如何使用tf.gather从列表中获取非Tensor类型元素?
解决tf.gather获取自定义ExtensionType元素的问题
问题原因
tf.gather会尝试将输入序列转换为Tensor对象,但直接把自定义ExtensionType实例放在Python列表中时,TensorFlow无法自动将其转为支持的张量结构,因此抛出类型转换错误。
解决方案
方法1:将自定义对象列表转为ExtensionType张量
tf.experimental.ExtensionType支持被TensorFlow包装为张量,先把列表转成对应张量,再用tf.gather即可正常操作:
import tensorflow as tf class Point(tf.experimental.ExtensionType): xx: tf.Tensor def __init__(self, xx): self.xx = tf.convert_to_tensor(xx) super().__init__() # 将Point实例列表转换为ExtensionType张量 point_tensor = tf.convert_to_tensor([Point(1), Point(2), Point(3), Point(4)]) # 使用tf.gather获取指定索引的元素 out2 = tf.gather(point_tensor, [0, 2]) print('Second gather ', out2)
运行后输出示例:
Second gather <tf.Tensor: shape=(2,), dtype=variant, numpy=[Point(xx=<tf.Tensor: shape=(), dtype=int32, numpy=1>), Point(xx=<tf.Tensor: shape=(), dtype=int32, numpy=3>)]>
方法2:Python列表索引结合Tensor索引(适合Eager模式)
如果不需要TensorFlow图追踪支持,可以把tf.gather的索引转为numpy数组,直接用Python列表索引提取元素:
list2 = [Point(1), Point(2), Point(3), Point(4)] indices = tf.constant([0, 2]).numpy() out2 = [list2[i] for i in indices] print('Second gather ', out2)
说明
- 方法1是推荐方案,保留了TensorFlow的张量操作特性,支持图模式和自动微分。
- 自定义
ExtensionType类的__init__方法中建议显式将输入转为Tensor,避免隐式转换带来的潜在问题。
内容的提问来源于stack exchange,提问作者NNN
相关产品推荐
相关产品推荐

