如何使用Python的geoopt库在庞加莱圆盘中绘制直线?
使用geoopt在双曲平面绘制测地线(直线)的正确方法
我正在同时学习geoopt的API和Python,写了一段尝试绘制双曲平面直线的代码,但无法正常运行:
import geoopt import torch import matplotlib.pyplot as plt # Create the Poincare ball model poincare = geoopt.PoincareBall() # Define two points in the hyperbolic space point1 = torch.tensor([0.1, 0.2]) point2 = torch.tensor([0.3, 0.4]) #Map the points to the tangent space at the identity element point1_tangent = poincare.expmap(point1, torch.tensor([1.0,0.0,0.0,1.0])) point2_tangent = poincare.expmap(point2, torch.tensor([1.0,0.0,0.0,1.0])) # Map the points back to the hyperbolic space point1_hyperbolic = poincare.logmap(point1_tangent, torch.tensor([1.0,0.0,0.0,1.0])) point2_hyperbolic = poincare.logmap(point2_tangent, torch.tensor([1.0,0.0,0.0,1.0])) # Transform the points using the Poincare ball model transformed_point1 = poincare.mobius_add(torch.tensor([1.0,0.0,0.0,1.0]), point1_hyperbolic) transformed_point2 = poincare.mobius_add(torch.tensor([1.0,0.0,0.0,1.0]), point2_hyperbolic) # Plot the line connecting the two points plt.plot([transformed_point1[0], transformed_point2[0]], [transformed_point1[1], transformed_point2[1]]) plt.show()
代码存在的核心问题
- 维度不匹配:PoincareBall默认是2维双曲空间,你传入的4维张量作为基点完全不符合要求,2维空间的单位元是
torch.tensor([0.0, 0.0])。 - API用法颠倒:
expmap是切空间→双曲空间,logmap是双曲空间→切空间,你搞反了两者的输入输出逻辑。 - 冗余错误操作:
mobius_add的参数错误且完全多余,4维张量不属于当前2维双曲空间的元素。 - 绘制对象错误:直接连接两点的是欧氏直线,不是双曲空间的测地线(双曲直线)。
正确实现代码
import geoopt import torch import matplotlib.pyplot as plt # 创建2维Poincare球模型(双曲平面) poincare = geoopt.PoincareBall() # 定义双曲空间中的两个有效点(范数必须小于1) point1 = torch.tensor([0.1, 0.2]) point2 = torch.tensor([0.3, 0.4]) # 生成两点间的双曲测地线 # 1. 将point2映射到point1的切空间,得到切向量 tangent_vec = poincare.logmap(point2, point1) # 2. 在切空间线性插值,再映射回双曲空间 t = torch.linspace(0, 1, 100) # 生成100个插值点 geodesic_points = poincare.expmap(t.unsqueeze(1) * tangent_vec, point1) # 可视化 plt.figure(figsize=(6,6)) # 绘制测地线 plt.plot(geodesic_points[:,0], geodesic_points[:,1], label='双曲测地线') # 绘制Poincare球边界(单位圆) circle = plt.Circle((0,0), 1, color='black', fill=False) plt.gca().add_patch(circle) # 标记原始点 plt.scatter(point1[0], point1[1], color='red', label='点1') plt.scatter(point2[0], point2[1], color='blue', label='点2') # 设置绘图参数 plt.xlim(-1.1, 1.1) plt.ylim(-1.1, 1.1) plt.gca().set_aspect('equal', adjustable='box') plt.legend() plt.title('Poincare球模型中的双曲测地线') plt.show()
关键逻辑说明
- 测地线生成:双曲空间的测地线不能直接用欧氏插值,必须通过
logmap将终点转换到起点的切空间,插值后再用expmap映射回双曲空间,这样得到的才是双曲直线。 - 边界约束:Poincare球模型中所有有效点的范数都小于1,绘制单位圆边界能清晰展示双曲空间的范围。
内容的提问来源于stack exchange,提问作者Lance Pollard
相关产品推荐
相关产品推荐

