如何修改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
相关产品推荐
相关产品推荐

