#!/usr/bin/env python3 # ruff: noqa: E501 """Stage 1 of the GPU port of the Marin/Levanter Qwen3 self-instill SFT recipe. HF dataset (pinned revision) -> Marin conversation transform (DEFAULT_TEXT_REPLACEMENTS) -> Qwen3 chat template with {% generation %} assistant masking -> Levanter greedy *contiguous* packing (max 64 segments / 32768 tokens) -> local .npy memmaps + manifest.json Every step here is a deliberate port of the TPU pipeline. Where the TPU pipeline has a bug or a surprising behaviour, this script reproduces it by default and exposes a flag to turn it off. Read the "FIDELITY NOTES" block before changing anything. FIDELITY NOTES -------------- 1. TEXT REPLACEMENTS (reproduced by default; `--text-replacements none` to disable). marin/transform/conversation/transform_conversation.py applies DEFAULT_TEXT_REPLACEMENTS = {"": "<|start_think|>", "": "<|end_think|>"} to every message because multi_turn_adapter() leaves `replacements=None`. Those two strings are Marin/Llama3 tokens, NOT Qwen3 tokens, so under the Qwen3 tokenizer they encode as ordinary multi-token text. Consequence: the chat template's `'' in content` reasoning-split branch never fires, `reasoning_content` stays empty, and every FINAL assistant turn is rendered as <|im_start|>assistant\n\n\n\n\n<|start_think|>...<|end_think|>...<|im_end|>\n i.e. an EMPTY Qwen3 think block followed by the real reasoning in non-special text. This is what the TPU checkpoints were trained on. Disabling it changes the token layout and the assistant mask, and the resulting model will not match the TPU runs. 2. DOCUMENT ORDER. Packing is greedy over *adjacent* documents in cache order. We use HF dataset row order. The TPU cache order came from the zephyr shard writer; if that interleaved shards, pack boundaries differ. Impact is limited: cross-document attention is blocked and per-token loss weights are unchanged, so only the grouping of documents into 32768-token windows (and hence the microbatch denominators) moves. 3. VOCAB. We do NOT resize embeddings. The HF checkpoint already has vocab_size=151936 while len(tokenizer) is 151669; Levanter padded the *tokenizer* up to the model. We emit the 267 <|padding_i|> tokens only at HF-export time (see train_sft_qwen3.py) so the exported tokenizer matches the TPU exports. Token ids are unaffected. 4. NO SHUFFLE. The TPU experiment passes `shuffle=DATASET_SIZE` (an int) where levanter expects bool | BlockShuffleConfig, which makes shuffling a silent no-op. The trainer therefore walks packed examples in order and wraps around. Nothing to do here. Usage ----- python prepare_sft_data.py \\ --dataset-id fzzhang/qwen3_8b_hs_competition_depth2_probval_instill_n8_valredundancy5_round1 \\ --revision <7-char-sha> \\ --tokenizer Qwen/Qwen3-8B \\ --max-seq-len 32768 \\ --out /data/sft/hs_competition_depth2 \\ --num-proc 16 # prove the encoder matches Levanter's segment-split encoder on 200 sampled docs python prepare_sft_data.py ... --verify 200 # packing statistics only (answers "how many epochs is 2000 steps?") python prepare_sft_data.py ... --stats-only """ from __future__ import annotations import argparse import hashlib import json import os import re import sys import time from itertools import chain from typing import Any import numpy as np # ---------------------------------------------------------------------------------- # Verbatim copy of marin experiments/qwen3_chat_template.py :: QWEN_3_CHAT_TEMPLATE # (upstream Qwen3 template + {% generation %} tags around assistant content, tool-call # bodies and the trailing '<|im_end|>\n'). # ---------------------------------------------------------------------------------- QWEN3_CHAT_TEMPLATE = r"""{%- if tools is defined and tools %} {{- '<|im_start|>system\n' }} {%- if messages[0].role == 'system' %} {{- messages[0].content + '\n\n' }} {%- endif %} {{- "# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within XML tags:\n" }} {%- for tool in tools %} {{- "\n" }} {{- tool | tojson }} {%- endfor %} {{- "\n\n\nFor each function call, return a json object with function name and arguments within XML tags:\n\n{\"name\": , \"arguments\": }\n<|im_end|>\n" }} {%- else %} {%- if messages[0].role == 'system' %} {{- '<|im_start|>system\n' + messages[0].content + '<|im_end|>\n' }} {%- endif %} {%- endif %} {%- set ns = namespace(multi_step_tool=true, last_query_index=messages|length - 1) %} {%- for message in messages[::-1] %} {%- set index = (messages|length - 1) - loop.index0 %} {%- if ns.multi_step_tool and message.role == "user" and message.content is string and not(message.content.startswith('') and message.content.endswith('')) %} {%- set ns.multi_step_tool = false %} {%- set ns.last_query_index = index %} {%- endif %} {%- endfor %} {%- for message in messages %} {%- if message.content is string %} {%- set content = message.content %} {%- else %} {%- set content = '' %} {%- endif %} {%- if (message.role == "user") or (message.role == "system" and not loop.first) %} {{- '<|im_start|>' + message.role + '\n' + content + '<|im_end|>' + '\n' }} {%- elif message.role == "assistant" %} {%- set reasoning_content = '' %} {%- if message.reasoning_content is string %} {%- set reasoning_content = message.reasoning_content %} {%- else %} {%- if '' in content %} {%- set reasoning_content = content.split('')[0].rstrip('\n').split('')[-1].lstrip('\n') %} {%- set content = content.split('')[-1].lstrip('\n') %} {%- endif %} {%- endif %} {%- if loop.index0 > ns.last_query_index %} {%- if loop.last or (not loop.last and reasoning_content) %} {{- '<|im_start|>' + message.role + '\n' }}{% generation %}{{- '\n' + reasoning_content.strip('\n') + '\n\n\n' + content.lstrip('\n') }}{% endgeneration %} {%- else %} {{- '<|im_start|>' + message.role + '\n' }}{% generation %}{{- content }}{% endgeneration %} {%- endif %} {%- else %} {{- '<|im_start|>' + message.role + '\n' }}{% generation %}{{- content }}{% endgeneration %} {%- endif %} {%- if message.tool_calls %} {%- for tool_call in message.tool_calls %} {%- if (loop.first and content) or (not loop.first) %} {% generation %}{{- '\n' }}{% endgeneration %} {%- endif %} {%- if tool_call.function %} {%- set tool_call = tool_call.function %} {%- endif %} {% generation %}{{- '\n{"name": "' }} {{- tool_call.name }} {{- '", "arguments": ' }} {%- if tool_call.arguments is string %} {{- tool_call.arguments }} {%- else %} {{- tool_call.arguments | tojson }} {%- endif %} {{- '}\n' }}{% endgeneration %} {%- endfor %} {%- endif %} {% generation %}{{- '<|im_end|>\n' }}{% endgeneration %} {%- elif message.role == "tool" %} {%- if loop.first or (messages[loop.index0 - 1].role != "tool") %} {{- '<|im_start|>user' }} {%- endif %} {{- '\n\n' }} {{- content }} {{- '\n' }} {%- if loop.last or (messages[loop.index0 + 1].role != "tool") %} {{- '<|im_end|>\n' }} {%- endif %} {%- endif %} {%- endfor %} {%- if add_generation_prompt %} {{- '<|im_start|>assistant\n' }} {%- if enable_thinking is defined and enable_thinking is false %} {% generation %}{{- '\n\n\n\n' }}{% endgeneration %} {%- endif %} {%- endif %}""" DEFAULT_TEXT_REPLACEMENTS = {"": "<|start_think|>", "": "<|end_think|>"} # levanter/tokenizers.py sentinels, used only by --verify. _GEN_START = "__MARIN_GEN_START_7f3a9c__" _GEN_END = "__MARIN_GEN_END_7f3a9c__" MAX_SEGMENTS_PER_EXAMPLE = 64 # levanter datasets.py: 64 when pack is True FORMAT_VERSION = 2 # ---------------------------------------------------------------------------------- # message normalisation (port of multi_turn_adapter + transform_row) # ---------------------------------------------------------------------------------- def normalize_messages( row: dict[str, Any], conversation_column: str, role_key: str, content_key: str, replacements: dict[str, str] | None, ) -> list[dict[str, Any]]: raw = row[conversation_column] out: list[dict[str, Any]] = [] for m in raw: role = m[role_key] content = m.get(content_key) if isinstance(content, str) and replacements: for old, new in replacements.items(): content = content.replace(old, new) msg: dict[str, Any] = {"role": role, "content": content} # marin's OpenAIChatMessage carries these; the template only reads tool_calls. if m.get("tool_calls"): msg["tool_calls"] = m["tool_calls"] out.append(msg) return out # ---------------------------------------------------------------------------------- # encoding # ---------------------------------------------------------------------------------- def encode_conversation(tokenizer, messages: list[dict[str, Any]]) -> tuple[list[int], list[int]]: """HF path: render once, tokenize once, map {% generation %} char spans to tokens.""" enc = tokenizer.apply_chat_template( messages, chat_template=QWEN3_CHAT_TEMPLATE, tokenize=True, add_generation_prompt=False, return_dict=True, return_assistant_tokens_mask=True, ) return list(enc["input_ids"]), list(enc["assistant_masks"]) def _levanter_env(): """Recreate levanter._make_jinja_env([_GenerationSentinelExtension]).""" import time as _time import jinja2 import jinja2.ext import jinja2.sandbox class _GenerationSentinelExtension(jinja2.ext.Extension): tags = {"generation"} def parse(self, parser): lineno = next(parser.stream).lineno body = parser.parse_statements(("name:endgeneration",), drop_needle=True) call = self.call_method("_wrap_generation", []) return jinja2.nodes.CallBlock(call, [], [], body).set_lineno(lineno) @staticmethod def _wrap_generation(caller): return _GEN_START + caller() + _GEN_END def _raise(message): raise jinja2.exceptions.TemplateError(message) # Match transformers.apply_chat_template, which renders with the DEFAULT lenient # Undefined: the Qwen3 template references optional fields (message.tool_calls, # message.reasoning_content) that plain {role, content} messages lack, and under # StrictUndefined `{% if message.tool_calls %}` raises UndefinedError. Lenient # Undefined evaluates missing attrs as falsy, exactly as the real encoding path does. env = jinja2.sandbox.SandboxedEnvironment( undefined=jinja2.Undefined, trim_blocks=True, lstrip_blocks=True, extensions=[_GenerationSentinelExtension], ) env.globals.update({"raise_exception": _raise, "strftime_now": lambda fmt: _time.strftime(fmt)}) return env def encode_conversation_levanter(tokenizer, messages: list[dict[str, Any]]) -> tuple[list[int], list[int]]: """Reference path: levanter's segment-split encoder, used only by --verify.""" env = _levanter_env() rendered = env.from_string(QWEN3_CHAT_TEMPLATE).render( messages=messages, add_generation_prompt=False, bos_token=tokenizer.bos_token or "", eos_token=tokenizer.eos_token or "", ) parts = re.split(f"({re.escape(_GEN_START)}|{re.escape(_GEN_END)})", rendered) ids: list[int] = [] mask: list[int] = [] is_assistant = False for part in parts: if part == _GEN_START: is_assistant = True continue if part == _GEN_END: is_assistant = False continue if not part: continue seg = tokenizer.encode(part, add_special_tokens=False) ids.extend(seg) mask.extend([1 if is_assistant else 0] * len(seg)) return ids, mask # ---------------------------------------------------------------------------------- # packing (line-by-line port of levanter/data/packing.py::pack_documents, # slice_strategy="left", pad_with_zeros=True, max_segments_per_example=64) # ---------------------------------------------------------------------------------- def pack_document_ranges(lengths: np.ndarray, max_length: int, max_segments: int) -> list[tuple[int, int]]: n = int(len(lengths)) ranges: list[tuple[int, int]] = [] i = 0 while i < n: start = i total_segments = 0 running = 0 while i < n: if total_segments + 1 > max_segments: break candidate = running + int(lengths[i]) end_pack_after_this = False if candidate > max_length: if i == start: # single oversized document: keep it and slice from the left end_pack_after_this = True else: break running = candidate total_segments += 1 i += 1 if end_pack_after_this: break if i == start: # defensive; unreachable with max_segments >= 1 i = start + 1 ranges.append((start, i)) return ranges # ---------------------------------------------------------------------------------- def build(args) -> None: from datasets import load_dataset from transformers import AutoTokenizer t0 = time.time() replacements = DEFAULT_TEXT_REPLACEMENTS if args.text_replacements == "marin" else None print(f"[prepare] text replacements: {replacements}", flush=True) tokenizer = AutoTokenizer.from_pretrained(args.tokenizer) if not tokenizer.is_fast: raise RuntimeError("A fast tokenizer is required for return_assistant_tokens_mask.") ds = load_dataset(args.dataset_id, revision=args.revision, split=args.split) if args.limit: ds = ds.select(range(min(args.limit, len(ds)))) print(f"[prepare] loaded {len(ds)} rows from {args.dataset_id}@{args.revision}", flush=True) conversation_column = args.conversation_column role_key, content_key = args.role_key, args.content_key def _map(batch, indices): tok = tokenizer ids_out, mask_out, len_out = [], [], [] for k in range(len(indices)): row = {conversation_column: batch[conversation_column][k]} messages = normalize_messages(row, conversation_column, role_key, content_key, replacements) ids, mask = encode_conversation(tok, messages) if not any(mask): raise ValueError( f"Row {indices[k]} produced no assistant tokens. levanter's ChatProcessor raises here too." ) ids_out.append(ids) mask_out.append(mask) len_out.append(len(ids)) return {"input_ids": ids_out, "assistant_mask": mask_out, "length": len_out} enc = ds.map( _map, batched=True, batch_size=16, with_indices=True, num_proc=args.num_proc, remove_columns=ds.column_names, desc="tokenize", ) lengths = np.asarray(enc["length"], dtype=np.int64) L = args.max_seq_len ranges = pack_document_ranges(lengths, L, MAX_SEGMENTS_PER_EXAMPLE) n_packs = len(ranges) stats = { "n_documents": int(len(lengths)), "doc_tokens_total": int(lengths.sum()), "doc_len_mean": float(lengths.mean()), "doc_len_p50": float(np.percentile(lengths, 50)), "doc_len_p90": float(np.percentile(lengths, 90)), "doc_len_p99": float(np.percentile(lengths, 99)), "doc_len_max": int(lengths.max()), "docs_over_max_seq_len": int((lengths > L).sum()), "tokens_truncated": int(np.clip(lengths - L, 0, None).sum()), "n_packed_examples": n_packs, "docs_per_pack_mean": float(len(lengths) / max(n_packs, 1)), "pack_fill_ratio": float(min(lengths.sum(), n_packs * L) / max(n_packs * L, 1)), } for gb in (64, 128, 256): stats[f"epochs_at_2000_steps_bs{gb}"] = round(2000 * gb / max(n_packs, 1), 3) print("[prepare] stats: " + json.dumps(stats, indent=2), flush=True) if args.stats_only: return os.makedirs(args.out, exist_ok=True) tokens = np.lib.format.open_memmap( os.path.join(args.out, "tokens.npy"), mode="w+", dtype=np.uint32, shape=(n_packs, L) ) amask = np.lib.format.open_memmap( os.path.join(args.out, "assistant_mask.npy"), mode="w+", dtype=np.uint8, shape=(n_packs, L) ) doc_lens_flat: list[int] = [] pack_offsets = np.zeros(n_packs + 1, dtype=np.int64) assistant_counts = np.zeros(n_packs, dtype=np.int64) real_tokens = np.zeros(n_packs, dtype=np.int64) for p, (s, e) in enumerate(ranges): rows = enc[s:e] ids = list(chain.from_iterable(rows["input_ids"])) msk = list(chain.from_iterable(rows["assistant_mask"])) dlens = [len(x) for x in rows["input_ids"]] if len(ids) > L: assert e - s == 1, "levanter never packs two docs when one alone overflows" ids, msk, dlens = ids[:L], msk[:L], [L] n = len(ids) tokens[p, :n] = np.asarray(ids, dtype=np.uint32) tokens[p, n:] = 0 # pad_with_zeros=True amask[p, :n] = np.asarray(msk, dtype=np.uint8) amask[p, n:] = 0 doc_lens_flat.extend(dlens) pack_offsets[p + 1] = len(doc_lens_flat) # levanter loss weight = roll(mask, -1) * not_last_mask => sum(mask[1:]) assistant_counts[p] = int(sum(msk)) - int(msk[0]) real_tokens[p] = n if p % 2000 == 0: print(f"[prepare] packed {p}/{n_packs}", flush=True) tokens.flush() amask.flush() np.save(os.path.join(args.out, "doc_lens.npy"), np.asarray(doc_lens_flat, dtype=np.int32)) np.save(os.path.join(args.out, "pack_offsets.npy"), pack_offsets) np.save(os.path.join(args.out, "assistant_counts.npy"), assistant_counts) np.save(os.path.join(args.out, "real_tokens.npy"), real_tokens) stats["assistant_tokens_total"] = int(assistant_counts.sum()) stats["assistant_token_fraction"] = float(assistant_counts.sum() / max(real_tokens.sum(), 1)) cfg = { "format_version": FORMAT_VERSION, "dataset_id": args.dataset_id, "revision": args.revision, "split": args.split, "tokenizer": args.tokenizer, "max_seq_len": L, "max_segments_per_example": MAX_SEGMENTS_PER_EXAMPLE, "text_replacements": replacements, "chat_template_sha256": hashlib.sha256(QWEN3_CHAT_TEMPLATE.encode()).hexdigest(), "tokenizer_len_unpadded": len(tokenizer), "n_packs": n_packs, "stats": stats, "built_at": time.strftime("%Y-%m-%dT%H:%M:%S"), } with open(os.path.join(args.out, "manifest.json"), "w") as f: json.dump(cfg, f, indent=2) print(f"[prepare] wrote {n_packs} packed examples to {args.out} in {time.time() - t0:.0f}s", flush=True) # ---------------------------------------------------------------------------------- def verify(args) -> None: """Assert the HF encoder == levanter's segment-split encoder on sampled documents.""" from datasets import load_dataset from transformers import AutoTokenizer replacements = DEFAULT_TEXT_REPLACEMENTS if args.text_replacements == "marin" else None tokenizer = AutoTokenizer.from_pretrained(args.tokenizer) ds = load_dataset(args.dataset_id, revision=args.revision, split=args.split) rng = np.random.default_rng(0) idx = rng.choice(len(ds), size=min(args.verify, len(ds)), replace=False) bad = 0 for j, i in enumerate(idx): messages = normalize_messages( ds[int(i)], args.conversation_column, args.role_key, args.content_key, replacements ) a_ids, a_mask = encode_conversation(tokenizer, messages) b_ids, b_mask = encode_conversation_levanter(tokenizer, messages) if a_ids != b_ids or a_mask != b_mask: bad += 1 n = min(len(a_ids), len(b_ids)) first = next((k for k in range(n) if a_ids[k] != b_ids[k] or a_mask[k] != b_mask[k]), n) print( f"[verify] MISMATCH row={int(i)} len_hf={len(a_ids)} len_lev={len(b_ids)} first_diff={first}", file=sys.stderr, ) print(f" hf : {a_ids[max(0, first - 5) : first + 5]}", file=sys.stderr) print(f" lev: {b_ids[max(0, first - 5) : first + 5]}", file=sys.stderr) if j == 0: txt = tokenizer.decode(a_ids[:120]) print(f"[verify] sample head: {txt!r}", flush=True) print(f"[verify] assistant tokens: {sum(a_mask)}/{len(a_mask)}", flush=True) if bad: raise SystemExit(f"[verify] {bad}/{len(idx)} documents disagree with the levanter encoder") print(f"[verify] OK: {len(idx)} documents byte-identical to the levanter encoder", flush=True) def main() -> None: p = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) p.add_argument("--dataset-id", required=True) p.add_argument("--revision", required=True, help="7-char HF sha, pinned exactly like instruction_datasets.py") p.add_argument("--split", default="train") p.add_argument("--tokenizer", default="Qwen/Qwen3-8B") p.add_argument("--max-seq-len", type=int, default=32768) p.add_argument("--out", default=None) p.add_argument("--num-proc", type=int, default=max(1, (os.cpu_count() or 8) // 2)) p.add_argument("--limit", type=int, default=0) p.add_argument("--conversation-column", default="messages") p.add_argument("--role-key", default="role") p.add_argument("--content-key", default="content") p.add_argument("--text-replacements", choices=("marin", "none"), default="marin") p.add_argument("--verify", type=int, default=0, help="verify N sampled docs against the levanter encoder and exit") p.add_argument("--stats-only", action="store_true") args = p.parse_args() if args.verify: verify(args) return if not args.out and not args.stats_only: p.error("--out is required unless --stats-only/--verify") build(args) if __name__ == "__main__": main()