Download prepare_assistant_masks.py from Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150: direct link, hf CLI and curl.
- Browser
- Download file 3.1 kB
-
https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/prepare_assistant_masks.py
- Command line
-
hf download hf://Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/prepare_assistant_masks.py
-
curl -L -o prepare_assistant_masks.py https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/prepare_assistant_masks.py
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())})) | |