如何在d维球冠表面均匀采样点?现有方法存在非均匀问题
在d维球冠表面均匀采样满足角距约束的点
需求描述
给定d维球面(S^D)上的点X,以及最大角距θ,生成满足angle(X,Y) ≤ θ且均匀分布在球冠表面的点Y。
现有代码问题
已知生成S^D上随机点的正确实现:
import numpy as np def get_point(N, D): X = np.random.randn(N, D + 1) X /= np.linalg.norm(X, axis=1)[:, np.newaxis] return X
尝试通过旋转矩阵实现球冠采样的代码如下,但采样结果非均匀(可通过3D可视化、角距直方图验证):
def get_sample_around_point(point, theta_max): D = point.shape[0] - 1 theta = np.random.uniform(0, theta_max) rotation_axis = np.random.normal(size=(D + 1,)) rotation_axis /= np.linalg.norm(rotation_axis) cos_theta = np.cos(theta) sin_theta = np.sin(theta) rotation_matrix = ( cos_theta * np.eye(D + 1) + sin_theta * np.outer(rotation_axis, rotation_axis) + (1 - cos_theta) * np.eye(D + 1) ) new_point = np.dot(rotation_matrix, point.T).reshape(1, -1) new_point /= np.linalg.norm(new_point, axis=1)[:, np.newaxis] return new_point.flatten()
示例调用代码:
N = 1 D = 2 theta = np.pi/4 x = get_point(N, D)[0] get_sample_around_point(x, theta)
问题根源
- 角度采样逻辑错误:直接用均匀分布采样θ会导致点集中在靠近中心点X的区域,因为球冠的面积随θ的三角函数(高维下为更高阶形式)递增,必须按面积加权采样θ。
- 旋转矩阵构造错误:原代码的旋转矩阵公式完全不符合Rodrigues旋转规则,会导致旋转后的点角度偏离预期。
正确实现方法
方法1:坐标变换法(高效通用)
将中心点X转换为球面坐标系的极点,在球冠内生成均匀点后再转换回原坐标系:
import numpy as np def uniform_ball_cap_sample(center_point, theta_max, num_samples=1): dim = center_point.shape[0] # 1. 按面积加权采样theta对应的余弦值 cos_theta = np.random.uniform(np.cos(theta_max), 1.0, num_samples) theta = np.arccos(cos_theta) # 2. 生成球冠内的随机点(以原点为极点) if dim == 1: points = np.array([np.cos(theta), np.sin(theta)]).T else: # 生成垂直于极轴的随机方向 random_dirs = np.random.randn(num_samples, dim - 1) random_dirs /= np.linalg.norm(random_dirs, axis=1, keepdims=True) # 构造球冠内的点坐标 r = np.sin(theta).reshape(-1, 1) z = cos_theta.reshape(-1, 1) points = np.hstack([r * random_dirs, z]) # 3. 将点从极点坐标系转换回原中心点的坐标系 if np.allclose(center_point, np.array([0]*(dim-1) + [1])): return points else: # 用Gram-Schmidt构造正交基,将中心点映射为极轴 e_d = center_point / np.linalg.norm(center_point) e_1 = np.zeros(dim) e_1[0] = 1.0 e_1 -= e_1.dot(e_d) * e_d e_1 /= np.linalg.norm(e_1) basis = [e_1] for i in range(1, dim-1): e_i = np.zeros(dim) e_i[i] = 1.0 for b in basis + [e_d]: e_i -= e_i.dot(b) * b e_i /= np.linalg.norm(e_i) basis.append(e_i) basis.append(e_d) rotation_mat = np.array(basis).T # 完成坐标变换 transformed_points = points @ rotation_mat.T return transformed_points
方法2:修正旋转矩阵法
若坚持使用旋转思路,需修正角度采样和旋转矩阵构造:
import numpy as np def rodrigues_rotation_matrix(axis, theta): # 高维通用旋转矩阵构造 dim = axis.shape[0] axis = axis / np.linalg.norm(axis) if dim == 3: # 3维下直接用Rodrigues公式 K = np.array([ [0, -axis[2], axis[1]], [axis[2], 0, -axis[0]], [-axis[1], axis[0], 0] ]) return np.eye(dim) + np.sin(theta)*K + (1 - np.cos(theta))*np.dot(K,K) else: # 高维下用Householder变换实现绕任意轴旋转 v = axis - np.array([1] + [0]*(dim-1)) v /= np.linalg.norm(v) H = np.eye(dim) - 2 * np.outer(v, v) # 构造绕e1轴旋转theta的矩阵 R_e1 = np.eye(dim) R_e1[1:,1:] = np.array([[np.cos(theta), -np.sin(theta)], [np.sin(theta), np.cos(theta)]]) return H.T @ R_e1 @ H def corrected_sample_around_point(point, theta_max): dim = point.shape[0] # 按面积加权采样theta cos_theta = np.random.uniform(np.cos(theta_max), 1.0) theta = np.arccos(cos_theta) # 生成垂直于中心点的随机旋转轴 random_vec = np.random.randn(dim) rotation_axis = random_vec - random_vec.dot(point)*point # 处理共线情况 if np.linalg.norm(rotation_axis) < 1e-10: random_vec = np.random.randn(dim) rotation_axis = random_vec - random_vec.dot(point)*point rotation_axis /= np.linalg.norm(rotation_axis) # 构造正确的旋转矩阵并旋转 rotation_mat = rodrigues_rotation_matrix(rotation_axis, theta) new_point = rotation_mat @ point # 数值误差修正 new_point /= np.linalg.norm(new_point) return new_point
均匀性验证代码
import matplotlib.pyplot as plt from mpl_toolkits.mplot3d import Axes3D # 生成采样点 D = 2 # 对应3维球面 center = get_point(1, D)[0] theta_max = np.pi/4 samples = uniform_ball_cap_sample(center, theta_max, 1000) # 绘制角距分布直方图 cos_distances = samples @ center angles = np.arccos(cos_distances) plt.figure() plt.hist(angles, bins=30) plt.xlabel('角距') plt.ylabel('采样点数量') plt.title('球冠采样角距分布') plt.show() # 3D可视化采样结果 fig = plt.figure() ax = fig.add_subplot(111, projection='3d') ax.scatter(samples[:,0], samples[:,1], samples[:,2], s=5, label='采样点') ax.scatter(center[0], center[1], center[2], s=50, c='red', label='中心点') ax.set_xlabel('X') ax.set_ylabel('Y') ax.set_zlabel('Z') ax.legend() plt.show()
内容的提问来源于stack exchange,提问作者RobJan
相关产品推荐
相关产品推荐

