顯示具有 python 標籤的文章。 顯示所有文章
顯示具有 python 標籤的文章。 顯示所有文章

2023年8月23日 星期三

Python + SAPGUI Script = 自動化重覆性事務性作業

Python + SAPGUI Script = 自動化重覆性事務性作業

reference

今天因為其他部門的資訊作業上的問題導致我這邊必需協助刪除庫存批屬性,雖然才八十幾筆,但也是讓我感到憤怒,憤怒可以提升技術力,所以我決定找方法自動化處理。

兩年前的今天,k-weiming寫下這個,真的是緣份,我的環境是python 3.11,需要的工具就是pywin32:

pip install pywin32

sap的設置也要確認,不過這取決於公司的管理策略,如果剛好沒被管控的話那就可以使用:

  1. Enable scripting:啟用scripting
  2. Notify when a script attaches to SAP GUI:對接上的時候是不是會跳出訊息,另一個也是一樣的設置

準備工作完成,再來就是要先取得執行的GUI的session:

import win32com.client def get_client(): #: 從com服務器中取得SAPGUI的物件 sap_gui_auto = win32com.client.GetObject('SAPGUI') if not type(sap_gui_auto) == win32com.client.CDispatch: print("SAPGUI is not running") return #: 取得SAPGUI的腳本引擎,所以接上的時候GUI會出現警示訊息 #: 這可以從gui設罝上取消 application = sap_gui_auto.GetScriptingEngine if not type(application) == win32com.client.CDispatch: print("SAPGUI is not running") return #: 這邊取得的是實際啟動的SAPGUI的物件 #: 以實際登入的使用者為條件 for conn in range(application.Children.Count): connection = application.Children(conn) #: 這邊取得的是登入的帳號所開啟的視窗(session) for sess in range(connection.Children.Count): session = connection.Children(sess) if session.Info.Transaction == "SESSION_MANAGER": return session

接下來的動作我們可以先利用錄製的方式來取得腳本,畢竟欄位框那麼多,即使看著官方的文件也很難一個一個找出來,最好的方法就是讓系統幫我們找出來。

開啟錄製程式:

錄製程式開啟之後如下:

中間那個圓點就是腳本錄製,腳本的保存地方則是根據你的『Save To』而定。

錄製完成之後在指定資料夾就會有vbstxt檔,打開txt檔,裡面會有你的腳本資訊:

紅框處的部份就是操作過程的相關物件的定位資訊以及操作的動作,像是按下去或是選擇某個分頁之類的。

以我自己的作業目標為例,我是想要自動化刪除很多庫存批屬性資料,所以就可以這麼做:

if __name__ == '__main__': #: reference #: https://k-weiming.github.io/2021-08-23-sap-connection-python/ session = get_client() #: 到SAP的交易代碼輸入欄位輸入MSC2N session.findById("wnd[0]/tbar[0]/okcd").text = "/nMSC2N" #: 按下Enter session.findById("wnd[0]").sendVKey(0) #: 輸入物料編號 session.findById("wnd[0]/usr/subSUBSCR_BATCH_MASTER:SAPLCHRG:1111/subSUBSCR_HEADER:SAPLCHRG:1500/ctxtDFBATCH-MATNR").text = "CA09B2002-1Z01" #: 輸入工廠 session.findById("wnd[0]/usr/subSUBSCR_BATCH_MASTER:SAPLCHRG:1111/subSUBSCR_HEADER:SAPLCHRG:1500/ctxtDFBATCH-WERKS").text = "A101" #: 輸入批號 session.findById("wnd[0]/usr/subSUBSCR_BATCH_MASTER:SAPLCHRG:1111/subSUBSCR_HEADER:SAPLCHRG:1500/ctxtDFBATCH-CHARG").text = "A120069722" #: 按下Enter session.findById("wnd[0]").sendVKey(0) #: 選擇批屬性分頁 session.findById("wnd[0]/usr/subSUBSCR_BATCH_MASTER:SAPLCHRG:1111/subSUBSCR_TABSTRIP:SAPLCHRG:2000/tabsTS_BODY/tabpCLAS").select() #: 按下刪除 session.findById("wnd[0]/usr/subSUBSCR_BATCH_MASTER:SAPLCHRG:1111/subSUBSCR_TABSTRIP:SAPLCHRG:2000/tabsTS_BODY/tabpCLAS/ssubSUBSCR_BODY:SAPLCHRG:2300/btnCLASS_DEL").press()

我的範例中少了按下保存的動作,因為只是測試,所以就沒有錄保存的動作。

這接口很方便,因為我不需要去寫abap來做錄製,不過這種作法應該是比較適用於重覆性的事務性作業,又或者剛好沒有BAPI、RFC可以使用,那就可以用這種方式來搜集資料。

2023年2月20日 星期一

Flask 2.x實作開工

hackmd的book連結

在職修課期間有幾位網友問我會不會針對2.x的部份再寫些新的東西,當初真的是心有餘而力不足,現在其實也還是,因為我在玩天堂w...

不過如何總算是開始動手了,除了重寫所用到的extension之外,題目的部份也暫定是天堂w的血盟分鑽管理,如果這個題目走不下去的話就會調整成自己寫一個記帳網頁。

會選擇血盟分鑽的題目是因為,有在玩的就知道,每一筆交易都需要手續費,過於頻繁的交易會付出過多的手續費。如果可以月結的話,那也許無形成也可以省下一筆不少的費用,因為角色之間可以互相沖帳。

當然如果真的不行的話變成是一個記帳網頁也不錯,生活中很多開銷如果可以好好記錄,回頭可以看到不少亮點。雖然現在也已經有很多的手機app可以使用,而且有更好的服務,但是寫一個專屬於自己的工具那種感覺是很不錯的。

最重要的當然就是我們可以從實作中瞭解到Flask帶給我們的便利性,然後從一個lib延伸學習到更多的lib。舉例來說,flask就單純的是flask,你想要跟資料庫對接也許就會考慮使用flask-sqlalchemy,使用flask-sqlalchemy也許你會採用orm,那你就學到orm於sqlalchemy的整合應用。

總之,還是那句話,flask是自由的,自由的意思就是很多事情你要自己處理,而不是由flask這個框架幫你決定一切。

PS: 目前已重置完成flask-sqlalchemy、flask-login


2020年10月27日 星期二

Algorithm_Sort_Quick sort

Algorithm_Sort_Quick sort

tags: Algorithms Sort python
import random
def quick_sort(input_array, is_random=False):
    """Quick Sort
    parameters:
        input_array: 輸入的未排序陣列
        is_random: 是否隨機取pivot point
             預設情況下是取陣列最後一個元素做pivot point      
    
    在worst case的情況下(資料已排序過),如果沒有亂數取pivot無法成功執行
    """
    # print(f'input array: {input_array}')
    array_size = len(input_array)
    if array_size < 2:
        # 不處理
        return input_array
    
    # 取得pivot
    if is_random:
        pivot_point = random.randint(0, array_size - 1)
    else:
        pivot_point = array_size - 1
    
    pivot = input_array.pop(pivot_point)
    # print(f'pivot: {pivot}, pivot_point: {pivot_point}')
        
    # 寫入比pivot大以及小的資料
    greater_pivot = [element for element in input_array if element > pivot]
    # print(f'greater_pivot: {greater_pivot}')
    lesser_pivot = [element for element in input_array if element < pivot]
    # print(f'lesser_pivot: {lesser_pivot}')
    # print(f'total: {lesser_pivot} + {[pivot]} + {greater_pivot}')
    # 利用python list相加的特性做遞迴處理
    return quick_sort(lesser_pivot, is_random) + [pivot] + quick_sort(greater_pivot, is_random)

下面給出五個資料集在不同資料量下的結果,其中

  • l1: 已經排序好的,由小到大
  • l2: 已經排序好的,但是逆排序,由大到小
  • l3: 亂數生成
  • l4: 亂數生成
  • l5: 亂數生成
Quick Sort(no random) 20000 40000 60000 80000 100000
l1
l2
l3 89 ms 172 ms 238 ms 438 ms 434 ms
l4 130 ms 153 ms 248 ms 330 ms 416 ms
l5 81 ms 154 ms 256 ms 335 ms 439 ms
Quick Sort(random) 20000 40000 60000 80000 100000
l1 154 ms 219 ms 315 ms 468 ms 583 ms
l2 107 ms 225 ms 333 ms 475 ms 539 ms
l3 103 ms 241 ms 370 ms 526 ms 592 ms
l4 121 ms 246 ms 349 ms 511 ms 656 ms
l5 111 ms 226 ms 403 ms 503 ms 627 ms

Algorithm_Sort_Heap srot

Algorithm_Sort_Heap srot

tags: Algorithms Sort python
def heapify(input_array, index, bound):
    """將資料結構化為heap
    parameters:
        input_array: 輸入的未排序陣列
        index: 要排序的root
        bound: 比較邊界,如果某個root的左、右索引已超過陣列長度就不處理
        
    假設,root的index=n,那其節點索引為
        root: n
        root_left_node: 2n + 1
        root_right_node: 2n + 2
        
    
    舉例來說,[0, 1, 2, 3, 4, 5, 6]
          0
        1   2
       3 4 5 6
    0為root,左邊節點為1,右邊則為2
    1為root,左邊節點為3,右邊則為4
    
    heap sort的一半是最後的節點,因此只會比較到陣列一半取上高斯
    其餘皆會是最終節點,不需要再比較,所以3456本身是不需要比較的
    我們會從2開始比,也就是7/2 = 3 - 1 = 2(因為python的索引初始值為0)
        
    要記得,python的索引是由0開始
    
    root永遠是三個節點中最大的(或最小的),但節點不考慮大小      
      5            5
    1   2        2   1 左邊兩種情況都可以
    """
    # 先預設最大值為root index
    largest_idx = index
    # print(f'init_array: {input_array}')
    # print(f'init_larget_idx: {largest_idx}, value: {input_array[largest_idx]}')
    left_idx = index * 2 + 1
    # print(f'left_idx: {left_idx}')
    right_idx = index * 2 + 2 
    # print(f'right_idx: {right_idx}')
    # 下面的比較,因為我們只關心root是最大的,因此不需要再比較左右兩邊
    # 注意到與陣列長度的比較條件必需在前面,不然會引發索引異常
    if  left_idx < bound  and input_array[largest_idx] < input_array[left_idx]:
        largest_idx = left_idx
        # print(f'l_larget_idx: {largest_idx}, value: {input_array[largest_idx]}')
    
    if right_idx < bound  and input_array[largest_idx] < input_array[right_idx]:        
        largest_idx = right_idx
        # print(f'r_larget_idx: {largest_idx}, value: {input_array[largest_idx]}')
    
    # 比較之後,如果最大值的異動,那就調整array的順序
    if largest_idx != index:        
        input_array[index], input_array[largest_idx] = input_array[largest_idx], input_array[index]
        # print(f'new array: {input_array}')
        # 遞迴處理heapify
        heapify(input_array, largest_idx, bound)
        
 def heap_sort(input_array):
    """heap sort"""
    array_size = len(input_array)
    # heap sort會從底層開始排序上來
    # 而且並不會所有的元素都排序,因為n//2 + 1之後的節點都是最終節點
    # 只有n//2 + 1之前是root,因此只做這邊的迴圈比較
    # 這邊只是做heap的結構
    for i in range(array_size // 2 - 1, -1, -1):        
        heapify(input_array, i, array_size)
    
    # 處理排序
    # 我們知道,經過heap過的結構,最大的一定是在最上面(假設是max heap)
    # 因此,實作上會將root的部份跟最後的結點做交換,再將n-1做一次heap
    for i in range(array_size - 1, 0, -1):
        # 每次heap之後都做頭尾的交換
        # print(f'heap_sort: {input_array}')
        input_array[0], input_array[i] = input_array[i], input_array[0]
        # 特別注意到,每一次迭代都會減少一個,因此這邊給的是迴圈中的i,而不是array_size
        heapify(input_array, 0, i)
    return input_array   

下面給出五個資料集在不同資料量下的結果,其中

  • l1: 已經排序好的,由小到大
  • l2: 已經排序好的,但是逆排序,由大到小
  • l3: 亂數生成
  • l4: 亂數生成
  • l5: 亂數生成
Heap Sort 20000 40000 60000 80000 100000
l1 416 ms 606 ms 864 ms 1.23 s 1.51 s
l2 255 ms 721 ms 876 ms 1.15 s 1.64 s
l3 265 ms 603 ms 923 ms 1.23 s 1.8 s
l4 297 ms 562 ms 946 m 1.16 s 1.59 s
l5 279 ms 553 ms 966 ms 1.22 s 1.66 s

2020年10月26日 星期一

Algorithm_Sort_Merge sort

Algorithm_Sort_Merge sort

tags: Algorithms Sort python
def merge(left_a, right_a):
    """合併兩個陣列"""
    def _merge():        
        while left_a and right_a:
#             print('_merge:', left_a, right_a)
            # 比較兩個陣列的第一個索引的值那一個比較小就優先回傳    
            yield (left_a if left_a[0] < right_a[0] else right_a).pop(0)
        yield from left_a
        yield from right_a
    return list(_merge())
        
def divide(input_array):
    """分割陣列"""
    # 取上高斯點做為中點切掉兩個陣列
    divide_point = len(input_array) // 2
    left_a = input_array[: divide_point]
    right_a = input_array[divide_point: ]
    return left_a, right_a  
    
def merge_sort(input_array):
    """利用遞迴呼執行分割、合併、排序"""
#     print('merge_sort_array:', input_array)
    if len(input_array) == 1:
#         print('return array:', input_array)
        return input_array
    left_a, right_a = divide(input_array)
    return merge(merge_sort(left_a), merge_sort(right_a))
    

下面給出五個資料集在不同資料量下的結果,其中

  • l1: 已經排序好的,由小到大
  • l2: 已經排序好的,但是逆排序,由大到小
  • l3: 亂數生成
  • l4: 亂數生成
  • l5: 亂數生成
Merge Sort 20000 40000 60000 80000 100000
l1 261 ms 585 ms 805 ms 1.35 s 2.18 s
l2 199 ms 528 ms 880 ms 1.47 s 2.09 s
l3 294 ms 844 ms 1.52 s 2.37 s 3.64 s
l4 282 ms 761 ms 1.45 s 2.39 s 3.42 s
l5 331 ms 735 ms 1.5 s 2.27 s 3.48 s

Algorithm_Sort_Insrtion sort

Algorithm_Sort_Insrtion sort

tags: Algorithms Sort python
def insertion_sort(input_array):
    """insertion sort
    input_array: list
                 list內的元素需為數值
                 
    return 
        由小大到排序後的list
    """
    if len(input_array) <= 1:
        return
    
    # 從第2個element開始處理
    for idx, idx_value in enumerate(input_array[1: ]):        
        # 內圈迴圈使用        
        inner_idx = idx        
        while inner_idx >= 0 and idx_value < input_array[inner_idx]:            
            # 這邊的idx剛好就是j-1的索引
            # 判斷兩值的大小,如果索引值(右)比前一個元素(左)小
            # 那就將前一個元素複製到索引處                                       
            input_array[inner_idx + 1] = input_array[inner_idx]                
            inner_idx -= 1            
        
        # 當兩個索引不同代表異動排序過
        if inner_idx != idx:                
            # 將原始索引值寫入最終停止的索引處            
            input_array[inner_idx + 1] = idx_value
    
    return input_array

下面給出五個資料集在不同資料量下的結果,其中

  • l1: 已經排序好的,由小到大
  • l2: 已經排序好的,但是逆排序,由大到小
  • l3: 亂數生成
  • l4: 亂數生成
  • l5: 亂數生成
Insertion sort 20000 40000 60000 80000 100000
l1 7ms 16ms 25ms 22ms 37ms
l2 1min 31s 5min 47s 13min 27s 24min 18s 37min 34s
l3 43.4s 2min 59s 6min 52s 12min 27s 20min 6s
l4 43.1s 3min 2s 6min 50s 13min 1s 19min 52s
l5 42.9s 2min 57s 6min 51s 12min 42s 19min 47s

2020年1月16日 星期四

Coroutines-and-Tasks翻譯

Coroutines and Tasks(翻譯)

這章節概略說明用於協程、任務的高階asyncio的APIs。

Coroutines

使用async/await語法宣告的協程是寫asyncio應用程式最好的方法。舉例來說,下面片段程式碼(需要Python 3.7+)列印"hello",等待1秒,然後列印"world":

>>> import asyncio

>>> async def main():
...     print('hello')
...     await asyncio.sleep(1)
...     print('world')

>>> asyncio.run(main())
hello
world

注意到,單純的呼叫協程並不會調度它被執行:

>>> main()
<coroutine object main at 0x1053bb7c8>

要實際的執行協程,asyncio提供三種機制:

  • 使用函數asyncio.run()來執行頂層的入口"main()"函數(見上範例)
  • 協程中的等待。下面程式碼片段會在等待1秒之後列印"hello",,然後在等待另外2秒之後列印"world":
import asyncio
import time

async def say_after(delay, what):
    await asyncio.sleep(delay)
    print(what)

async def main():
    print(f"started at {time.strftime('%X')}")

    await say_after(1, 'hello')
    await say_after(2, 'world')

    print(f"finished at {time.strftime('%X')}")

asyncio.run(main())

預期輸出如下:

started at 17:13:52
hello
world
finished at 17:13:55

讓我們調整上面範例,並且同時執行兩個協程say_after

async def main():
    task1 = asyncio.create_task(
        say_after(1, 'hello'))

    task2 = asyncio.create_task(
        say_after(2, 'world'))

    print(f"started at {time.strftime('%X')}")

    # Wait until both tasks are completed (should take
    # around 2 seconds.)
    await task1
    await task2

    print(f"finished at {time.strftime('%X')}")

注意到,預期輸出現在顯示出這程式碼片段比之前還要快一秒:

started at 17:14:32
hello
world
finished at 17:14:34

Awaitables

我們說,如果一個物件可以用於await表達示中,那它就是一個可等待物件。許多asyncio APIs被設計為接受等待。

有三種類型的可等待物件:coroutines, Tasks, and Futures

Coroutines

Python協程是可等待的,因此可以從其它協程中等待:

import asyncio

async def nested():
    return 42

async def main():
    # Nothing happens if we just call "nested()".
    # A coroutine object is created but not awaited,
    # so it *won't run at all*.
    nested()

    # Let's do it differently now and await it:
    print(await nested())  # will print "42".

asyncio.run(main())

重要:這文件中的術語"協程"可以用於兩個緊密相關的概念:

  • 協程函數:一個async def函數
  • 協程物件:透過呼叫協程函數回傳的物件

asyncio亦支援傳統基於生程的協程。

Tasks

Tasks用於並行調度協程。

當協程被包裝到具有像是asyncio.create_task()函數的任務內時,這個協程會很快的自動調度來執行:

import asyncio

async def nested():
    return 42

async def main():
    # Schedule nested() to run soon concurrently
    # with "main()".
    task = asyncio.create_task(nested())

    # "task" can now be used to cancel "nested()", or
    # can simply be awaited to wait until it is complete:
    await task

asyncio.run(main())

Futures

Future是一個特別的低階可等待物件,代表一個非同步(異步)操作的事件結果。

當Future物件為awaited的時候,意味著協程會等待,一直到Future在其它位置被解析。

asyncio中的Future物件需要允許基於回呼程式碼與async/await一起使用。

一般來說,不需要在應用程式級別程式碼中建立Future物件。

Future objects, sometimes exposed by libraries and some asyncio APIs, can be awaited:

Future物件(有些時候會由套件或asyncio APIs公開)可以是awaited:

    await function_that_returns_a_future_object()

    # this is also valid:
    await asyncio.gather(
        function_that_returns_a_future_object(),
        some_python_coroutine()
    )

一個回傳Future物件的低階函數範例為loop.run_in_executor()

Running an asyncio Program

  • asyncio.run(coro, *, debug=False)
    執行協程並回傳結果。

    這個函數執行一個傳遞過來的協程,並負責管理asyncio事件迴圈以及完成非同步生成器。

    當另一個asyncio事件迴圈在相同thread(執行緒)執行的時候,這函數無法被呼叫。

    如果debug=True,那事件迴圈會在除錯模式中執行。

    這函數總是建立一個新的事件迴圈,並在最後關閉它。它應該用做為asyncio程式的主要入口點,而且理想情況下應該只調用一次。
    範例:

    async def main():
        await asyncio.sleep(1)
        print('hello')
    
    asyncio.run(main())
    

    New in version 3.7.

    原始碼可以在Lib/asyncio/runners.py找到

Creating Tasks

  • asyncio.create_task(coro, *, name=None)
    將協程包裝到Task中,並調度它的執行。回傳Task物件。

    如果它的name不為None,則使用Task.set_name()來設置task名稱。

    task在get_running_loop()回傳的迴圈中執行,如果當前的執行緒(thread)沒有正在執行中的迴圈,那就拋出RuntimeError

    這個函數在Python 3.7中被加入。Python 3.7之前的版本,以低階函數asyncio.ensure_future()替代:

    async def coro():
    ...
    
    # In Python 3.7+
    task = asyncio.create_task(coro())
    ...
    
    # This works in all Python versions but is less readable
    task = asyncio.ensure_future(coro())
    ...
    

    New in version 3.7.

    Changed in version 3.8:增加參數name

Sleeping

  • coroutine asyncio.sleep(delay, result=None, *, loop=None)
    延遲秒數的阻塞。

    如果提供結果,那在協程完成的時候會回傳給調用者。

    sleep()總是暫停當前task,允許其它tasks執行。

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

    協程範例,每秒顯示當前日期5秒:

    import asyncio
    import datetime
    
    async def display_date():
        loop = asyncio.get_running_loop()
        end_time = loop.time() + 5.0
        while True:
            print(datetime.datetime.now())
            if (loop.time() + 1.0) >= end_time:
                break
            await asyncio.sleep(1)
    
    asyncio.run(display_date())
    

Running Tasks Concurrently

  • awaitable asyncio.gather(*aws, loop=None, return_exceptions=False)
    在asw序列中同時執行awaitable物件

    如果任一個awaitable在aws中是一個協程,那就視為Task自動調度。

    如果所有的awaitables都成功的完成,彙總清單(list)。其順序相對應於aws中的awaitables的順序。

    如果return_exceptions為False(預設),那第一個拋出的異常會立即被傳送到在gather()上等待的task。

    如果return_exceptions為True,那異常會跟成功的結果一樣的處理方法,並彙總於結果清單。

    如果gather()被取消,所有提交的awaitables(未完成的)也會一併被取消。

    如果任一來自aws序列的Task或Future被取消,那就將它視為引發CancelledError,這種情況下不會取消gather()的呼叫。

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

    範例:

    import asyncio
    
    async def factorial(name, number):
        f = 1
        for i in range(2, number + 1):
            print(f"Task {name}: Compute factorial({i})...")
            await asyncio.sleep(1)
            f *= i
        print(f"Task {name}: factorial({number}) = {f}")
        
    async def main():
        # Schedule three calls *concurrently*:
        await asyncio.gather(
            factorial("A", 2),
            factorial("B", 3),
            factorial("C", 4),
        )
        
    asyncio.run(main())
    
    # Expected output:
    #
    #     Task A: Compute factorial(2)...
    #     Task B: Compute factorial(2)...
    #     Task C: Compute factorial(2)...
    #     Task A: factorial(2) = 2
    #     Task B: Compute factorial(3)...
    #     Task C: Compute factorial(3)...
    #     Task B: factorial(3) = 6
    #     Task C: Compute factorial(4)...
    #     Task C: factorial(4) = 24
    

    Changed in version 3.7:如果gather本身被取消,那取消傳播與return_exceptions無關。

Shielding From Cancellation

  • awaitable asyncio.shield(aw, *, loop=None)
    保護一個awaitable物件不被取消

    如果aw是一個協程,那就視為Task自動調度。

    語法:

    res = await shield(something())
    

    等價於:

    res = await something()
    

    除非包含它的協程被取消,那在something()中執行的Task就不會被取消。從something()的觀點來看,取消沒有發生。儘管其調用者仍然被取消,但"await"表達式依然會拋出CancelledError

    如果something()被其它方式取消(由自身內部),那會同時取消shield()

    如果希望完全忽略取消(不建議),那函數shield()應該結合try/except,如下:

    try:
        res = await shield(something())
    except CancelledError:
        res = None
    

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

Timeouts

  • coroutine asyncio.wait_for(aw, timeout, *, loop=None)
    等待aw awaitable以超時完成。

    如果aw是一個協程,那就視為Task自動調度。

    timeout可以是None或是float或是int做為等待的秒數。如果timeoutNone,那阻塞會一直到future完成。

    避免task取消,可以將task包裝在shield()

    函數會等待,一直到future確實的取消,因此總等待時間也許會超過timeout

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

    範例:

    async def eternity():
        # Sleep for one hour
        await asyncio.sleep(3600)
        print('yay!')
    async def main():
        # Wait for at most 1 second
        try:
            await asyncio.wait_for(eternity(), timeout=1.0)
        except asyncio.TimeoutError:
            print('timeout!') 
            
    asyncio.run(main())
    
    # Expected output:
    #
    #     timeout!    
    

    Changed in version 3.7:當aw因為timeout而取消,則wait_for等待aw被取消。之前的版本會立即拋出asyncio.TimeoutError

Waiting Primitives

  • coroutine asyncio.wait(aws, *, loop=None, timeout=None, return_when=ALL_COMPLETED)
    同時在aws上執行awaitable物件,並阻塞直到return_when指定的條件為止。

    回傳兩組Tasks/Futures:(done, pending)

    用法:

    done, pending = await asyncio.wait(aws)
    

    timeout(float or ing),如果指定,可以用來控制回傳前等待的最大秒數。

    注意,這個函數並不會拋出asyncio.TimeoutError。發生timeout未完成的Futures或Tasks僅在第二組中回傳。

    return_when指示這函數何時應該回傳。它必須為下面項目之一:

    Constant Description
    FIRST_COMPLETED The function will return when any future finishes or is cancelled.
    FIRST_EXCEPTION The function will return when any future finishes by raising an exception. If no future raises an exception then it is equivalent to ALL_COMPLETED.
    ALL_COMPLETED The function will return when all futures finish or are cancelled.

    不像wait_for(),當發生timeout的時候,wait並不會取消futures。

    Python 3.8之後不建議使用:如果任一個awaitable在aws中是一個協程,那就視為Task自動調度。不建議將協程物件直接傳遞給wait(),因為它會導致混亂的行為。

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

    Note:wait()自動的將協程視為Task調度,稍後會回傳隱式建立的Task物件(done, pending)。因此,下面程式碼將無法如預期般執行:

    async def foo():
        return 42
    
    coro = foo()
    done, pending = await asyncio.wait({coro})
    
    if coro in done:
        # This branch will never be run!
    

    修復上面片段程式碼如下:

    async def foo():
        return 42
    
    task = asyncio.create_task(foo())
    done, pending = await asyncio.wait({task})
    
    if task in done:
        # Everything will work as expected now.
    

    Python 3.8之後不建議使用:不建議直接將協程物件直接傳遞給wait()

  • asyncio.as_completed(aws, *, loop=None, timeout=None)
    同時在aws集中執行awaitable物件。回傳一個Future物件的迭代器。回傳的每一個Future物件表示剩餘的awaitables集的最早結果。

    如果在所有的Futures完成之前發生timeout,則拋出asyncio.TimeoutError

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

    範例:

    for f in as_completed(aws):
        earliest_result = await f
        # ... 
    

Scheduling From Other Threads

  • asyncio.run_coroutine_threadsafe(coro, loop)
    提交協程到給定的事件迴圈。線程安全。

    回傳concurrent.futures.Future以等待其它OS線程的結果。

    這個函數意思是,從不同於執行事件循環的OS線程中調用。範例:

    # Create a coroutine
    coro = asyncio.sleep(1, result=3)
    
    # Submit the coroutine to a given loop
    future = asyncio.run_coroutine_threadsafe(coro, loop)
    
    # Wait for the result with an optional timeout argument
    assert future.result(timeout) == 3   
    

    如果協程中拋出異常,那將通知回傳的Future。這也可以用來取消事件迴圈中的task:

    try:
        result = future.result(timeout)
    except asyncio.TimeoutError:
        print('The coroutine took too long, cancelling the task...')
        future.cancel()
    except Exception as exc:
        print(f'The coroutine raised an exception: {exc!r}')
    else:
        print(f'The coroutine returned: {result!r}')
    

    參考concurrency and multithreading

    不像其它asyncio函數,這個函數要求顯式的傳遞參數loop

    New in version 3.5.1.

Introspection

  • asyncio.current_task(loop=None)
    回傳目前執行中的Task實例,如果沒有Task在執行,則回傳None

    如果loop is None,那就用get_running_loop()來取得當前迴圈。

    New in version 3.7.

  • asyncio.all_tasks(loop=None)
    回傳由迴圈執行的一組未完成的Task物件。

    如果loop is None,那就用get_running_loop()來取得當前迴圈。

    New in version 3.7.

Task Object

  • class asyncio.Task(coro, *, loop=None, name=None)
    執行Python協程的Future-like物件。非線程安全。

    Tasks用於在事件迴圈中執行協程。如果協程等待Future,那Task會暫停協程的執行並等待Future完成。當Future完成,將繼續執行包裝的協程。

    事件迴圈使用協同調度:事件迴圈一次執行一個Task。當Task等待協程完成的時候,事件迴圈會執行其它Task、回呼或執行IO操作。

    使用高階函數asyncio.create_task(),或低階函數loop.create_task()ensure_future()建立Tasks。不建議手動實例化Tasks。

    要取消執行中的Task,請使用cancal()方法。呼叫它將導致Task向包裝的協程中拋出CancelledError異常。如果在取消過程中某個Future物件正在等待協程,那這個Future物件將被取消。

    calcelled()可以用來確認Task是否已經被取消。如果包裝的協程沒有抑制CancelledError異常並確實的取消,那就回傳True。

    asyncio.Task繼承Future所有的APIs,除了Future.set_result()Future.set_exception()

    Tasks支援contextvars模組。建立Task的時候,它會複製當前的上下文,然後在複製的上下文中執行協程。

    Changed in version 3.7: 增加支援contextvars模組。

    Changed in version 3.8: 增加參數name

    Python 3.8之後不建議使用,將在Python 3.10的時候移除:參數loop

cancel()

請求取消Task。

這安排在事件迴圈的下一個週期將CancelledError異常拋出到包裝的協程中。

然後,協程有機會清楚,或甚至拒絕請求,透過利用try except CancelledError finally區塊來抑制異常。

然而,不同於Future.cancel()Task.cancel()並不保證Task一定被取消,儘管完全抑制取消並不常見,並且主動阻止。

下面範例說明協程如何攔截取消請求:

async def cancel_me():
    print('cancel_me(): before sleep')

    try:
        # Wait for 1 hour
        await asyncio.sleep(3600)
    except asyncio.CancelledError:
        print('cancel_me(): cancel sleep')
        raise
    finally:
        print('cancel_me(): after sleep')

async def main():
    # Create a "cancel_me" Task
    task = asyncio.create_task(cancel_me())

    # Wait for 1 second
    await asyncio.sleep(1)

    task.cancel()
    try:
        await task
    except asyncio.CancelledError:
        print("main(): cancel_me is cancelled now")

asyncio.run(main())

# Expected output:
#
#     cancel_me(): before sleep
#     cancel_me(): cancel sleep
#     cancel_me(): after sleep
#     main(): cancel_me is cancelled now

cancelled()

如果Task已經取消,則回傳True。

當使用calcel()請求取消,Task將被取消,並且包裝的協程將CancelledError異常傳播到Task中。

done()

如果Task完成,則回傳True。

當包裝的協程回傳值,拋出異常、或Task被取消,Task就算完成了。

result()

回傳Task的結果。

如果Task完成,那就回傳包裝的協程的結果(或者如何協程拋出異常,則重新拋出該異常)

如果Task已經被取消,這個方法將拋出CancelledError異常。

如果Task的結果尚不可用,這個方法會拋出InvalidStateError異常。

exception()

回傳Task的異常。

如果包裝的協程拋出異常,則回傳該異常。如果包裝的協程正常回傳,這個方法將回傳None

如果Task已經被取消,這個方法將拋出CancelledError異常。

如果Task尚未完成,這個方法會拋出InvalidStateError異常。

add_done_callback(callback, *, context=None)

增加一個callback function給Task完成的時候執行。

這個方法只能在基於低階回呼的程式碼中使用。

更多細節請參考Future.add_done_callback()

remove_done_callback(callback)

從callback清單中取消callback。

這個方法只能在基於低階回呼的程式碼中使用。

更多細節請參考Future.remove_done_callback()

get_stack(*, limit=None)

回傳該Task的堆疊框的清單(list of stack frames)。

如果包裝的協程還沒有完成,那就回傳掛起它的堆疊。如果協程已經成功地完成或被取消,那就回傳空清單(empty list)。如果協程被異常終止,那就回傳回溯框清單(list of traceback frames)。

框(frames)總是依著舊到新的排序。

對懸置的協程,只會回傳一個堆疊框(stack frame)。

選項參數limit設置回傳的框(frame)的最大數量;預設情況下回傳所有可用的框(frame)。回傳的清單順序依回傳是堆疊或回溯而有所不同:回傳最新的堆疊,但回傳最舊的回溯。(這與回溯模組的行為匹配。)

列印該Task的堆棧或回溯。

對於get_stack()檢索到的frames,將生成與回溯模組類似的輸出。

get_coro()

回傳Task包裝的協程。

New in version 3.8.

get_name()

回傳Task的名稱。

如果沒有為Task明確的分配名稱,則預設asyncio Task實現在實例化過程中生成一個預設名稱。

New in version 3.8.

set_name(value)

設置Task的名稱。

參數value可以是任意物件,然後將它轉為字串。

在預設的Task實現中,名稱將在Task物件的repr()的輸出看的見。

New in version 3.8.

classmethod all_tasks(loop=None)

回傳事件迴圈的所有Tasks集合。

預設情況下,回傳當前事件迴圈的所有Tasks。如果loop is None,則使用函數get_event_loop()取得當前迴圈。

版本3.7之後不建議使用,將在版本3.9中移除:不要將此做為task方法調用。改用asyncio.all_tasks()

classmethod current_task(loop=None)

回傳當前的task或None。

如果loop is None,則使用函數get_event_loop()取得當前迴圈。

版本3.7之後不建議使用,將在版本3.9中移除:不要將此做為task方法調用。改用asyncio.current_task()

Generator-based Coroutines

2019年12月25日 星期三

Release-Highlights-for-scikit-learn-022翻譯

Release Highlights for scikit-learn 0.22(翻譯)

原文連結

我們很高興宣佈scikit-learn 0.22的發布,其中包含許多bug的修復以及新功能!下面我們詳細說明這版本的一些主要功能。關於完整的修正清單,請參閱發行說明。

使用pip安裝最新版本:

pip install --upgrade scikit-learn

或使用conda

conda install scikit-learn

New plotting API

新的plotting API可用於建立可視化。這個新的API允許在不涉及任何重新計算情況下快速調整繪圖的視覺效果。也可以在同一個圖(figure)上加入不同圖形。下面範例說明plot_roc_curve,但支援其它繪圖工具,像是plot_partial_dependenceplot_precision_recall_curveplot_confusion_matrix。關於這個API可參閱使用者指南

from sklearn.model_selection import train_test_split
from sklearn.svm import SVC
from sklearn.metrics import plot_roc_curve
from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import make_classification
import matplotlib.pyplot as plt

# 生成測試資料
X, y = make_classification(random_state=0)
# 資料集分為訓練與測試資料集
X_train, X_test, y_train, y_test = train_test_split(X, y, random_state=42)

svc = SVC(random_state=42)
svc.fit(X_train, y_train)

rfc = RandomForestClassifier(random_state=42)
rfc.fit(X_train, y_train)

# 利用plot_roc_curve計算roc
svc_disp = plot_roc_curve(svc, X_test, y_test)
rfc_disp = plot_roc_curve(rfc, X_test, y_test, ax=svc_disp.ax_)
rfc_disp.figure_.suptitle("ROC curve comparison")

plt.show()

plot_roc_curve很快速的幫我們計算出ROC曲線,並且將兩個不同模型的結果繪製在同一張圖(figure)上,這對我們瞭解不同模型的比較非常實用。

Stacking Classifier and Regressor

StackingClassifierStackingRegressor允許你有一個帶有最終分類器或回歸器的估計器堆疊。堆疊概括包括堆疊各別估計器的輸出,並使用一個分類器來完全最終的預測。堆疊允許使用每個各別估計器的強度,透過使用它們的輸出做為最終估計器的輸入。基礎估計器在完整的X上擬合,而最終估計器的訓練則使用cross_val_predict對基礎估計器做交叉驗證的預測。

更多資訊可參閱使用者指南

from sklearn.datasets import load_iris
from sklearn.svm import LinearSVC
from sklearn.linear_model import LogisticRegression
from sklearn.preprocessing import StandardScaler
from sklearn.pipeline import make_pipeline
from sklearn.ensemble import StackingClassifier
from sklearn.model_selection import train_test_split

X, y = load_iris(return_X_y=True)

# 利用list加入多個分類器,分類器中一樣可以設置pipeline
estimators = [
    ('rf', RandomForestClassifier(n_estimators=10, random_state=42)),
    ('svr', make_pipeline(StandardScaler(),
                          LinearSVC(random_state=42)))
]

# 實作StackingClassifier,指定最終分類器
clf = StackingClassifier(
    estimators=estimators, final_estimator=LogisticRegression()
)
# 分割資料集
X_train, X_test, y_train, y_test = train_test_split(
    X, y, stratify=y, random_state=42
)

clf.fit(X_train, y_train).score(X_test, y_test)

Permutation-based feature importance

對任一個擬合的估計器,inspection.permutation_importance可以用來得到每一個特徵的重要性的估計。

from sklearn.ensemble import RandomForestClassifier
from sklearn.inspection import permutation_importance

# 取得測試資料
X, y = make_classification(random_state=0, n_features=5, n_informative=3)

# 實作estimator
rf = RandomForestClassifier(random_state=0).fit(X, y)
# 實作permutation_importance
result = permutation_importance(rf, X, y, n_repeats=10, random_state=0,
                                n_jobs=-1)

# 定義圖表 
fig, ax = plt.subplots()
# 取得重要性均值排序
sorted_idx = result.importances_mean.argsort()
# 繪製圖表
ax.boxplot(result.importances[sorted_idx].T,
           vert=False, labels=range(X.shape[1]))
# 細部調整圖表
ax.set_title("Permutation Importance of each feature")
ax.set_ylabel("Features")
fig.tight_layout()
plt.show()

Native support for missing values for gradient boosting

ensemble.HistGradientBoostingClassifierensemble.HistGradientBoostingRegressor現在對缺失值(NaNs)有本機支援。這意味著在訓練或預測的時候不需要再做缺值插補。

from sklearn.experimental import enable_hist_gradient_boosting  # noqa
from sklearn.ensemble import HistGradientBoostingClassifier
import numpy as np

# 故意放一個缺值
X = np.array([0, 1, 2, np.nan]).reshape(-1, 1)
y = [0, 0, 1, 1]

gbdt = HistGradientBoostingClassifier(min_samples_leaf=1).fit(X, y)
print(gbdt.predict(X))

Precomputed sparse nearest neighbors graph

大多基於最近鄰圖的估計器現在接受預先計算的稀疏圖做為輸入,以便重覆使用同一圖進行多個估計器擬合。要在pipeline中使用這個功能,可以使用記憶體參數,以及兩個新的轉換器neighbors.KNeighborsTransformerneighbors.RadiusNeighborsTransformer。預先計算也可以透過自定義估計器來執行做為替代的實現,像是近似最近鄰方法。更多資訊可參閱使用者指南

from tempfile import TemporaryDirectory
from sklearn.neighbors import KNeighborsTransformer
from sklearn.manifold import Isomap
from sklearn.pipeline import make_pipeline

X, y = make_classification(random_state=0)

with TemporaryDirectory(prefix="sklearn_cache_") as tmpdir:
    estimator = make_pipeline(
        KNeighborsTransformer(n_neighbors=10, mode='distance'),
        Isomap(n_neighbors=10, metric='precomputed'),
        memory=tmpdir)
    estimator.fit(X)

    # We can decrease the number of neighbors and the graph will not be
    # recomputed.
    estimator.set_params(isomap__n_neighbors=5)
    estimator.fit(X)

KNN Based Imputation

我們現在支援使用KNN來做缺值插補。

每個樣本的缺失值都用訓練集中發現的n_neighbors個最近鄰的平均值來插補。如果兩個不缺值的特徵都很接近,則兩個樣本是接近的。預設情況下,使用支援缺失值nan_euclidean_distances的歐氏距離度量來找出最近鄰。

更多資訊可參閱使用者指南

import numpy as np
from sklearn.impute import KNNImputer

# 故意插入空值
X = [[1, 2, np.nan], [3, 4, 3], [np.nan, 6, 5], [8, 8, 7]]
imputer = KNNImputer(n_neighbors=2)
print(imputer.fit_transform(X))

Tree pruning

在樹(tree)建構完成之後,現在可以修剪多數基於樹的估計器。修剪法是基於最小化成本複雜。更多資訊可參閱使用者指南

X, y = make_classification(random_state=0)

rf = RandomForestClassifier(random_state=0, ccp_alpha=0).fit(X, y)
print("Average number of nodes without pruning {:.1f}".format(
    np.mean([e.tree_.node_count for e in rf.estimators_])))

rf = RandomForestClassifier(random_state=0, ccp_alpha=0.05).fit(X, y)
print("Average number of nodes with pruning {:.1f}".format(
    np.mean([e.tree_.node_count for e in rf.estimators_])))

Retrieve dataframes from OpenML

datasets.fetch_openml現在可以回傳pandas dataframe,從而正確處理帶有異質資料的資料集。

from sklearn.datasets import fetch_openml

titanic = fetch_openml('titanic', version=1, as_frame=True)
print(titanic.data.head()[['pclass', 'embarked']])

Checking scikit-learn compatibility of an estimator

開發人員可以使用check_estimator檢查他們的scikit-learn相容估計器的相容性。例如,check_estimator(LinearSVC)通過。

我們現在提供一個pytest特定的裝飾器,它允許pytest獨立執行所有檢查並報告失敗的檢查。

from sklearn.linear_model import LogisticRegression
from sklearn.tree import DecisionTreeRegressor
from sklearn.utils.estimator_checks import parametrize_with_checks


@parametrize_with_checks([LogisticRegression, DecisionTreeRegressor])
def test_sklearn_compatible_estimator(estimator, check):
    check(estimator)

ROC AUC now supports multiclass classification

函數roc_auc_score也可以應用在多類別的分類。目前支援兩種平均策略:一對一演算法計算成對的ROC AUC分數的平均值,一對多演算法計算每個類別相對於其它類別的ROC AUC分數的平均值。這兩種情況下,多類別的ROC AUC分數都是根據模型從樣本屬於特定類別的機率估計計算而來。OvO與OvR演算法皆支援均勻加權(average='macro'),與依盛行率加權(average='weighted')。

更多資訊可參閱使用者指南

from sklearn.datasets import make_classification
from sklearn.svm import SVC
from sklearn.metrics import roc_auc_score

X, y = make_classification(n_classes=4, n_informative=16)
clf = SVC(decision_function_shape='ovo', probability=True).fit(X, y)
print(roc_auc_score(y, clf.predict_proba(X), multi_class='ovo'))

2019年8月7日 星期三

Socket-Programming-HOWTO翻譯

Socket Programming HOWTO(翻譯)

官方文件

Author: Gordon McMillan

Abstract
Sockets的應用無所不在,但卻是最被誤會的技術之一。這只是關於sockets的概述。它並不是一個真正的教程,你仍然需要下點工夫。它並沒有涵蓋的很精確(絕大部份如此),但我希望可以給你足夠的背景來正確的使用它。

Sockets

我只想談談INET(i.e. IPv4)sockets,但它們佔了將近99%使用中的sockets。並且我將只談STREAM(i.e. TCP)sockets,除了你真的知道你在做什麼(這種情況下這個HOWTO並不適合你),你將從STREAM socket得到比其它模型更好的做法及效能。

理解這些事的部份麻煩在於,socket可以代表許多不同細微的東西,這取決於應用上下文。首先,我們先區分client端的socket-一個會話的端點,與server端的socket,它更像是一個交換機的操作。client應用程式(如,瀏灠器)單純的使用client sockets;而與它談的web server則同時使用著server sockets與client sockets。

History

各種形式的IPC中,sockets是目前為止最受歡迎的。在任何給定平台上,可能其它形式的IPC是更快的,但對於跨平台通信,sockets是唯一的選擇。

它們在Berkeley發明,做為Unix BSD的一部份。他們在網路上迅速傳播。有充份的理由,sockets與INET的結合使得與世界各地的任何機器交談變的異常的簡單(至少與其它方案相比)

Creating a Socket

大致而言,當你點擊連結,帶你到這個網頁,你的瀏灠器做了下面事情:

# create an INET, STREAMing socket
s = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
# now connect to the web server on port 80 - the normal http port
s.connect(("www.python.org", 80))

connect完成的時候,socket-s可以用來發送對頁面文字的請求。相同的socket將會讀取回覆,然後被銷毀。沒錯,被銷毀了。Client sockets通常只用於一次交換(或一小組序列交換)

在web server上發生的事有點複雜。首先,web server建立一個server socket

# create an INET, STREAMing socket
serversocket = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
# bind the socket to a public host, and a well-known port
serversocket.bind((socket.gethostname(), 80))
# become a server socket
serversocket.listen(5)

有幾點需要注意:我們使用socket.gethostname(),這樣外部就可以看到socket。如果我們使用s.bind(('localhost', 80))s.bind(('127.0.0.1', 80)),我們依然會有一個server socket,但那只會在相同機器上可以看的見。s.bind(('', 80))指定socket是可以被機器碰巧有的位址訪問。

第二件要注意的事:數值較小的port通常保留給"已知"的服務(HTTP、SNMP等)。如果你只是玩玩,請記得用四位數以上的port號。

最後,listen的參數告訴socket library,我們希望它在拒絕外部連線之前,隊列內應該要有五個連線(正常最大值)。如果其餘的程式碼正確的話,那應該是足夠的。

現在,我們擁有server socket,監聽80port,我們可以進入web server的主要迴圈:

while True:
    # accept connections from outside
    (clientsocket, address) = serversocket.accept()
    # now do something with the clientsocket
    # in this case, we'll pretend this is a threaded server
    ct = client_thread(clientsocket)
    ct.run()

實際上,有三種通用方法是這個迴圈內能做的事-調度一個線程處理clientsocket,建立一個進程處理clientsocket,或重構這個應用程式,使用非阻塞式socket,以及使用select在我們的server socket及任一活動的clientstockets進行多工。稍後會詳細介紹。現在需要理解的一個重點是:這就是server socket做的全部的事了。它沒有發送任何的資料。它沒有接收任何的資料。它只生成client sockets。每個clientsocket都是為了回應某些其它client socket執行connect()到我們所綁定的主機與port而建立。一但我們建立clientsocket,我們就會回頭監聽更多的連線。兩個clients可以自由交握-它們使用一些動態分配的port,當對話結束的時候,這些port會被回收。

IPC

如果你需要在一台機器上的兩個進程之間快速的IPC,你應該調查pipes或共享記憶體。如果你決定使用AF_INETsockets,那就綁定server socket到localhost。在多數平台上,這將圍繞兩層網絡程式碼走一條捷徑,而且速度要快得多。

See also multiprocessing整合跨平台IPC到高階API中。

Using a Socket

首先要注意的是web browser的client socket與web server的client socket是完全相同的。意思是,這是點對點的通訊。或者換種方式,做為設計人員,你必須決定通訊的禮儀規則。通常,連接socket通過發送一個需求或登錄來啟動通訊。但它是設計決定-它並不是sockets的規則。

現在有兩組動詞用於通訊。你可以使用sendrecv,或首可以轉換你的client socket變為file-like,然後使用readwrite。後者是Java呈現socket的方式。除了警告你需要在socket上使用flush之外,我並不會在這邊討論它。這些是緩衝檔案,常見錯誤是寫入一些東西,然後讀取回覆。如果沒有flush,你也許會一直等待回覆,因為request依然在你的輸出緩衝區中。

現在我們回到socket的主要癥結-sendrecv操作在網路緩衝區中。它們並不需要處理所有你傳送給它們的位元組,因為它們主要關注的是處理網路緩衝區。通常,它們在相關的網路緩衝被填滿(send)或清空(recv)的時候回傳。然後它們會告訴你它們處理多少位元組。在你的訊息完全的被處理之前,你有責任再次調用它們。

recv回傳0位元組的時候,這代表另一端已經關閉(或者正在關閉)連線。你將不會在這連線上再接收到任何資料。曾經,你也許能夠成功的傳送資料,稍後我會再詳細討論這點。

像HTTP這樣的協定使用socket只能進行一次傳輸。client送出一個request,然後讀取回覆。就這樣。然後socket就被丟掉。這意味著client可以透過接收到0位元組來偵測回覆的結束。

但是,如果你計劃重新使用你的socket做進一步的傳輸,你必需意識到,socket沒有EOT(End-of-Transmission)。再說一遍,如果socket在處理0位元組之後sendrecv回傳,那這個連線已經斷了。如果連線沒有斷,你也許會一直等著recv,因為socket不會告訴你(現在)沒東西可讀了。現在,如果你稍微思考一下,你就會瞭解sockets的基本原理:訊息必需固定長度、或是被分割、或者指定它們的長度、或通過關閉連線結束。這都取決於你的決定(但有些方法比其它方法更正確)。

假設你沒有要斷開連線,最簡單的解決方式就是固定訊息長度:

class MySocket:
    """demonstration class only
      - coded for clarity, not efficiency
    """

    def __init__(self, sock=None):
        if sock is None:
            self.sock = socket.socket(
                            socket.AF_INET, socket.SOCK_STREAM)
        else:
            self.sock = sock

    def connect(self, host, port):
        self.sock.connect((host, port))

    def mysend(self, msg):
        totalsent = 0
        while totalsent < MSGLEN:
            sent = self.sock.send(msg[totalsent:])
            if sent == 0:
                raise RuntimeError("socket connection broken")
            totalsent = totalsent + sent

    def myreceive(self):
        chunks = []
        bytes_recd = 0
        while bytes_recd < MSGLEN:
            chunk = self.sock.recv(min(MSGLEN - bytes_recd, 2048))
            if chunk == b'':
                raise RuntimeError("socket connection broken")
            chunks.append(chunk)
            bytes_recd = bytes_recd + len(chunk)
        return b''.join(chunks)

這裡的發送程式碼幾乎可以用於任何消息傳遞方式-在Python你發送字串,你可以使用len()來確認它的長(即使它包含0個字元)。主要是接收程式碼變得更複雜了(在C語言中,它並不會更糟,除非你不能使用strlen)。

最簡單的增強方式就是讓訊息的第一個字元做為訊息的指示器,並且確定長度。現在你有兩個recvs-第一個取得(最少)第一個字元,你可以查詢長度,第二個在迴圈中取得其餘的字元。如果你決定使用分離路由,你會接收到一些任意chunk size,(4096或8192通常很適合網路緩衝區大小),並掃描你接收到的分隔符號。

需要注意的一個複雜因素是:如果你的通訊協定允許多個訊息被回送(back to back)(不需某種回覆),然後你傳遞recv任意chunk size,你可能最終會讀取到後面的訊息的起始字元。你需要把它放一邊並保持住,直到需要它為止。

訊息前綴長度(假設是五個字元)變的更複雜,因為(信不信由你)你可能無法一次recv得到五個字元。在遊戲中你可以擺脫它,但在高網路負載中,除非你可以使用兩個recv迴圈,否則你的程式碼會很快中斷-第一個迴圈用於確定長度,第二個用於取得訊息的資料部份。雅打。這也是當你發現send並不總是可以在一次傳輸處理所有事情的時候。儘管已經讀過這篇文章了,你最終還是會被它反咬一口。

Binary Data

通過socket發送二進制資料是完全有可能的。主要的問題是並不是所有的機器都可以使用相同格式的二進制資料。舉例來說,一個Motorola芯片將代表一個十六位元整數,值為1,即兩個十六進制位元組00 01,Intel與DEC,然而,位元組是倒轉的,相同的1,其十六進制為01 00。socket套件要求轉換十六與三十二位元整數-ntohl, htonl, ntohs, htonsn代表networkh代表hosts代表shortl代表long。如果網路順序是host順序,那它們什麼也不會做,但如果是位元組反轉的話,那它們會適當的交換位元組。

在現在32bit的機器中,二進制資料的ascii表示通常小於二進制表示。這是因為驚人的時間量,所有這些長整數型的值都是0,或者可能是1。字串0是2bytes,而二進制是4bytes。當然,這不是適固定長度的訊息。

Disconnecting

嚴格來說,你應該在closesocket之前先執行shutdownshutdown是給另一端socket的提示。取決於你傳送的參數,它可是"我沒有要傳送資料了,但我依然監聽中",或"我沒監聽了,清除!"。然而,多數的socket libraries都是習慣程式設計師忽略這規則,通常closeshutdown()相同。close()。因此,多數情況下不需要直接執行shutdown

一種有效使用shutdown的方法是在HTTP-like exchange。client端發送一個request,然後執行shutdown(1)。這告訴server端"這個client完成發送,但依然可以接收"。server可以透過接收0位元組來偵測"EOF"。它可以假設它已經完成request。server發送一個回覆。如果send成功完成,那實際上client依然在接收。

Python將自動關閉更進一步,並表明當socket被垃圾回收時,如果需要,它將自動執行close。但依賴這個是一個壞習慣。如果你的socket沒有close情況下消失,那另一端的socket也許會無限期的掛著,你還會以為只是變慢了。當你完成的時候請記得一定要close你的sockets。

When Sockets Die

使用blocking sockets最糟糕的事情,就是另一邊掛掉的時候會發生什麼事(沒有執行close)。你的socket會掛著。TCP是一種可靠的協議,在放棄連線之前它會等待非常長的一段時間。如果你使用thread,那基本上整個線程已經掛了。對此,你無能為力。只要你沒有做一些很愚蠢的事情,像是在blocking read的時候拿著鎖,那thread並不會消耗太多資源。不要試著去終止thread-threads比processes更有效率部份原因在於它們避免資源自動回收的開銷。換句話說,如果你試著去終止thread,你的整個過程很可能被搞砸。

Non-blocking Sockets

如果你有瞭解上述內容,那你已經瞭解使用sockets機制的大致內容。你仍然會使用大致相同的方式調用。就是這樣,如果你做對了,那你的應用程式幾乎是內而外。

在Python,你使用socket.setblocking(0)來使它non-blocking。在C,它會更複雜,(首先,你需要在BSD風格O_NONBLOCK與幾乎無法區分的Posix風格O_NDELAY之間做個選擇,這與TCP_NODELAY完全不同)這是完全相同的概念。在建立socket之後執行這個操作, 但在使用它之前。(事實上,如果你瘋了,你可以來回切換。)

主要的機械式差異在於sendrecvconnectaccept可以在不做任何事情況下回傳。你有(當然)多種選擇。你可以檢查回傳的程式碼與異常程式碼,這通常讓你抓狂。如果你不相信我,試一下。你的應用程式會變的愈來愈大、愈來愈多蟲,吸光CPU。因此,讓我們跳過腦死的解決方案,把它做對吧。

使用select

在C,寫一個select非常複雜。在Python,輕而易舉,但是它非常接近C的版本,如果你瞭解Python的select,那在C裡面你幾乎不會有任何困擾:

ready_to_read, ready_to_write, in_error = \
               select.select(
                  potential_readers,
                  potential_writers,
                  potential_errs,
                  timeout)

你傳給select三個lists:第一個包含了所有你想試著讀取的sockets;第二個是你想試著寫入的sockets,最後一個(通常空的)是你想檢查的錯誤。你應該注意到,socket可以進入多個lists。select的調用是blocking,但你可以設置timeout,這通常是一個明智的作法-給它一個很長的timeout(假設一分鐘),除非你有一個很好的理由不這麼做。

回傳的部份也將得到三個lists。他們包含了實際可讀、可寫以及錯誤的sockets。這些清單中的每一個都是你傳入的相對應清單的子集(可能是空的)。

如果socket在可讀的list中,你可以盡可能的像我們在這業務中所得的那樣,以便在該socket上的recv將回傳某些東西。對可寫list是相同的想法,你可以發送一些東西。這也許不是你想要的,但有一些東西總比沒東西好。(事實上,任何正常的socket都將以可寫模式回傳-這意味著輸出網路緩衝空間是可用的)。

如果你有一個server socket,把它放進potential_readers。如果它出現在可讀的list中,你的accept(幾乎確定)會有作用。如果你已經建立新的socket,並且connect到其它人身上了,把它放到potential_writers list。如果它出現在可寫的list中,那就有機會它是已經連接的了。

事實上,即使使用blocking sockets,select依然是非常方便的。這是一個確認你是否阻塞的方法-當有些東西在緩衝區的時候,socket回傳為可讀。然而,這對判斷另一端是否完成是沒有幫助的,或者它只是忙著其它事情。

跨平台警示:
在Unix上,select可以同時適用於sockets與files。不要在Windows上這麼做。在Windows上,select僅支援sockets。還有C語言,socket很多進階選項在Windows上是不同的。事實上,在Windows我通常使用執行緒(執行狀況非常、非常好)。

2019年7月16日 星期二

DCGAN

DCGAN

DCGAN,即deep convolutional GAN的簡寫

數據集

CIFAR10為32x32x3的照片,並擁有10個輸出類別,每個類別5,000張,共計50,000張照片

範例說明

為了簡化範例,單純的使用一個類別做測試
這邊會著重在實作,不會有過深的理論說明

作業開始

GAN由Generator與Discriminator兩個nn結合而成,迭代過程大致如下說明:
  1. 初始化Generator-v1與Discriminator-v1
  2. 空間中sample出一筆資料,經過Generator-v1生成
  3. Discriminator-v1驗證真假,發現是假
  4. Generator-v1升級為Generator-v2
  5. 成功騙過Discriminator
  6. Discriminator-v1升級為Discriminator-v2
  7. Discriminator-v2驗證真假,發現是假
作業開始之前定義使用的GPU資源
import os
# 限制gpu資源
os.environ["CUDA_VISIBLE_DEVICES"] = "0"
首先我們建置Generator
載入建置Generator的需求套件
import keras
from keras import layers
import numpy as np
這是後面保有照片會用到的套件
from keras.preprocessing import image
定義dimension,CIFAR10的照片維度為32x32x3
Generator會從空間中sample出一個點,那個點是一個Vector,在下面範例中即是laten_dim的設置
# CIFAR10資料維度
height = 32
width = 32
channel =3

# 預計生成資料維度
laten_dim = 32
  • W:輸入維度
  • F:filter size
  • S:stride
  • P:padding size
  • N:輸入維度
layers.Conv2D:卷積,經過卷積之後的維度計算如下:
N=(WF+2P)/S+1
layers.Conv2DTranspose:反卷積,經過反卷積之後的維度計算如下:
W=(N1)S2P+F
generator_input = keras.Input(shape=(laten_dim, ))
x = layers.Dense(128 * 16 * 16)(generator_input)
x = layers.LeakyReLU()(x)
x = layers.Reshape((16, 16, 128))(x)

x = layers.Conv2D(256, 5, padding='same')(x)
x = layers.LeakyReLU()(x)

# 輸出為32x32
x = layers.Conv2DTranspose(256, 4, strides=2, padding='same')(x)
x = layers.LeakyReLU()(x)

x = layers.Conv2D(256, 5, padding='same')(x)
x = layers.LeakyReLU()(x)
x = layers.Conv2D(256, 5, padding='same')(x)
x = layers.LeakyReLU()(x)

# 圖片channel設置為3,即輸出為32x32x3
x = layers.Conv2D(channel, 7, activation='tanh', padding='same')(x)

generator = keras.models.Model(generator_input, x)
利用Model.summary()來確認模型的資料維度變化
generator.summary()
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
input_2 (InputLayer)         (None, 32)                0         
_________________________________________________________________
dense_2 (Dense)              (None, 32768)             1081344   
_________________________________________________________________
leaky_re_lu_6 (LeakyReLU)    (None, 32768)             0         
_________________________________________________________________
reshape_2 (Reshape)          (None, 16, 16, 128)       0         
_________________________________________________________________
conv2d_5 (Conv2D)            (None, 16, 16, 256)       819456    
_________________________________________________________________
leaky_re_lu_7 (LeakyReLU)    (None, 16, 16, 256)       0         
_________________________________________________________________
conv2d_transpose_2 (Conv2DTr (None, 32, 32, 256)       1048832   
_________________________________________________________________
leaky_re_lu_8 (LeakyReLU)    (None, 32, 32, 256)       0         
_________________________________________________________________
conv2d_6 (Conv2D)            (None, 32, 32, 256)       1638656   
_________________________________________________________________
leaky_re_lu_9 (LeakyReLU)    (None, 32, 32, 256)       0         
_________________________________________________________________
conv2d_7 (Conv2D)            (None, 32, 32, 256)       1638656   
_________________________________________________________________
leaky_re_lu_10 (LeakyReLU)   (None, 32, 32, 256)       0         
_________________________________________________________________
conv2d_8 (Conv2D)            (None, 32, 32, 3)         37635     
=================================================================
Total params: 6,264,579
Trainable params: 6,264,579
Non-trainable params: 0
_________________________________________________________________
Generator完成之後我們要架構Discriminator,Discriminator的作用就是判斷給定的資料是真或假,因此它的輸入維度即照片維度(範例為32x32x3)
discriminator_input = layers.Input(shape=(height, width, channel))
x = layers.Conv2D(128, 3)(discriminator_input)
x = layers.LeakyReLU()(x)
x = layers.Conv2D(128, 4, strides=2)(x)
x = layers.LeakyReLU()(x)
x = layers.Conv2D(128, 4, strides=2)(x)
x = layers.LeakyReLU()(x)
x = layers.Conv2D(128, 4, strides=2)(x)
x = layers.LeakyReLU()(x)
x = layers.Flatten()(x)
x = layers.Dropout(0.5)(x)
# 判斷真假,因此輸出為1個unit,並搭配sigmoid
x = layers.Dense(1, activation='sigmoid')(x)

discriminator = keras.models.Model(discriminator_input, x)
利用Model.summary()確認模型維度變化
discriminator.summary()
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
input_3 (InputLayer)         (None, 32, 32, 3)         0         
_________________________________________________________________
conv2d_9 (Conv2D)            (None, 30, 30, 128)       3584      
_________________________________________________________________
leaky_re_lu_11 (LeakyReLU)   (None, 30, 30, 128)       0         
_________________________________________________________________
conv2d_10 (Conv2D)           (None, 14, 14, 128)       262272    
_________________________________________________________________
leaky_re_lu_12 (LeakyReLU)   (None, 14, 14, 128)       0         
_________________________________________________________________
conv2d_11 (Conv2D)           (None, 6, 6, 128)         262272    
_________________________________________________________________
leaky_re_lu_13 (LeakyReLU)   (None, 6, 6, 128)         0         
_________________________________________________________________
conv2d_12 (Conv2D)           (None, 2, 2, 128)         262272    
_________________________________________________________________
leaky_re_lu_14 (LeakyReLU)   (None, 2, 2, 128)         0         
_________________________________________________________________
flatten_1 (Flatten)          (None, 512)               0         
_________________________________________________________________
dropout_1 (Dropout)          (None, 512)               0         
_________________________________________________________________
dense_3 (Dense)              (None, 1)                 513       
=================================================================
Total params: 790,913
Trainable params: 790,913
Non-trainable params: 0
_________________________________________________________________
定義模型最佳化方式
範例使用RMSprop做為最佳化的方式,透過clipvalue來限制梯度範圍
discriminator_optimizer = keras.optimizers.RMSprop(
    lr = 0.0008,
    clipvalue=1.0,
    decay=1e-8
)
discriminator.compile(optimizer=discriminator_optimizer,
                      loss='binary_crossentropy')
Generator與Discriminator都已經設置好了,現在我們要讓他們對抗,幾個點簡單說明:
  1. 原始GAN中雖然提到Generator要訓練多次,但實際上Ian Goodfellow只訓練一次而以
  2. 訓練Generator的時候Discriminator是凍結的
首先,將discriminator凍結
discriminator.trainable = False
gan model的input是一開始所設置的laten_dim,而output的部份則是generator model的output給discriminator model的input所做的判斷,即判斷generator model所生成的資料是真還是假
gan_input = keras.Input(shape=(laten_dim, ))
gan_output = discriminator(generator(gan_input))
gan = keras.models.Model(gan_input, gan_output)
設置gan model的最佳化方式
gan_optimizer = keras.optimizers.RMSprop(
    lr=0.0004,
    clipvalue=1.0,
    decay=1e-8
)
gan.compile(optimizer=gan_optimizer,
            loss='binary_crossentropy')
gan.summary()
_________________________________________________________________
Layer (type)                 Output Shape              Param #   
=================================================================
input_5 (InputLayer)         (None, 32)                0         
_________________________________________________________________
model_2 (Model)              (None, 32, 32, 3)         6264579   
_________________________________________________________________
model_3 (Model)              (None, 1)                 790913    
=================================================================
Total params: 7,055,492
Trainable params: 6,264,579
Non-trainable params: 790,913
_________________________________________________________________
現在我們可以開始來訓練模型,上面提到,我們會從空間隨機sample出資料,經過generator model來生成,這筆資料當然假資料,其標記為Negative,另外我們也會有真實的資料,其標記為Positive
首先我們先下載keras自帶的資料集,CIFAR10
  • 0: airplane
  • 1: automobile
  • 2: bird
  • 3: cat
  • 4: deer
  • 5: dog
  • 6: frog
  • 7: horse
  • 8: ship
  • 9: truck
(x_train, y_train), (_, _) = keras.datasets.cifar10.load_data()
Downloading data from https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz
170500096/170498071 [==============================] - 109s 1us/step
任何資料集都請務必先行驗證確認資料維度
x_train.shape, y_train.shape
((50000, 32, 32, 3), (50000, 1))
下載下來的照片資料,我們只會採用其中一個類別來簡化範例的進行
y_train.flatten() == 5為mask來取得資料
x_train = x_train[y_train.flatten() == 5]
調整之後確認資料維度
x_train.shape
(5000, 32, 32, 3)
後面不需要太過在意y_train,因為對我們目前需求來說,所有真實資料皆為Positive
x_train.dtype
dtype('uint8')
資料標準化
x_train = x_train / 255.
x_train.dtype
dtype('float64')
轉換型別,理論上使用float32足矣
x_train = x_train.astype(np.float32)
x_train.dtype
dtype('float32')
設置幾個簡單參數
# 預計執行迭代次數
iterations = 10000
# 每次處理數量
batch_size = 20
# 預計保存生成照片資料集,依個人需求設置
folder_path = '/tf/GAN/DCGAN' 
前置作業完畢之後就可以正式的來訓練模型
# 記錄起始索引的參數
idx_start = 0
for i in range(iterations):
    # 從高斯分佈空間中隨機sample出資料點
    random_vectors = np.random.normal(size=(batch_size, laten_dim))
    # 從generator model取得output
    # 依我們所設置的模型,會回傳Nx32x32x3的資料
    generator_img = generator.predict(random_vectors)
    # 目前的結束索引
    idx_stop = idx_start + batch_size
    # 利用索引取得相對應的真實照片
    true_imgs = x_train[idx_start: idx_stop]
    # 將真假堆疊在一起
    combian_imgs = np.concatenate([generator_img, true_imgs])
    # 產生相對應的label,生成為假:1,真實照片:0
    labels = np.concatenate([np.ones((batch_size, 1)), 
                             np.zeros((batch_size, 1))])
    
    # 在label中加入隨機的噪點,很多事不能說明為什麼
    # 跟著打牌就不會放槍就是了
    labels += 0.01 * np.random.random(labels.shape)
    
    # 目前已經隨機生成照片,要拿這些照片來訓練discriminator
    discriminator_loss = discriminator.train_on_batch(combian_imgs, labels)
    
    # 現在,discriminator已經知道這些是假照片,因此要更新generator
    # 接下來的生成照片,印象中李弘毅老師的線上課程中有提到,要不要重新生成都可以    
    random_vectors = np.random.normal(size=(batch_size, laten_dim))
    # 產生上面新生成照片的labels
    # 這邊設置的標籤是『Positive』,這是因為我們要欺騙discriminator
    # 讓discriminator感覺這個generator生成的照片是真的
    # 這個過程之後generator就升級了
    generator_labels = np.zeros((batch_size, 1))
    # 訓練gan model,記得這時候訓練的是我們凍結discriminator的模型-gan
    gan_loss = gan.train_on_batch(random_vectors, generator_labels)
    
    # 更新索引
    idx_start = idx_start + batch_size 
    # 判斷索引是否超過資料集索引,超過的話就重新計算
    # 這邊也要特別注意到資料集數量與batch_size的設置關係,要可以整除
    if idx_start > len(x_train) - batch_size:
        idx_start = 0
        
    # 這邊每100次迭代記錄一次
    if i % 100 == 0:
        gan.save_weights('gan.h5')
        
        print('discriminator loss: ', discriminator_loss)
        print('gan loss: ', gan_loss)
        
        img = image.array_to_img(generator_img[0] * 255., scale=False)
        img.save(os.path.join(folder_path, 'epoch-' + str(i) + '-generator.jpg'))
        
        img = image.array_to_img(true_imgs[0] * 255., scale=False)
        img.save(os.path.join(folder_path, 'epoch-' + str(i) + '-true_image.jpg'))        
    
discriminator loss:  5.1067543
gan loss:  13.483853
discriminator loss:  0.55595607
gan loss:  0.9773754
discriminator loss:  0.6778943
gan loss:  0.7218493
discriminator loss:  0.7140575
gan loss:  0.74051535
discriminator loss:  0.6997445
gan loss:  0.74154454
discriminator loss:  0.6921107
gan loss:  0.7020267
discriminator loss:  0.69686
gan loss:  0.7030369
discriminator loss:  0.6800611
gan loss:  0.6471882
discriminator loss:  0.6967187
gan loss:  0.77145576
discriminator loss:  0.7020233
gan loss:  0.71897966
discriminator loss:  0.7126138
gan loss:  0.73591757
discriminator loss:  0.7110381
gan loss:  0.7163498
discriminator loss:  0.6996112
gan loss:  0.71271735
discriminator loss:  0.73551977
gan loss:  0.6986626
discriminator loss:  0.67326313
gan loss:  0.71304166
discriminator loss:  0.7034636
gan loss:  1.0788002
discriminator loss:  0.7141048
gan loss:  0.7726129
discriminator loss:  0.696671
gan loss:  0.720789
discriminator loss:  0.7359967
gan loss:  0.7729639
discriminator loss:  0.69505966
gan loss:  0.65915006
discriminator loss:  0.67265576
gan loss:  0.89056253
discriminator loss:  0.70829356
gan loss:  0.6985985
discriminator loss:  0.6924041
gan loss:  0.6630027
discriminator loss:  0.71791756
gan loss:  0.6948942
discriminator loss:  0.7060956
gan loss:  0.70852387
discriminator loss:  0.69083965
gan loss:  0.7589138
discriminator loss:  0.7217103
gan loss:  0.7043828
discriminator loss:  0.69030714
gan loss:  0.7687561
discriminator loss:  0.7466472
gan loss:  0.83582276
discriminator loss:  0.6909427
gan loss:  0.75084436
discriminator loss:  0.692063
gan loss:  0.66303813
discriminator loss:  0.6578157
gan loss:  0.7915683
discriminator loss:  0.660645
gan loss:  0.717038
discriminator loss:  0.6897063
gan loss:  0.71162593
discriminator loss:  0.72813225
gan loss:  0.7520748
discriminator loss:  0.70937884
gan loss:  0.9043194
discriminator loss:  0.6839577
gan loss:  0.72316694
discriminator loss:  0.68821317
gan loss:  0.6758722
discriminator loss:  0.66450214
gan loss:  0.7818069
discriminator loss:  0.68543863
gan loss:  0.7772101
discriminator loss:  0.72352195
gan loss:  1.1552292
discriminator loss:  0.66245687
gan loss:  0.8431295
discriminator loss:  0.64920413
gan loss:  0.6762968
discriminator loss:  0.6851862
gan loss:  0.6802823
discriminator loss:  0.6546106
gan loss:  0.8981015
discriminator loss:  0.66509783
gan loss:  0.8244335
discriminator loss:  0.71177804
gan loss:  0.64901894
discriminator loss:  0.6947624
gan loss:  0.6984755
discriminator loss:  0.70206606
gan loss:  0.8204377
discriminator loss:  0.7477706
gan loss:  0.9481209
discriminator loss:  0.67449886
gan loss:  0.83607197
discriminator loss:  0.6896846
gan loss:  0.6858331
discriminator loss:  0.6856981
gan loss:  0.71470636
discriminator loss:  0.7168747
gan loss:  0.58466285
discriminator loss:  0.69816226
gan loss:  0.76343626
discriminator loss:  0.7020701
gan loss:  0.8184387
discriminator loss:  0.64522004
gan loss:  0.6753376
discriminator loss:  0.6865121
gan loss:  0.7008316
discriminator loss:  0.7476796
gan loss:  0.87338126
discriminator loss:  0.7263964
gan loss:  0.6889228
discriminator loss:  0.70218945
gan loss:  0.6831338
discriminator loss:  0.6857745
gan loss:  0.6565727
discriminator loss:  0.68626153
gan loss:  0.6991853
discriminator loss:  0.68162763
gan loss:  0.74716157
discriminator loss:  0.7115513
gan loss:  0.7860044
discriminator loss:  0.7063864
gan loss:  0.74256736
discriminator loss:  0.6947535
gan loss:  0.73066074
discriminator loss:  0.70494205
gan loss:  0.64904773
discriminator loss:  0.69508123
gan loss:  0.72680366
discriminator loss:  0.6899625
gan loss:  0.81384104
discriminator loss:  0.7159114
gan loss:  0.7416827
discriminator loss:  0.6912014
gan loss:  0.7360595
discriminator loss:  0.72471696
gan loss:  0.6513818
discriminator loss:  0.6992035
gan loss:  0.6941215
discriminator loss:  0.6863073
gan loss:  0.7213627
discriminator loss:  0.6924833
gan loss:  0.71849513
discriminator loss:  0.6969613
gan loss:  0.72678196
discriminator loss:  0.685106
gan loss:  0.7152079
discriminator loss:  0.68887657
gan loss:  0.73665506
discriminator loss:  0.71015817
gan loss:  0.732918
discriminator loss:  0.67974985
gan loss:  0.8060713
discriminator loss:  0.693206
gan loss:  0.67660487
discriminator loss:  0.68673164
gan loss:  0.7380837
discriminator loss:  0.6858455
gan loss:  0.7205802
discriminator loss:  0.69764405
gan loss:  0.7232281
discriminator loss:  0.7109472
gan loss:  0.66223055
discriminator loss:  0.707724
gan loss:  1.1661718
discriminator loss:  0.6913794
gan loss:  0.6941385
discriminator loss:  0.69227374
gan loss:  0.7601951
discriminator loss:  0.6878415
gan loss:  0.7479371
discriminator loss:  0.71728295
gan loss:  0.83168066
discriminator loss:  0.6864918
gan loss:  0.7533535
discriminator loss:  0.7127117
gan loss:  0.6897801
discriminator loss:  0.69414794
gan loss:  0.7193168
discriminator loss:  0.67149776
gan loss:  0.7141426
discriminator loss:  0.6828767
gan loss:  0.7632409
discriminator loss:  0.67732584
gan loss:  0.7958939
discriminator loss:  0.69132864
gan loss:  0.721274
discriminator loss:  0.71267766
gan loss:  0.71191424
discriminator loss:  0.6940938
gan loss:  0.8811827

確認生成照片

import matplotlib.pyplot as plt
import matplotlib.image as mpimg
%matplotlib inline
img_path = 'epoch-9600-generator.jpg'
gen_img = mpimg.imread(img_path)
plt.imshow(gen_img)

像哈巴狗?

柯基與柴犬?

結論

範例可以看的出來,在Keras的高階api協助之下,實作GAN並不是那麼樣子的困難,但GAN的訓練有許多比程式碼還要困難的部份(上面結果可以發現,優化的並不是那麼好),不管是Generator還是Discriminator,過強或過弱對模型來說都不是好事,許多情況下GAN難以訓練,需要的是參數上的不斷調校,唯有不斷的踩雷才能擁有足夠的經驗。

後記:後面繼續訓練約50,000次迭代之後就整個壞掉了!

延伸

單純保存一張照片似乎較難以判斷生成狀況,下面function將圖片串接起來

直接使用keras自帶工具來處理照片保存的作業

import numpy as np
from keras.preprocessing import image
def save_generator_img(generator_img, h_int, w_int, epoch, save_file_name='generator', save_dir='.'):
    """保存generator生成的照片
    function:
        利用keras自帶工具keras.preprocessing.Image保存照片
    
    parameters:
        generator_img: generator生成的numpy object
        h_int: 湊成一張照片的時候高要幾張
        w_ing: 湊成一張照片的時候寬要幾張
        epoch: 第幾次的迭代,保存照片的時候插入檔案名稱
        save_file_name: 檔案名稱,預設為generator
        save_dir: 保存的路徑,預設在執行程式的根目錄    
        
    remark:
        h_int x w_int 不能超過generator_img的長度,一定只能等於
    
    example:
        save_generator_img(generator_img, 4, 5, 9900)
    """
    # 取得資料維度
    N, H, W, C = generator_img.shape
    
    # 驗證拼湊照片數量相符
    assert int(h_int) * int(w_int) == N        
    
    # 開一張全zero的大陣列
    target_img = np.zeros(shape=(h_int * H, w_int * W, C))
    
    # 索引
    now_idx = 0    
    for h_idx in range(h_int):
        for w_idx in range(w_int):
            # 取得照片
            _img = generator_img[now_idx]
            # 取代相對應陣列位置
            target_img[h_idx * H: h_idx * H + H, w_idx * W: w_idx * W + W] = _img                        
            now_idx += 1
    
    
    file_name = os.path.join(save_dir, save_file_name + str(epoch) + '.png')
    save_img = image.array_to_img(target_img * 255., scale=False)
    save_img.save(file_name)