Keras + RNN の文章生成チュートリアルをリファクタしてみた
概要
過去にやった文章生成のチュートリアルをリファクタして更に理解を深めてみました
環境
- macOS 14.0
- Python 3.11.6
- keras 2.14.0
- tensorflow 2.14.0
- matplotlib 3.8.1
サンプルコード
import os
from dataclasses import dataclass
from typing import Optional
import numpy as np
import tensorflow as tf
from keras.src.callbacks import History
from tensorflow.keras.callbacks import ModelCheckpoint
from tensorflow.keras.layers import GRU, Dense, Embedding
from tensorflow.keras.losses import sparse_categorical_crossentropy
from tensorflow.keras.models import Sequential
from tensorflow.keras.utils import get_file
class TestData:
SEQ_LENGTH = 100
BUFFER_SIZE = 10000
BATCH_SIZE = 64
def __init__(self) -> None:
self.__download()
self.char2idx = self.__gen_char2idx()
self.idx2char = self.__gen_idx2char()
self.text_as_int = np.array([self.char2idx[c] for c in self.text])
self.char_dataset = tf.data.Dataset.from_tensor_slices(self.text_as_int)
self.sequences = self.char_dataset.batch(
self.SEQ_LENGTH + 1, drop_remainder=True
)
dataset = self.sequences.map(self.shift)
self.dataset = dataset.shuffle(self.BUFFER_SIZE).batch(
self.BATCH_SIZE, drop_remainder=True
)
@property
def vocab(self) -> list:
return sorted(set(self.text))
def __download(self):
path_to_file = get_file(
"shakespeare.txt",
"https://storage.googleapis.com/download.tensorflow.org/data/shakespeare.txt",
)
self.text = open(path_to_file, "rb").read().decode(encoding="utf-8")
def __gen_char2idx(self) -> dict:
return {u: i for i, u in enumerate(self.vocab)}
def __gen_idx2char(self) -> np.ndarray:
return np.array(self.vocab)
def shift(self, chunk: tf.raw_ops.BatchDataset):
independent = chunk[:-1]
dependent = chunk[1:]
return independent, dependent
class RNNModel:
EMBEDDING_DIM = 256
RUN_UNITS = 1024
model: Sequential
def __init__(self, vocab_size: int, batch_size) -> None:
self.model = Sequential()
self.vocab_size = vocab_size
self.batch_size = batch_size
def build(self, batch_size: Optional[int] = None, new: bool = False):
if batch_size is None:
bs = self.batch_size
else:
bs = batch_size
if new:
self.model = Sequential()
self.model.add(
Embedding(
self.vocab_size,
self.EMBEDDING_DIM,
batch_input_shape=[bs, None],
)
)
self.model.add(
GRU(
self.RUN_UNITS,
return_sequences=True,
stateful=True,
recurrent_initializer="glorot_uniform",
)
)
self.model.add(Dense(self.vocab_size))
def compile(self) -> None:
self.model.compile(optimizer="adam", loss=Loss().keras_scc)
def show(self):
self.model.summary()
def train(
self,
dataset: tf.raw_ops.BatchDataset,
epochs=10,
callbacks=[],
) -> History:
history = self.model.fit(dataset, epochs=epochs, callbacks=callbacks)
return history
def save(self, file_name="my_model"):
self.model.save(file_name)
def rebuild(self, checkpoint_dir: str = "./training_checkpoints"):
self.build(batch_size=1, new=True)
self.model.load_weights(tf.train.latest_checkpoint(checkpoint_dir))
self.model.build(tf.TensorShape([1, None]))
class Loss:
def keras_scc(self, y_true, y_pred):
return sparse_categorical_crossentropy(y_true, y_pred, from_logits=True)
class Callback:
@classmethod
def save_model(cls) -> ModelCheckpoint:
checkpoint_dir = "./training_checkpoints"
checkpoint_prefix = os.path.join(checkpoint_dir, "ckpt_{epoch}")
return ModelCheckpoint(filepath=checkpoint_prefix, save_weights_only=True)
@dataclass
class Result:
text: str
def show(self):
print(self.text)
class Test:
TEMPERRATURE = 1.0
NUM_GENERATE_CHAR = 1000
def __init__(
self, model: RNNModel, test_data: TestData, start_string: str = "ROMEO: "
) -> None:
self.model = model
self.test_data = test_data
self.start_string = start_string
input_eval = [self.test_data.char2idx[char] for char in start_string]
self.input_eval = tf.expand_dims(input_eval, 0)
def run(self) -> Result:
self.model.model.reset_states()
generated_text = []
for _ in range(self.NUM_GENERATE_CHAR):
predictions = self.model.model(self.input_eval)
predictions = tf.squeeze(predictions, 0)
predictions = predictions / self.TEMPERRATURE
predicted_id = tf.random.categorical(predictions, num_samples=1)[
-1, 0
].numpy()
self.input_eval = tf.expand_dims([predicted_id], 0)
generated_text.append(self.test_data.idx2char[predicted_id])
return Result(text=self.start_string + "".join(generated_text))
if __name__ == "__main__":
test_data = TestData()
model = RNNModel(len(test_data.vocab), TestData.BATCH_SIZE)
model.build()
model.compile()
model.train(test_data.dataset, callbacks=[Callback.save_model()])
model.rebuild()
model.show()
test = Test(model, test_data)
result = test.run()
result.show()
ちょっと解説
テストデータは単純なテキストを使います
テキスト情報を1文字ずつ分解し数字とのマッピング情報を作成しテキストをすべて数値化することで学習データを作成します
文字情報の数値化は自然言語ではよくある手法になります
データをシャッフルしたりバッチサイズで分割するのは Keras で学習させる際のフォーマットに変換していると理解していますがこの辺りの手法を自然に使いこなせるレベルにならないとダメかなと思います
モデルの生成は Embedding -> GRU -> Dense とシンプルです
Embedding は自然言語の学習ではほぼ必須です
GRU (Gated Recurrent Unit) は RNN の一種です
モデル学習時にチェックポイントを保存しています
評価する際に入力をシンプルにしたいので学習時のバッチサイズ64を再度チェックポイントからモデルをビルドしてバッチサイズを1にするためです
予測した値はそのまま使わずカテゴリー分布という手法を使って再度計算させています (正直理由は不明です
リファクタリング不足点
- TestData をもう少しリファクタリングしたほうがいい
- データの生成を init ではなくそれぞれの関数にしたほうがいい
- 説明変数と目的変数を別のクラスで管理した方がいい
- Model の model の管理をリファクタリングしたほうがいい
- build で新規 or 既存の書き換えを制御するよりかは ModelFactory を作って model を個別に生成できるような仕組みにするといいかも
- Test を柔軟にしたほうがいい
- いろいろなデータでテストできるようにしたほうがいい
- 生成文字列サイズなども可変にできるといい
- 何をやっているかわからないところがあるのでコメントで補完したい
最後に
Keras を使った学習方法や流れはだいたい把握できましたがフルスクラッチでゼロから numpy や keras を使ってモデルを作るにはまだまだ学習が足りないと感じました
numpy で生成したベクトルの扱い方や keras のレイヤーの生成方法やレイヤーに対する入出力の制御、最適化あたりを駆使できるようにならないとオリジナルなモデルを作るのは難しいなと感じました
参考サイト