TF 2.6运行文本生成RNN示例时tf.sparse.to_dense报索引重复错误
错误原因
tf.sparse.to_dense接口要求输入的SparseTensor索引不能重复,你提供的日志显示构造的SparseTensor的indices字段存在两个完全相同的[0]值,直接触发参数校验报错。- 索引重复的根源是
ids_from_chars(['[UNK]'])返回了两个重复的id:你构造ids_from_chars(tf.keras.layers.StringLookup层)时,要么手动将[UNK]加入了自定义词汇表,同时层默认开启了oov_token=[UNK]自动添加逻辑,导致词汇表存在重复的[UNK],查询时返回两个相同的id;要么是层参数配置异常导致单条查询返回重复结果。
修复方案
方案1:快速去重(修改代码最少)
在生成skip_ids的代码后新增去重逻辑,直接消除重复索引:
skip_ids = self.ids_from_chars(['[UNK]'])[:, None] # 新增去重代码 skip_ids = tf.unique(tf.squeeze(skip_ids))[0][:, None]
处理后skip_ids只会保留唯一id,不会再出现重复索引问题。
方案2:直接构造稠密mask(稳定性更高)
完全替换原有的稀疏张量生成逻辑,绕开稀疏张量的索引校验规则:
# 替换原有的sparse_mask相关全部代码 vocab_size = len(ids_from_chars.get_vocabulary()) self.prediction_mask = tf.zeros(vocab_size, dtype=tf.float32) # 获取[UNK]对应的id unk_id = ids_from_chars('[UNK]').numpy() # 更新[UNK]位置的mask值为-inf self.prediction_mask = tf.tensor_scatter_nd_update(self.prediction_mask, [[unk_id]], [-float('inf')])
方案3:修正词汇表构造逻辑(根源修复)
检查ids_from_chars层的构造代码:
- 如果你手动将
[UNK]加入了传入的词汇列表,同时设置了oov_token='[UNK]',删除手动加入的[UNK]即可; - 如果你不需要OOV逻辑,构造层时显式设置
oov_token=None即可。
内容的提问来源于stack exchange,提问作者James
相关产品推荐
相关产品推荐

