复现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
相关产品推荐
相关产品推荐

