Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150 / prepare_assistant_masks.py
Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
3.1 kB
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())}))