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

495 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",
"循環神經網絡RNNs及其門控單元變體例如長短期記憶單元LSTMs和門控循環單元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": []
},
{
"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' 的範例圖片。](../../../../../translated_images/rnn-generate.56c54afb52f9781d.mo.png)\n",
"\n",
"對於序列中的最後一個字元,我們會要求網路生成 `<eos>` 標記。\n",
"\n",
"這裡使用的生成式 RNN 與其他的主要區別在於,我們會從 RNN 的每一步輸出中取結果,而不僅僅是從最後一個單元格中取結果。這可以通過為 RNN 單元指定 `return_sequences` 參數來實現。\n",
"\n",
"因此,在訓練過程中,網路的輸入將是一個特定長度的編碼字元序列,而輸出則是一個相同長度的序列,但向後偏移一個元素並以 `<eos>` 結尾。小批次minibatch將由多個這樣的序列組成我們需要使用**填充padding**來對齊所有序列。\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",
"由於網路以字元作為輸入,且詞彙表的大小相對較小,我們不需要嵌入層,直接將 one-hot 編碼的輸入傳遞給 LSTM 單元即可。輸出層將是一個 `Dense` 分類器,它會將 LSTM 的輸出轉換為 one-hot 編碼的標記編號。\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` 函數;然而,這個方法在字元級別的標記化中效果並不好。因此,我們將從 tokenizer 中取得一個標記的字典(稱為 `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-28T12:01:00+00:00",
"source_file": "lessons/5-NLP/17-GenerativeNetworks/GenerativeTF.ipynb",
"language_code": "mo"
}
},
"nbformat": 4,
"nbformat_minor": 4
}