如何将带嵌套不规则维度的标准张量转回RaggedTensor?
嵌套RaggedTensor调用from_tensor还原失败的问题
当RaggedTensor存在嵌套不规则维度时,直接调用tf.RaggedTensor.from_tensor(ragged_tensor.to_tensor(), padding=0)会抛出维度不兼容错误,无法还原出原始的RaggedTensor。
示例代码
import tensorflow as tf data = tf.ragged.constant([ [[4,35,6,33], [7,2], [89,56,12]], [[2,11], [9]] ]) # 直接调用会触发报错 tf.RaggedTensor.from_tensor(data.to_tensor(), padding=0)
报错信息
Traceback (most recent call last): File "./src/tppmodel.py", line 34, in <module> tf.RaggedTensor.from_tensor(data.to_tensor(), padding=0) File "/opt/conda/lib/python3.8/site-packages/tensorflow/python/util/traceback_utils.py", line 153, in error_handler raise e.with_traceback(filtered_tb) from None File "/opt/conda/lib/python3.8/site-packages/tensorflow/python/framework/tensor_shape.py", line 1307, in assert_is_compatible_with raise ValueError("Shapes %s and %s are incompatible" % (self, other)) ValueError: Shapes () and (4,) are incompatible
问题原因
tf.RaggedTensor.from_tensor默认仅处理1个不规则维度(ragged_rank=1),但示例中的data是嵌套结构的RaggedTensor,存在2个不规则维度(外层列表长度不一致,内层子列表长度也不一致)。若不指定ragged_rank参数,TensorFlow无法正确识别嵌套的不规则结构,从而触发维度不兼容错误。
解决方案
调用from_tensor时,显式指定ragged_rank参数为原始RaggedTensor的不规则维度数(可通过data.ragged_rank直接获取):
# 获取原始RaggedTensor的不规则维度数 ragged_rank = data.ragged_rank # 指定ragged_rank完成还原 restored_data = tf.RaggedTensor.from_tensor(data.to_tensor(), padding=0, ragged_rank=ragged_rank) # 验证还原结果与原数据一致 print(tf.equal(data, restored_data)) # 输出:<tf.RaggedTensor [[[True True True True], [True True], [True True True]], [[True True], [True]]]>
最终效果
执行上述代码后,restored_data与原始data完全一致,实现了嵌套RaggedTensor的正确还原。
内容的提问来源于stack exchange,提问作者abes
相关产品推荐
相关产品推荐

