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
相关产品推荐
相关产品推荐

