PyTorch实现SOM批量权重更新:移除循环并支持批量输入
自组织映射(SOM)批量权重更新优化实现
原实现通过双重循环逐个更新SOM单元权重,且仅处理单个样本,运算效率极低。以下是针对需求的优化实现:
核心优化思路
- 利用PyTorch广播机制与向量化运算替代双重循环,大幅提升计算效率
- 遍历所有输入样本,对每个样本执行向量化的权重更新操作
完整优化代码
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
相关产品推荐
相关产品推荐

