TensorFlow拼接Ragged Tensor返回形状异常问题
问题描述
我需要拼接两个ragged tensor,保持最后一个维度固定为2。
查看model.output_shape时我得到了预期的(None, None, 2)形状,但调用模型进行推理时,得到的输出形状为(batch_size, None, None),需要获取正确的输出形状。
复现代码如下:
import tensorflow as tf a_input = tf.keras.layers.Input([None, 2], ragged=True) b_input = tf.keras.layers.Input([None, 2], ragged=True) output = tf.concat([a_input, b_input], axis=1) model = tf.keras.Model([a_input, b_input], output) a = tf.ragged.constant([ [[1, 2], [3, 4], [5, 6]], [[1, 2], [3, 4]], [[1, 2]], ]) b = tf.ragged.constant([ [[1, 2]], [[1, 2], [3, 4], [5, 6], [7, 8]], [[1, 2], [3, 4]], ]) print(model.output_shape) # (None, None, 2) print(model([a, b]).shape) # (3, None, None)
问题原因
这是TensorFlow Keras处理ragged tensor静态形状推导的固有问题:直接将原生tf.concat算子的返回值作为模型输出时,框架不会自动继承输入张量最后一维固定为2的形状信息,推理阶段就会把最后一维错误标记为动态维度None。
解决方法
把原生tf.concat替换为Keras内置的Concatenate层,同时显式声明输出的静态形状即可,修改后的可运行代码如下:
import tensorflow as tf a_input = tf.keras.layers.Input([None, 2], ragged=True) b_input = tf.keras.layers.Input([None, 2], ragged=True) # 用Keras内置拼接层替代tf.concat output = tf.keras.layers.Concatenate(axis=1)([a_input, b_input]) # 显式补全静态形状信息,不会修改张量实际值 output.set_shape([None, None, 2]) model = tf.keras.Model([a_input, b_input], output) a = tf.ragged.constant([ [[1, 2], [3, 4], [5, 6]], [[1, 2], [3, 4]], [[1, 2]], ]) b = tf.ragged.constant([ [[1, 2]], [[1, 2], [3, 4], [5, 6], [7, 8]], [[1, 2], [3, 4]], ]) print(model.output_shape) # (None, None, 2) print(model([a, b]).shape) # (3, None, 2)
操作说明:
- Keras内置层自带适配ragged tensor的形状推导逻辑,比直接调用TensorFlow原生算子的形状识别更准确
set_shape仅做静态元信息补全,不会带来任何运行时性能损耗,也不会改变张量的实际计算结果- 如果其他ragged算子出现同类静态形状丢失问题,都可以用这个方式手动补全已知的固定维度信息,是这类场景的通用解法
内容的提问来源于stack exchange,提问作者yolisses
相关产品推荐
相关产品推荐

