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 }