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

如何用PyTorch向量化高效计算批量输入的SOM获胜单元?

批量计算SOM最佳匹配单元的PyTorch向量化方案

实现思路

通过PyTorch的广播机制和批量张量运算,替代逐样本循环,全程在张量层面完成L2距离计算与BMU定位,既避免numpy/torch数据转换开销,又能利用GPU加速提升效率。

完整代码实现

import torch

# 输入批次:batch-size=512,input-dim=84
z = torch.randn(512, 84)
# SOM结构:(height, width, input-dim)
som = torch.randn(40, 40, 84)

# 调整张量形状以支持广播
z_expanded = z.unsqueeze(1).unsqueeze(1)  # 形状变为 (512, 1, 1, 84)
som_expanded = som.unsqueeze(0)           # 形状变为 (1, 40, 40, 84)

# 批量计算所有样本与SOM单元的L2距离
dist_l2_batch = torch.norm(z_expanded - som_expanded, p=2, dim=-1)
# 结果形状:(512, 40, 40),对应每个样本的40x40距离矩阵

# 批量定位BMU的二维索引
# 先将每个样本的40x40距离矩阵展平,找到最小值的一维索引
flat_bmu_indices = torch.argmin(dist_l2_batch, dim=(-2, -1))
# 将一维索引还原为(row, col)坐标
bmu_rows, bmu_cols = torch.unravel_index(flat_bmu_indices, som.shape[:2])

# 验证第一个样本的结果(与原单样本代码对齐)
print(f"BMU for z[0]; row = {bmu_rows[0].item()}, col = {bmu_cols[0].item()}")

关键细节

  • 广播对齐:通过unsqueeze为输入批次添加两个维度,与SOM的高度、宽度维度对齐,实现每个样本与所有SOM单元的并行计算。
  • 无数据转换:全程使用PyTorch张量运算,无需切换到numpy,减少内存拷贝损耗,天然支持GPU加速。
  • 批量索引转换:torch.argmin在dim=(-2,-1)维度上一次性找到每个样本的最小距离索引,再通过torch.unravel_index批量转换为二维坐标,一步完成所有样本的BMU定位。

内容的提问来源于stack exchange,提问作者Arun

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 11:36:05