run_direction9_experiments.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389
  1. #!/usr/bin/env python
  2. """
  3. 方向 9 实验运行脚本 - 整合所有实验流程
  4. 实验列表:
  5. - E1: 基础谱收敛验证
  6. - E2: 跨域谱分层验证
  7. - E3: 能力涌现谱突变检测
  8. - E4: 对齐即谱工程
  9. - LCP: 逐层谱收敛轮廓
  10. 用法:
  11. # 运行所有实验
  12. python experiments/run_direction9_experiments.py --all
  13. # 运行单个实验
  14. python experiments/run_direction9_experiments.py --experiment E1
  15. # 生成可视化
  16. python experiments/run_direction9_experiments.py --visualize
  17. """
  18. import os
  19. import sys
  20. import json
  21. import argparse
  22. from pathlib import Path
  23. from datetime import datetime
  24. # 设置 HF 镜像
  25. os.environ["HF_ENDPOINT"] = "https://hf-mirror.com"
  26. # 项目根目录
  27. PROJECT_ROOT = Path(__file__).parent.parent
  28. sys.path.insert(0, str(PROJECT_ROOT))
  29. def check_prerequisites():
  30. """检查实验前置条件"""
  31. print("=" * 60)
  32. print("检查实验前置条件")
  33. print("=" * 60)
  34. issues = []
  35. # 检查模型
  36. print("\n[模型检查]")
  37. required_models = {
  38. "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
  39. "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
  40. "llama3-8b": "model/weights/llama3-8b",
  41. "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
  42. "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
  43. }
  44. for name, path in required_models.items():
  45. model_path = PROJECT_ROOT / path
  46. config_file = model_path / "config.json"
  47. if config_file.exists():
  48. print(f" ✓ {name}")
  49. else:
  50. print(f" ✗ {name} - 缺失")
  51. issues.append(f"模型缺失:{name}")
  52. # 检查数据集
  53. print("\n[数据集检查]")
  54. from database.corpus import CorpusLoader
  55. loader = CorpusLoader(PROJECT_ROOT / "database/corpus")
  56. required_datasets = ["gsm8k", "math", "humaneval", "alpaca", "bigbench", "flores200"]
  57. available = loader.list_available_domains()
  58. for ds in required_datasets:
  59. if ds in available:
  60. stats = loader.get_corpus_stats(ds)
  61. print(f" ✓ {ds}: {stats['num_texts']} 条")
  62. else:
  63. print(f" ✗ {ds} - 缺失")
  64. issues.append(f"数据集缺失:{ds}")
  65. # 检查依赖
  66. print("\n[依赖检查]")
  67. try:
  68. import ruptures
  69. print(" ✓ ruptures")
  70. except ImportError:
  71. print(" ✗ ruptures - 未安装")
  72. issues.append("依赖缺失:ruptures")
  73. try:
  74. import seaborn
  75. print(" ✓ seaborn")
  76. except ImportError:
  77. print(" ✗ seaborn - 未安装")
  78. issues.append("依赖缺失:seaborn")
  79. # GPU 检查
  80. print("\n[GPU 检查]")
  81. import torch
  82. if torch.cuda.is_available():
  83. print(f" ✓ CUDA {torch.version.cuda}")
  84. print(f" ✓ GPU: {torch.cuda.get_device_name(0)}")
  85. print(f" ✓ 显存:{torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")
  86. else:
  87. print(" ✗ CUDA 不可用 - 将使用 CPU (慢)")
  88. print("\n" + "=" * 60)
  89. if issues:
  90. print(f"发现 {len(issues)} 个问题:")
  91. for issue in issues:
  92. print(f" - {issue}")
  93. print("\n建议先解决上述问题再运行实验")
  94. return False
  95. else:
  96. print("✓ 所有前置条件满足")
  97. return True
  98. def run_experiment_e1(output_dir: str):
  99. """运行 E1 实验"""
  100. from experiments.cross_model_convergence import CrossModelConvergence, ExperimentConfig
  101. config = ExperimentConfig(
  102. models=[
  103. "llama3.2-3b-instruct",
  104. "Mistral-7B-v0.3",
  105. "llama3-8b",
  106. "Qwen2.5-7B-Instruct",
  107. "gemma-2-9b-it"
  108. ],
  109. model_paths={
  110. "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
  111. "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
  112. "llama3-8b": "model/weights/llama3-8b",
  113. "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
  114. "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
  115. },
  116. datasets=["gsm8k", "math", "humaneval", "alpaca", "flores200"],
  117. output_dir=output_dir,
  118. batch_size=4, # 降低 batch size 以防止 OOM
  119. max_length=512,
  120. pooling="mean"
  121. )
  122. analyzer = CrossModelConvergence(config)
  123. return analyzer.run_experiment_e1()
  124. def run_experiment_e2(output_dir: str):
  125. """运行 E2 实验 - 跨域谱分层验证"""
  126. from experiments.cross_model_convergence import CrossModelConvergence, ExperimentConfig
  127. # E2 复用 E1 的配置和数据
  128. config = ExperimentConfig(
  129. models=[
  130. "llama3.2-3b-instruct",
  131. "Mistral-7B-v0.3",
  132. "llama3-8b",
  133. "Qwen2.5-7B-Instruct",
  134. "gemma-2-9b-it"
  135. ],
  136. model_paths={
  137. "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
  138. "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
  139. "llama3-8b": "model/weights/llama3-8b",
  140. "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
  141. "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
  142. },
  143. datasets=["gsm8k", "math", "humaneval", "alpaca", "flores200"],
  144. output_dir=output_dir,
  145. batch_size=4,
  146. max_length=512,
  147. pooling="mean"
  148. )
  149. analyzer = CrossModelConvergence(config)
  150. return analyzer.run_experiment_e2()
  151. def run_experiment_e3(output_dir: str):
  152. """运行 E3 实验"""
  153. from experiments.emergence_detection import EmergenceDetector
  154. MODEL_SEQUENCE = [
  155. "llama3.2-3b-instruct",
  156. "Mistral-7B-v0.3",
  157. "llama3-8b",
  158. "Qwen2.5-7B-Instruct",
  159. "gemma-2-9b-it"
  160. ]
  161. MODEL_PATHS = {
  162. "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
  163. "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
  164. "llama3-8b": "model/weights/llama3-8b",
  165. "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
  166. "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
  167. }
  168. MODEL_MMLU = {
  169. "llama3.2-3b-instruct": 58.0,
  170. "Mistral-7B-v0.3": 62.5,
  171. "llama3-8b": 68.4,
  172. "Qwen2.5-7B-Instruct": 86.3,
  173. "gemma-2-9b-it": 82.0
  174. }
  175. detector = EmergenceDetector(
  176. model_sequence=MODEL_SEQUENCE,
  177. model_paths=MODEL_PATHS,
  178. model_mmlu=MODEL_MMLU,
  179. output_dir=output_dir
  180. )
  181. return detector.run_full_analysis()
  182. def run_experiment_e5(output_dir: str, dataset: str = "bigbench", sample_size: int = 100):
  183. """运行 E5 实验 - RS_cross 作为文本难度代理"""
  184. from experiments.text_difficulty_proxy import TextDifficultyProxy, E5Config
  185. config = E5Config(
  186. models=[
  187. "llama3.2-3b-instruct",
  188. "Mistral-7B-v0.3",
  189. "llama3-8b",
  190. "Qwen2.5-7B-Instruct",
  191. "gemma-2-9b-it"
  192. ],
  193. model_paths={
  194. "llama3.2-3b-instruct": "model/weights/llama3.2-3b-instruct",
  195. "Mistral-7B-v0.3": "model/weights/Mistral-7B-v0.3",
  196. "llama3-8b": "model/weights/llama3-8b",
  197. "Qwen2.5-7B-Instruct": "model/weights/Qwen2.5-7B-Instruct",
  198. "gemma-2-9b-it": "model/weights/gemma-2-9b-it"
  199. },
  200. model_mmlu={
  201. "llama3.2-3b-instruct": 58.0,
  202. "Mistral-7B-v0.3": 62.5,
  203. "llama3-8b": 68.4,
  204. "Qwen2.5-7B-Instruct": 86.3,
  205. "gemma-2-9b-it": 82.0
  206. },
  207. dataset=dataset,
  208. output_dir=output_dir,
  209. sample_size=sample_size,
  210. batch_size=4,
  211. max_length=512
  212. )
  213. analyzer = TextDifficultyProxy(config)
  214. return analyzer
  215. def run_visualization(input_dir: str, output_dir: str):
  216. """运行可视化"""
  217. from experiments.visualization_direction9 import Direction9Visualizer
  218. visualizer = Direction9Visualizer(output_dir)
  219. # 加载 E1 结果
  220. e1_result_file = Path(input_dir) / "cross_model/e1_result.json"
  221. if e1_result_file.exists():
  222. with open(e1_result_file) as f:
  223. e1_result = json.load(f)
  224. print("加载 E1 结果成功")
  225. else:
  226. e1_result = {}
  227. print("未找到 E1 结果文件")
  228. # 加载 E3 结果
  229. e3_result_file = Path(input_dir) / "emergence/emergence_results.json"
  230. if e3_result_file.exists():
  231. with open(e3_result_file) as f:
  232. e3_result = json.load(f)
  233. print("加载 E3 结果成功")
  234. else:
  235. e3_result = {}
  236. print("未找到 E3 结果文件")
  237. # 生成可视化
  238. model_mmlu = {
  239. "llama3.2-3b-instruct": 58.0,
  240. "Mistral-7B-v0.3": 62.5,
  241. "llama3-8b": 68.4,
  242. "Qwen2.5-7B-Instruct": 86.3,
  243. "gemma-2-9b-it": 82.0
  244. }
  245. if e1_result.get('r_effs'):
  246. visualizer.plot_rs_cross_comparison(e1_result.get('rs_cross', {}))
  247. # 尝试绘制 MMLU vs r_eff (需要重新组织数据结构)
  248. if e1_result.get('spd_matrices'):
  249. for dataset, spd in e1_result['spd_matrices'].items():
  250. model_names = list(e1_result.get('r_effs', {}).get(dataset, {}).keys())
  251. if model_names:
  252. visualizer.plot_spd_heatmap(spd, model_names, dataset)
  253. if e3_result:
  254. visualizer.plot_emergence_curve(e3_result)
  255. # 生成摘要仪表盘
  256. visualizer.create_summary_dashboard(e1_result, e3_result)
  257. def main():
  258. parser = argparse.ArgumentParser(description="方向 9 实验运行脚本")
  259. parser.add_argument("--all", action="store_true", help="运行所有实验")
  260. parser.add_argument("--experiment", type=str,
  261. choices=["E1", "E2", "E3", "E5", "LCP"],
  262. help="运行指定实验")
  263. parser.add_argument("--dataset", type=str, default="bigbench",
  264. help="数据集名称 (E5 实验用)")
  265. parser.add_argument("--sample", type=int, default=100,
  266. help="采样文本数量 (E5 实验用)")
  267. parser.add_argument("--visualize", action="store_true",
  268. help="生成可视化")
  269. parser.add_argument("--output-dir", type=str,
  270. default="experiments/output",
  271. help="输出目录")
  272. parser.add_argument("--skip-checks", action="store_true",
  273. help="跳过前置检查")
  274. args = parser.parse_args()
  275. print("=" * 60)
  276. print("方向 9 实验运行脚本")
  277. print(f"时间:{datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
  278. print("=" * 60)
  279. # 前置检查
  280. if not args.skip_checks:
  281. if not check_prerequisites():
  282. print("\n前置检查未通过,请解决上述问题")
  283. sys.exit(1)
  284. # 创建输出目录
  285. output_dir = Path(args.output_dir)
  286. output_dir.mkdir(parents=True, exist_ok=True)
  287. # 运行实验
  288. if args.all:
  289. print("\n" + "=" * 60)
  290. print("运行 E1: 基础谱收敛验证")
  291. print("=" * 60)
  292. run_experiment_e1(str(output_dir / "cross_model"))
  293. print("\n" + "=" * 60)
  294. print("运行 E3: 能力涌现谱突变检测")
  295. print("=" * 60)
  296. run_experiment_e3(str(output_dir / "emergence"))
  297. print("\n" + "=" * 60)
  298. print("生成可视化")
  299. print("=" * 60)
  300. run_visualization(str(output_dir), str(output_dir / "figures"))
  301. elif args.experiment:
  302. if args.experiment == "E1":
  303. run_experiment_e1(str(output_dir / "cross_model"))
  304. elif args.experiment == "E2":
  305. run_experiment_e2(str(output_dir / "cross_model"))
  306. elif args.experiment == "E3":
  307. run_experiment_e3(str(output_dir / "emergence"))
  308. elif args.experiment == "E5":
  309. analyzer = run_experiment_e5(
  310. str(output_dir / "e5"),
  311. dataset=args.dataset,
  312. sample_size=args.sample
  313. )
  314. # 加载文本并运行
  315. from database.corpus import CorpusLoader
  316. loader = CorpusLoader(PROJECT_ROOT / "database/corpus")
  317. texts = loader.load_texts(args.dataset, limit=args.sample)
  318. analyzer.run_analysis(texts)
  319. analyzer.save_results()
  320. analyzer.generate_report()
  321. elif args.visualize:
  322. run_visualization(str(output_dir), str(output_dir / "figures"))
  323. else:
  324. parser.print_help()
  325. if __name__ == "__main__":
  326. main()