如何在TensorFlow模型中按坐标选取像素并输出其值?
问题:从TensorFlow模型的2D特征图中提取指定坐标像素,输出形状不符合预期
我要构建的模型包含两个输出:一个是2D数组,另一个是从第一个输出分支中选取预定义坐标的像素,再经位置相关的特定学习函数得到的数值集合(该函数随位置变化,无法全局学习)。
我需要实现的功能对应NumPy代码如下:
import numpy as np m = np.random.randint(0, 100, size=[3, 3]) pixels = [[0, 0], [1, 0], [2, 1]] rows = [p[0] for p in pixels] cols = [p[1] for p in pixels] print('Original array: \n') print(m) print('Selected pixel values: \n') print(m[rows, cols])
运行输出:
Original array: [[89 35 0] [27 74 93] [96 13 13]] Selected pixel values: [89 27 13]
我用TensorFlow尝试的代码如下:
import tensorflow as tf from tensorflow import keras import numpy as np input_array = keras.layers.Input(shape=(3, 3, 1)) x = keras.layers.Conv2D(1, 3, padding='same')(input_array) print(f'array shape: {keras.backend.int_shape(x)}') pixels_np = np.array([[0, 0], [1, 0], [2, 1]]).reshape([3, 2, 1]) pixels = tf.constant(pixels_np) print(f'pixels shape: {keras.backend.int_shape(pixels)}') output1 = keras.layers.Lambda(lambda x:tf.gather_nd(x, pixels))(x) model = keras.models.Model(inputs=input_array, outputs=[output1]) model.compile(loss='mse', optimizer=keras.optimizers.Adam()) print('model summary: \n') model.summary()
运行结果:
array shape: (None, 3, 3, 1) pixels shape: (3, 2, 1) model summary: _________________________________________________________________ input_1 (InputLayer) [(None, 3, 3, 1)] 0 conv2d (Conv2D) (None, 3, 3, 1) 10 lambda (Lambda) (3, 2, 3, 3, 1) 0 ================================================================= Total params: 10 Trainable params: 10 Non-trainable params: 0
但输出形状不符合预期(预期形状:[None, 3, 1])。
解决方案
问题出在tf.gather_nd的索引格式错误:输入特征图是(batch_size, height, width, channels)的批量数据,但提供的索引没有考虑batch维度,导致tf.gather_nd错误解析索引,输出形状混乱。
以下是两种可行的修正方法:
方法1:显式处理batch维度的tf.gather_nd实现
import tensorflow as tf from tensorflow import keras import numpy as np input_array = keras.layers.Input(shape=(3, 3, 1)) x = keras.layers.Conv2D(1, 3, padding='same')(input_array) print(f'array shape: {keras.backend.int_shape(x)}') # 预定义像素坐标,保持(3,2)的形状:每个元素是[row, col] pixels_np = np.array([[0, 0], [1, 0], [2, 1]]) print(f'pixels shape: {pixels_np.shape}') def gather_specific_pixels(x): batch_size = tf.shape(x)[0] # 生成每个像素对应的batch索引,形状为(3*batch_size, 1) batch_indices = tf.tile(tf.expand_dims(tf.range(batch_size), 1), [1, 3]) batch_indices = tf.reshape(batch_indices, [-1, 1]) # 重复像素坐标batch_size次,形状为(3*batch_size, 2) pixel_indices = tf.tile(pixels_np, [batch_size, 1]) # 拼接得到完整索引:[batch_idx, row, col],形状(3*batch_size, 3) indices = tf.concat([batch_indices, pixel_indices], axis=1) # 提取像素值并调整形状为(batch_size, 3, 1) gathered = tf.gather_nd(x, indices) return tf.reshape(gathered, [batch_size, 3, 1]) output1 = keras.layers.Lambda(gather_specific_pixels)(x) model = keras.models.Model(inputs=input_array, outputs=[output1]) model.compile(loss='mse', optimizer=keras.optimizers.Adam()) print('model summary: \n') model.summary()
方法2:模仿NumPy索引的tf.gather实现(更简洁)
import tensorflow as tf from tensorflow import keras import numpy as np input_array = keras.layers.Input(shape=(3, 3, 1)) x = keras.layers.Conv2D(1, 3, padding='same')(input_array) print(f'array shape: {keras.backend.int_shape(x)}') pixels_np = np.array([[0, 0], [1, 0], [2, 1]]) rows = pixels_np[:, 0] cols = pixels_np[:, 1] def gather_specific_pixels(x): # 先提取所有batch中指定行的特征,形状(None, 3, 3, 1) row_selected = tf.gather(x, rows, axis=1) # 再从选中的行中提取指定列,形状(None, 3, 1) return tf.gather(row_selected, cols, axis=2) output1 = keras.layers.Lambda(gather_specific_pixels)(x) model = keras.models.Model(inputs=input_array, outputs=[output1]) model.compile(loss='mse', optimizer=keras.optimizers.Adam()) print('model summary: \n') model.summary()
两种方法的输出形状都会符合预期的(None, 3, 1),模型摘要如下:
array shape: (None, 3, 3, 1) pixels shape: (3, 2) model summary: _________________________________________________________________ Layer (type) Output Shape Param # ================================================================= input_1 (InputLayer) [(None, 3, 3, 1)] 0 conv2d (Conv2D) (None, 3, 3, 1) 10 lambda (Lambda) (None, 3, 1) 0 ================================================================= Total params: 10 Trainable params: 10 Non-trainable params: 0
内容的提问来源于stack exchange,提问作者Amit Oved
相关产品推荐
相关产品推荐

