497 lines
26 KiB
Plaintext
497 lines
26 KiB
Plaintext
{
|
||
"cells": [
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"# 生成式網絡\n",
|
||
"\n",
|
||
"循環神經網絡(Recurrent Neural Networks, RNNs)及其門控單元變體,例如長短期記憶單元(Long Short Term Memory Cells, LSTMs)和門控循環單元(Gated Recurrent Units, GRUs),提供了一種語言建模的機制,也就是說,它們可以學習單詞的排列順序,並對序列中的下一個單詞進行預測。這使得我們可以使用 RNNs 來完成**生成任務**,例如普通文本生成、機器翻譯,甚至是圖像描述生成。\n",
|
||
"\n",
|
||
"在我們上一單元討論的 RNN 架構中,每個 RNN 單元會生成下一個隱藏狀態作為輸出。然而,我們也可以為每個循環單元添加另一個輸出,這樣就可以輸出一個**序列**(其長度等於原始序列的長度)。此外,我們還可以使用不在每一步接受輸入的 RNN 單元,而僅僅接受一個初始狀態向量,然後生成一個輸出序列。\n",
|
||
"\n",
|
||
"在這個筆記本中,我們將專注於幫助我們生成文本的簡單生成模型。為了簡化起見,我們將構建一個**字符級網絡**,逐字母生成文本。在訓練過程中,我們需要使用一些文本語料庫,並將其拆分為字母序列。\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 1,
|
||
"metadata": {},
|
||
"outputs": [],
|
||
"source": [
|
||
"import tensorflow as tf\n",
|
||
"from tensorflow import keras\n",
|
||
"import tensorflow_datasets as tfds\n",
|
||
"import numpy as np\n",
|
||
"\n",
|
||
"ds_train, ds_test = tfds.load('ag_news_subset').values()"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"## 建立字元詞彙表\n",
|
||
"\n",
|
||
"為了建立字元級別的生成網絡,我們需要將文本拆分為單個字元,而不是單詞。我們之前使用的 `TextVectorization` 層無法做到這一點,因此我們有以下兩個選擇:\n",
|
||
"\n",
|
||
"* 手動加載文本並自行進行分詞,如 [這個官方 Keras 範例](https://keras.io/examples/generative/lstm_character_level_text_generation/) 中所示\n",
|
||
"* 使用 `Tokenizer` 類進行字元級別的分詞。\n",
|
||
"\n",
|
||
"我們將選擇第二種方法。`Tokenizer` 也可以用於將文本分詞為單詞,因此可以很輕鬆地從字元級分詞切換到單詞級分詞。\n",
|
||
"\n",
|
||
"要進行字元級分詞,我們需要傳遞參數 `char_level=True`:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 2,
|
||
"metadata": {},
|
||
"outputs": [],
|
||
"source": [
|
||
"def extract_text(x):\n",
|
||
" return x['title']+' '+x['description']\n",
|
||
"\n",
|
||
"def tupelize(x):\n",
|
||
" return (extract_text(x),x['label'])\n",
|
||
"\n",
|
||
"tokenizer = keras.preprocessing.text.Tokenizer(char_level=True,lower=False)\n",
|
||
"tokenizer.fit_on_texts([x['title'].numpy().decode('utf-8') for x in ds_train])"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"我們還希望使用一個特殊的標記來表示**序列結束**,我們將其稱為`<eos>`。讓我們手動將其添加到詞彙表中:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 3,
|
||
"metadata": {},
|
||
"outputs": [],
|
||
"source": [
|
||
"eos_token = len(tokenizer.word_index)+1\n",
|
||
"tokenizer.word_index['<eos>'] = eos_token\n",
|
||
"\n",
|
||
"vocab_size = eos_token + 1"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"現在,要將文本編碼為數字序列,我們可以使用:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 4,
|
||
"metadata": {},
|
||
"outputs": [
|
||
{
|
||
"data": {
|
||
"text/plain": [
|
||
"[[48, 2, 10, 10, 5, 44, 1, 25, 5, 8, 10, 13, 78]]"
|
||
]
|
||
},
|
||
"execution_count": 4,
|
||
"metadata": {},
|
||
"output_type": "execute_result"
|
||
}
|
||
],
|
||
"source": [
|
||
"tokenizer.texts_to_sequences(['Hello, world!'])"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"## 訓練生成式 RNN 來生成標題\n",
|
||
"\n",
|
||
"我們將以以下方式訓練 RNN 來生成新聞標題。在每一步中,我們會選取一個標題,將其輸入到 RNN 中,並且對於每個輸入的字元,我們會要求網路生成下一個輸出的字元:\n",
|
||
"\n",
|
||
"\n",
|
||
"\n",
|
||
"對於序列中的最後一個字元,我們會要求網路生成 `<eos>` 標記。\n",
|
||
"\n",
|
||
"我們在此使用的生成式 RNN 的主要不同之處在於,我們會從 RNN 的每一步輸出中提取結果,而不僅僅是從最後一個單元提取。這可以通過向 RNN 單元指定 `return_sequences` 參數來實現。\n",
|
||
"\n",
|
||
"因此,在訓練過程中,網路的輸入將是一段編碼字元的序列,而輸出將是相同長度的序列,但向右偏移一個元素並以 `<eos>` 結束。小批次將由多個這樣的序列組成,我們需要使用**填充**來對齊所有序列。\n",
|
||
"\n",
|
||
"接下來,我們來建立一些函數,用於轉換數據集。由於我們希望在小批次層級進行序列填充,我們會先通過調用 `.batch()` 將數據集分批,然後使用 `map` 來進行轉換。因此,轉換函數將以整個小批次作為參數:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 5,
|
||
"metadata": {},
|
||
"outputs": [],
|
||
"source": [
|
||
"def title_batch(x):\n",
|
||
" x = [t.numpy().decode('utf-8') for t in x]\n",
|
||
" z = tokenizer.texts_to_sequences(x)\n",
|
||
" z = tf.keras.preprocessing.sequence.pad_sequences(z)\n",
|
||
" return tf.one_hot(z,vocab_size), tf.one_hot(tf.concat([z[:,1:],tf.constant(eos_token,shape=(len(z),1))],axis=1),vocab_size)"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"一些我們在這裡執行的重要步驟:\n",
|
||
"* 我們首先從字串張量中提取實際的文字\n",
|
||
"* `text_to_sequences` 將字串列表轉換為整數張量列表\n",
|
||
"* `pad_sequences` 將這些張量填充到它們的最大長度\n",
|
||
"* 最後,我們對所有字符進行獨熱編碼,並執行位移和 `<eos>` 附加操作。我們很快就會了解為什麼需要獨熱編碼的字符\n",
|
||
"\n",
|
||
"然而,這個函數是 **Pythonic** 的,也就是說它無法自動轉換為 Tensorflow 的計算圖。如果我們直接在 `Dataset.map` 函數中使用這個函數,會出現錯誤。我們需要使用 `py_function` 包裝器來封裝這個 Pythonic 調用:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 6,
|
||
"metadata": {},
|
||
"outputs": [],
|
||
"source": [
|
||
"def title_batch_fn(x):\n",
|
||
" x = x['title']\n",
|
||
" a,b = tf.py_function(title_batch,inp=[x],Tout=(tf.float32,tf.float32))\n",
|
||
" return a,b"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"> **注意**:區分 Pythonic 和 Tensorflow 的轉換函數可能看起來有些複雜,您可能會疑惑為什麼我們不在將數據集傳遞給 `fit` 之前使用標準的 Python 函數進行轉換。雖然這確實是可行的,但使用 `Dataset.map` 有一個巨大的優勢,因為數據轉換管道是通過 Tensorflow 的計算圖執行的,這可以利用 GPU 的計算能力,並且減少了在 CPU 和 GPU 之間傳遞數據的需求。\n",
|
||
"\n",
|
||
"現在我們可以構建生成器網絡並開始訓練。它可以基於我們在上一單元中討論的任何循環單元(簡單、LSTM 或 GRU)。在我們的示例中,我們將使用 LSTM。\n",
|
||
"\n",
|
||
"由於網絡以字符作為輸入,且詞彙表的大小相對較小,我們不需要嵌入層,經過一次熱編碼的輸入可以直接進入 LSTM 單元。輸出層將是一個 `Dense` 分類器,它會將 LSTM 的輸出轉換為一次熱編碼的標記編號。\n",
|
||
"\n",
|
||
"此外,因為我們處理的是可變長度的序列,我們可以使用 `Masking` 層來創建一個掩碼,忽略字符串中填充的部分。這並不是絕對必要的,因為我們對超過 `<eos>` 標記的部分並不太感興趣,但我們會使用它來獲得一些使用這類型層的經驗。`input_shape` 將是 `(None, vocab_size)`,其中 `None` 表示可變長度的序列,而輸出形狀也是 `(None, vocab_size)`,正如您可以從 `summary` 中看到的那樣:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 7,
|
||
"metadata": {},
|
||
"outputs": [
|
||
{
|
||
"name": "stdout",
|
||
"output_type": "stream",
|
||
"text": [
|
||
"Model: \"sequential\"\n",
|
||
"_________________________________________________________________\n",
|
||
"Layer (type) Output Shape Param # \n",
|
||
"=================================================================\n",
|
||
"masking (Masking) (None, None, 84) 0 \n",
|
||
"_________________________________________________________________\n",
|
||
"lstm (LSTM) (None, None, 128) 109056 \n",
|
||
"_________________________________________________________________\n",
|
||
"dense (Dense) (None, None, 84) 10836 \n",
|
||
"=================================================================\n",
|
||
"Total params: 119,892\n",
|
||
"Trainable params: 119,892\n",
|
||
"Non-trainable params: 0\n",
|
||
"_________________________________________________________________\n",
|
||
"15000/15000 [==============================] - 229s 15ms/step - loss: 1.5385\n"
|
||
]
|
||
},
|
||
{
|
||
"data": {
|
||
"text/plain": [
|
||
"<tensorflow.python.keras.callbacks.History at 0x7fa40c1245e0>"
|
||
]
|
||
},
|
||
"execution_count": 7,
|
||
"metadata": {},
|
||
"output_type": "execute_result"
|
||
}
|
||
],
|
||
"source": [
|
||
"model = keras.models.Sequential([\n",
|
||
" keras.layers.Masking(input_shape=(None,vocab_size)),\n",
|
||
" keras.layers.LSTM(128,return_sequences=True),\n",
|
||
" keras.layers.Dense(vocab_size,activation='softmax')\n",
|
||
"])\n",
|
||
"\n",
|
||
"model.summary()\n",
|
||
"model.compile(loss='categorical_crossentropy')\n",
|
||
"\n",
|
||
"model.fit(ds_train.batch(8).map(title_batch_fn))"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"## 生成輸出\n",
|
||
"\n",
|
||
"現在我們已經訓練了模型,接下來我們希望使用它來生成一些輸出。首先,我們需要一種方法來解碼由一系列標記數字表示的文本。為此,我們可以使用 `tokenizer.sequences_to_texts` 函數;然而,該函數在字元級標記化時效果不佳。因此,我們將從標記器中提取標記字典(稱為 `word_index`),建立一個反向映射,並撰寫自己的解碼函數:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 10,
|
||
"metadata": {},
|
||
"outputs": [],
|
||
"source": [
|
||
"reverse_map = {val:key for key, val in tokenizer.word_index.items()}\n",
|
||
"\n",
|
||
"def decode(x):\n",
|
||
" return ''.join([reverse_map[t] for t in x])"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"現在,我們開始進行生成。我們將以某個字串 `start` 作為起點,將其編碼為序列 `inp`,接著在每一步中呼叫我們的網路來推斷下一個字元。\n",
|
||
"\n",
|
||
"網路的輸出 `out` 是一個包含 `vocab_size` 元素的向量,代表每個標記的概率。我們可以使用 `argmax` 找出最可能的標記編號。然後,我們將這個字元附加到已生成的標記列表中,並繼續進行生成。這個生成一個字元的過程會重複執行 `size` 次,以生成所需的字元數量。如果在過程中遇到 `eos_token`,我們會提前終止生成。\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 12,
|
||
"metadata": {},
|
||
"outputs": [
|
||
{
|
||
"data": {
|
||
"text/plain": [
|
||
"'Today #39;s lead to strike for the strike for the strike for the strike (AFP)'"
|
||
]
|
||
},
|
||
"execution_count": 12,
|
||
"metadata": {},
|
||
"output_type": "execute_result"
|
||
}
|
||
],
|
||
"source": [
|
||
"def generate(model,size=100,start='Today '):\n",
|
||
" inp = tokenizer.texts_to_sequences([start])[0]\n",
|
||
" chars = inp\n",
|
||
" for i in range(size):\n",
|
||
" out = model(tf.expand_dims(tf.one_hot(inp,vocab_size),0))[0][-1]\n",
|
||
" nc = tf.argmax(out)\n",
|
||
" if nc==eos_token:\n",
|
||
" break\n",
|
||
" chars.append(nc.numpy())\n",
|
||
" inp = inp+[nc]\n",
|
||
" return decode(chars)\n",
|
||
" \n",
|
||
"generate(model)"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"## 在訓練期間抽樣輸出\n",
|
||
"\n",
|
||
"由於我們沒有任何像 *準確率* 這樣的有用指標,我們唯一能夠判斷模型是否有所改善的方法就是在訓練期間通過**抽樣**生成的字串來觀察。為了實現這一點,我們將使用**回調函數**,即可以傳遞給 `fit` 函數的函數,並在訓練過程中定期被調用。\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 13,
|
||
"metadata": {},
|
||
"outputs": [
|
||
{
|
||
"name": "stdout",
|
||
"output_type": "stream",
|
||
"text": [
|
||
"Epoch 1/3\n",
|
||
"15000/15000 [==============================] - 226s 15ms/step - loss: 1.2703\n",
|
||
"Today #39;s a lead in the company for the strike\n",
|
||
"Epoch 2/3\n",
|
||
"15000/15000 [==============================] - 227s 15ms/step - loss: 1.2057\n",
|
||
"Today #39;s the Market Service on Security Start (AP)\n",
|
||
"Epoch 3/3\n",
|
||
"15000/15000 [==============================] - 226s 15ms/step - loss: 1.1752\n",
|
||
"Today #39;s a line on the strike to start for the start\n"
|
||
]
|
||
},
|
||
{
|
||
"data": {
|
||
"text/plain": [
|
||
"<tensorflow.python.keras.callbacks.History at 0x7fa40c74e3d0>"
|
||
]
|
||
},
|
||
"execution_count": 13,
|
||
"metadata": {},
|
||
"output_type": "execute_result"
|
||
}
|
||
],
|
||
"source": [
|
||
"sampling_callback = keras.callbacks.LambdaCallback(\n",
|
||
" on_epoch_end = lambda batch, logs: print(generate(model))\n",
|
||
")\n",
|
||
"\n",
|
||
"model.fit(ds_train.batch(8).map(title_batch_fn),callbacks=[sampling_callback],epochs=3)"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"這個範例已經生成了一些相當不錯的文本,但仍有多種方式可以進一步改進:\n",
|
||
"\n",
|
||
"* **更多文本**。我們僅使用了標題作為任務內容,但您可能希望嘗試使用完整文本。請記住,RNN在處理長序列方面表現不佳,因此可以將文本拆分成較短的句子,或者始終在固定的序列長度(例如,`num_chars`,假設為256)上進行訓練。您可以嘗試將上述範例改為這種架構,並參考[官方 Keras 教學](https://keras.io/examples/generative/lstm_character_level_text_generation/)作為靈感。\n",
|
||
"\n",
|
||
"* **多層 LSTM**。嘗試使用2或3層的LSTM單元是有意義的。如我們在前一單元提到的,每層LSTM會從文本中提取特定的模式,而在字元級生成器的情況下,我們可以預期較低層的LSTM負責提取音節,而較高層則負責提取單詞及單詞組合。這可以通過向LSTM構造函數傳遞層數參數來簡單實現。\n",
|
||
"\n",
|
||
"* 您可能還希望嘗試使用**GRU單元**,看看哪種表現更好,以及嘗試**不同的隱藏層大小**。隱藏層過大可能導致過擬合(例如,網絡會學習精確的文本),而過小的隱藏層可能無法生成良好的結果。\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"## 軟性文本生成與溫度\n",
|
||
"\n",
|
||
"在之前 `generate` 的定義中,我們總是選擇機率最高的字元作為生成文本的下一個字元。這導致文本經常在相同的字元序列之間不斷「循環」,如下例所示:\n",
|
||
"```\n",
|
||
"today of the second the company and a second the company ...\n",
|
||
"```\n",
|
||
"\n",
|
||
"然而,如果我們觀察下一個字元的機率分佈,可能會發現幾個最高機率之間的差距並不大,例如一個字元的機率是 0.2,另一個是 0.19,等等。例如,在尋找序列 '*play*' 的下一個字元時,下一個字元可能同樣有可能是空格,或者是 **e**(如單詞 *player* 中的情況)。\n",
|
||
"\n",
|
||
"這讓我們得出一個結論:選擇機率最高的字元並不總是「公平」的,因為選擇第二高的字元仍然可能生成有意義的文本。更明智的做法是從網路輸出的機率分佈中**抽樣**字元。\n",
|
||
"\n",
|
||
"這種抽樣可以使用 `np.multinomial` 函數來完成,該函數實現了所謂的**多項分佈**。下面定義了一個實現這種**軟性**文本生成的函數:\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "code",
|
||
"execution_count": 33,
|
||
"metadata": {
|
||
"scrolled": true
|
||
},
|
||
"outputs": [
|
||
{
|
||
"name": "stdout",
|
||
"output_type": "stream",
|
||
"text": [
|
||
"\n",
|
||
"--- Temperature = 0.3\n",
|
||
"Today #39;s strike #39; to start at the store return\n",
|
||
"On Sunday PO to Be Data Profit Up (Reuters)\n",
|
||
"Moscow, SP wins straight to the Microsoft #39;s control of the space start\n",
|
||
"President olding of the blast start for the strike to pay <b>...</b>\n",
|
||
"Little red riding hood ficed to the spam countered in European <b>...</b>\n",
|
||
"\n",
|
||
"--- Temperature = 0.8\n",
|
||
"Today countie strikes ryder missile faces food market blut\n",
|
||
"On Sunday collores lose-toppy of sale of Bullment in <b>...</b>\n",
|
||
"Moscow, IBM Diffeiting in Afghan Software Hotels (Reuters)\n",
|
||
"President Ol Luster for Profit Peaced Raised (AP)\n",
|
||
"Little red riding hood dace on depart talks #39; bank up\n",
|
||
"\n",
|
||
"--- Temperature = 1.0\n",
|
||
"Today wits House buiting debate fixes #39; supervice stake again\n",
|
||
"On Sunday arling digital poaching In for level\n",
|
||
"Moscow, DS Up 7, Top Proble Protest Caprey Mamarian Strike\n",
|
||
"President teps help of roubler stepted lessabul-Dhalitics (AFP)\n",
|
||
"Little red riding hood signs on cash in Carter-youb\n",
|
||
"\n",
|
||
"--- Temperature = 1.3\n",
|
||
"Today wits flawer ro, pSIA figat's co DroftwavesIs Talo up\n",
|
||
"On Sunday hround elitwing wint EU Powerburlinetien\n",
|
||
"Moscow, Bazz #39;s sentries olymen winnelds' next for Olympite Huc?\n",
|
||
"President lost securitys from power Elections in Smiltrials\n",
|
||
"Little red riding hood vides profit, exponituity, profitmainalist-at said listers\n",
|
||
"\n",
|
||
"--- Temperature = 1.8\n",
|
||
"Today #39;It: He deat: N.KA Asside\n",
|
||
"On Sunday i arry Par aldeup patient Wo stele1\n"
|
||
]
|
||
},
|
||
{
|
||
"ename": "KeyError",
|
||
"evalue": "0",
|
||
"output_type": "error",
|
||
"traceback": [
|
||
"\u001b[0;31m---------------------------------------------------------------------------\u001b[0m",
|
||
"\u001b[0;31mKeyError\u001b[0m Traceback (most recent call last)",
|
||
"\u001b[0;32m<ipython-input-33-db32367a0feb>\u001b[0m in \u001b[0;36m<module>\u001b[0;34m\u001b[0m\n\u001b[1;32m 18\u001b[0m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34mf\"\\n--- Temperature = {i}\"\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 19\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mj\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mrange\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;36m5\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 20\u001b[0;31m \u001b[0mprint\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mgenerate_soft\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mmodel\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0msize\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0;36m300\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0mstart\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mwords\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mj\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0mtemperature\u001b[0m\u001b[0;34m=\u001b[0m\u001b[0mi\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||
"\u001b[0;32m<ipython-input-33-db32367a0feb>\u001b[0m in \u001b[0;36mgenerate_soft\u001b[0;34m(model, size, start, temperature)\u001b[0m\n\u001b[1;32m 11\u001b[0m \u001b[0mchars\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mappend\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mnc\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 12\u001b[0m \u001b[0minp\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0minp\u001b[0m\u001b[0;34m+\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mnc\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m---> 13\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0mdecode\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mchars\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m\u001b[1;32m 14\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 15\u001b[0m \u001b[0mwords\u001b[0m \u001b[0;34m=\u001b[0m \u001b[0;34m[\u001b[0m\u001b[0;34m'Today '\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m'On Sunday '\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m'Moscow, '\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m'President '\u001b[0m\u001b[0;34m,\u001b[0m\u001b[0;34m'Little red riding hood '\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n",
|
||
"\u001b[0;32m<ipython-input-10-3f5fa6130b1d>\u001b[0m in \u001b[0;36mdecode\u001b[0;34m(x)\u001b[0m\n\u001b[1;32m 2\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 3\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0mdecode\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 4\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0;34m''\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mjoin\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mreverse_map\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mt\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mt\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mx\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||
"\u001b[0;32m<ipython-input-10-3f5fa6130b1d>\u001b[0m in \u001b[0;36m<listcomp>\u001b[0;34m(.0)\u001b[0m\n\u001b[1;32m 2\u001b[0m \u001b[0;34m\u001b[0m\u001b[0m\n\u001b[1;32m 3\u001b[0m \u001b[0;32mdef\u001b[0m \u001b[0mdecode\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0mx\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m:\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0;32m----> 4\u001b[0;31m \u001b[0;32mreturn\u001b[0m \u001b[0;34m''\u001b[0m\u001b[0;34m.\u001b[0m\u001b[0mjoin\u001b[0m\u001b[0;34m(\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mreverse_map\u001b[0m\u001b[0;34m[\u001b[0m\u001b[0mt\u001b[0m\u001b[0;34m]\u001b[0m \u001b[0;32mfor\u001b[0m \u001b[0mt\u001b[0m \u001b[0;32min\u001b[0m \u001b[0mx\u001b[0m\u001b[0;34m]\u001b[0m\u001b[0;34m)\u001b[0m\u001b[0;34m\u001b[0m\u001b[0;34m\u001b[0m\u001b[0m\n\u001b[0m",
|
||
"\u001b[0;31mKeyError\u001b[0m: 0"
|
||
]
|
||
}
|
||
],
|
||
"source": [
|
||
"def generate_soft(model,size=100,start='Today ',temperature=1.0):\n",
|
||
" inp = tokenizer.texts_to_sequences([start])[0]\n",
|
||
" chars = inp\n",
|
||
" for i in range(size):\n",
|
||
" out = model(tf.expand_dims(tf.one_hot(inp,vocab_size),0))[0][-1]\n",
|
||
" probs = tf.exp(tf.math.log(out)/temperature).numpy().astype(np.float64)\n",
|
||
" probs = probs/np.sum(probs)\n",
|
||
" nc = np.argmax(np.random.multinomial(1,probs,1))\n",
|
||
" if nc==eos_token:\n",
|
||
" break\n",
|
||
" chars.append(nc)\n",
|
||
" inp = inp+[nc]\n",
|
||
" return decode(chars)\n",
|
||
"\n",
|
||
"words = ['Today ','On Sunday ','Moscow, ','President ','Little red riding hood ']\n",
|
||
" \n",
|
||
"for i in [0.3,0.8,1.0,1.3,1.8]:\n",
|
||
" print(f\"\\n--- Temperature = {i}\")\n",
|
||
" for j in range(5):\n",
|
||
" print(generate_soft(model,size=300,start=words[j],temperature=i))"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"我們引入了一個名為 **溫度** 的參數,用於指示我們應該多大程度地堅持最高概率。如果溫度為 1.0,我們進行公平的多項式抽樣,而當溫度趨於無窮大時,所有概率變得相等,我們隨機選擇下一個字符。在下面的例子中,我們可以觀察到當溫度增加過多時,文本變得毫無意義,而當溫度接近 0 時,則類似於「循環」的硬生成文本。\n"
|
||
]
|
||
},
|
||
{
|
||
"cell_type": "markdown",
|
||
"metadata": {},
|
||
"source": [
|
||
"\n---\n\n**免責聲明**: \n本文件使用 AI 翻譯服務 [Co-op Translator](https://github.com/Azure/co-op-translator) 進行翻譯。我們致力於提供準確的翻譯,但請注意,自動翻譯可能包含錯誤或不準確之處。應以原文文件作為權威來源。對於關鍵資訊,建議尋求專業人工翻譯。我們對因使用此翻譯而引起的任何誤解或誤讀概不負責。\n"
|
||
]
|
||
}
|
||
],
|
||
"metadata": {
|
||
"interpreter": {
|
||
"hash": "16af2a8bbb083ea23e5e41c7f5787656b2ce26968575d8763f2c4b17f9cd711f"
|
||
},
|
||
"kernelspec": {
|
||
"display_name": "Python 3.8.12 ('py38')",
|
||
"language": "python",
|
||
"name": "python3"
|
||
},
|
||
"language_info": {
|
||
"codemirror_mode": {
|
||
"name": "ipython",
|
||
"version": 3
|
||
},
|
||
"file_extension": ".py",
|
||
"mimetype": "text/x-python",
|
||
"name": "python",
|
||
"nbconvert_exporter": "python",
|
||
"pygments_lexer": "ipython3",
|
||
"version": "3.8.12"
|
||
},
|
||
"coopTranslator": {
|
||
"original_hash": "9fbb7d5fda708537649f71f5f646fcde",
|
||
"translation_date": "2025-08-31T10:31:48+00:00",
|
||
"source_file": "lessons/5-NLP/17-GenerativeNetworks/GenerativeTF.ipynb",
|
||
"language_code": "tw"
|
||
}
|
||
},
|
||
"nbformat": 4,
|
||
"nbformat_minor": 4
|
||
} |