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

TensorFlow2的@tf.function装饰函数中拼接不同batch size张量报错如何解决

问题原因分析

两个报错的根源都是@tf.function构建静态计算图时的静态结构校验要求:

  • 第一种方案用Python列表存储中间结果:静态图转换时会将tf.range驱动的循环转为tf.while_loop,要求循环前后所有变量的嵌套结构完全一致,Python列表每次append元素结构都会变化,无法通过校验。
  • 第二种方案用张量存储结果:循环前results的静态形状是(0, 2),第一次循环后变成了(None, 2),静态图要求循环变量的形状在迭代前后保持一致,因此触发形状不匹配报错。
解决方案

方案1:显式声明循环形状不变量

按照报错提示添加tf.autograph.experimental.set_loop_options指定循环变量的形状不变量即可,修改后代码如下:

import tensorflow as tf

@tf.function  # indispensable
def fun2(tensors, indices):
    results = tf.reshape(tf.constant([], dtype=tf.float32), (0, 2))  # make empty tensors

    for i in tf.range(2):  # batch size = 2
        # 声明循环变量results的形状不变量:第0维动态可变,第1维固定为2
        tf.autograph.experimental.set_loop_options(
            shape_invariants=[(results, tf.TensorShape([None, 2]))]
        )
        pos = tf.where(indices==i)
        emb = tf.gather_nd(tensors, pos)
        # 可在这里添加对emb的自定义处理逻辑
        results = tf.concat([results, emb], axis=0)

    return results

tensors = tf.random.uniform((5, 2))
print(fun2(tensors, indices=[0, 0, 1, 1, 1]))

方案2(更推荐):用TensorFlow原生分组算子替代显式循环

你的需求本质是按indices对tensors分组后拼接,直接用tf.dynamic_partition就能实现,完全避免自定义循环带来的静态图校验问题,性能也更优:

import tensorflow as tf

@tf.function
def fun3(tensors, indices):
    # 按indices的值将tensors分成2组
    grouped = tf.dynamic_partition(tensors, indices, num_partitions=2)
    # 如有需要可以在这里对每组张量做自定义处理
    # for i in range(len(grouped)):
    #     grouped[i] = your_process(grouped[i])
    return tf.concat(grouped, axis=0)

tensors = tf.random.uniform((5, 2))
print(fun3(tensors, indices=[0, 0, 1, 1, 1]))

内容的提问来源于stack exchange,提问作者Wilson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.03 03:27:03