AI-For-Beginners/translations/zh-TW/lessons/5-NLP/17-GenerativeNetworks/GenerativeTF.ipynb

497 lines
26 KiB
Plaintext
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

{
"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",
"![顯示 RNN 生成單詞 'HELLO' 的示例圖像。](../../../../../lessons/5-NLP/17-GenerativeNetworks/images/rnn-generate.png)\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 &lt;b&gt;...&lt;/b&gt;\n",
"Little red riding hood ficed to the spam countered in European &lt;b&gt;...&lt;/b&gt;\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 &lt;b&gt;...&lt;/b&gt;\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
}