#!/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()