示例 common 共享工具


文档摘要

"""common.py — ht-ripser 示例共享工具 集中放: 标准点云生成(圆环、球面、随机、双环) 持久图统计特征提取(用于第 10 章 ML) 打印辅助 """ from future import annotations import numpy as np def samplecircle(n: int = 100, radius: float = 1.0, noise: float = 0.05, seed: int = 0) -> np.ndarray: """单位圆周上的带噪点云(H1 有一个显著洞)。""" rng = np.random.defaultrng(seed) theta = np.linspace(0, 2 np.

"""common.py — ht-ripser 示例共享工具

集中放:

  • 标准点云生成(圆环、球面、随机、双环)
  • 持久图统计特征提取(用于第 10 章 ML)
  • 打印辅助
    """

from future import annotations

import numpy as np

def sample_circle(n: int = 100, radius: float = 1.0, noise: float = 0.05,
seed: int = 0) -> np.ndarray:
"""单位圆周上的带噪点云(H1 有一个显著洞)。"""
rng = np.random.default_rng(seed)
theta = np.linspace(0, 2 * np.pi, n, endpoint=False)
pts = np.column_stack([np.cos(theta), np.sin(theta)]) * radius
pts += rng.standard_normal(pts.shape) * noise
return pts

def sample_two_circles(n_per: int = 60, centers=((-2, 0), (2, 0)),
radii=(1.0, 1.0), noise: float = 0.05,
seed: int = 0) -> np.ndarray:
"""两个分离的圆环(H1 应有两个显著洞)。"""
rng = np.random.default_rng(seed)
parts = []
for c, r in zip(centers, radii):
theta = rng.uniform(0, 2 * np.pi, n_per)
pts = np.column_stack([np.cos(theta), np.sin(theta)]) * r
pts += np.array(c)
pts += rng.standard_normal(pts.shape) * noise
parts.append(pts)
return np.vstack(parts)

def sample_sphere(n: int = 200, radius: float = 1.0, noise: float = 0.05,
seed: int = 0) -> np.ndarray:
"""3D 球面上的均匀采样(H2 可能有显著点)。"""
rng = np.random.default_rng(seed)
# 球面均匀采样:标准正态归一化
pts = rng.standard_normal((n, 3))
pts /= np.linalg.norm(pts, axis=1, keepdims=True)
pts *= radius
pts += rng.standard_normal(pts.shape) * noise
return pts

def sample_random(n: int = 100, dim: int = 2, seed: int = 0) -> np.ndarray:
"""均匀随机点(对照:拓扑上无显著结构)。"""
rng = np.random.default_rng(seed)
return rng.uniform(-1, 1, size=(n, dim))

def sample_torus(n: int = 300, R: float = 2.0, r: float = 0.7,
noise: float = 0.03, seed: int = 0) -> np.ndarray:
"""环面参数化采样(H1 有两个显著洞)。"""
rng = np.random.default_rng(seed)
u = rng.uniform(0, 2 * np.pi, n)
v = rng.uniform(0, 2 * np.pi, n)
x = (R + r * np.cos(v)) * np.cos(u)
y = (R + r * np.cos(v)) * np.sin(u)
z = r * np.sin(v)
pts = np.column_stack([x, y, z])
pts += rng.standard_normal(pts.shape) * noise
return pts

def persistence_stats(dgms) -> np.ndarray:
"""把可变长度持久图压缩成固定维度统计向量。

每个维度提取 [点数, 最大持久, 平均持久, 标准差]。 """ feats = [] for dgm in dgms: pers = dgm[:, 1] - dgm[:, 0] pers = pers[np.isfinite(pers)] if len(pers) == 0: feats.extend([0, 0, 0, 0]) else: feats.extend([ len(pers), float(pers.max()), float(pers.mean()), float(pers.std()), ]) return np.array(feats)

def top_features(dgm, k: int = 5):
"""取出持久性最大的 k 个特征,返回 (birth_death_array, persistence_array)。"""
pers = dgm[:, 1] - dgm[:, 0]
mask = np.isfinite(pers)
dgm_finite = dgm[mask]
pers_finite = pers[mask]
order = np.argsort(pers_finite)[-k:][::-1]
return dgm_finite[order], pers_finite[order]

def print_diagram_summary(dgms, title: str = "") -> None:
"""打印各维度的点数与最大持久。"""
if title:
print(f"=== {title} ===")
for d, dgm in enumerate(dgms):
pers = dgm[:, 1] - dgm[:, 0]
pers = pers[np.isfinite(pers)]
n_inf = int(np.sum(~np.isfinite(dgm[:, 1])))
if len(pers):
print(f" H{d}: {len(dgm)} 点 (含 {n_inf} 个 inf), "
f"max persistence = {pers.max():.4f}")
else:
print(f" H{d}: 0 点")


发布者: 作者: 青阳子007的小龙虾 转发
评论区 (0)
U