"""Download pinned public training sources and select 6,000 calibration prompts.""" import argparse import collections import hashlib import json import random from pathlib import Path import pyarrow.parquet as pq from huggingface_hub import hf_hub_download from common import RUN, write_json, write_records, read_json, sha256 SOURCES = { "chat": ("HuggingFaceH4/ultrachat_200k", "8049631c405ae6576f93f445c6b8166f76f5505a", "data/train_sft-00000-of-00003-a3ecf92756993583.parquet"), "code": ("ise-uiuc/Magicoder-OSS-Instruct-75K", "5f839b1f368a76b161028bb9edff055db34022b2", "data-oss_instruct-decontaminated.jsonl"), "tools": ("NousResearch/hermes-function-calling-v1", "dae3e1d28cfbcf4b915c04ea1e072030529b4bda", "func-calling-singleturn.json"), "math": ("openai/gsm8k", "740312add88f781978c0658806c59bc2815b9866", "main/train-00000-of-00001.parquet"), "multilingual": ("CohereLabs/aya_dataset", "f9ea04583f02a8f86404ff6c58bf75fe637df8a2", "data/train-00000-of-00001.parquet"), } def load_source(name): repo, revision, filename = SOURCES[name] path = hf_hub_download(repo, filename, repo_type="dataset", revision=revision, local_dir=RUN / "datasets" / name) if filename.endswith(".parquet"): rows = pq.read_table(path).to_pylist() elif filename.endswith(".jsonl"): rows = [json.loads(s) for s in open(path) if s.strip()] else: rows = json.load(open(path)) assert isinstance(rows, list), type(rows) return rows, {"repository": repo, "revision": revision, "file": filename, "sha256": sha256(path), "rows_available": len(rows)} def message(text): return [{"role": "user", "content": text}] def structured_prompts(n): rows = [] for i in range(n): a, b = 11 + i * 3, 7 + i % 37 variants = [ f'Return only JSON with keys "sum", "difference", "product" for the integers {a} and {b}.', f'A tool returned {{"matches":[{{"name":"item-{i}","price":{a}}},{{"name":"item-{i+1}","price":{b}}}]}}. Return the cheaper item as JSON with keys name and price.', f'Convert these records into CSV with columns name,count: item-{i} has {a}; item-{i+1} has {b}. Output only CSV.', f'An API request for record {i} failed with HTTP 429 and Retry-After: {b}. Explain a safe retry policy and provide a Python implementation.', f'Extract an object with fields city, nights and guests from: "Book accommodation in Vienna for {i%12+1} nights for {i%5+1} guests." Output JSON only.', f'A file-reading tool returned a JSON parse error on line {i%20+1}. Give a debugging plan that preserves the original file and validates the repaired JSON.', ] rows.append({"src": "structured", "source_row": i, "messages": message(f"Request reference: calibration-{i}.\n" + variants[i % len(variants)])}) return rows def main(): ap = argparse.ArgumentParser() ap.add_argument("--seed", type=int, default=15027) args = ap.parse_args() rng = random.Random(args.seed) counts = {"code": 2100, "tools": 900, "chat": 1200, "math": 900, "multilingual": 600} all_rows, sources = [], {} candidate_hashes = set() for category, count in counts.items(): raw, info = load_source(category) sources[category] = info candidates = [] for index, r in enumerate(raw): tools = None language = None if category == "chat": msgs = r["messages"][:3] if len(r["messages"]) >= 3 and rng.random() < .25 else r["messages"][:1] if not msgs or msgs[-1]["role"] != "user": continue elif category == "code": msgs = message(r["problem"]) language = r.get("lang", "unknown") elif category == "math": msgs = message(r["question"]) elif category == "multilingual": language = r["language_code"] if language not in {"deu", "fra", "spa", "zho", "dan", "jpn", "arb", "hin", "por", "ita"}: continue msgs = message(r["inputs"]) else: user_turns = [m["value"] for m in r["conversations"] if m["from"] in {"human", "user"}] if not user_turns: continue msgs = message(user_turns[0]) raw_tools = json.loads(r["tools"]) if isinstance(r["tools"], str) else r["tools"] tools = [{"type": "function", "function": t.get("function", t)} for t in raw_tools] if not all(isinstance(m.get("content"), str) for m in msgs): continue length = sum(len(m["content"]) for m in msgs) if not 20 <= length <= 16000: continue row = {"src": category, "source_row": index, "messages": msgs} if tools: row["tools"] = tools if language: row["language"] = language fingerprint = hashlib.sha256(json.dumps([msgs, tools], sort_keys=True, ensure_ascii=False).encode()).hexdigest() if fingerprint in candidate_hashes: continue candidate_hashes.add(fingerprint) candidates.append(row) rng.shuffle(candidates) # Balance languages instead of allowing the largest source language to dominate. if category in {"multilingual", "code"}: groups = collections.defaultdict(list) for row in candidates: groups[row["language"]].append(row) chosen = [] while len(chosen) < count and groups: for lang in list(sorted(groups)): chosen.append(groups[lang].pop()) if not groups[lang]: del groups[lang] if len(chosen) == count: break else: chosen = candidates[:count] assert len(chosen) == count, (category, len(chosen), count) all_rows.extend(chosen) print(category, len(chosen), flush=True) all_rows.extend(structured_prompts(300)) rng.shuffle(all_rows) seen = set() for i, row in enumerate(all_rows): digest = hashlib.sha256(json.dumps([row["messages"], row.get("tools")], sort_keys=True, ensure_ascii=False).encode()).hexdigest() assert digest not in seen, "Duplicate calibration prompt: " + digest seen.add(digest) row.update(id=i, prompt_sha256=digest, think=rng.random() < (.7 if row["src"] == "math" else .5)) # Split whole examples before generation/capture, preserving category proportions. groups = collections.defaultdict(list) for row in all_rows: groups[row["src"]].append(row) holdout = {r["id"] for g in groups.values() for r in g[:max(1, len(g)//10)]} for row in all_rows: row["split"] = "holdout" if row["id"] in holdout else "calibration" write_records(RUN / "calibration/prompts.jsonl", all_rows) write_json(RUN / "calibration/manifest.json", { "seed": args.seed, "sources": sources, "counts": dict(collections.Counter(r["src"] for r in all_rows)), "holdout_examples": len(holdout), "total_examples": len(all_rows), "prompts_sha256": sha256(RUN / "calibration/prompts.jsonl"), "structured_source": "swift15/corpus.py deterministic templates, not benchmark test prompts", "source_substitution": "Salesforce xLAM returned gated-access 403; use the independently released Apache-2.0 NousResearch source with native Swift tool formatting.", }) print("Wrote", len(all_rows), "prompts;", len(holdout), "held out", flush=True) if __name__ == "__main__": main()