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
        }