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

复现TensorFlow官网Transformer翻译模型遇AttributeError错误求助

错误原因及解决方法

错误原因

这个报错的核心是训练数据中的标签(label)为RaggedTensor类型,但masked_loss函数中存在针对RaggedTensor专属属性nested_row_splits的调用,而实际传入loss函数的却是普通Tensor。可能的触发场景:

  • 数据管道中某一步(比如错误使用普通batch而非ragged_batch)将RaggedTensor强制转为了带padding的普通Tensor,但loss函数仍保留了处理RaggedTensor的逻辑;
  • 官网默认代码的masked_loss是为密集Tensor(padding后的)设计的,你的数据是RaggedTensor,直接传入后loss函数内部的掩码逻辑错误调用了RaggedTensor的属性。

解决方法

方法1:修改masked_loss函数,兼容RaggedTensor和普通Tensor

在loss函数中先判断输入label的类型,针对性处理掩码生成逻辑,无需修改原始数据类型:

def masked_loss(label, pred):
    # 处理RaggedTensor类型的label
    if isinstance(label, tf.RaggedTensor):
        # 基于RaggedTensor的行长度生成掩码
        row_lengths = label.row_lengths()
        mask = tf.sequence_mask(row_lengths, maxlen=pred.shape[1])
        # 将RaggedTensor转为带padding的密集Tensor,用于计算损失
        label = label.to_tensor(default_value=0)
    else:
        # 处理普通Tensor的情况(原官网逻辑)
        mask = tf.math.logical_not(tf.math.equal(label, 0))
    
    loss_object = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction='none')
    loss = loss_object(label, pred)
    
    mask = tf.cast(mask, dtype=loss.dtype)
    loss *= mask
    return tf.reduce_sum(loss) / tf.reduce_sum(mask)

方法2:检查并修正数据管道的批处理方式

确保数据批处理时保留RaggedTensor类型,将普通的batch()替换为ragged_batch():

# 替换之前的batch逻辑
train_batches = train_dataset.ragged_batch(batch_size=你的批次大小).prefetch(tf.data.AUTOTUNE)

这样可以保证label始终是RaggedTensor类型,匹配loss函数中针对nested_row_splits的调用逻辑(如果官网原始loss是为RaggedTensor设计的)。

方法3:调整loss函数的掩码生成逻辑(针对RaggedTensor原生处理)

如果确认label是RaggedTensor,直接利用其内置属性生成掩码,避免类型转换:

def masked_loss(label, pred):
    # 直接从RaggedTensor获取掩码
    mask = label.to_tensor(default_value=0) != 0
    loss_object = tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True, reduction='none')
    # 使用RaggedTensor的flat_values获取真实标签值,匹配模型输出的形状
    loss = loss_object(label.flat_values, tf.gather(pred, label.value_rowids(), axis=0))
    
    mask = tf.cast(mask, dtype=loss.dtype)
    loss *= mask.flat_values
    return tf.reduce_sum(loss) / tf.reduce_sum(mask.flat_values)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 09:07:29