yetrun's picture
ver3: 将源码迁入 src/deep_learning 包,重塑训练流水线,规范 data/model 契约
07cb7d3
Raw
History Blame Contribute Delete
2.64 kB
from dataclasses import dataclass
from functools import partial
from pathlib import Path
import keras
from keras import layers
from .gpt_components import PositionalEmbedding, TransformerDecoder
from deep_learning.pipeline.impl.text_generation.generation import generate_with_training_model
from deep_learning.models.spec import TextGenerationModel, TextGenerationModelBuilder
@dataclass
class GptModelBuilder(TextGenerationModelBuilder):
hidden_dim: int
intermediate_dim: int
num_heads: int
num_layers: int
def build_training_artifact(
self,
vocab_size: int,
sequence_length: int
) -> TextGenerationModel:
inputs = keras.Input(shape=(None,), dtype="int32", name="inputs")
embedding = PositionalEmbedding(
sequence_length,
vocab_size,
self.hidden_dim,
name="embedding"
)
x = embedding(inputs)
x = layers.LayerNormalization(name="input_layer_norm")(x)
for i in range(self.num_layers):
decoder = TransformerDecoder(
self.hidden_dim,
self.intermediate_dim,
self.num_heads,
name=f"decoder_{i}"
)
x = decoder(x)
# 这里 reverse 复用 embedding 权重生成输出 logits
outputs = embedding(x, reverse=True)
model = keras.Model(inputs, outputs, name="mini_gpt")
return TextGenerationModel(
model=model,
generate=partial(generate_with_training_model, model)
)
def compile_training_model(self, model: keras.Model) -> None:
from deep_learning.pipeline.impl.text_generation.pipeline import WarmupSchedule
model.compile(
optimizer=keras.optimizers.Adam(WarmupSchedule()),
loss=keras.losses.SparseCategoricalCrossentropy(from_logits=True),
metrics=["accuracy"]
)
def load_inference_artifact(self, model_path: Path) -> TextGenerationModel:
model = keras.models.load_model(
str(model_path),
custom_objects=self._custom_objects()
)
return TextGenerationModel(
model=model,
generate=partial(generate_with_training_model, model)
)
def _custom_objects(self) -> dict:
# ASK: 这个组件要规范下位置
from deep_learning.pipeline.impl.text_generation.pipeline import WarmupSchedule
return {
"WarmupSchedule": WarmupSchedule,
"PositionalEmbedding": PositionalEmbedding,
"TransformerDecoder": TransformerDecoder
}