spectrum.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251
  1. """
  2. 谱计算核心模块 - 基于 SVCCA 对偶性理论
  3. 核心功能:
  4. - 计算 Gram 矩阵的谱(利用对偶性处理不同维度 d)
  5. - 计算有效秩 (Effective Rank)
  6. - 计算两个谱之间的 Wasserstein-1 距离
  7. 数学基础:
  8. H^T H 与 HH^T 有相同的非零特征值(对偶性定理)
  9. 有效秩定义:
  10. r_eff = exp(S), 其中 S = -Σ σ'_i log σ'_i
  11. σ'_i = σ_i / Σ σ_j (归一化奇异值)
  12. """
  13. import torch
  14. import numpy as np
  15. from typing import Tuple, Optional, Union
  16. def normalize_representations(H: torch.Tensor, method: str = "per_sample_l2") -> torch.Tensor:
  17. """
  18. 对表示矩阵进行归一化,消除 scale 影响
  19. Args:
  20. H: 表示矩阵,shape = (d, N)
  21. method: 归一化方法
  22. - "per_sample_l2": 对每个样本(列)做 L2 归一化(推荐)
  23. - "global_frobenius": 全局 Frobenius 范数归一化
  24. - "none": 不归一化
  25. Returns:
  26. 归一化后的表示矩阵
  27. """
  28. if method == "none":
  29. return H
  30. if method == "global_frobenius":
  31. return H / (H.norm(p="fro") + 1e-10)
  32. if method == "per_sample_l2":
  33. # 对每列(每个样本)做 L2 归一化
  34. norm = H.norm(dim=0, keepdim=True) + 1e-10
  35. return H / norm
  36. raise ValueError(f"Unknown normalize method: {method}")
  37. def compute_gram_spectrum(
  38. H: torch.Tensor,
  39. normalize: str = "per_sample_l2",
  40. return_eigvals: bool = False
  41. ) -> Union[float, Tuple[float, np.ndarray]]:
  42. """
  43. 利用对偶性,通过 Gram 矩阵计算谱和有效秩
  44. 对于 H ∈ R^{d×N}:
  45. - 若 d <= N: 计算 d×d 协方差矩阵 C = HH^T / N
  46. - 若 d > N: 计算 N×N Gram 矩阵 G = H^T H / N(对偶性保证谱相同)
  47. Args:
  48. H: 表示矩阵,shape = (d, N),d=特征维度,N=样本数
  49. normalize: 归一化方法,见 normalize_representations
  50. return_eigvals: 是否返回特征值数组
  51. Returns:
  52. r_eff: 有效秩(标量)
  53. eigvals: (可选) 归一化后的特征值分布
  54. """
  55. # Step 1: 归一化
  56. H_norm = normalize_representations(H, method=normalize)
  57. d, N = H_norm.shape
  58. # Step 2: 选择较小的矩阵计算(对偶性应用)
  59. if d <= N:
  60. # 使用 d×d 协方差矩阵
  61. C = (H_norm @ H_norm.T) / N
  62. else:
  63. # 使用 N×N Gram 矩阵
  64. C = (H_norm.T @ H_norm) / N
  65. # Step 3: 特征值分解(对称矩阵用 eigvalsh)
  66. eigvals_raw = torch.linalg.eigvalsh(C)
  67. # Step 4: 过滤数值噪声,只保留正特征值
  68. # 按降序排列
  69. eigvals_raw = eigvals_raw.flip(0)
  70. mask = eigvals_raw > 1e-9
  71. eigvals = eigvals_raw[mask]
  72. if len(eigvals) == 0:
  73. # 所有特征值都是数值噪声
  74. return (1.0, np.array([])) if return_eigvals else 1.0
  75. # Step 5: 归一化为概率分布
  76. prob = eigvals / eigvals.sum()
  77. # Step 6: 计算谱熵
  78. entropy = -(prob * torch.log(prob + 1e-10)).sum()
  79. # Step 7: 有效秩 = exp(熵)
  80. r_eff = torch.exp(entropy).item()
  81. if return_eigvals:
  82. return r_eff, prob.cpu().numpy()
  83. return r_eff
  84. def compute_spectrum_array(
  85. H: torch.Tensor,
  86. normalize: str = "per_sample_l2"
  87. ) -> Tuple[np.ndarray, np.ndarray]:
  88. """
  89. 计算完整的谱信息(不仅有效秩)
  90. Returns:
  91. eigvals: 原始特征值(降序)
  92. prob: 归一化概率分布
  93. """
  94. H_norm = normalize_representations(H, method=normalize)
  95. d, N = H_norm.shape
  96. # 对偶性选择
  97. if d <= N:
  98. C = (H_norm @ H_norm.T) / N
  99. else:
  100. C = (H_norm.T @ H_norm) / N
  101. eigvals_raw = torch.linalg.eigvalsh(C).flip(0)
  102. mask = eigvals_raw > 1e-9
  103. eigvals = eigvals_raw[mask].cpu().numpy()
  104. prob = eigvals / eigvals.sum()
  105. return eigvals, prob
  106. def wasserstein1_distance(p: np.ndarray, q: np.ndarray) -> float:
  107. """
  108. 计算两个谱概率分布之间的 Wasserstein-1 距离(Earth Mover's Distance)
  109. 对于 1D 分布,W1 距离 = CDF 差的 L1 范数
  110. Args:
  111. p, q: 概率分布数组(会自动归一化)
  112. Returns:
  113. W1 距离
  114. """
  115. # 截断/补零到相同长度
  116. n = max(len(p), len(q))
  117. p_ = np.pad(p, (0, n - len(p)), mode='constant').astype(np.float64)
  118. q_ = np.pad(q, (0, n - len(q)), mode='constant').astype(np.float64)
  119. # 重新归一化(确保是概率分布)
  120. p_ = p_ / (p_.sum() + 1e-10)
  121. q_ = q_ / (q_.sum() + 1e-10)
  122. # 1D Wasserstein = CDF 差的积分
  123. cdf_p = np.cumsum(p_)
  124. cdf_q = np.cumsum(q_)
  125. return float(np.sum(np.abs(cdf_p - cdf_q)))
  126. def compute_spectrum_distance_matrix(
  127. prob_dict: dict,
  128. distance_fn: str = "wasserstein1"
  129. ) -> Tuple[np.ndarray, list]:
  130. """
  131. 计算多个模型谱之间的距离矩阵
  132. Args:
  133. prob_dict: {model_name: prob_distribution}
  134. distance_fn: 距离函数名
  135. Returns:
  136. distance_matrix: n×n 距离矩阵
  137. model_names: 模型名称列表
  138. """
  139. model_names = list(prob_dict.keys())
  140. n = len(model_names)
  141. dist_matrix = np.zeros((n, n))
  142. dist_fn = wasserstein1_distance if distance_fn == "wasserstein1" else None
  143. for i in range(n):
  144. for j in range(i + 1, n):
  145. d = dist_fn(prob_dict[model_names[i]], prob_dict[model_names[j]])
  146. dist_matrix[i, j] = d
  147. dist_matrix[j, i] = d
  148. return dist_matrix, model_names
  149. def effective_rank_from_eigvals(eigvals: np.ndarray) -> float:
  150. """
  151. 从特征值直接计算有效秩
  152. Args:
  153. eigvals: 特征值数组(无需预先归一化)
  154. Returns:
  155. 有效秩
  156. """
  157. eigvals = np.array(eigvals)
  158. eigvals = eigvals[eigvals > 1e-10] # 过滤零值
  159. if len(eigvals) == 0:
  160. return 1.0
  161. # 归一化为概率分布
  162. prob = eigvals / eigvals.sum()
  163. # 谱熵
  164. entropy = -np.sum(prob * np.log(prob + 1e-10))
  165. return float(np.exp(entropy))
  166. # ============= 可视化辅助函数 =============
  167. def compute_top_k_ratio(prob: np.ndarray, k: int = 10) -> float:
  168. """
  169. 计算前 k 个特征值的能量占比
  170. Args:
  171. prob: 归一化特征值分布(降序)
  172. k: 前 k 个
  173. Returns:
  174. 能量占比
  175. """
  176. return float(prob[:k].sum())
  177. def spectrum_summary(r_eff: float, prob: np.ndarray) -> dict:
  178. """
  179. 生成谱的摘要统计
  180. Returns:
  181. 包含各种统计指标的字典
  182. """
  183. return {
  184. "r_eff": r_eff,
  185. "num_eigvals": len(prob),
  186. "top1_ratio": float(prob[0]) if len(prob) > 0 else 0,
  187. "top5_ratio": float(prob[:5].sum()) if len(prob) >= 5 else float(prob.sum()),
  188. "top10_ratio": float(prob[:10].sum()) if len(prob) >= 10 else float(prob.sum()),
  189. "entropy": float(np.log(r_eff)), # S = log(r)
  190. }