自定义Transformer三级分类训练:Decoder输入与数据集适配问题
自定义Transformer层级分类Decoder训练输入解决方案
核心问题拆解
你需要实现基于Transformer Decoder的prompt层级划分任务:Encoder编码完整prompt,三个独立Decoder分别输出lvl1/lvl2/lvl3对应的文本序列。训练阶段的核心是正确给每个Decoder喂入teacher forcing输入序列,并对齐标签计算损失。
数据集预处理调整
你的数据集每个层级包含sequence(目标词元序列),训练前需对每个层级做以下处理:
- 统一序列长度:对所有层级的
sequence做padding/truncation到预设的max_seq_len,lvl3的空序列用特殊<PAD>token(词元ID可设为0)填充。 - 构造Decoder输入序列:对每个层级的
sequence执行左移一位+起始token操作:- 若
start字段是该层级的起始token(如示例中的321),则Decoder输入序列为[start] + sequence[:-1] - 若
sequence为空(如lvl3),则输入序列全为<PAD>或仅保留起始token+pad
- 若
- 标签序列:直接使用原
sequence(padding后)作为该Decoder的预测标签。
示例处理后的数据格式(以lvl1为例):
{ "prompt_input": [123, 456, ...], # 编码后的prompt词元序列(padding后) "lvl1_decoder_input": [321, 321, 269, 539, 1131], # start + sequence[:-1] "lvl1_label": [321, 269, 539, 1131, 419], # 原sequence "lvl2_decoder_input": [1013, 1013, 269, 493, 877, 1072, 1003], "lvl2_label": [1013, 269, 493, 877, 1072, 1003, 1026], "lvl3_decoder_input": [0, 0, ..., 0], # <PAD>填充 "lvl3_label": [0, 0, ..., 0] }
Decoder输入逻辑改造
你的Decoder实现已支持接收target输入,训练时需为三个Decoder分别传入对应层级的decoder_input,并利用Encoder输出作为memory:
- 掩码处理:
decoder_mask:使用下三角掩码(防止Decoder看到未来token)+ padding掩码,屏蔽<PAD>和未来位置memory_mask:仅使用padding掩码,屏蔽prompt中的<PAD>部分
- 多Decoder整合:构建主模型类,将Encoder与三个Decoder串联,每个Decoder独立处理对应层级的输入。
训练流程整合
1. 主模型实现示例
class HierarchicalTransformer(tf.keras.Model): def __init__(self, encoder_params, decoder_params): super().__init__() self.encoder = Encoder(**encoder_params) # 三个独立Decoder,参数可共享或独立,此处用独立参数 self.decoder_lvl1 = Decoder(**decoder_params) self.decoder_lvl2 = Decoder(**decoder_params) self.decoder_lvl3 = Decoder(**decoder_params) # 输出映射层,每个Decoder对应一个 self.dense_lvl1 = tf.keras.layers.Dense(decoder_params['target_vocab_size']) self.dense_lvl2 = tf.keras.layers.Dense(decoder_params['target_vocab_size']) self.dense_lvl3 = tf.keras.layers.Dense(decoder_params['target_vocab_size']) def call(self, inputs, training=False): prompt_input, lvl1_input, lvl2_input, lvl3_input = inputs # 计算掩码 encoder_mask = self.create_padding_mask(prompt_input) decoder_mask_lvl1 = self.create_decoder_mask(lvl1_input) decoder_mask_lvl2 = self.create_decoder_mask(lvl2_input) decoder_mask_lvl3 = self.create_decoder_mask(lvl3_input) # Encoder编码prompt encoder_output, _ = self.encoder(prompt_input, training, encoder_mask) # 三个Decoder分别处理 lvl1_output, _ = self.decoder_lvl1(encoder_output, lvl1_input, training, decoder_mask_lvl1, encoder_mask) lvl2_output, _ = self.decoder_lvl2(encoder_output, lvl2_input, training, decoder_mask_lvl2, encoder_mask) lvl3_output, _ = self.decoder_lvl3(encoder_output, lvl3_input, training, decoder_mask_lvl3, encoder_mask) # 映射到词汇表 lvl1_logits = self.dense_lvl1(lvl1_output) lvl2_logits = self.dense_lvl2(lvl2_output) lvl3_logits = self.dense_lvl3(lvl3_output) return [lvl1_logits, lvl2_logits, lvl3_logits] # 辅助函数:创建padding掩码 def create_padding_mask(self, seq): seq = tf.cast(tf.math.equal(seq, 0), tf.float32) return seq[:, tf.newaxis, tf.newaxis, :] # (batch_size, 1, 1, seq_len) # 辅助函数:创建Decoder掩码(下三角+padding) def create_decoder_mask(self, seq): padding_mask = self.create_padding_mask(seq) seq_len = tf.shape(seq)[1] look_ahead_mask = 1 - tf.linalg.band_part(tf.ones((seq_len, seq_len)), -1, 0) look_ahead_mask = tf.cast(look_ahead_mask, tf.float32) return tf.maximum(padding_mask, look_ahead_mask)
2. 损失函数与训练步骤
def hierarchical_loss(labels, logits): # 对每个层级计算交叉熵,屏蔽padding部分 loss_fn = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction='none') total_loss = 0.0 for label, logit in zip(labels, logits): loss = loss_fn(label, logit) # 掩码掉padding的损失 mask = tf.cast(tf.math.not_equal(label, 0), tf.float32) loss *= mask total_loss += tf.reduce_mean(loss) return total_loss # 模型初始化示例 encoder_params = { "num_blocks": 2, "dimension_model": 512, "num_heads": 8, "hidden_dimension": 2048, "src_vocab_size": 15000, "max_seq_len": 50, "dropout_rate": 0.1 } decoder_params = { "num_blocks": 2, "d_model": 512, "num_heads": 8, "hidden_dim": 2048, "target_vocab_size": 15000, "max_seq_len": 50, "dropout_rate": 0.1 } model = HierarchicalTransformer(encoder_params, decoder_params) optimizer = tf.keras.optimizers.Adam(learning_rate=1e-4) # 训练步骤 @tf.function def train_step(inputs, labels): with tf.GradientTape() as tape: logits = model(inputs, training=True) loss = hierarchical_loss(labels, logits) gradients = tape.gradient(loss, model.trainable_variables) optimizer.apply_gradients(zip(gradients, model.trainable_variables)) return loss
关键注意点
- Encoder位置编码修复:你的Encoder中
positional_index计算有误,替换为TensorFlow原生操作避免兼容性问题:# 替换原positional_index计算代码 positional_index = tf.range(self.max_sql_len, dtype=tf.int32)[tf.newaxis, :] positional_index = tf.repeat(positional_index, repeats=tf.shape(input)[0], axis=0) - lvl3特殊处理:若lvl3对应无内容,可分配特殊
<NULL>token,让模型学习输出该token表示无对应层级内容,减少全padding噪声。 - 参数共享:若三个层级任务相似,可共享Decoder参数,只需将三个Decoder实例改为同一个对象即可减少参数量。
内容的提问来源于stack exchange,提问作者Osama Elsherif
相关产品推荐
相关产品推荐

