extract_representations.py 7.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261
  1. #!/usr/bin/env python
  2. """
  3. 从预训练模型提取文本表示
  4. 使用现有的模型权重,为每个领域的文本提取最后一层隐藏状态,
  5. 计算 mean pooling 后保存为 H ∈ R^{d×N} 矩阵。
  6. """
  7. import os
  8. # 设置环境变量以减少显存碎片化
  9. os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True"
  10. import gc
  11. import torch
  12. import json
  13. from pathlib import Path
  14. from typing import List, Dict
  15. from transformers import AutoTokenizer, AutoModel
  16. import warnings
  17. warnings.filterwarnings("ignore")
  18. def load_texts(filepath: str) -> List[str]:
  19. """从 JSONL 文件加载文本"""
  20. texts = []
  21. with open(filepath, "r", encoding="utf-8") as f:
  22. for line in f:
  23. record = json.loads(line.strip())
  24. texts.append(record.get("text", ""))
  25. return texts
  26. def extract_representations(
  27. model_path: str,
  28. texts: List[str],
  29. batch_size: int = 4, # 减小 batch_size 以避免 OOM
  30. max_length: int = 512
  31. ) -> torch.Tensor:
  32. """
  33. 从文本列表提取模型表示
  34. Args:
  35. model_path: 模型路径(本地目录)
  36. texts: 文本列表
  37. batch_size: batch 大小
  38. max_length: 最大序列长度
  39. Returns:
  40. H: 表示矩阵,shape = (d, N)
  41. """
  42. print(f" 加载模型:{model_path}")
  43. tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
  44. # 修复:某些模型(如 Mistral)缺少 pad_token,需要手动设置
  45. if tokenizer.pad_token is None:
  46. tokenizer.pad_token = tokenizer.eos_token
  47. print(f" 设置 pad_token = eos_token: {repr(tokenizer.pad_token)}")
  48. model = AutoModel.from_pretrained(
  49. model_path,
  50. torch_dtype=torch.float32,
  51. device_map="auto",
  52. trust_remote_code=True,
  53. output_hidden_states=True
  54. )
  55. model.eval()
  56. print(f" 模型设备:{model.device}")
  57. print(f" 文本数量:{len(texts)}")
  58. all_embeddings = []
  59. with torch.no_grad():
  60. for i in range(0, len(texts), batch_size):
  61. batch_texts = texts[i:i + batch_size]
  62. # Tokenize
  63. inputs = tokenizer(
  64. batch_texts,
  65. return_tensors="pt",
  66. padding=True,
  67. truncation=True,
  68. max_length=max_length
  69. )
  70. inputs = inputs.to(model.device)
  71. # Forward pass
  72. outputs = model(**inputs)
  73. # Mean pooling over sequence length
  74. # hidden_states: (batch, seq_len, hidden_dim)
  75. hidden = outputs.last_hidden_state
  76. attention_mask = inputs["attention_mask"].unsqueeze(-1).float()
  77. # 加权平均(考虑 padding)
  78. sum_embeddings = (hidden * attention_mask).sum(dim=1)
  79. valid_tokens = attention_mask.sum(dim=1)
  80. mean_embeddings = sum_embeddings / (valid_tokens + 1e-9)
  81. all_embeddings.append(mean_embeddings.cpu())
  82. print(f" 进度:{min(i + batch_size, len(texts))}/{len(texts)}")
  83. # 拼接所有 embedding: (N, d)
  84. embeddings = torch.cat(all_embeddings, dim=0)
  85. # 转置为 (d, N) 格式
  86. H = embeddings.T
  87. # 清理 GPU 内存
  88. del model, tokenizer
  89. if torch.cuda.is_available():
  90. torch.cuda.empty_cache()
  91. return H
  92. def save_representation(
  93. H: torch.Tensor,
  94. model_name: str,
  95. domain: str,
  96. output_dir: str = "model/representations"
  97. ) -> str:
  98. """保存表示矩阵"""
  99. output_path = Path(output_dir) / f"{model_name}_{domain}.pt"
  100. output_path.parent.mkdir(parents=True, exist_ok=True)
  101. torch.save(H, output_path)
  102. # 保存元数据
  103. meta_path = output_path.with_suffix(".json")
  104. meta = {
  105. "model_name": model_name,
  106. "domain": domain,
  107. "d": int(H.shape[0]),
  108. "N": int(H.shape[1]),
  109. "dtype": str(H.dtype)
  110. }
  111. with open(meta_path, "w") as f:
  112. json.dump(meta, f, indent=2)
  113. return str(output_path)
  114. def get_model_name(model_path: str) -> str:
  115. """从路径提取模型简称"""
  116. path = Path(model_path)
  117. name = path.name.lower()
  118. # 简化命名
  119. if "qwen" in name:
  120. return "qwen_7b"
  121. elif "mistral" in name:
  122. return "mistral_7b"
  123. elif "llama3" in name or "llama-3" in name:
  124. return "llama3_8b"
  125. elif "gemma" in name:
  126. return "gemma_9b"
  127. elif "bert" in name:
  128. return "bert"
  129. elif "gpt2" in name:
  130. return "gpt2"
  131. else:
  132. return path.name.replace("-", "_")
  133. def main():
  134. print("=" * 60)
  135. print("模型表示提取")
  136. print("=" * 60)
  137. print()
  138. # 可用模型
  139. models = [
  140. "model/weights/Qwen2.5-7B-Instruct",
  141. "model/weights/Mistral-7B-v0.3",
  142. "model/weights/llama3-8b",
  143. ]
  144. # 可用领域(6 个原始领域 + 5 个 GLUE 领域)
  145. domains = [
  146. # 原始领域
  147. ("news_en", "database/corpus/news_en/texts.jsonl"),
  148. ("news_zh", "database/corpus/news_zh/texts.jsonl"),
  149. ("academic", "database/corpus/academic/texts.jsonl"),
  150. ("code", "database/corpus/code/texts.jsonl"),
  151. ("dialogue", "database/corpus/dialogue/texts.jsonl"),
  152. ("literature", "database/corpus/literature/texts.jsonl"),
  153. # GLUE 领域
  154. ("glue_mnli", "database/corpus/glue_mnli/texts.jsonl"),
  155. ("glue_qnli", "database/corpus/glue_qnli/texts.jsonl"),
  156. ("glue_sst2", "database/corpus/glue_sst2/texts.jsonl"),
  157. ("glue_stsb", "database/corpus/glue_stsb/texts.jsonl"),
  158. ("glue_cola", "database/corpus/glue_cola/texts.jsonl"),
  159. ]
  160. # 检查哪些文本文件存在
  161. available_domains = []
  162. for domain, path in domains:
  163. if Path(path).exists():
  164. available_domains.append((domain, path))
  165. else:
  166. print(f"跳过(文件不存在):{path}")
  167. if not available_domains:
  168. print("错误:没有找到任何文本语料文件!")
  169. return
  170. print(f"可用领域:{[d[0] for d in available_domains]}")
  171. print(f"模型数量:{len(models)}")
  172. print()
  173. # 对每个模型 × 领域组合提取表示
  174. for model_path in models:
  175. if not Path(model_path).exists():
  176. print(f"跳过(模型不存在):{model_path}")
  177. continue
  178. model_name = get_model_name(model_path)
  179. print(f"\n[模型:{model_name}]")
  180. for domain, texts_path in available_domains:
  181. print(f"\n [{domain}]")
  182. # 加载文本
  183. texts = load_texts(texts_path)
  184. print(f" 加载了 {len(texts)} 条文本")
  185. if len(texts) == 0:
  186. print(" 跳过:没有文本")
  187. continue
  188. # 提取表示
  189. H = extract_representations(model_path, texts)
  190. print(f" 表示矩阵形状:{H.shape}")
  191. print(f" dtype: {H.dtype}")
  192. # 保存
  193. output_path = save_representation(H, model_name, domain)
  194. print(f" 已保存:{output_path}")
  195. # 每个领域结束后清理 GPU 缓存
  196. del H
  197. gc.collect()
  198. if torch.cuda.is_available():
  199. torch.cuda.empty_cache()
  200. print()
  201. print("=" * 60)
  202. print("提取完成!")
  203. print("=" * 60)
  204. print()
  205. print("下一步:运行 MVP 实验")
  206. print(" python -m experiments.verify --mode mvp")
  207. if __name__ == "__main__":
  208. main()