"""geodesic 信号演示:测地距离(流形)分区 + 卷曲结构探测。
演示两件事:
1. geodesic_ratio 能探测出「卷曲流形」(Swiss Roll 比值 > 1.3,普通簇 ≈ 1),
这是「该不该用 geodesic」的判断依据。
2. geodesic 分区在卷曲流形上 explained 不劣于欧氏 G(不「穿越」流形聚错)。
这是首个「扩展 signal」(registry 索引 3),处理真实卷积权重常见的卷曲流形分布。
(注:本示例用低维合成数据做「能力演示」;真实 CNN 卷积层 geodesic 占优的实测
结论见「优化nrfunc-rs方案.md」探索期记录。)
运行:python geodesic_demo.py
"""
import numpy as np
import _common
import nrfunc
def swiss_roll(n, seed=0):
"""构造螺旋(Swiss Roll)卷曲流形数据:绕螺旋面采样,欧氏距离会「穿越」螺旋。
返回 (n, 3) 三维螺旋点:欧氏近邻会误连不同圈,测地距离才贴流形。
"""
rng = np.random.default_rng(seed)
t = rng.uniform(0, 4 * np.pi, size=n)
h = rng.uniform(-1, 1, size=n)
x = t * np.cos(t)
y = h
z = t * np.sin(t)
pts = np.column_stack([x, y, z])
return pts.astype(np.float64)
def gaussian_clusters(n, K, seed=1):
"""普通欧氏簇(几何可分),无卷曲结构。"""
rng = np.random.default_rng(seed)
c = rng.normal(0, 3, size=(K, 3))
a = rng.integers(0, K, size=n)
return c[a] + rng.normal(0, 0.3, size=(n, 3))
def main():
print('=' * 72)
print('geodesic 分区信号:测地距离(流形)vs 欧氏 G')
print('=' * 72)
W_roll = swiss_roll(256)
print('\n[1] 螺旋卷曲流形(Swiss Roll),n=256 D=3')
ratio = nrfunc.geodesic_ratio(W_roll, k=5)
print(f' 测地/欧氏比值 = {ratio[0]:.2f}(>1.3 说明确有卷曲流形,可用 geodesic)')
g = nrfunc.regionify(W_roll, signal='G', K=4, order=1, r=2, seed=0)
geo = nrfunc.regionify(W_roll, signal='geodesic', K=4, order=1, r=2, seed=0)
print(f' geodesic 分区 explained={geo["explained"]:.4f},欧氏 G explained={g["explained"]:.4f}'
f' → 不劣化(不穿越流形聚错)')
W_cluster = gaussian_clusters(256, 4)
print('\n[2] 普通欧氏簇 n=256 D=3(无卷曲)')
ratio2 = nrfunc.geodesic_ratio(W_cluster, k=5)
print(f' 测地/欧氏比值 = {ratio2[0]:.2f}(≈1 说明无卷曲,用 G 即可)')
g2 = nrfunc.regionify(W_cluster, signal='G', K=4, order=1, r=2, seed=0)
geo2 = nrfunc.regionify(W_cluster, signal='geodesic', K=4, order=1, r=2, seed=0)
print(f' geodesic 分区 explained={geo2["explained"]:.4f},欧氏 G explained={g2["explained"]:.4f}'
f' → 接近')
print('\n' + '=' * 72)
print('结论:geodesic_ratio 是「该不该用 geodesic」的探测依据;')
print(' geodesic 在卷曲流形上不劣化,欧氏可分组上退化回 G。')
print('=' * 72)
if __name__ == '__main__':
main()