如何在Keras中基于列索引构建生成176x208二进制掩码的网络
实现方案:从列索引生成二进制掩码的Keras网络
一、明确目标掩码的定义
首先要确定35个列索引对应掩码的规则,常见两种场景:
- 场景1:35个列是掩码中需全部设为1的列(即这35列的所有176行像素为1,其余列全为0)
- 场景2:35个索引对应35个特定行的列位置(比如第i个索引对应第i行的列位置,该行仅该列设为1,其余行保持0)
下面以场景1为例给出实现,场景2只需调整目标掩码的生成逻辑即可。
二、数据预处理:列索引转目标掩码
先实现函数将输入的35个索引转换成176×208的二进制掩码:
import numpy as np def indices_to_mask(indices, img_shape=(176, 208)): mask = np.zeros(img_shape, dtype=np.float32) # 将指定列的所有行设为1 for col_idx in indices: mask[:, col_idx] = 1.0 # 增加batch和通道维度,适配Keras输入格式 return mask[np.newaxis, ..., np.newaxis] # 最终形状:(1, 176, 208, 1) # 测试示例 sample_indices = [121,55,115,82,59,84,85,77,155,15,29,105,48,97,158,32,104,39,111,110,47,1,45,0,120,154,130,98,118,95,160,22,63,86,80] sample_mask = indices_to_mask(sample_indices) print(sample_mask.shape) # 输出:(1, 176, 208, 1)
三、构建生成掩码的Keras网络
网络核心是把35维的索引向量映射到176×208的特征图,这里采用「全连接+转置卷积」的轻量结构:
from tensorflow import keras from tensorflow.keras import layers def build_mask_generator(input_dim=35, img_shape=(176, 208)): # 输入层:接收35个列索引 inputs = keras.Input(shape=(input_dim,)) # 索引归一化到[-1,1]区间,适配后续网络层 x = layers.Lambda(lambda x: x / (img_shape[1]-1) * 2 - 1)(inputs) # 全连接层将向量映射到高维特征,对应目标尺寸8倍缩小的特征图 x = layers.Dense(128 * 22 * 26, activation='relu')(x) x = layers.Reshape((22, 26, 128))(x) # 转置卷积逐步上采样到目标尺寸 x = layers.Conv2DTranspose(64, kernel_size=3, strides=2, padding='same', activation='relu')(x) # 形状:(44, 52, 64) x = layers.Conv2DTranspose(32, kernel_size=3, strides=2, padding='same', activation='relu')(x) # 形状:(88, 104, 32) x = layers.Conv2DTranspose(1, kernel_size=3, strides=2, padding='same', activation='sigmoid')(x) # 形状:(176, 208, 1) # 将sigmoid输出转为0/1的二进制掩码 outputs = layers.Lambda(lambda x: keras.backend.round(x))(x) model = keras.Model(inputs=inputs, outputs=outputs) return model # 初始化模型并查看结构 model = build_mask_generator() model.summary()
四、训练配置
针对二进制掩码生成任务,采用二元交叉熵作为损失函数,Adam作为优化器:
model.compile(optimizer=keras.optimizers.Adam(learning_rate=1e-4), loss=keras.losses.BinaryCrossentropy(), metrics=['accuracy'])
训练时,输入为批量的35维索引数组(形状:(batch_size, 35)),目标为对应的批量掩码数组(形状:(batch_size, 176, 208, 1))。
五、场景2的调整(索引对应特定行的列)
如果35个索引对应35行的列位置(比如第i个索引对应第i行),只需修改掩码生成函数:
def indices_to_mask_scenario2(indices, img_shape=(176, 208)): mask = np.zeros(img_shape, dtype=np.float32) # 仅指定行的对应列设为1 for row_idx, col_idx in enumerate(indices): if row_idx < img_shape[0]: # 避免索引超出图像行数 mask[row_idx, col_idx] = 1.0 return mask[np.newaxis, ..., np.newaxis]
网络结构无需改动,直接沿用相同训练流程即可。
内容的提问来源于stack exchange,提问作者Razor
相关产品推荐
相关产品推荐

