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

PyTorch实现SOM批量权重更新:移除循环并支持批量输入

自组织映射(SOM)批量权重更新优化实现

原实现通过双重循环逐个更新SOM单元权重,且仅处理单个样本,运算效率极低。以下是针对需求的优化实现:

核心优化思路

  1. 利用PyTorch广播机制与向量化运算替代双重循环,大幅提升计算效率
  2. 遍历所有输入样本,对每个样本执行向量化的权重更新操作

完整优化代码

import torch

# 输入批次:批量大小=512,输入维度=84
z = torch.randn(512, 84)

# SOM形状:(高度, 宽度, 输入维度)
som = torch.randn(40, 40, 84)

# 假设row和col为已计算得到的每个样本对应的BMU坐标,形状均为[512]
row = torch.randint(0, 40, (512,))
col = torch.randint(0, 40, (512,))

# 定义初始邻域半径和学习率
neighb_rad = torch.tensor(2.0)
lr = 0.5

# 处理所有输入样本,内部无双重循环
for i in range(z.shape[0]):
    # 获取当前样本对应的BMU权重
    bmu_weights = som[row[i], col[i]]  # 形状:(84,)
    # 向量化计算所有SOM单元与BMU的权重L2距离
    distances = torch.norm(som - bmu_weights.unsqueeze(0).unsqueeze(0), dim=-1)  # 形状:(40,40)
    # 计算邻域衰减系数
    neigh_coeff = torch.exp(-distances / (2 * torch.pow(neighb_rad, 2)))  # 形状:(40,40)
    # 向量化更新所有SOM单元权重
    som += lr * neigh_coeff.unsqueeze(-1) * (z[i] - som)

# 验证更新后的SOM形状
print(f"更新后SOM形状: {som.shape}")
# 输出:更新后SOM形状: torch.Size([40, 40, 84])

优化说明

  • 移除双重循环:通过将BMU权重扩展为(1,1,84)形状,与整个SOM张量(40,40,84)进行广播运算,一次性计算所有单元到BMU的距离矩阵,替代原有的逐元素循环计算。
  • 批量处理所有样本:仅保留外层对输入样本的单循环(因SOM为在线学习,每个样本的更新会改变SOM状态,必须顺序处理),内部完全使用向量化操作完成权重更新,覆盖全部512个输入样本。

内容的提问来源于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:42:54