File size: 3,102 Bytes
12f320c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
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())}))