emergence_detection.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402
  1. #!/usr/bin/env python
  2. """
  3. 能力涌现谱突变检测 - 方向 9 实验 E3
  4. 功能:
  5. - 检测 r_eff 曲线中的突变点
  6. - 使用 Ruptures 库进行变化点检测
  7. - 关联突变点与模型能力涌现
  8. 用法:
  9. python experiments/emergence_detection.py --output-dir output/e3
  10. """
  11. import os
  12. import json
  13. import torch
  14. import numpy as np
  15. from pathlib import Path
  16. from typing import Dict, List, Tuple, Optional, Any
  17. from dataclasses import dataclass, asdict
  18. from datetime import datetime
  19. from transformers import AutoTokenizer, AutoModelForCausalLM
  20. # 导入谱计算模块
  21. import sys
  22. sys.path.insert(0, str(Path(__file__).parent.parent))
  23. from model.spectrum import compute_gram_spectrum
  24. from model.extractor import save_H, load_H
  25. from experiments.cross_model_convergence import CrossModelConvergence, ExperimentConfig
  26. try:
  27. import ruptures as rpt
  28. HAS_RUPTURES = True
  29. except ImportError:
  30. HAS_RUPTURES = False
  31. print("警告:ruptures 未安装,变化点检测功能不可用")
  32. @dataclass
  33. class EmergenceResult:
  34. """能力涌现检测结果"""
  35. timestamp: str
  36. ability_name: str
  37. r_eff_curve: List[float]
  38. model_order: List[str]
  39. model_mmlu: List[float]
  40. max_jump: dict
  41. changepoints: List[int]
  42. jump_significance: float
  43. class EmergenceDetector:
  44. """能力涌现检测器"""
  45. def __init__(
  46. self,
  47. model_sequence: List[str],
  48. model_paths: Dict[str, str],
  49. model_mmlu: Dict[str, float],
  50. output_dir: str = "experiments/output/emergence"
  51. ):
  52. self.model_sequence = model_sequence
  53. self.model_paths = model_paths
  54. self.model_mmlu = model_mmlu
  55. self.output_dir = Path(output_dir)
  56. self.output_dir.mkdir(parents=True, exist_ok=True)
  57. # 按 MMLU 排序模型
  58. self.sorted_models = sorted(
  59. model_sequence,
  60. key=lambda m: model_mmlu.get(m, 0)
  61. )
  62. def load_texts_for_ability(
  63. self,
  64. ability_name: str
  65. ) -> List[str]:
  66. """根据能力类型加载探针文本"""
  67. from database.corpus import CorpusLoader
  68. loader = CorpusLoader("database/corpus")
  69. # 能力 - 数据集映射
  70. ability_mapping = {
  71. "cot_reasoning": ["bigbench"],
  72. "arithmetic": ["gsm8k", "math"],
  73. "code_generation": ["humaneval", "code"],
  74. "instruction_following": ["alpaca"],
  75. }
  76. datasets = ability_mapping.get(ability_name, ["gsm8k"])
  77. texts = []
  78. for ds in datasets:
  79. try:
  80. texts.extend(loader.load_texts(ds, limit=200))
  81. except FileNotFoundError:
  82. continue
  83. return texts[:500]
  84. def compute_r_eff_curve(
  85. self,
  86. texts: List[str],
  87. ability_name: str
  88. ) -> Dict[str, float]:
  89. """计算每个模型的 r_eff"""
  90. r_effs = {}
  91. for model_name in self.sorted_models:
  92. print(f" 处理 {model_name}...")
  93. # 构建唯一的数据集名称
  94. dataset_name = f"emergence_{ability_name}"
  95. # 提取表示
  96. H = self._extract_hidden_states(model_name, texts)
  97. # 计算 r_eff
  98. H_tensor = torch.tensor(H.numpy() if hasattr(H, "cpu") else H) if not isinstance(H, torch.Tensor) else H
  99. r_eff = compute_gram_spectrum(H_tensor)
  100. r_effs[model_name] = r_eff
  101. # 保存表示
  102. rep_dir = self.output_dir.parent / "representations"
  103. rep_dir.mkdir(parents=True, exist_ok=True)
  104. save_H(H, model_name, dataset_name,
  105. representations_dir=str(rep_dir))
  106. return r_effs
  107. def _extract_hidden_states(
  108. self,
  109. model_name: str,
  110. texts: List[str]
  111. ) -> torch.Tensor:
  112. """提取隐层表示"""
  113. model_path = self.model_paths.get(
  114. model_name,
  115. f"model/weights/{model_name}"
  116. )
  117. tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
  118. # 设置 pad_token
  119. if tokenizer.pad_token is None:
  120. tokenizer.pad_token = tokenizer.eos_token
  121. if tokenizer.eos_token_id is None:
  122. tokenizer.eos_token_id = 50256
  123. model = AutoModelForCausalLM.from_pretrained(
  124. model_path,
  125. torch_dtype=torch.float16,
  126. device_map="auto",
  127. output_hidden_states=True,
  128. trust_remote_code=True
  129. )
  130. model.eval()
  131. all_features = []
  132. batch_size = 8
  133. with torch.no_grad():
  134. for i in range(0, len(texts), batch_size):
  135. batch = texts[i:i + batch_size]
  136. inputs = tokenizer(
  137. batch,
  138. return_tensors="pt",
  139. padding=True,
  140. truncation=True,
  141. max_length=512
  142. ).to(model.device)
  143. outputs = model(**inputs)
  144. hidden = outputs.hidden_states[-1]
  145. attention_mask = inputs["attention_mask"]
  146. # Mean pooling
  147. mask_expanded = attention_mask.unsqueeze(-1).float()
  148. feat = (hidden * mask_expanded).sum(1) / (mask_expanded.sum(1) + 1e-10)
  149. all_features.append(feat.cpu().float())
  150. H = torch.cat(all_features, dim=0).T
  151. return H
  152. def detect_changepoints(
  153. self,
  154. r_effs: Dict[str, float],
  155. ability_name: str
  156. ) -> EmergenceResult:
  157. """检测变化点"""
  158. # 按能力排序的 r_eff 曲线
  159. r_eff_curve = [r_effs[m] for m in self.sorted_models]
  160. mmlu_scores = [self.model_mmlu.get(m, 0) for m in self.sorted_models]
  161. # 方法 1:一阶差分寻找最大跃变
  162. if len(r_eff_curve) > 1:
  163. diffs = np.abs(np.diff(r_eff_curve))
  164. jump_idx = int(np.argmax(diffs))
  165. jump_magnitude = float(diffs[jump_idx])
  166. else:
  167. jump_idx = 0
  168. jump_magnitude = 0
  169. # 方法 2:Ruptures 变化点检测
  170. changepoints = []
  171. if HAS_RUPTURES and len(r_eff_curve) >= 3:
  172. try:
  173. # PELT 算法
  174. algo = rpt.Pelt(model="rbf").fit(np.array(r_eff_curve).reshape(-1, 1))
  175. changepoints = algo.predict(pen=2)
  176. except Exception as e:
  177. print(f" Ruptures 检测失败:{e}")
  178. # 降级:使用二阶差分
  179. if len(r_eff_curve) >= 3:
  180. second_diff = np.diff(np.diff(r_eff_curve))
  181. changepoints = [int(np.argmax(np.abs(second_diff)) + 1) + 1]
  182. # 计算跃变显著性
  183. if len(r_eff_curve) >= 2:
  184. mean_r_eff = np.mean(r_eff_curve)
  185. std_r_eff = np.std(r_eff_curve) if len(r_eff_curve) > 1 else 1
  186. jump_significance = jump_magnitude / (std_r_eff + 1e-10)
  187. else:
  188. jump_significance = 0
  189. result = EmergenceResult(
  190. timestamp=datetime.now().isoformat(),
  191. ability_name=ability_name,
  192. r_eff_curve=r_eff_curve,
  193. model_order=self.sorted_models.copy(),
  194. model_mmlu=mmlu_scores,
  195. max_jump={
  196. "model_index": jump_idx,
  197. "model_name": self.sorted_models[jump_idx] if jump_idx < len(self.sorted_models) else "N/A",
  198. "mmlu_at_jump": mmlu_scores[jump_idx] if jump_idx < len(mmlu_scores) else 0,
  199. "jump_magnitude": jump_magnitude
  200. },
  201. changepoints=changepoints,
  202. jump_significance=jump_significance
  203. )
  204. return result
  205. def run_full_analysis(
  206. self,
  207. abilities: Optional[List[str]] = None
  208. ) -> Dict[str, EmergenceResult]:
  209. """运行完整的能力涌现分析"""
  210. if abilities is None:
  211. abilities = [
  212. "cot_reasoning",
  213. "arithmetic",
  214. "code_generation",
  215. "instruction_following"
  216. ]
  217. results = {}
  218. for ability in abilities:
  219. print(f"\n{'='*60}")
  220. print(f"分析能力:{ability}")
  221. print(f"{'='*60}")
  222. # 加载探针文本
  223. texts = self.load_texts_for_ability(ability)
  224. if not texts:
  225. print(f" 跳过:无可用文本")
  226. continue
  227. print(f" 文本数:{len(texts)}")
  228. # 计算 r_eff 曲线
  229. r_effs = self.compute_r_eff_curve(texts, ability)
  230. # 检测变化点
  231. result = self.detect_changepoints(r_effs, ability)
  232. results[ability] = result
  233. # 打印摘要
  234. print(f"\n 结果摘要:")
  235. print(f" 最大跃变模型:{result.max_jump['model_name']}")
  236. print(f" MMLU@跃变点:{result.max_jump['mmlu_at_jump']:.1f}%")
  237. print(f" 跃变幅度:{result.max_jump['jump_magnitude']:.3f}")
  238. print(f" 跃变显著性:{result.jump_significance:.2f}σ")
  239. if result.changepoints:
  240. print(f" 变化点索引:{result.changepoints}")
  241. # 保存结果
  242. self._save_results(results)
  243. return results
  244. def _save_results(self, results: Dict[str, EmergenceResult]):
  245. """保存结果"""
  246. results_dict = {}
  247. for name, result in results.items():
  248. results_dict[name] = asdict(result)
  249. result_file = self.output_dir / "emergence_results.json"
  250. with open(result_file, "w") as f:
  251. json.dump(results_dict, f, indent=2)
  252. print(f"\n结果已保存到:{result_file}")
  253. # 生成摘要报告
  254. self._generate_summary(results)
  255. def _generate_summary(self, results: Dict[str, EmergenceResult]):
  256. """生成摘要报告"""
  257. # 获取第一个结果来访问 model_order 和 model_mmlu
  258. first_result = list(results.values())[0]
  259. model_order = first_result.model_order
  260. model_mmlu = first_result.model_mmlu
  261. summary = [
  262. "# E3 能力涌现谱突变检测摘要",
  263. "",
  264. f"分析时间:{datetime.now().strftime('%Y-%m-%d %H:%M')}",
  265. "",
  266. "## 检测到的涌现点",
  267. "",
  268. "| 能力 | 涌现模型 | MMLU@涌现 | 跃变幅度 | 显著性 |",
  269. "|------|---------|-----------|----------|--------|"
  270. ]
  271. for name, result in results.items():
  272. summary.append(
  273. f"| {name} | {result.max_jump['model_name']} | "
  274. f"{result.max_jump['mmlu_at_jump']:.1f}% | "
  275. f"{result.max_jump['jump_magnitude']:.3f} | "
  276. f"{result.jump_significance:.2f}σ |"
  277. )
  278. summary.extend([
  279. "",
  280. "## 模型序列 (按 MMLU 排序)",
  281. ""
  282. ])
  283. for i, model in enumerate(model_order):
  284. mmlu = model_mmlu[i]
  285. summary.append(f"{i+1}. {model} (MMLU: {mmlu:.1f}%)")
  286. summary_file = self.output_dir / "emergence_summary.md"
  287. with open(summary_file, "w") as f:
  288. f.write("\n".join(summary))
  289. print(f"摘要报告已保存到:{summary_file}")
  290. def main():
  291. import argparse
  292. parser = argparse.ArgumentParser(description="能力涌现谱突变检测")
  293. parser.add_argument("--output-dir", type=str,
  294. default="experiments/output/emergence",
  295. help="输出目录")
  296. parser.add_argument("--abilities", type=str, nargs="+",
  297. default=["cot_reasoning", "arithmetic", "code_generation", "instruction_following"],
  298. help="要分析的能力列表")
  299. args = parser.parse_args()
  300. # 模型序列和 MMLU 分数(文献值)
  301. MODEL_SEQUENCE = [
  302. "llama3.2-3b-instruct",
  303. "Mistral-7B-v0.3",
  304. "llama3-8b",
  305. "Qwen2.5-7B-Instruct",
  306. "gemma-2-9b-it"
  307. ]
  308. MODEL_PATHS = {
  309. "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
  310. "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
  311. "llama3-8b": "model/weights/llama3-8b",
  312. "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
  313. "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
  314. }
  315. # 近似 MMLU 分数(用于排序)
  316. MODEL_MMLU = {
  317. "llama3.2-3b-instruct": 58.0,
  318. "Mistral-7B-v0.3": 62.5,
  319. "llama3-8b": 68.4,
  320. "Qwen2.5-7B-Instruct": 86.3,
  321. "gemma-2-9b-it": 82.0
  322. }
  323. detector = EmergenceDetector(
  324. model_sequence=MODEL_SEQUENCE,
  325. model_paths=MODEL_PATHS,
  326. model_mmlu=MODEL_MMLU,
  327. output_dir=args.output_dir
  328. )
  329. results = detector.run_full_analysis(args.abilities)
  330. if __name__ == "__main__":
  331. main()