File size: 6,707 Bytes
580cb69 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 | #!/usr/bin/env python3
"""Prepare compact, assistant-only code SFT binaries for a roughly 2B model."""
from __future__ import annotations
import argparse
import hashlib
import json
from pathlib import Path
from typing import Any, Iterable
import numpy as np
from tokenizers import Tokenizer
def arguments() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
source = parser.add_mutually_exclusive_group(required=True)
source.add_argument("--input", type=Path, help="Local JSON or JSONL records")
source.add_argument("--dataset", help="Hugging Face dataset ID, for example HuggingFaceH4/CodeAlpaca_20K")
parser.add_argument("--dataset-config")
parser.add_argument("--split", default="train")
parser.add_argument("--tokenizer", type=Path, required=True)
parser.add_argument("--out", type=Path, required=True)
parser.add_argument("--max-records", type=int, default=30000)
parser.add_argument("--max-length", type=int, default=1024)
parser.add_argument("--max-answer-tokens", type=int, default=768)
parser.add_argument("--validation-fraction", type=float, default=0.02)
return parser.parse_args()
def records(args: argparse.Namespace) -> Iterable[dict[str, Any]]:
if args.input:
text = args.input.read_text(encoding="utf-8")
value = json.loads(text) if args.input.suffix == ".json" else [json.loads(line) for line in text.splitlines() if line.strip()]
if isinstance(value, dict):
value = value.get("data", value.get("records"))
if not isinstance(value, list):
raise ValueError("local input must be a JSON array/JSONL, or contain data/records array")
yield from value
return
try:
from datasets import load_dataset
except ImportError as exc:
raise SystemExit("Hugging Face input requires: pip install datasets") from exc
dataset = load_dataset(args.dataset, args.dataset_config, split=args.split, streaming=True)
yield from dataset
def instruction_answer(row: dict[str, Any]) -> tuple[str, str] | None:
instruction = row.get("instruction") or row.get("prompt") or row.get("question")
answer = row.get("output") or row.get("completion") or row.get("response") or row.get("answer")
if not isinstance(instruction, str) or not isinstance(answer, str):
messages = row.get("messages")
if isinstance(messages, list):
users = [m.get("content") for m in messages if m.get("role") == "user"]
assistants = [m.get("content") for m in messages if m.get("role") == "assistant"]
if users and assistants:
instruction, answer = users[-1], assistants[-1]
if not isinstance(instruction, str) or not isinstance(answer, str):
return None
extra = row.get("input")
prompt = instruction.strip()
if isinstance(extra, str) and extra.strip():
prompt += "\n\nInput:\n" + extra.strip()
answer = answer.strip()
if len(prompt) < 8 or len(answer) < 8 or "\x00" in prompt or "\x00" in answer:
return None
return prompt, answer
def write_split(path: Path, examples: list[tuple[np.ndarray, np.ndarray]]) -> dict[str, int]:
path.mkdir(parents=True, exist_ok=True)
lengths = np.asarray([len(tokens) for tokens, _ in examples], dtype="<i4")
tokens = np.concatenate([item[0] for item in examples]) if examples else np.asarray([], dtype="<u2")
masks = np.concatenate([item[1] for item in examples]) if examples else np.asarray([], dtype="u1")
tokens.astype("<u2", copy=False).tofile(path / "sft_tokens.u16")
masks.tofile(path / "sft_mask.u8")
lengths.tofile(path / "sft_lens.i32")
return {"records": len(examples), "tokens": int(tokens.size), "assistant_tokens": int(masks.sum())}
def main() -> None:
args = arguments()
if not 0 < args.validation_fraction < 0.5:
raise SystemExit("--validation-fraction must be between 0 and 0.5")
tokenizer = Tokenizer.from_file(str(args.tokenizer))
if tokenizer.get_vocab_size() > 65536:
raise SystemExit("uint16 output requires tokenizer vocabulary <= 65,536")
special = {name: tokenizer.token_to_id(name) for name in ("<s>", "</s>", "<|user|>", "<|assistant|>")}
missing = [name for name, token_id in special.items() if token_id is None]
if missing:
raise SystemExit("tokenizer lacks required tokens: " + ", ".join(missing))
train: list[tuple[np.ndarray, np.ndarray]] = []
validation: list[tuple[np.ndarray, np.ndarray]] = []
seen: set[str] = set()
rejected = 0
for row in records(args):
pair = instruction_answer(row)
if pair is None:
rejected += 1
continue
prompt, answer = pair
digest = hashlib.sha256((prompt + "\0" + answer).encode()).hexdigest()
if digest in seen:
rejected += 1
continue
answer_ids = tokenizer.encode("\n" + answer).ids[: args.max_answer_tokens]
prefix = [special["<s>"], special["<|user|>"]] + tokenizer.encode("\n" + prompt).ids + [special["</s>"], special["<|assistant|>"]]
room = args.max_length - len(prefix) - 1
if room < 16 or len(answer_ids) < 4:
rejected += 1
continue
answer_ids = answer_ids[:room]
sequence = prefix + answer_ids + [special["</s>"]]
mask = [0] * len(prefix) + [1] * (len(answer_ids) + 1)
item = (np.asarray(sequence, dtype="<u2"), np.asarray(mask, dtype="u1"))
bucket = int(digest[:8], 16) / 0xFFFFFFFF
(validation if bucket < args.validation_fraction else train).append(item)
seen.add(digest)
if len(train) + len(validation) >= args.max_records:
break
if not train or not validation:
raise SystemExit("dataset filtering produced an empty train or validation split")
summary = {
"format": "v25-assistant-only-u16-v1",
"source": str(args.input) if args.input else args.dataset,
"dataset_config": args.dataset_config,
"source_split": args.split,
"tokenizer": str(args.tokenizer),
"max_length": args.max_length,
"max_answer_tokens": args.max_answer_tokens,
"split_method": "sha256(prompt + NUL + answer)",
"validation_fraction": args.validation_fraction,
"train": write_split(args.out / "train", train),
"validation": write_split(args.out / "validation", validation),
"rejected_or_duplicate_records": rejected,
}
(args.out / "meta.json").write_text(json.dumps(summary, indent=2) + "\n", encoding="utf-8")
print(json.dumps(summary, sort_keys=True))
if __name__ == "__main__":
main()
|