You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.30 00:31:02