import argparse import json import os import pickle import torch from tqdm import tqdm from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer from lib.tokenize_data_mask import apply_chat_with_tokenize_with_mark CUSTOM_MASK_IDENTIFIER = { "data": ["", ""], "inst": ["", ""], } class SepDataset: name = "sep" def __init__(self, path: str): self.path = path self.records = [] self._load() def _load(self): with open(self.path, "r", encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue prompts = json.loads(line) if not isinstance(prompts, list) or len(prompts) < 2: raise ValueError("Each SEP line must be a JSON list with two prompts.") self.records.append(prompts) def iter_records(self): for idx, prompts in enumerate(self.records): yield idx, prompts @staticmethod def build_metric_masks(instruction_mask, segment_type, is_normal_token, custom_mask): base = [ (seg == "usr") and norm and (cust == "inst") for seg, norm, cust in zip(segment_type, is_normal_token, custom_mask) ] instr = [b and im for b, im in zip(base, instruction_mask)] return { "sep_native": base, "sep_instrtive": instr, } def load_model(model_path): cfg = AutoConfig.from_pretrained(model_path, trust_remote_code=True) tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) if tok.pad_token_id is None: tok.pad_token = tok.eos_token tok.pad_token_id = tok.eos_token_id tok.padding_side = "right" model = AutoModelForCausalLM.from_pretrained( model_path, config=cfg, device_map="auto", torch_dtype=torch.bfloat16, trust_remote_code=True, attn_implementation="eager", ) model.eval() return model, tok @torch.inference_mode() def main(): parser = argparse.ArgumentParser() parser.add_argument("--model", required=True) parser.add_argument("--dataset-sep", default=None) parser.add_argument("--output-dir", required=True) parser.add_argument("--split-size", type=int, default=100) parser.add_argument("--max-data", type=int, default=None) args = parser.parse_args() model, tok = load_model(args.model) n_layers = getattr(model.config, "num_hidden_layers", None) n_heads = getattr(model.config, "num_attention_heads", None) if n_layers is None or n_heads is None: raise ValueError("Model config missing num_hidden_layers or num_attention_heads.") datasets = [] if args.dataset_sep: datasets.append(SepDataset(args.dataset_sep)) os.makedirs(args.output_dir, exist_ok=True) split_idx = 0 raw_output = None total_prompts = sum(len(prompts) for dataset in datasets for _rid, prompts in dataset.iter_records()) if total_prompts == 0: raise ValueError("No prompts loaded. Please provide at least one dataset.") remaining_records = args.max_data for dataset in datasets: record_iter = list(dataset.iter_records()) if remaining_records is not None: record_iter = record_iter[:remaining_records] all_prompts = [ (record_id, prompt) for record_id, prompts in record_iter for prompt in prompts ] for record_id, prompt in tqdm( all_prompts, desc=f"Processing {dataset.name} dataset", leave=True, ): if raw_output is None: raw_output = { "prompts": [], "heads": { f"L{l}_H{h}": { "attn_weight": [], } for l in range(n_layers) for h in range(n_heads) }, } messages = [ {"role": "system", "content": ""}, {"role": "user", "content": prompt}, ] ( input_ids, instruction_mask, _data_mask, segment_type, is_normal_token, custom_mask, _rendered, ) = apply_chat_with_tokenize_with_mark( messages, tok, custom_mask_identifier=CUSTOM_MASK_IDENTIFIER, add_generation_prompt=True, ) input_ids_tensor = torch.tensor([input_ids], dtype=torch.long, device=model.device) attention_mask = torch.ones_like(input_ids_tensor) out = model( input_ids=input_ids_tensor, attention_mask=attention_mask, output_attentions=True, ) attn = out.attentions # tuple layers: (B, H, T, S) user_mask = [seg == "usr" for seg in segment_type] inst_mask = [cust == "inst" for cust in custom_mask] instr_mask = instruction_mask raw_output["prompts"].append( { "token_ids": input_ids, "user_mask": user_mask, "inst_mask": inst_mask, "instr_mask": instr_mask, } ) for l in range(n_layers): layer_attn = attn[l][0] for h in range(n_heads): head_key = f"L{l}_H{h}" head_store = raw_output["heads"][head_key] head_store["attn_weight"].append( layer_attn[h].to(torch.float16).cpu().numpy()[-1, :] ) del out, attn, input_ids_tensor, attention_mask, layer_attn if len(raw_output["prompts"]) >= args.split_size: output_path = os.path.join( args.output_dir, f"head_scoring_raw_split_{split_idx}.pkl" ) with open(output_path, "wb") as f: pickle.dump(raw_output, f) split_idx += 1 raw_output = None if remaining_records is not None: remaining_records -= len(record_iter) if remaining_records <= 0: break if raw_output is not None and raw_output["prompts"]: output_path = os.path.join( args.output_dir, f"head_scoring_raw_split_{split_idx}.pkl" ) with open(output_path, "wb") as f: pickle.dump(raw_output, f) if __name__ == "__main__": main()