TensorFlow中如何从维度为[n,k,d]的张量生成自定义对象列表
错误原因分析
OperatorNotAllowedInGraphError报错:代码运行在TensorFlow图执行模式(通常是被@tf.function修饰的函数内部),图模式不支持Python原生迭代语法遍历Tensor,列表推导属于Python原生迭代操作,因此被AutoGraph拦截报错。ValueError报错:tf.map_fn仅支持返回可转换为Tensor类型的映射函数,自定义类Myobject的实例无法被转换为Tensor,因此触发报错;构造函数中拿到的x维度为[None, d]是图模式下的动态形状推断特性,第一维为动态批次维度,因此显示为None。
正确实现方案
场景1:代码运行在eager执行模式(无@tf.function修饰的普通Python上下文)
使用tf.unstack将输入张量按第一维拆分,再遍历生成实例即可:
# 将shape为[n, k, d]的X沿第0维拆分为n个shape为[k, d]的子张量 sub_tensors = tf.unstack(X, axis=0) # 若Myobject构造函数要求输入为numpy数组,可将x替换为x.numpy() listofobject = [Myobject(x) for x in sub_tensors]
场景2:必须在@tf.function修饰的图模式函数内处理
TensorFlow图模式核心为张量运算,无法兼容任意Python自定义对象的生成逻辑,建议调整逻辑:
- 先在图内完成所有张量运算逻辑,输出处理后的子张量集合
- 在图外的eager模式下,遍历子张量生成
Myobject实例列表
若必须在图内实现Myobject相关功能,需将Myobject的所有逻辑改造为纯张量运算实现,取消Python类封装。
内容的提问来源于stack exchange,提问作者jajamaharaja
相关产品推荐
相关产品推荐

