基于PyTorch梯度下降的单位单纯形投影结果和不为1问题排查
单位单纯形投影结果和不为1的问题排查
问题背景
在Boyd教授的单位单纯形投影作业解决方案中,推导得到的公式被写成:
g_of_nu = (1/2)*torch.norm(-relu(-(x-nu)))**2 + nu*(torch.sum(x) -1) - x.size()[0]*nu**2
(对应的公式逻辑:将x各元素减去nu后取负,再经过relu取负,计算其二范数平方的一半,加上nu乘以(x元素和减1),再减去元素个数乘以nu的平方)
按照推导,最优值nu*确定后,单位单纯形投影结果为y*=relu(x-nu*1)。由于g_of_nu是严格凹函数,将其取负得到f_of_nu,通过梯度下降寻找全局最小值,但最终得到的y*元素和不为1(示例中为1.6652),问题出在公式推导错误。
复现代码
torch.manual_seed(1) x = torch.randn(10)#.view(-1, 1) x_list = x.tolist() print(list(map(lambda x: round(x, 4), x_list))) nu_0 = torch.tensor(0., requires_grad = True) nu = nu_0 optimizer = torch.optim.SGD([nu], lr=1e-1) nu_old = torch.tensor(float('inf')) steps = 100 eps = 1e-6 i = 1 while torch.norm(nu_old-nu) > eps: nu_old = nu.clone() optimizer.zero_grad() f_of_nu = -( (1/2)*torch.norm(-relu(-(x-nu)))**2 + nu*(torch.sum(x) -1) - x.size()[0]*nu**2 ) f_of_nu.backward() optimizer.step() print(f'At step {i+1:2} the function value is {f_of_nu.item(): 1.4f} and nu={nu: 0.4f}' ) i += 1 y_star = relu(x-nu).cpu().detach() print(list(map(lambda x: round(x, 4), y_star.tolist()))) print(y_star.sum())
运行输出:
[0.6614, 0.2669, 0.0617, 0.6213, -0.4519, -0.1661, -1.5228, 0.3817, -1.0276, -0.5631] At step 1 the function value is -1.9618 and nu= 0.0993 . . . At step 14 the function value is -1.9947 and nu= 0.0665 [0.5948, 0.2004, 0.0, 0.5548, 0.0, 0.0, 0.0, 0.3152, 0.0, 0.0] tensor(1.6652)
函数可视化
torch.manual_seed(1) x = torch.randn(10) nu = torch.linspace(-1, 1, steps=10000) f = lambda x, nu: -( (1/2)*torch.norm(-relu(-(x-nu)))**2 + nu*(torch.sum(x) -1) - x.size()[0]*nu**2 ) f_value_list = np.asarray( [f(x, i) for i in nu.tolist()] ) i_min = np.argmin(f_value_list) print(nu[i_min]) fig, ax = plt.subplots() ax.plot(nu.cpu().detach().numpy(), f_value_list);
可视化得到的最小值点与梯度下降结果一致:
tensor(0.0665)
(函数可视化图显示,f_of_nu的最小值点确实在nu=0.0665处)
错误原因与修正
1. 公式核心错误
单位单纯形投影的对偶函数推导中,正确的g(ν)应该包含投影结果relu(x-ν)的元素和,而非原向量x的元素和,同时不存在-x.size()[0]*nu²这一项。用户的公式存在两处关键错误:
- 错误地将
nu*(torch.sum(relu(x-nu)) - 1)写成了nu*(torch.sum(x) - 1) - 额外添加了错误的
-x.size()[0]*nu²项
正确的对偶函数g(ν)应为:
g_of_nu = (1/2)*torch.norm(relu(x - nu))**2 + nu*(torch.sum(relu(x - nu)) - 1) # 等价于:(1/2)*torch.norm(-relu(-(x-nu)))**2 + nu*(torch.sum(relu(x - nu)) - 1)
2. 修正后的代码
将f_of_nu的计算替换为正确的表达式:
torch.manual_seed(1) x = torch.randn(10) x_list = x.tolist() print(list(map(lambda x: round(x, 4), x_list))) nu_0 = torch.tensor(0., requires_grad = True) nu = nu_0 optimizer = torch.optim.SGD([nu], lr=1e-1) nu_old = torch.tensor(float('inf')) eps = 1e-6 i = 1 while torch.norm(nu_old-nu) > eps: nu_old = nu.clone() optimizer.zero_grad() relu_x_nu = relu(x - nu) g_of_nu = (1/2)*torch.norm(relu_x_nu)**2 + nu*(torch.sum(relu_x_nu) - 1) f_of_nu = -g_of_nu f_of_nu.backward() optimizer.step() print(f'At step {i+1:2} the function value is {f_of_nu.item(): 1.4f} and nu={nu: 0.4f}' ) i += 1 y_star = relu(x-nu).cpu().detach() print(list(map(lambda x: round(x, 4), y_star.tolist()))) print(y_star.sum())
3. 修正后的结果
运行修正后的代码,最终y_star的元素和会趋近于1,符合单位单纯形的约束要求。
内容的提问来源于stack exchange,提问作者Saeed
相关产品推荐
相关产品推荐

