如何无循环为超光谱图像分类的contrastive loss准备超像素正样本对?
无循环实现超光谱图像对比损失的输入准备
核心思路
利用张量的向量化分组与采样操作替代循环遍历逻辑,核心通过掩码矩阵+多项式采样完成每个超像素内的双像素随机选取,全程依赖框架原生张量运算,无显式循环。
实现步骤(PyTorch为例)
假设输入张量:
logit: 形状为[1, num_classes, H, W],超光谱分类模型输出的logitsegments: 形状为[H, W],SLIC超像素分割结果,每个元素为对应像素的超像素ID
1. 张量形状预处理
将logit和超像素结果展平为一维像素序列,方便后续索引操作:
import torch # 展平logit:从[1, C, H, W]转为[H*W, C] logit_flat = logit.squeeze(0).permute(1, 2, 0).reshape(-1, logit.shape[1]) # 展平超像素ID:从[H, W]转为[H*W] segments_flat = segments.reshape(-1)
2. 筛选有效超像素
过滤掉像素数量不足2的超像素(无法生成正样本对):
# 获取所有超像素ID及其对应的像素数量 seg_ids, counts = torch.unique(segments_flat, return_counts=True) # 保留像素数≥2的超像素 valid_mask = counts >= 2 valid_seg_ids = seg_ids[valid_mask]
3. 构建超像素掩码矩阵
生成掩码矩阵,标记每个像素属于哪个有效超像素:
# 掩码矩阵形状:[H*W, N],N为有效超像素数量 # mask[i, k] = 1 表示第i个像素属于第k个有效超像素,否则为0 mask = (segments_flat.unsqueeze(1) == valid_seg_ids.unsqueeze(0)).float()
4. 向量化采样正样本对
利用torch.multinomial对每个超像素完成双像素随机采样:
# 对每个超像素(掩码矩阵转置后每行对应一个超像素),无放回选取2个像素 # samples形状:[N, 2],每行对应一个超像素的两个像素索引 samples = torch.multinomial(mask.T, num_samples=2, replacement=False) # 提取对应logit作为正样本对 emb_i = logit_flat[samples[:, 0]] # 形状[N, num_classes] emb_j = logit_flat[samples[:, 1]] # 形状[N, num_classes]
关键说明
torch.multinomial自动对每个超像素的有效像素位置进行无放回采样,确保两个像素来自同一超像素且不重复。- 所有操作均为PyTorch原生张量运算,可自动适配GPU加速,效率远高于循环实现。
- 若需处理批量输入(logit形状为
[B, C, H, W]),只需在预处理阶段添加批量维度的拆分与拼接逻辑,核心采样逻辑保持不变。
内容的提问来源于stack exchange,提问作者niu yuanzhuo
相关产品推荐
相关产品推荐

