# -*- coding: utf-8 -*-
"""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  # noqa: F401  # 先导入:完成 src 路径引导
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)

    # —— 1. 螺旋卷曲分布:geodesic_ratio 能探测出卷曲 ——
    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'  → 不劣化(不穿越流形聚错)')

    # —— 2. 普通欧氏簇:比值 ≈1,用 G 即可 ——
    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()