Transformer模型MRR计算时张量重塑矛盾报错求助
Transformer模型MRR计算的Shape错误修复
你的核心问题是完全忽略了批次维度,硬把带批次的输入张量强行reshape成单样本形状,导致后续shape完全混乱。从报错看,你的输入y_true实际是[32, 39](32个样本,每个样本39个token),总元素数1248=32*39,但你硬把它reshape成[1,39],直接丢掉了批次信息,后续操作自然会出现元素数不匹配的矛盾。
下面是修正后的_count_mrr函数,完全适配批次输入:
def _count_mrr(self, y_true: tf.Tensor, y_pred: tf.Tensor): # 确认输入shape:y_true [batch_size, seq_len], y_pred [batch_size, seq_len, vocab_size] batch_size, seq_len = tf.shape(y_true)[0], tf.shape(y_true)[1] # 把y_true转成[batch_size, seq_len, 1],方便和y_pred的排序结果做广播匹配 y_true = tf.cast(y_true, tf.int32) y_true = tf.expand_dims(y_true, axis=-1) # 对每个位置的预测概率取全量排序(取top_k为vocab_size等价于全排序) # y_pred_ranks shape: [batch_size, seq_len, vocab_size],每个位置是token的排序索引(0是概率最高的) y_pred_ranks = tf.math.top_k(y_pred, k=self._data_controller._vocab_size).indices # 找到真实token在排序中的位置,注意top_k的indices是从0开始的,所以要+1得到实际rank(第1名是1,不是0) match_indices = tf.where(tf.equal(y_pred_ranks, y_true)) # 提取每个匹配对应的rank值,还原回原批次序列的shape ranks = tf.scatter_nd( indices=match_indices[:, :2], updates=tf.cast(match_indices[:, 2] + 1, tf.float64), shape=[batch_size, seq_len] ) # 计算每个位置的1/rank,然后求全局平均(所有样本所有位置的平均) mrr_per_position = tf.divide(tf.ones_like(ranks, dtype=tf.float64), ranks) return tf.math.reduce_mean(mrr_per_position)
关键修正点
- 保留批次维度,不再强行reshape成单样本
- 用
tf.expand_dims给y_true增加最后一维,和y_pred的排序结果做广播匹配 - 用
tf.scatter_nd把匹配到的rank值还原回原批次序列的shape,避免遗漏任何样本或位置 - 完全适配
[batch_size, seq_len]的输入shape,不管batch_size是多少都能正确计算
原函数报错原因
- 第一次报错:你把原本1248个元素的张量(32*39)reshape成(39,1),总元素数只有39,自然不匹配
- 第二次报错:你又把经过
squeeze后变成39个元素的张量(因为之前硬转成(1,39),squeeze后成了(39,))reshape成(1248,1),元素数39≠1248,所以报错
内容的提问来源于stack exchange,提问作者zer0deck
相关产品推荐
相关产品推荐

