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

自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.03 17:41:06