You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 20:05:30