如何在TensorFlow自定义层中实现张量列的重排操作
TensorFlow 2.x自定义层实现张量列交换的解决方案
你之前的写法报错有两个核心原因:
- 自定义层的输入是计算图中的符号张量,不需要封装为
ResourceVariable做运算,且TensorFlow原生不支持变量/张量的Python风格切片原位赋值 scatter_nd_update的入参格式错误,且这类简单列重排操作不需要用更新变量的复杂方案实现
最简实现方案
直接用tf.gather实现任意列重排,适配动态batch大小,兼容计算图构建逻辑:
import tensorflow as tf from tensorflow.keras.layers import Input from tensorflow.keras.models import Model class CustomLayer(tf.keras.layers.Layer): def __init__(self, **kwargs): super(CustomLayer, self).__init__(**kwargs) def call(self, inputs): # indices=[1,0]表示将原最后一维的第1位放到输出第0位,原第0位放到输出第1位,即交换两列 return tf.gather(inputs, indices=[1, 0], axis=-1)
模型构建与验证
# 构建模型 input_1 = Input(shape=(2, 2), name='inp') output = CustomLayer()(input_1) model = Model(input_1, output) # 效果验证 test_input = tf.constant([[[1.,2.],[3.,4.]]]) print(model(test_input)) # 输出:tf.Tensor([[[2. 1.] [4. 3.]]], shape=(1, 2, 2), dtype=float32)
可选方案:tf.concat实现
如果需要更直观的列拼接逻辑,也可以用tf.concat直接拼接交换后的列:
def call(self, inputs): # 分别取第1列、第0列,沿最后一维拼接 return tf.concat([inputs[..., 1:2], inputs[..., 0:1]], axis=-1)
内容的提问来源于stack exchange,提问作者Boris Salimov
相关产品推荐
相关产品推荐

