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

PyTorch优化器几乎保留初始值,双曲路径优化失效求助

Poincaré圆盘模型双曲路径优化问题排查

背景

刚接触PyTorch,完成x²+y²简单优化示例后,尝试基于Poincaré圆盘模型实现双曲度量路径优化:输入一组y坐标相同的初始点,期望优化后点向上移动形成曲线。

1. 定义Poincaré距离函数

import numpy
import torch
import torch.optim as optim
import pandas

# poincare distance between two points
def poincare_distance(x1, y1, x2, y2, R=1):
  r = torch.sqrt(x1**2 + y1**2)
  matg = torch.diag(torch.tensor([2.0, 2.0])) * (4 * R**4/(R**2-r**2)**2)
  delta = torch.tensor( [[x2 - x1], [y2 - y1]] )
  f = torch.sqrt( delta.T @ matg @ delta )
  return f

2. 带约束的损失函数

引入两个Lagrange乘数,约束点间等距且首尾点保持初始值:

# loss function: arc length + equidistance between all points + equality 
# between the fist and last point and the original (should remain the same)
def poincare_loss(coords, k, kequality, eqfirst, eqlast):
    res = torch.zeros( (coords.shape[0]-1) )
    for i in range(0, coords.shape[0]-1):
        res[i] = poincare_distance(coords[i][0], coords[i][1],
                                   coords[i+1][0], coords[i+1][1])
    
    diffs = torch.zeros( (res.shape[0]-1) )
    for i in range(0, res.shape[0]-1):
        diffs[i] = res[i] - res[i+1]

    eq = kequality * (torch.sum((coords[0] - eqfirst)**2) + torch.sum((coords[-1] - eqlast)**2))

    d = torch.sum(res) + k * torch.sum(diffs**2) + eq
    return d

3. 读取初始点数据

# handy
usedevice = 'cpu'
usedtype = torch.float32

# load the file in one float per line: x x x x ... y y y y y y
arc_csv = pandas.read_csv("../mesh/poincare/arc.txt", header=None, delimiter=' ')
arc = torch.tensor(arc_csv.values.astype(numpy.float64), dtype=usedtype, device=usedevice)

# coordinates as tensor [ [x,y], [x,y], ... ]
coords = arc.reshape([2,40]).T

4. 优化循环实现

# initial guess
inicoords = coords.clone().detach().requires_grad_(True)

# lagrange multipliers for the equidistance, and equality of first and last point
k = torch.tensor([100.0], dtype=usedtype, device=usedevice, requires_grad=True)
keq = torch.tensor([1000.0], dtype=usedtype, device=usedevice, requires_grad=True)

# optimizer
learning_rate = 1e-3
optimizer = optim.Adam([inicoords, k, keq], lr=learning_rate)

num_steps = 1000

for step in range(num_steps):
    optimizer.zero_grad()  
    loss = poincare_loss(inicoords, k, keq, coords[0], coords[-1])    
    # print("step", step, "loss", loss)     
    loss.backward()        
    optimizer.step()

optimal_inicoords = inicoords.detach().numpy()
optimal_k = k.detach().numpy()
optimal_keq = keq.detach().numpy()
optimal_value = poincare_loss(inicoords, k, keq, coords[0], coords[-1]).item()

print("Optimal coords:\n", optimal_inicoords)
print("Optimal k:\n", optimal_k)
print("Optimal keq:\n", optimal_keq)
print("Optimal Value:\n", optimal_value)

print("FIRST:\n", optimal_inicoords[0])
print("LAST:\n", optimal_inicoords[-1])

问题

调整Lagrange乘数k、keq及学习率后,优化器几乎保留初始点,y坐标与乘数几乎无变化,怀疑对PyTorch优化器的理解有误,需要排查思路。

排查思路

  • 梯度有效性检查:在loss.backward()后,打印inicoords.grad、k.grad、keq.grad的数值,确认梯度是否为0或极小值。如果梯度接近0,说明当前损失在初始点附近是局部极小,或者距离函数/损失函数的梯度计算有问题。
  • Poincaré距离函数正确性验证:手动计算两个点的Poincaré距离,对比函数输出结果,确认距离函数实现无误。另外检查matg的计算是否符合Poincaré圆盘的度量张量定义,注意当点接近圆盘边界(r→R)时,度量张量会趋于无穷大,可能导致数值不稳定。
  • 损失函数权重平衡性:尝试固定k和keq为常数(不加入优化器),只优化inicoords,看是否能得到变化的结果。如果此时有变化,说明同时优化乘数导致了目标函数的不稳定;如果仍无变化,说明约束权重(keq=1000)过大,完全压制了路径长度的优化,可尝试降低keq的初始值(比如100、10)。
  • 优化器与学习率调整:尝试更换优化器为SGD(配合较大的学习率如1e-2),观察是否有变化。Adam优化器对学习率较敏感,初始学习率1e-3可能过小,可尝试逐步调高到1e-2、5e-2,同时监控损失变化,防止发散。
  • 张量操作的可导性检查:确认poincare_distance中的所有操作都支持自动微分。比如torch.diag(torch.tensor([2.0, 2.0]))创建的张量是否在计算图中,可改为torch.diag(torch.tensor([2.0, 2.0], dtype=usedtype, device=usedevice))确保设备和 dtype 一致,避免隐式类型转换破坏计算图。另外,delta的创建应使用输入张量的设备和 dtype,改为delta = torch.stack([x2 - x1, y2 - y1]).unsqueeze(1),避免创建新的CPU张量。
  • 初始点边界检查:确认初始点的r(x²+y²的平方根)是否接近R=1,如果初始点靠近圆盘边界,度量张量会异常大,导致梯度爆炸或消失,可尝试将初始点向圆盘中心移动,验证优化是否生效。
  • 损失函数输出监控:取消print("step", step, "loss", loss)的注释,观察每一步损失的变化情况。如果损失几乎不变,说明梯度为0或优化器步长不足以更新参数;如果损失突变,说明数值不稳定。

内容的提问来源于stack exchange,提问作者senseiwa

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 03:35:53