| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261 |
- #!/usr/bin/env python
- """
- 从预训练模型提取文本表示
- 使用现有的模型权重,为每个领域的文本提取最后一层隐藏状态,
- 计算 mean pooling 后保存为 H ∈ R^{d×N} 矩阵。
- """
- import os
- # 设置环境变量以减少显存碎片化
- os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True"
- import gc
- import torch
- import json
- from pathlib import Path
- from typing import List, Dict
- from transformers import AutoTokenizer, AutoModel
- import warnings
- warnings.filterwarnings("ignore")
- def load_texts(filepath: str) -> List[str]:
- """从 JSONL 文件加载文本"""
- texts = []
- with open(filepath, "r", encoding="utf-8") as f:
- for line in f:
- record = json.loads(line.strip())
- texts.append(record.get("text", ""))
- return texts
- def extract_representations(
- model_path: str,
- texts: List[str],
- batch_size: int = 4, # 减小 batch_size 以避免 OOM
- max_length: int = 512
- ) -> torch.Tensor:
- """
- 从文本列表提取模型表示
- Args:
- model_path: 模型路径(本地目录)
- texts: 文本列表
- batch_size: batch 大小
- max_length: 最大序列长度
- Returns:
- H: 表示矩阵,shape = (d, N)
- """
- print(f" 加载模型:{model_path}")
- tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
- # 修复:某些模型(如 Mistral)缺少 pad_token,需要手动设置
- if tokenizer.pad_token is None:
- tokenizer.pad_token = tokenizer.eos_token
- print(f" 设置 pad_token = eos_token: {repr(tokenizer.pad_token)}")
- model = AutoModel.from_pretrained(
- model_path,
- torch_dtype=torch.float32,
- device_map="auto",
- trust_remote_code=True,
- output_hidden_states=True
- )
- model.eval()
- print(f" 模型设备:{model.device}")
- print(f" 文本数量:{len(texts)}")
- all_embeddings = []
- with torch.no_grad():
- for i in range(0, len(texts), batch_size):
- batch_texts = texts[i:i + batch_size]
- # Tokenize
- inputs = tokenizer(
- batch_texts,
- return_tensors="pt",
- padding=True,
- truncation=True,
- max_length=max_length
- )
- inputs = inputs.to(model.device)
- # Forward pass
- outputs = model(**inputs)
- # Mean pooling over sequence length
- # hidden_states: (batch, seq_len, hidden_dim)
- hidden = outputs.last_hidden_state
- attention_mask = inputs["attention_mask"].unsqueeze(-1).float()
- # 加权平均(考虑 padding)
- sum_embeddings = (hidden * attention_mask).sum(dim=1)
- valid_tokens = attention_mask.sum(dim=1)
- mean_embeddings = sum_embeddings / (valid_tokens + 1e-9)
- all_embeddings.append(mean_embeddings.cpu())
- print(f" 进度:{min(i + batch_size, len(texts))}/{len(texts)}")
- # 拼接所有 embedding: (N, d)
- embeddings = torch.cat(all_embeddings, dim=0)
- # 转置为 (d, N) 格式
- H = embeddings.T
- # 清理 GPU 内存
- del model, tokenizer
- if torch.cuda.is_available():
- torch.cuda.empty_cache()
- return H
- def save_representation(
- H: torch.Tensor,
- model_name: str,
- domain: str,
- output_dir: str = "model/representations"
- ) -> str:
- """保存表示矩阵"""
- output_path = Path(output_dir) / f"{model_name}_{domain}.pt"
- output_path.parent.mkdir(parents=True, exist_ok=True)
- torch.save(H, output_path)
- # 保存元数据
- meta_path = output_path.with_suffix(".json")
- meta = {
- "model_name": model_name,
- "domain": domain,
- "d": int(H.shape[0]),
- "N": int(H.shape[1]),
- "dtype": str(H.dtype)
- }
- with open(meta_path, "w") as f:
- json.dump(meta, f, indent=2)
- return str(output_path)
- def get_model_name(model_path: str) -> str:
- """从路径提取模型简称"""
- path = Path(model_path)
- name = path.name.lower()
- # 简化命名
- if "qwen" in name:
- return "qwen_7b"
- elif "mistral" in name:
- return "mistral_7b"
- elif "llama3" in name or "llama-3" in name:
- return "llama3_8b"
- elif "gemma" in name:
- return "gemma_9b"
- elif "bert" in name:
- return "bert"
- elif "gpt2" in name:
- return "gpt2"
- else:
- return path.name.replace("-", "_")
- def main():
- print("=" * 60)
- print("模型表示提取")
- print("=" * 60)
- print()
- # 可用模型
- models = [
- "model/weights/Qwen2.5-7B-Instruct",
- "model/weights/Mistral-7B-v0.3",
- "model/weights/llama3-8b",
- ]
- # 可用领域(6 个原始领域 + 5 个 GLUE 领域)
- domains = [
- # 原始领域
- ("news_en", "database/corpus/news_en/texts.jsonl"),
- ("news_zh", "database/corpus/news_zh/texts.jsonl"),
- ("academic", "database/corpus/academic/texts.jsonl"),
- ("code", "database/corpus/code/texts.jsonl"),
- ("dialogue", "database/corpus/dialogue/texts.jsonl"),
- ("literature", "database/corpus/literature/texts.jsonl"),
- # GLUE 领域
- ("glue_mnli", "database/corpus/glue_mnli/texts.jsonl"),
- ("glue_qnli", "database/corpus/glue_qnli/texts.jsonl"),
- ("glue_sst2", "database/corpus/glue_sst2/texts.jsonl"),
- ("glue_stsb", "database/corpus/glue_stsb/texts.jsonl"),
- ("glue_cola", "database/corpus/glue_cola/texts.jsonl"),
- ]
- # 检查哪些文本文件存在
- available_domains = []
- for domain, path in domains:
- if Path(path).exists():
- available_domains.append((domain, path))
- else:
- print(f"跳过(文件不存在):{path}")
- if not available_domains:
- print("错误:没有找到任何文本语料文件!")
- return
- print(f"可用领域:{[d[0] for d in available_domains]}")
- print(f"模型数量:{len(models)}")
- print()
- # 对每个模型 × 领域组合提取表示
- for model_path in models:
- if not Path(model_path).exists():
- print(f"跳过(模型不存在):{model_path}")
- continue
- model_name = get_model_name(model_path)
- print(f"\n[模型:{model_name}]")
- for domain, texts_path in available_domains:
- print(f"\n [{domain}]")
- # 加载文本
- texts = load_texts(texts_path)
- print(f" 加载了 {len(texts)} 条文本")
- if len(texts) == 0:
- print(" 跳过:没有文本")
- continue
- # 提取表示
- H = extract_representations(model_path, texts)
- print(f" 表示矩阵形状:{H.shape}")
- print(f" dtype: {H.dtype}")
- # 保存
- output_path = save_representation(H, model_name, domain)
- print(f" 已保存:{output_path}")
- # 每个领域结束后清理 GPU 缓存
- del H
- gc.collect()
- if torch.cuda.is_available():
- torch.cuda.empty_cache()
- print()
- print("=" * 60)
- print("提取完成!")
- print("=" * 60)
- print()
- print("下一步:运行 MVP 实验")
- print(" python -m experiments.verify --mode mvp")
- if __name__ == "__main__":
- main()
|