You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于TensorFlow Keras子类化实现Transformer的编译与训练错误排查

Transformer实现问题排查与修复

错误现象

  1. 模型编译时抛出类型错误:

TypeError: Error converting shape to a TensorShape: Dimension value must be integer or None or have an index method, got value '(12, 1)' with type '<class 'tuple'>'

  1. 样本数据编译通过后,训练时抛出梯度错误:

ValueError: No gradients provided for any variable: [...]


问题定位与修复方案

1. 模型构建的形状定义错误

问题点:forecastor.build(((12, 1),(3, 1)))传入的输入形状不符合实际数据维度。实际输入trainX是(batch_size, 12),trainXt_in是(batch_size, 3),build方法需要传入不带batch维度的输入形状。

修复代码:

forecastor.build(((12,), (3,)))

2. 位置编码层的梯度中断与形状计算错误

问题点:

  • 原PositionalEncodingLayer使用numpy操作+tf.map_fn,默认不跟踪梯度,导致梯度无法回传
  • 手动循环计算位置编码效率低,且形状获取逻辑错误

修复后的位置编码层:

class PositionalEncodingLayer(Layer):
    def __init__(self, **kwargs):
        super().__init__()
        self.add = Add()

    def call(self, x):
        seq_len = tf.shape(x)[1]
        d_model = tf.shape(x)[2]
        
        # 生成位置索引与分母项
        position = tf.range(seq_len, dtype=tf.float32)[:, tf.newaxis]
        div_term = tf.pow(10000, tf.range(0, d_model, 2, dtype=tf.float32) / d_model)
        
        # 构建位置编码张量
        pos_enc = tf.zeros((seq_len, d_model))
        pos_enc = tf.tensor_scatter_nd_update(
            pos_enc,
            indices=tf.stack([tf.range(seq_len), tf.range(0, d_model, 2)], axis=1),
            updates=tf.sin(position / div_term)
        )
        pos_enc = tf.tensor_scatter_nd_update(
            pos_enc,
            indices=tf.stack([tf.range(seq_len), tf.range(1, d_model, 2)], axis=1),
            updates=tf.cos(position / div_term)
        )
        
        # 扩展batch维度后与输入相加
        pos_enc = pos_enc[tf.newaxis, ...]
        pos_embeddings = self.add([x, pos_enc])
        return pos_embeddings

3. Decoder输出层的梯度中断

问题点:Decoder的tf.round(tf.abs(output))是不可导操作,直接切断了梯度传递链路,导致训练时无梯度可用。

修复代码:

### Output Dense Layer
output = self.output_dense_layer(sub_layer6_out)
# 训练阶段保留连续输出,离散化操作放到推理阶段执行
# output = tf.round(tf.abs(output))
return output

4. 学习率调度器的梯度兼容问题

问题点:原调度器使用step.numpy()和numpy计算,破坏了TensorFlow的梯度跟踪机制。

修复后的学习率调度器:

class MyLRSchedule(tf.keras.optimizers.schedules.LearningRateSchedule):
    def __init__(self, d_model, warmup_steps):
        self.d_model = tf.cast(d_model, tf.float32)
        self.warmup_steps = tf.cast(warmup_steps, tf.float32)

    def __call__(self, step):
        step = tf.cast(step, tf.float32)
        denom = tf.pow(self.d_model, -0.5)
        term1 = tf.pow(step, -0.5)
        term2 = step * tf.pow(self.warmup_steps, -1.5)
        numer = tf.minimum(term1, term2)
        lrate = numer / denom
        return lrate

5. Encoder中的变量名笔误

问题点:Encoder的call方法中存在变量名拼写错误:postitional_embedding多写了一个t,导致dropout未生效,可能引发后续逻辑错误。

修复代码:

positional_embedding = self.pos_embedding(embedding_output)
positional_embedding = self.dropout(positional_embedding)  # 修正变量名拼写

6. 训练参数合理性调整

问题点:样本数据仅5条,steps_per_epoch = trainX.shape[0]//batch_size会得到0,导致fit时出错。

修复代码:

history = forecastor.fit((trainX, trainXt_in), trainY,
                          batch_size=batch_size,
                          # 移除steps_per_epoch,由Keras自动计算训练步数
                          epochs=1,
                          validation_data=((valX, ValXt_in), valY),
                          callbacks=cb)

内容的提问来源于stack exchange,提问作者Krishnang K Dalal

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.05 11:02:02