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
相关产品推荐
相关产品推荐

