PyTorch中概率单纯形约束下的投影梯度下降矩阵优化实现
概率单纯形约束投影梯度下降实现方案
1. 实现逐行单纯形投影函数
先实现适配PyTorch的向量化投影函数,可高效处理你1000×70000尺寸的矩阵,确保每行满足和为1、所有元素非负的单纯形约束:
import torch def project_simplex(v, row_sum=1.0): """ 将张量的每行投影到概率单纯形 参数: v: 输入张量,形状为[行数, 维度] row_sum: 约束的行和值,默认1 返回: 投影后的张量,形状与输入一致 """ dim = v.shape[1] # 每行元素降序排序 sorted_v, _ = torch.sort(v, descending=True, dim=-1) # 计算前缀和减去约束行和 cumsum_v = torch.cumsum(sorted_v, dim=-1) - row_sum # 构造维度索引 idx = torch.arange(dim, device=v.device) + 1 # 筛选有效阈值位置 valid_mask = sorted_v - cumsum_v / idx > 0 # 取每行最大有效位置对应的阈值 rho = idx[valid_mask].view(v.shape[0], -1)[:, -1] theta = cumsum_v[valid_mask].view(v.shape[0], -1)[:, -1] / rho # 计算最终投影结果 return torch.maximum(v - theta[:, None], torch.tensor(0.0, device=v.device))
2. 修改训练循环加入投影逻辑
只需在optimizer.step()参数更新后,追加投影操作即可,所有投影操作不干扰梯度计算链路:
for epoch in range(500): y_pred=forward(X) y=model(torch.mm(A.float(),X)) l=loss(y,y_pred) l.backward() A.grad.data=-A.grad.data optimizer.step() # 新增:投影A到概率单纯形 with torch.no_grad(): A.data = project_simplex(A.data) optimizer.zero_grad() if epoch%2==0: print("Loss",l,"\n")
注意事项
- 投影操作全程包裹在
torch.no_grad()中,仅修改A的数值部分,不会产生额外计算图开销,也不会打断反向传播流程 - 该实现完全适配GPU运行,1000×70000的矩阵投影不会产生明显的性能损耗
- 兼容带动量的优化器(如Adam、RMSprop等),投影操作不会影响优化器内部的动量缓存
内容的提问来源于stack exchange,提问作者Pratik Karmakar
相关产品推荐
相关产品推荐

