Spaces:
Sleeping
Sleeping
File size: 2,639 Bytes
07cb7d3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 | 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
}
|