欢迎光临

TensorFlow 2.x 自定义层与模型子类化深度实战:从 Layer 到 Model 的完整工程化指南

在 TensorFlow 2.x 的 Keras API 中,

1
tf.keras.layers.Layer

1
tf.keras.Model

是构建一切深度学习模型的基础。虽然内置层(Dense、Conv2D、LSTM 等)覆盖了大部分常见需求,但在实际工程中,我们经常需要创建自定义层来实现特殊的计算逻辑、非标准的激活函数、或复杂的注意力机制。同时,模型子类化(Model Subclassing)赋予了开发者最大的灵活性,让我们能够突破 Sequential 和 Functional API 的限制,实现动态计算图、多输入输出、自定义训练步骤等高级功能。

本文将从零开始,系统讲解 TensorFlow 2.x 中自定义层和模型子类化的完整工程化实践,涵盖层的三件套(

1
__init__

1
build

1
call

)、序列化与配置恢复、模型子类化的最佳实践、自定义训练循环的集成,以及生产环境中常见的陷阱与解决方案。

深度学习模型架构

一、自定义层的核心:理解 Layer 生命周期

在 Keras 中,一个自定义层的完整生命周期包含三个核心方法。理解这三者的职责边界和调用时机,是写出高质量自定义层的前提。

1.1

1
__init__

:声明式配置

1
__init__

方法负责接收超参数并保存,不创建任何权重。这是一个纯粹的声明阶段——我们只记录层需要什么配置,但不知道输入的形状,因此无法在此创建权重张量。


1
2
3
4
5
6
7
8
9
10
11
12
import tensorflow as tf

class GatedLinearUnit(tf.keras.layers.Layer):
    """门控线性单元:将输入分为两路,一路过 sigmoid 生成门控信号,另一路线性变换后逐元素相乘。"""

    def __init__(self, units, dropout_rate=0.0, use_bias=True, **kwargs):
        super().__init__(**kwargs)
        self.units = units
        self.dropout_rate = dropout_rate
        self.use_bias = use_bias

    # build 和 call 将在后续实现

1.2

1
build

:延迟创建权重

1
build

方法在层第一次收到输入时被调用,此时我们知道了输入形状,可以根据它创建对应维度的权重。这种延迟创建(Lazy Instantiation)机制使得自定义层能自动适配不同形状的输入,而无需在

1
__init__

中硬编码维度。


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
    def build(self, input_shape):
        # input_shape 是 TensorShape 对象,例如 (None, 128)
        last_dim = input_shape[-1]

        # 线性变换权重
        self.w = self.add_weight(
            name="kernel",
            shape=(last_dim, self.units * 2),  # 2x units: 一路给门控,一路给线性
            initializer="glorot_uniform",
            trainable=True,
        )

        if self.use_bias:
            self.b = self.add_weight(
                name="bias",
                shape=(self.units * 2,),
                initializer="zeros",
                trainable=True,
            )

        # 创建 dropout 层
        self.dropout = tf.keras.layers.Dropout(self.dropout_rate)

        # 必须调用父类 build
        super().build(input_shape)
1
add_weight

是 Keras 提供的权重注册方法,它会自动将权重纳入

1
trainable_weights

1
non_trainable_weights

列表,确保在训练时被正确优化。你也可以使用

1
tf.Variable

手动创建权重,但必须通过

1
self._trainable_weights.append()

手动注册,否则优化器无法感知这些权重。

1.3

1
call

:前向计算逻辑

1
call

方法实现层的前向传播。它接收输入张量,执行计算,返回输出张量。这是自定义层的核心——你所有的计算逻辑都在这里。


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
    def call(self, inputs, training=None):
        # 线性变换
        output = tf.matmul(inputs, self.w)
        if self.use_bias:
            output = output + self.b

        # 分裂为门控信号和线性信号
        gate, linear = tf.split(output, 2, axis=-1)

        # 门控信号过 sigmoid
        gate = tf.sigmoid(gate)

        # training 模式下应用 dropout
        linear = self.dropout(linear, training=training)

        # 逐元素相乘
        return gate * linear

注意

1
training

参数:Keras 会在

1
model.fit()

时自动传入

1
training=True

,而在

1
model.predict()

1
model.evaluate()

时传入

1
training=False

。如果你的层包含 Dropout 或 BatchNormalization 等训练/推理行为不同的逻辑,必须正确处理这个参数。

神经网络层结构

二、序列化与配置恢复:让自定义层可保存

在生产环境中,模型需要被保存、加载、传输。如果你的自定义层不支持序列化,

1
model.save()

1
tf.saved_model.save()

都会失败。Keras 的序列化机制依赖两个方法:

1
get_config

1
from_config

2.1 实现

1
get_config

1
get_config

返回一个字典,包含重建该层所需的全部超参数。父类的配置也必须包含在内。


1
2
3
4
5
6
7
8
    def get_config(self):
        config = super().get_config()  # 获取父类配置(含 name, dtype 等)
        config.update({
            "units": self.units,
            "dropout_rate": self.dropout_rate,
            "use_bias": self.use_bias,
        })
        return config

2.2 实现

1
from_config

(可选)

默认的

1
from_config

实现将 config 字典作为

1
**kwargs

传入

1
__init__

。如果你的

1
__init__

参数名和 config 的 key 完全一致,就不需要覆盖这个方法。但如果需要额外的反序列化逻辑(比如将字符串转回对象),则需要自定义:


1
2
3
4
5
    @classmethod
    def from_config(cls, config):
        # 如果 config 中有需要特殊处理的字段,在这里做
        # 例如:config["activation"] = tf.keras.activations.get(config["activation"])
        return cls(**config)

2.3 注册自定义层(关键步骤)

即使实现了

1
get_config

1
model.save()

仍然可能报错——因为 Keras 在反序列化时需要通过类名字符串找到对应的类。你必须将自定义层注册到 Keras 的全局对象字典中:


1
2
3
4
5
6
7
# 方式一:使用装饰器注册
@tf.keras.utils.register_keras_serializable(package="MyLayers")
class GatedLinearUnit(tf.keras.layers.Layer):
    ...

# 方式二:手动注册
tf.keras.utils.get_custom_objects()["GatedLinearUnit"] = GatedLinearUnit

注册后,

1
model.save("model.keras")

1
tf.keras.models.load_model("model.keras")

就能正常工作了。注意:加载模型时,自定义类的 Python 定义必须在 Python 进程中可用——Keras 不会自动重建代码,只会用

1
from_config

重建对象。

三、模型子类化:突破 API 限制的终极武器

Keras 提供三种构建模型的 API:Sequential、Functional、Model Subclassing。前两者适合静态计算图,而 Model Subclassing 则赋予你完全的控制权——动态条件分支、循环、递归、多阶段计算,一切 Python 能做的,子类化模型都能做。

3.1 基础模型子类化


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
class TransformerBlock(tf.keras.Model):
    def __init__(self, embed_dim, num_heads, ff_dim, dropout_rate=0.1, **kwargs):
        super().__init__(**kwargs)
        self.att = tf.keras.layers.MultiHeadAttention(
            num_heads=num_heads, key_dim=embed_dim // num_heads
        )
        self.ffn = tf.keras.Sequential([
            tf.keras.layers.Dense(ff_dim, activation="gelu"),
            tf.keras.layers.Dropout(dropout_rate),
            tf.keras.layers.Dense(embed_dim),
        ])
        self.layernorm1 = tf.keras.layers.LayerNormalization(epsilon=1e-6)
        self.layernorm2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)
        self.dropout1 = tf.keras.layers.Dropout(dropout_rate)
        self.dropout2 = tf.keras.layers.Dropout(dropout_rate)

    def call(self, inputs, training=None, mask=None):
        # Multi-Head Attention + 残差 + LayerNorm
        attn_output = self.att(inputs, inputs, attention_mask=mask)
        attn_output = self.dropout1(attn_output, training=training)
        out1 = self.layernorm1(inputs + attn_output)

        # Feed-Forward + 残差 + LayerNorm
        ffn_output = self.ffn(out1, training=training)
        ffn_output = self.dropout2(ffn_output, training=training)
        return self.layernorm2(out1 + ffn_output)

3.2 动态计算图:条件分支与循环

模型子类化的真正威力在于可以在

1
call

中使用 Python 原生的控制流。以下是一个根据输入长度动态调整层数的模型:


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
class AdaptiveDepthModel(tf.keras.Model):
    """根据输入复杂度动态决定计算深度。"""

    def __init__(self, units, max_depth=6, threshold=0.1, **kwargs):
        super().__init__(**kwargs)
        self.max_depth = max_depth
        self.threshold = threshold

        self.shared_layer = tf.keras.layers.Dense(units, activation="relu")
        self.exit_classifiers = [
            tf.keras.layers.Dense(10, name=f"exit_{i}")
            for i in range(max_depth)
        ]
        self.confidence_gate = tf.keras.layers.Dense(1, activation="sigmoid")

    def call(self, inputs, training=None):
        x = self.shared_layer(inputs)
        outputs = []

        for depth in range(self.max_depth):
            logits = self.exit_classifiers[depth](x)
            outputs.append(logits)

            # 推理模式下:如果置信度足够高,提前退出
            if not training:
                confidence = self.confidence_gate(x)
                if confidence > self.threshold:
                    return logits  # 提前返回,节省计算

            x = tf.keras.layers.Dense(
                x.shape[-1], activation="relu"
            )(x)  # 继续加深

        # 训练模式:返回最深层的输出
        return outputs[-1]

这种自适应深度(Adaptive Depth / Early Exit)机制在边缘设备部署中非常有价值——简单样本快速返回,复杂样本深入计算,有效降低平均推理延迟。

模型架构设计

四、自定义训练步骤:与

1
model.fit

无缝集成

模型子类化的另一个核心场景是自定义

1
train_step

。这让你既能享受

1
model.fit()

的便利(进度条、回调、分布式策略),又能完全控制每一步的训练逻辑。

4.1 覆写

1
train_step


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
class ContrastiveModel(tf.keras.Model):
    """对比学习模型,包含自定义训练步骤。"""

    def __init__(self, encoder, projection_dim=128, temperature=0.07, **kwargs):
        super().__init__(**kwargs)
        self.encoder = encoder
        self.projection = tf.keras.Sequential([
            tf.keras.layers.Dense(512, activation="relu"),
            tf.keras.layers.Dense(projection_dim),
        ])
        self.temperature = temperature

    def call(self, inputs, training=None):
        features = self.encoder(inputs, training=training)
        projections = self.projection(features, training=training)
        # L2 归一化
        return tf.math.l2_normalize(projections, axis=-1)

    def train_step(self, data):
        # data 是 fit() 传入的每个 batch
        (view1, view2), _ = data  # 对比学习的两个视图

        with tf.GradientTape() as tape:
            z1 = self(view1, training=True)
            z2 = self(view2, training=True)

            # 计算对比损失(NT-Xent / SimCLR 风格)
            loss = self._contrastive_loss(z1, z2)

        # 计算梯度并更新
        trainable_vars = self.trainable_variables
        gradients = tape.gradient(loss, trainable_vars)
        self.optimizer.apply_gradients(zip(gradients, trainable_vars))

        # 返回指标字典,供 Keras 显示和记录
        return {"contrastive_loss": loss}

    def _contrastive_loss(self, z1, z2):
        batch_size = tf.shape(z1)[0]
        # 拼接两个视图
        representations = tf.concat([z1, z2], axis=0)  # (2N, dim)
        similarity = tf.matmul(representations, representations, transpose_b=True)
        similarity = similarity / self.temperature

        # 掩码:排除自身相似度
        mask = tf.eye(2 * batch_size)
        similarity = similarity - mask * 1e9

        # 正对标签:z1[i] 的正对是 z2[i],反之亦然
        labels = tf.concat([
            tf.range(batch_size, 2 * batch_size),
            tf.range(0, batch_size),
        ], axis=0)

        loss = tf.nn.sparse_softmax_cross_entropy_with_logits(
            labels=labels, logits=similarity
        )
        return tf.reduce_mean(loss)

4.2 覆写

1
test_step

1
predict_step

同样的模式适用于评估和推理。如果你在

1
train_step

中做了特殊处理,务必也覆写

1
test_step

以保持一致:


1
2
3
4
5
6
7
8
9
10
11
12
    def test_step(self, data):
        (view1, view2), _ = data
        z1 = self(view1, training=False)
        z2 = self(view2, training=False)
        loss = self._contrastive_loss(z1, z2)
        return {"contrastive_loss": loss}

    def predict_step(self, data):
        # 推理时只需要编码器输出
        if isinstance(data, tuple):
            data = data[0]
        return self.encoder(data, training=False)

完成上述覆写后,你就可以像普通模型一样调用

1
model.fit(train_dataset, validation_data=val_dataset, epochs=100)

,所有 Keras 回调(ModelCheckpoint、EarlyStopping、TensorBoard)都能正常工作。

五、多输入输出与复杂拓扑:Functional API vs 子类化

对于复杂的多输入多输出模型,Functional API 和子类化各有优劣。下面通过一个多任务学习模型来对比两种实现。

5.1 Functional API 实现


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
# Functional API:静态图,可序列化,结构清晰
shared_input = tf.keras.Input(shape=(128,), name="shared_features")
task_a_input = tf.keras.Input(shape=(32,), name="task_a_features")
task_b_input = tf.keras.Input(shape=(64,), name="task_b_features")

# 共享主干
x = tf.keras.layers.Dense(256, activation="relu")(shared_input)
x = tf.keras.layers.Dense(128, activation="relu")(x)

# 任务 A 分支
a = tf.keras.layers.concatenate([x, task_a_input])
a = tf.keras.layers.Dense(64, activation="relu")(a)
output_a = tf.keras.layers.Dense(5, activation="softmax", name="task_a")(a)

# 任务 B 分支
b = tf.keras.layers.concatenate([x, task_b_input])
b = tf.keras.layers.Dense(64, activation="relu")(b)
output_b = tf.keras.layers.Dense(3, activation="softmax", name="task_b")(b)

model = tf.keras.Model(
    inputs=[shared_input, task_a_input, task_b_input],
    outputs=[output_a, output_b]
)

model.compile(
    optimizer="adam",
    loss={"task_a": "categorical_crossentropy", "task_b": "categorical_crossentropy"},
    loss_weights={"task_a": 1.0, "task_b": 0.5},
)

5.2 子类化实现(更灵活)


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
class MultiTaskModel(tf.keras.Model):
    def __init__(self, **kwargs):
        super().__init__(**kwargs)
        self.shared_backbone = tf.keras.Sequential([
            tf.keras.layers.Dense(256, activation="relu"),
            tf.keras.layers.Dense(128, activation="relu"),
        ])
        self.task_a_head = tf.keras.Sequential([
            tf.keras.layers.Dense(64, activation="relu"),
            tf.keras.layers.Dense(5, activation="softmax"),
        ])
        self.task_b_head = tf.keras.Sequential([
            tf.keras.layers.Dense(64, activation="relu"),
            tf.keras.layers.Dense(3, activation="softmax"),
        ])
        self.task_a_loss_tracker = tf.keras.metrics.Mean(name="task_a_loss")
        self.task_b_loss_tracker = tf.keras.metrics.Mean(name="task_b_loss")

    def call(self, inputs, training=None):
        shared, task_a_feat, task_b_feat = inputs
        shared_repr = self.shared_backbone(shared, training=training)

        a = tf.concat([shared_repr, task_a_feat], axis=-1)
        output_a = self.task_a_head(a, training=training)

        b = tf.concat([shared_repr, task_b_feat], axis=-1)
        output_b = self.task_b_head(b, training=training)

        return output_a, output_b

    @property
    def metrics(self):
        return [self.task_a_loss_tracker, self.task_b_loss_tracker]

    def train_step(self, data):
        (shared, a_feat, b_feat), (a_labels, b_labels) = data

        with tf.GradientTape() as tape:
            pred_a, pred_b = self((shared, a_feat, b_feat), training=True)
            loss_a = tf.keras.losses.categorical_crossentropy(a_labels, pred_a)
            loss_b = tf.keras.losses.categorical_crossentropy(b_labels, pred_b)
            # 动态损失权重:根据各任务的损失大小自动平衡
            total_loss = loss_a + 0.5 * loss_b

        gradients = tape.gradient(total_loss, self.trainable_variables)
        self.optimizer.apply_gradients(zip(gradients, self.trainable_variables))

        self.task_a_loss_tracker.update_state(loss_a)
        self.task_b_loss_tracker.update_state(loss_b)
        return {m.name: m.result() for m in self.metrics}

六、生产环境中的关键陷阱与最佳实践

自定义层和模型子类化在带来灵活性的同时,也引入了一些容易踩的坑。以下是我们团队在多个生产项目中总结的经验。

6.1

1
build

被多次调用的陷阱

Keras 保证

1
build

只被调用一次——在层第一次收到输入时。但如果你的模型在

1
call

中多次调用同一个子层并传入不同形状的输入,第二次调用时

1
build

不会再被触发,可能导致维度不匹配的运行时错误。


1
2
3
4
5
6
7
8
9
10
# 错误示范:同一个层处理不同形状的输入
class BadModel(tf.keras.Model):
    def __init__(self):
        super().__init__()
        self.dense = tf.keras.layers.Dense(64)  # 只能适配一种输入形状

    def call(self, inputs):
        x = self.dense(inputs)        # 假设输入 (None, 128),build 时 shape=(128, 64)
        y = self.dense(inputs[:, :32])  # 输入 (None, 32),但 dense 已 built,权重还是 (128, 64)
        return x + y  # RuntimeError!

解决方案:为不同形状的输入创建独立的层实例,或者在

1
__init__

中预先指定

1
input_shape

让层提前 build。

6.2 Eager vs Graph 模式的行为差异

子类化模型在 Eager 模式下可以使用 Python 原生控制流(if、for、print),但在

1
tf.function

包装后,Python 副作用只在 tracing 时执行一次。常见的坑:


1
2
3
4
5
6
7
8
9
10
class BuggyModel(tf.keras.Model):
    def call(self, inputs):
        # 这个 print 只在 tracing 时执行一次,不是每次调用都执行!
        print(f"Input shape: {inputs.shape}")

        # Python if 在 tf.function 中被静态追踪
        if tf.reduce_sum(inputs) > 0:  # 每次 tracing 只走一个分支
            return self.dense_a(inputs)
        else:
            return self.dense_b(inputs)

解决方案:使用

1
tf.cond

1
tf.print

替代 Python 原生控制流和打印:


1
2
3
4
5
6
7
8
9
class CorrectModel(tf.keras.Model):
    def call(self, inputs):
        tf.print("Input shape:", tf.shape(inputs))  # 每次执行都会打印

        return tf.cond(
            tf.reduce_sum(inputs) > 0,
            lambda: self.dense_a(inputs),
            lambda: self.dense_b(inputs),
        )

6.3 自定义层的

1
compute_output_shape

如果你的自定义层改变了输入的形状(如 Flatten、Reshape),覆写

1
compute_output_shape

方法可以帮助 Functional API 正确推断模型结构:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
class SqueezeLayer(tf.keras.layers.Layer):
    def __init__(self, axis=-1, **kwargs):
        super().__init__(**kwargs)
        self.axis = axis

    def call(self, inputs):
        return tf.squeeze(inputs, axis=self.axis)

    def compute_output_shape(self, input_shape):
        input_shape = list(input_shape)
        del input_shape[self.axis]
        return tf.TensorShape(input_shape)

    def get_config(self):
        config = super().get_config()
        config["axis"] = self.axis
        return config

6.4 内存泄漏:忘记释放计算图

在自定义训练循环中(不使用

1
model.fit

而是手写 for 循环),如果每次迭代都创建新的

1
GradientTape

和张量但不清理,会导致 GPU 内存持续增长:


1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
# 错误:tape 和中间张量不会被及时释放
for batch in dataset:
    with tf.GradientTape() as tape:
        predictions = model(batch, training=True)
        loss = loss_fn(labels, predictions)
    grads = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    # tape 的引用在循环结束后仍然存在,内存逐渐增长

# 正确:使用 del 显式释放
for batch in dataset:
    with tf.GradientTape() as tape:
        predictions = model(batch, training=True)
        loss = loss_fn(labels, predictions)
    grads = tape.gradient(loss, model.trainable_variables)
    optimizer.apply_gradients(zip(grads, model.trainable_variables))
    del tape  # 显式释放 tape

七、完整实战:构建一个可保存、可部署的自定义模型

最后,我们将以上所有知识点整合,构建一个完整的自定义模型——包含自定义层、子类化模型、自定义训练步骤,并确保它能被正确保存和加载。


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
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
import tensorflow as tf

@tf.keras.utils.register_keras_serializable(package="custom")
class ScaledDotProductAttention(tf.keras.layers.Layer):
    """缩放点积注意力层。"""

    def __init__(self, dropout_rate=0.1, **kwargs):
        super().__init__(**kwargs)
        self.dropout_rate = dropout_rate

    def build(self, input_shape):
        self.dropout = tf.keras.layers.Dropout(self.dropout_rate)
        super().build(input_shape)

    def call(self, query, key, value, mask=None, training=None):
        d_k = tf.cast(tf.shape(key)[-1], tf.float32)
        scores = tf.matmul(query, key, transpose_b=True) / tf.math.sqrt(d_k)

        if mask is not None:
            scores += (mask * -1e9)

        weights = tf.nn.softmax(scores, axis=-1)
        weights = self.dropout(weights, training=training)

        return tf.matmul(weights, value)

    def get_config(self):
        config = super().get_config()
        config["dropout_rate"] = self.dropout_rate
        return config


@tf.keras.utils.register_keras_serializable(package="custom")
class SelfAttentionBlock(tf.keras.layers.Layer):
    """自注意力块,包含 QKV 投影和前馈网络。"""

    def __init__(self, embed_dim, num_heads, ff_dim, dropout_rate=0.1, **kwargs):
        super().__init__(**kwargs)
        self.embed_dim = embed_dim
        self.num_heads = num_heads
        self.ff_dim = ff_dim
        self.dropout_rate = dropout_rate

    def build(self, input_shape):
        self.query_dense = tf.keras.layers.Dense(self.embed_dim)
        self.key_dense = tf.keras.layers.Dense(self.embed_dim)
        self.value_dense = tf.keras.layers.Dense(self.embed_dim)
        self.attention = ScaledDotProductAttention(self.dropout_rate)
        self.projection = tf.keras.layers.Dense(self.embed_dim)
        self.norm1 = tf.keras.layers.LayerNormalization(epsilon=1e-6)
        self.norm2 = tf.keras.layers.LayerNormalization(epsilon=1e-6)
        self.ffn = tf.keras.Sequential([
            tf.keras.layers.Dense(self.ff_dim, activation="gelu"),
            tf.keras.layers.Dropout(self.dropout_rate),
            tf.keras.layers.Dense(self.embed_dim),
        ])
        self.dropout1 = tf.keras.layers.Dropout(self.dropout_rate)
        super().build(input_shape)

    def call(self, inputs, training=None, mask=None):
        q = self.query_dense(inputs)
        k = self.key_dense(inputs)
        v = self.value_dense(inputs)

        attn = self.attention(q, k, v, mask=mask, training=training)
        attn = self.projection(attn)
        attn = self.dropout1(attn, training=training)
        out1 = self.norm1(inputs + attn)

        ffn_out = self.ffn(out1, training=training)
        return self.norm2(out1 + ffn_out)

    def get_config(self):
        config = super().get_config()
        config.update({
            "embed_dim": self.embed_dim,
            "num_heads": self.num_heads,
            "ff_dim": self.ff_dim,
            "dropout_rate": self.dropout_rate,
        })
        return config


# 构建完整模型
@tf.keras.utils.register_keras_serializable(package="custom")
class TextClassifier(tf.keras.Model):
    def __init__(self, vocab_size, embed_dim, num_heads, ff_dim,
                 num_blocks=2, num_classes=2, **kwargs):
        super().__init__(**kwargs)
        self.embedding = tf.keras.layers.Embedding(vocab_size, embed_dim)
        self.pos_embedding = tf.keras.layers.Embedding(512, embed_dim)
        self.blocks = [
            SelfAttentionBlock(embed_dim, num_heads, ff_dim)
            for _ in range(num_blocks)
        ]
        self.global_pool = tf.keras.layers.GlobalAveragePooling1D()
        self.classifier = tf.keras.layers.Dense(num_classes, activation="softmax")
        self.loss_tracker = tf.keras.metrics.Mean(name="loss")
        self.acc_metric = tf.keras.metrics.CategoricalAccuracy(name="accuracy")

    def call(self, inputs, training=None):
        seq_len = tf.shape(inputs)[1]
        positions = tf.range(seq_len)
        x = self.embedding(inputs) + self.pos_embedding(positions)

        for block in self.blocks:
            x = block(x, training=training)

        x = self.global_pool(x)
        return self.classifier(x)

    @property
    def metrics(self):
        return [self.loss_tracker, self.acc_metric]

    def train_step(self, data):
        x, y = data
        with tf.GradientTape() as tape:
            y_pred = self(x, training=True)
            loss = self.compiled_loss(y, y_pred)

        grads = tape.gradient(loss, self.trainable_variables)
        self.optimizer.apply_gradients(zip(grads, self.trainable_variables))

        self.loss_tracker.update_state(loss)
        self.acc_metric.update_state(y, y_pred)
        return {m.name: m.result() for m in self.metrics}

    def get_config(self):
        return {
            "vocab_size": self.embedding.input_dim,
            "embed_dim": self.embedding.output_dim,
            "num_heads": self.blocks[0].num_heads,
            "ff_dim": self.blocks[0].ff_dim,
            "num_blocks": len(self.blocks),
        }

    @classmethod
    def from_config(cls, config):
        return cls(**config)

# 使用示例
model = TextClassifier(
    vocab_size=30000, embed_dim=128,
    num_heads=4, ff_dim=512, num_blocks=4, num_classes=2
)
model.compile(
    optimizer=tf.keras.optimizers.Adam(1e-4),
    loss=tf.keras.losses.CategoricalCrossentropy(),
)

# 训练
# model.fit(train_dataset, validation_data=val_dataset, epochs=10)

# 保存与加载
# model.save("text_classifier.keras")
# loaded = tf.keras.models.load_model("text_classifier.keras")

八、Functional API vs Model Subclassing 选择指南

最后,总结两种模型构建方式的选择依据:

维度 Functional API Model Subclassing
序列化 原生支持,开箱即用 需手动实现 get_config / from_config
模型结构检查 model.summary() 在构建前即可用 需先传入数据触发 build
动态控制流 不支持 完全支持(if/for/while)
自定义 train_step 不支持 支持,与 fit() 无缝集成
多输出多损失 声明式定义,简洁 灵活,可在 train_step 中动态调整
学习曲线 中高
调试难度 较低 需注意 Eager/Graph 差异

推荐策略:能用 Functional API 就用 Functional API(90% 的场景足够);当需要动态控制流、自定义训练步骤或与

1
model.fit()

的深度集成时,再切换到 Model Subclassing。两者也可以混合使用——子类化模型内部使用 Functional 构建的子模块。

自定义层和模型子类化是 TensorFlow 2.x 高级开发的核心能力。掌握它们后,你将不再受限于内置组件,而是能以 Python 原生的表达力构建任何你能想到的模型架构。希望本文的实战经验能帮助你在下一个项目中游刃有余。

【本站文章皆为原创,未经允许不得转载】:汤不热吧 » TensorFlow 2.x 自定义层与模型子类化深度实战:从 Layer 到 Model 的完整工程化指南
分享到: 更多 (0)