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

如何修改Transformer模型适配(30,)输入与6类onehot分类输出

报错核心原因

你的代码存在几处逻辑硬伤直接触发报错:

  • 在call方法里硬编码tf.reshape(x, [64, 29]):batch size是动态值,训练、验证、推理阶段的batch size不一定为64,同时原始输入是30维特征,强制reshape到29会直接丢失维度,触发切片越界错误
  • 输入维度不匹配Transformer结构要求:原生Transformer的自注意力层仅接受形状为(batch_size, seq_len, feature_dim)的三维张量,直接传入(30,)的一维输入时,取tf.shape(x)[1]不存在第二维,自然会报维度为None、切片越界的错误
  • 层定义逻辑错误:在自定义层里硬绑input_shape、错误定义位置编码长度,会导致Dense层接收到未知维度的输入,触发最后一维未定义的报错
  • 输出头未适配分类任务:原生Transformer解码器输出是面向词表的序列结果,没有针对6分类任务设计投影层,直接使用无法输出(6,)形状的结果
快速验证可用的修改方案

该方案不需要完整复刻NLP场景下的编码器-解码器交叉注意力结构,适配输入形状(batch_size, 30)、输出形状(batch_size, 6)的分类需求,改完可直接运行:

1. 重写Encoder层

去掉硬编码逻辑,增加输入投影和维度适配:

class Encoder(tf.keras.layers.Layer):
  def __init__(self,*, num_layers, d_model, num_heads, dff, input_dim=30,
               rate=0.1):
    super(Encoder, self).__init__()
    self.d_model = d_model
    self.num_layers = num_layers
    # 输入投影层:把30维输入特征映射到Transformer要求的d_model维度
    self.input_proj = tf.keras.layers.Dense(d_model, activation='relu')
    # 单序列场景下位置编码长度设为1即可
    self.pos_encoding = positional_encoding(1, self.d_model)

    self.enc_layers = [
        EncoderLayer(d_model=d_model, num_heads=num_heads, dff=dff, rate=rate)
        for _ in range(num_layers)]
    self.dropout = tf.keras.layers.Dropout(rate)

  def call(self, x, training, mask=None):
    # x输入形状: (batch_size, 30)
    seq_len = 1
    # 特征投影
    x = self.input_proj(x)
    # 升维到3维,适配自注意力层输入要求
    x = tf.expand_dims(x, axis=1)
    x *= tf.math.sqrt(tf.cast(self.d_model, tf.float32))
    x += self.pos_encoding[:, :seq_len, :]
    x = self.dropout(x, training=training)

    for i in range(self.num_layers):
      x = self.enc_layers[i](x, training, mask)
    # 输出形状: (batch_size, 1, d_model)
    return x

2. 重写Transformer模型类

去掉原生NLP任务的解码器适配逻辑,直接添加6分类头:

class Transformer(tf.keras.Model):
  def __init__(self,*, num_layers, d_model, num_heads, dff, input_dim=30, num_classes=6, rate=0.1):
    super(Transformer, self).__init__()
    self.encoder = Encoder(
        num_layers=num_layers,
        d_model=d_model,
        num_heads=num_heads,
        dff=dff,
        input_dim=input_dim,
        rate=rate
    )
    # 池化层压平序列维度
    self.global_pool = tf.keras.layers.GlobalAveragePooling1D()
    # 分类头输出6分类概率
    self.class_head = tf.keras.layers.Dense(num_classes, activation='softmax')

  def call(self, x, training=False):
    enc_out = self.encoder(x, training=training, mask=None)
    pooled = self.global_pool(enc_out)
    return self.class_head(pooled)

3. 实例化与测试

# 参考超参,可根据验证效果调整
transformer = Transformer(
    num_layers=2,
    d_model=64,
    num_heads=2,
    dff=128,
    input_dim=30,
    num_classes=6,
    rate=0.1
)

# 形状测试
test_x = tf.random.normal((32, 30)) # 32个样本,每个样本30维特征
test_y = transformer(test_x, training=False)
print(test_y.shape) # 输出(32, 6),符合要求
使用注意事项
  • 训练时如果标签是onehot格式,损失函数用tf.keras.losses.CategoricalCrossentropy();如果是整数类别标签,直接用tf.keras.losses.SparseCategoricalCrossentropy(),不需要手动转onehot
  • 如果需要更强的特征提取能力,可以把30维输入拆成30个单值token(即序列长度为30,每个token维度为1),只需要把输入reshape为(batch_size, 30, 1),位置编码长度改为30即可,其余逻辑不变
  • 不要在call方法里硬编码任何batch size、维度数值,所有维度通过张量动态形状推导,避免不同阶段batch size变化触发报错

内容的提问来源于stack exchange,提问作者Jigberto

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 04:21:34