import argparse import hashlib import json from pathlib import Path from transformers import AutoTokenizer p = argparse.ArgumentParser() p.add_argument('--model', required=True) p.add_argument('--input', type=Path, required=True) p.add_argument('--output', type=Path, required=True) a = p.parse_args() tok = AutoTokenizer.from_pretrained(a.model, local_files_only=True) template = tok.chat_template start = '{%- elif message.role == "assistant" %}' stop = '{%- elif message.role == "tool" %}' if template.count(start) != 1 or template.count(stop) != 1: raise ValueError('Pinned template assistant boundaries differ; cannot infer masks') tracked = template.replace(start, start + '\n{%- generation %}').replace(stop, '{%- endgeneration %}\n' + stop) header = tok.encode('<|im_start|>assistant\n', add_special_tokens=False) eos = tok.convert_tokens_to_ids('<|im_end|>') records = {} for row in map(json.loads, a.input.read_text().splitlines()): encoded = tok.apply_chat_template(row['messages'], tools=row.get('tools'), chat_template=tracked, tokenize=True, return_dict=True, return_assistant_tokens_mask=True, add_generation_prompt=False, enable_thinking=False) ids, mask = encoded['input_ids'], encoded['assistant_masks'] if ids != row['token_ids']: raise ValueError(f'Mask instrumentation changed actual tokens: {row["id"]}') index = 0 assistant_messages = 0 while index < len(mask): if not mask[index]: index += 1 continue begin = index while index < len(mask) and mask[index]: index += 1 end = index if ids[begin:begin+len(header)] != header: raise ValueError(f'Unexpected assistant header tokenization: {row["id"]}') assistant_messages += 1 mask[begin:begin+len(header)] = [0] * len(header) eos_positions = [j for j in range(begin+len(header), end) if ids[j] == eos] if len(eos_positions) != 1: raise ValueError(f'Ambiguous assistant termination: {row["id"]}') mask[eos_positions[0]+1:end] = [0] * (end-eos_positions[0]-1) if assistant_messages != sum(m['role'] == 'assistant' for m in row['messages']): raise ValueError(f'Assistant span count mismatch: {row["id"]}') if not any(mask[1:]): raise ValueError(f'No assistant targets: {row["id"]}') records[row['id']] = {'token_ids': ids, 'target_mask': mask[1:], 'domain': row['domain'], 'split': row['split']} result = {'method': 'Instrument pinned Jinja assistant branch with generation tags; verify exact original token equality; exclude supplied role headers and post-EOS newline; include assistant content/tool-call syntax and EOS', 'input_sha256': hashlib.sha256(a.input.read_bytes()).hexdigest(), 'template_sha256': hashlib.sha256(template.encode()).hexdigest(), 'records': records} a.output.write_text(json.dumps(result, ensure_ascii=False) + '\n') print(json.dumps({'records': len(records), 'assistant_targets': sum(sum(r['target_mask']) for r in records.values()), 'all_targets': sum(len(r['target_mask']) for r in records.values())}))