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

PyTorch张量元素等式约束的高效优化实现方案问询

PyTorch张量元素等式约束的高效优化实现方案问询

大家好,我现在遇到一个PyTorch优化的问题:我需要在优化过程中让张量的某些元素始终保持相等。举个例子,比如一个2×9的张量,相同颜色标记的元素必须一直相等。

我先做了个1×4的最小示例,初始时前两个元素相等,后两个元素也相等:

import torch

x1 = torch.tensor([1.2, 1.2, -0.3, -0.3], requires_grad=True)
print(x1)
# tensor([ 1.2000,  1.2000, -0.3000, -0.3000])

如果直接做普通的最小二乘优化,等式约束就失效了:

y = torch.arange(4)

opt_1 = torch.optim.SGD([x1], lr=0.1)
opt_1.zero_grad()
loss = (y - x1).pow(2).sum()
loss.backward()
opt_1.step()

print(x1)
# tensor([0.9600, 1.1600, 0.1600, 0.3600], requires_grad=True)

我试过用掩码加权和的方式来构造张量,这样优化后约束能保持:

def weighted_sum(c, masks):
    return torch.sum(torch.stack([c[0] * masks[0], c[1] * masks[1]]), axis=0)

c = torch.tensor([1.2, -0.3], requires_grad=True)
masks = torch.tensor([[1, 1, 0, 0], [0, 0, 1, 1]])

x2 = weighted_sum(c, masks)
print(x2)
# tensor([ 1.2000,  1.2000, -0.3000, -0.3000])

优化后的结果确实保持了元素相等:

opt_c = torch.optim.SGD([c], lr=0.1)
opt_c.zero_grad()
y = torch.arange(4)
x2 = weighted_sum(c, masks)
loss = (y - x2).pow(2).sum()
loss.backward()
opt_c.step()

print(c)
# tensor([0.9200, 0.8200], requires_grad=True)
print(weighted_sum(c, masks))
# tensor([0.9200, 0.9200, 0.8200, 0.8200], grad_fn=<SumBackward1>)

但这个方法有个致命问题——当输入张量维度很高时,需要维护大量的掩码,很容易内存溢出。比如如果输入张量形状是d_0 * d_1 * ... * d_m,等式块数量是k,那就要存一个k * d_0 * d_1 * ... * d_m的大掩码,完全不现实。

我还考虑过用低分辨率张量上采样的方式,但这种方法没法处理不规则的等式块,比如下面这个张量:

tensor([[ 1.2000,  1.2000,  1.2000, -3.1000, -3.1000],
        [-0.1000,  2.0000,  2.0000,  2.0000,  2.0000]])

所以想请教大家,有没有更聪明的方法能在PyTorch里实现这种元素等式约束呢?

备注:内容来源于stack exchange,提问作者Cheng

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.23 12:09:27