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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.24 02:54:04