示例13 持久图向量化做分类


文档摘要

"""ex13mlfeatures.py — 第 10 章 持久图向量化做分类 对应教程:tutorials/10-advanced.md、tutorials/05-visualization.md 5.5 节。 把持久图压缩成固定维度特征向量,训练一个简单分类器, 区分「圆环」与「随机点云」两类。展示 TDA → ML 的完整管线。 依赖:scikit-learn(pip install scikit-learn) 运行: python ex13mlfeatures.py """ import numpy as np from ripser import ripser from sklearn.

"""ex13_ml_features.py — 第 10 章 持久图向量化做分类

对应教程:tutorials/10-advanced.md、tutorials/05-visualization.md 5.5 节。

把持久图压缩成固定维度特征向量,训练一个简单分类器,
区分「圆环」与「随机点云」两类。展示 TDA → ML 的完整管线。

依赖:scikit-learn(pip install scikit-learn)

运行:
python ex13_ml_features.py
"""

import numpy as np
from ripser import ripser
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import cross_val_score

from common import sample_circle, sample_random, persistence_stats

def build_dataset(n_samples=40, seed=0):
"""构造 (X_features, y_label) 数据集。

每个样本:生成一个点云 → 算持久同调 → 提取统计特征。 标签:1 = 圆环,0 = 随机。 """ rng = np.random.default_rng(seed) X, y = [], [] for i in range(n_samples): # 随机化点数、噪声,让任务更有挑战 n_pts = int(rng.integers(60, 120)) noise = float(rng.uniform(0.03, 0.15)) # 圆环 circle = sample_circle(n=n_pts, noise=noise, seed=i) dgms = ripser(circle, maxdim=1)['dgms'] X.append(persistence_stats(dgms)) y.append(1) # 随机 rand = sample_random(n=n_pts, dim=2, seed=i) dgms = ripser(rand, maxdim=1)['dgms'] X.append(persistence_stats(dgms)) y.append(0) return np.array(X), np.array(y)

def main():
print("构建数据集(圆环 vs 随机点)...")
X, y = build_dataset(n_samples=40, seed=0)
print(f" X shape={X.shape}, 标签分布: {np.bincount(y)}")
print(f" 特征维度: {X.shape[1]}(每维 4 个统计量:点数/max/mean/std)")

# 看一下两类特征均值差异 print("\n=== 两类特征均值对比 ===") feat_names = ['H0_count', 'H0_max', 'H0_mean', 'H0_std', 'H1_count', 'H1_max', 'H1_mean', 'H1_std'] for j, name in enumerate(feat_names): print(f" {name:10s}: 圆环={X[y==1, j].mean():8.3f} " f"随机={X[y==0, j].mean():8.3f}") # 逻辑回归 + 5 折交叉验证 clf = LogisticRegression(max_iter=1000) scores = cross_val_score(clf, X, y, cv=5, scoring='accuracy') print(f"\n=== 逻辑回归 5 折交叉验证 ===") print(f" accuracy: {scores.mean():.3f} ± {scores.std():.3f}") print(f" 各折: {np.round(scores, 3)}") print("\n解读:H1_max(H1 最大持久性)是最具区分度的特征——") print(" 圆环的 H1 显著点持久性远大于随机点。")

if name == "main":
main()


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