{ "cells": [ { "cell_type": "markdown", "metadata": {}, "source": [ "# 生成ネットワーク\n", "\n", "リカレントニューラルネットワーク (RNN) とそのゲート付きセルのバリエーション(例えば、長短期記憶セル (LSTM) やゲート付きリカレントユニット (GRU))は、言語モデル化の仕組みを提供します。つまり、これらは単語の順序を学習し、シーケンス内の次の単語を予測することができます。この特性により、RNNを使って**生成タスク**を実行することが可能になります。例えば、通常のテキスト生成、機械翻訳、さらには画像キャプション生成などです。\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`レイヤーではそれができないため、以下の2つの選択肢があります:\n", "\n", "* テキストを手動で読み込み、[この公式Kerasの例](https://keras.io/examples/generative/lstm_character_level_text_generation/)のように手作業でトークン化を行う\n", "* 文字レベルのトークン化に`Tokenizer`クラスを使用する\n", "\n", "ここでは2つ目の選択肢を選びます。`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": [ "シーケンスの終了を示す特別なトークンとして****を使用したいと考えています。それを語彙に手動で追加しましょう。\n" ] }, { "cell_type": "code", "execution_count": 3, "metadata": {}, "outputs": [], "source": [ "eos_token = len(tokenizer.word_index)+1\n", "tokenizer.word_index[''] = 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をトレーニングする方法は以下の通りです。各ステップで1つのタイトルを取り出し、それをRNNに入力します。そして、各入力文字に対してネットワークに次の出力文字を生成させます。\n", "\n", "![単語 'HELLO' を生成するRNNの例を示す画像。](../../../../../translated_images/ja/rnn-generate.56c54afb52f9781d.webp)\n", "\n", "シーケンスの最後の文字に対しては、ネットワークに `` トークンを生成させます。\n", "\n", "ここで使用する生成型RNNの主な違いは、RNNの最終セルだけでなく、各ステップの出力を利用する点です。これを実現するには、RNNセルに `return_sequences` パラメータを指定します。\n", "\n", "したがって、トレーニング中のネットワークへの入力は、ある長さのエンコードされた文字列のシーケンスであり、出力は同じ長さのシーケンスですが、1つの要素分シフトされ、最後に `` で終了します。ミニバッチは複数のそのようなシーケンスで構成され、すべてのシーケンスを揃えるために**パディング**を使用する必要があります。\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", "* 最後に、すべての文字をワンホットエンコードし、シフトと `` の追加も行います。なぜワンホットエンコードされた文字が必要なのかは、すぐにわかるでしょう\n", "\n", "ただし、この関数は **Pythonic** であり、つまり Tensorflow の計算グラフに自動的に変換することはできません。この関数を直接 `Dataset.map` 関数で使用しようとするとエラーが発生します。この Pythonic な呼び出しを `py_function` ラッパーを使用して囲む必要があります:\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": [ "> **注**: Python的な変換関数とTensorflowの変換関数を区別するのは少し複雑に感じるかもしれません。また、なぜデータセットを`fit`に渡す前に標準的なPython関数を使って変換しないのか疑問に思うかもしれません。確かにそれも可能ですが、`Dataset.map`を使用することには大きな利点があります。それは、データ変換パイプラインがTensorflowの計算グラフを使用して実行されるため、GPU計算を活用でき、CPUとGPU間でデータをやり取りする必要性が最小限に抑えられる点です。\n", "\n", "さて、ジェネレーターネットワークを構築し、トレーニングを開始しましょう。これは、前のユニットで議論した任意の再帰型セル(シンプル、LSTM、またはGRU)に基づくことができます。この例ではLSTMを使用します。\n", "\n", "ネットワークは文字を入力として受け取り、語彙サイズが比較的小さいため、埋め込み層は必要ありません。一つのホットエンコードされた入力を直接LSTMセルに渡すことができます。出力層は`Dense`分類器となり、LSTMの出力を一つのホットエンコードされたトークン番号に変換します。\n", "\n", "さらに、可変長のシーケンスを扱うため、`Masking`層を使用して、文字列のパディング部分を無視するマスクを作成することができます。これは厳密には必要ではありません。なぜなら、``トークン以降の内容にはあまり興味がないからです。しかし、この層タイプの経験を得るために使用してみます。`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": [ "" ] }, "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` を使用して最も確率の高いトークン番号を見つけ、その文字を生成されたトークンのリストに追加します。そして、生成を続けます。このようにして1文字を生成するプロセスを `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": [ "" ] }, "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**。LSTMセルを2層または3層にするのも試す価値があります。前のユニットで述べたように、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", "このことから、常に確率が最も高い文字を選ぶのが「公平」とは限らないという結論に至ります。2番目に高い確率の文字を選んでも、意味のあるテキストにつながる可能性があるからです。ネットワークの出力による確率分布から文字を**サンプリング**する方が賢明です。\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\u001b[0m in \u001b[0;36m\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\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\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\u001b[0m in \u001b[0;36m\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つのパラメーターを導入しました。これは、最高確率にどれだけ厳密に従うべきかを示すために使用されます。温度が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-30T10:09:51+00:00", "source_file": "lessons/5-NLP/17-GenerativeNetworks/GenerativeTF.ipynb", "language_code": "ja" } }, "nbformat": 4, "nbformat_minor": 4 }