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

如何在PyTorch中实现3D张量到单位球切面的可微2D投影?

3D张量到单位球相切平面的可微投影(PyTorch实现)

原理推导

首先明确单位球的相切平面定义:设切点为单位向量 $\mathbf{p}$(满足 $|\mathbf{p}|=1$),则该点处的相切平面方程为 $\mathbf{p} \cdot \mathbf{x} = 1$(平面上所有点 $\mathbf{x}$ 与 $\mathbf{p}$ 的点积等于1)。

对于任意3D点 $\mathbf{v}$,其到该平面的正交投影公式推导如下:

  1. 计算点 $\mathbf{v}$ 到平面的带符号距离:$d = \mathbf{p} \cdot \mathbf{v} - 1$
  2. 投影点即为原向量减去沿法线方向($\mathbf{p}$)的距离分量:$\mathbf{v}_{\text{proj}} = \mathbf{v} - d \cdot \mathbf{p}$

这个公式完全由线性运算和点积构成,天然支持PyTorch的自动微分机制,满足可微性要求。

PyTorch实现代码

import torch

def project_to_tangent_plane(v, p=None):
    """
    将3D张量投影到与单位球相切的平面上,支持批量处理且可微。
    
    参数:
        v: 输入3D张量,形状为 (..., 3),支持任意批量维度
        p: 切点的单位向量,形状为 (3,) 或与v匹配的批量形状 (..., 3)。
           若未指定,默认使用 (1, 0, 0) 作为切点。
    
    返回:
        proj_v: 投影后的张量,形状与输入v相同
    """
    # 默认切点为(1,0,0)
    if p is None:
        p = torch.tensor([1.0, 0.0, 0.0], device=v.device, dtype=v.dtype)
    else:
        # 确保切点是单位向量(防止输入非单位向量导致平面定义错误)
        p = torch.nn.functional.normalize(p, dim=-1)
    
    # 计算点积:p · v,保持批量维度
    dot_product = torch.sum(p * v, dim=-1, keepdim=True)
    # 计算投影向量
    proj_v = v - (dot_product - 1.0) * p
    
    return proj_v

使用示例

# 单个3D点
v_single = torch.tensor([2.0, 3.0, 4.0], requires_grad=True)
proj_single = project_to_tangent_plane(v_single)
print("单个点投影结果:", proj_single)

# 批量3D点
v_batch = torch.tensor([[2.0, 3.0, 4.0], [0.5, -1.0, 2.0]], requires_grad=True)
proj_batch = project_to_tangent_plane(v_batch)
print("批量点投影结果:\n", proj_batch)

# 验证可微性:计算梯度
proj_batch.sum().backward()
print("输入张量的梯度:\n", v_batch.grad)

关键说明

  • 可微性保证:所有运算(点积、乘法、减法)都是PyTorch的可微操作,自动微分系统可以正常计算梯度。
  • 批量支持:函数支持任意前置批量维度(如(batch_size, 3)、(num_samples, batch_size, 3)等),通过dim=-1确保对最后一维(3D坐标)进行运算。
  • 切点灵活性:可以自定义任意单位向量作为切点,函数会自动归一化输入的p,避免因非单位向量导致平面定义错误。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.14 13:31:04