Download gpu-sft/scripts/gpu_sft/prepare_sft_data.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 23.2 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/gpu-sft/scripts/gpu_sft/prepare_sft_data.py
- Command line
-
hf download hf://fzzhang/svd-code/gpu-sft/scripts/gpu_sft/prepare_sft_data.py
-
curl -L -o prepare_sft_data.py https://huggingface.co/fzzhang/svd-code/resolve/main/gpu-sft/scripts/gpu_sft/prepare_sft_data.py
23.2 kB
| #!/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 = {"<think>": "<|start_think|>", "</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 | |
| `'</think>' in content` reasoning-split branch never fires, `reasoning_content` stays | |
| empty, and every FINAL assistant turn is rendered as | |
| <|im_start|>assistant\n<think>\n\n</think>\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 <tools></tools> XML tags:\n<tools>" }} | |
| {%- for tool in tools %} | |
| {{- "\n" }} | |
| {{- tool | tojson }} | |
| {%- endfor %} | |
| {{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|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('<tool_response>') and message.content.endswith('</tool_response>')) %} | |
| {%- 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 '</think>' in content %} | |
| {%- set reasoning_content = content.split('</think>')[0].rstrip('\n').split('<think>')[-1].lstrip('\n') %} | |
| {%- set content = content.split('</think>')[-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 %}{{- '<think>\n' + reasoning_content.strip('\n') + '\n</think>\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 %}{{- '<tool_call>\n{"name": "' }} | |
| {{- tool_call.name }} | |
| {{- '", "arguments": ' }} | |
| {%- if tool_call.arguments is string %} | |
| {{- tool_call.arguments }} | |
| {%- else %} | |
| {{- tool_call.arguments | tojson }} | |
| {%- endif %} | |
| {{- '}\n</tool_call>' }}{% 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<tool_response>\n' }} | |
| {{- content }} | |
| {{- '\n</tool_response>' }} | |
| {%- 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 %}{{- '<think>\n\n</think>\n\n' }}{% endgeneration %} | |
| {%- endif %} | |
| {%- endif %}""" | |
| DEFAULT_TEXT_REPLACEMENTS = {"<think>": "<|start_think|>", "</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) | |
| 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() | |