如何在PyTorch中实现3D张量到单位球切面的可微2D投影?
3D张量到单位球相切平面的可微投影(PyTorch实现)
原理推导
首先明确单位球的相切平面定义:设切点为单位向量 $\mathbf{p}$(满足 $|\mathbf{p}|=1$),则该点处的相切平面方程为 $\mathbf{p} \cdot \mathbf{x} = 1$(平面上所有点 $\mathbf{x}$ 与 $\mathbf{p}$ 的点积等于1)。
对于任意3D点 $\mathbf{v}$,其到该平面的正交投影公式推导如下:
- 计算点 $\mathbf{v}$ 到平面的带符号距离:$d = \mathbf{p} \cdot \mathbf{v} - 1$
- 投影点即为原向量减去沿法线方向($\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
相关产品推荐
相关产品推荐

