Spaces:
Sleeping
Sleeping
| 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 | |
| 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 | |
| } | |