Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150 / prepare_instruction_corpus.py
Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
20.3 kB
#!/usr/bin/env python3
"""Fetch revision-checked HF row snapshots, then prepare exact-template PTQ corpora.
Run --fetch using stdlib Python; run --prepare separately in the isolated pinned
Transformers container. No model is instantiated. Existing snapshots are verified
and reused, making preparation reproducible even if the upstream branch moves.
"""
import argparse
import collections
import concurrent.futures
import hashlib
import importlib.metadata
import json
import re
import time
import unicodedata
import urllib.parse
import urllib.request
from pathlib import Path
REVISION = "c202236235762e1c871ad0ccb60c8ee5ba337b9a"
SOURCES = {
"chat": {"dataset": "HuggingFaceH4/ultrachat_200k", "revision": "8049631c405ae6576f93f445c6b8166f76f5505a", "license": "mit", "config": "default", "splits": {"train_sft": 20, "test_sft": 10}},
"code": {"dataset": "m-a-p/CodeFeedback-Filtered-Instruction", "revision": "a08c213a9748c66c15d0225814be80a2e77adf4a", "license": "apache-2.0", "config": "default", "splits": {"train": 24}},
"multilingual": {"dataset": "CohereLabs/aya_dataset", "revision": "f9ea04583f02a8f86404ff6c58bf75fe637df8a2", "license": "apache-2.0", "config": "default", "splits": {"train": 40, "test": 10}},
"tools": {"dataset": "NousResearch/hermes-function-calling-v1", "revision": "dae3e1d28cfbcf4b915c04ea1e072030529b4bda", "license": "apache-2.0", "config": "func_calling", "splits": {"train": 16}},
}
TARGETS = {"chat": 90000, "code": 90000, "multilingual": 90000, "tools": 30000}
def digest(data):
return hashlib.sha256(data).hexdigest()
def save_json(path, value):
path.write_text(json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n")
def request(url):
for attempt in range(5):
try:
with urllib.request.urlopen(urllib.request.Request(url, headers={"User-Agent": "qwen9b-ptq-corpus/1.0"}), timeout=90) as response:
return response.read(), dict(response.headers)
except (OSError, TimeoutError):
if attempt == 4:
raise
time.sleep(2 ** attempt)
def fetch(root):
folder = root / "instruction-sources"
folder.mkdir(exist_ok=True)
manifest_path = folder / "snapshots.json"
if manifest_path.exists():
manifest = json.loads(manifest_path.read_text())
if manifest["sources"] != SOURCES:
raise ValueError("Existing source manifest does not match pinned configuration")
for item in manifest["files"]:
if digest((folder / item["file"]).read_bytes()) != item["sha256"]:
raise ValueError(f"Snapshot hash mismatch: {item['file']}")
print(json.dumps({"status": "verified_existing_snapshots", "files": len(manifest["files"])}))
return
files = []
jobs = []
for domain, source in SOURCES.items():
api_url = f"https://huggingface.co/api/datasets/{source['dataset']}/revision/{source['revision']}"
raw, headers = request(api_url)
info = json.loads(raw)
if info["sha"] != source["revision"] or info["cardData"]["license"] != source["license"]:
raise ValueError(f"Revision/license mismatch: {domain}")
name = f"{domain}-repository.json"
(folder / name).write_bytes(raw)
files.append({"file": name, "url": api_url, "sha256": digest(raw), "domain": domain, "kind": "repository_metadata"})
card_url = f"https://huggingface.co/datasets/{source['dataset']}/resolve/{source['revision']}/README.md"
card, _ = request(card_url)
name = f"{domain}-dataset-card.txt"
(folder / name).write_bytes(card)
files.append({"file": name, "url": card_url, "sha256": digest(card), "domain": domain, "kind": "license_dataset_card"})
for split, pages in source["splits"].items():
base = "https://datasets-server.huggingface.co/rows?" + urllib.parse.urlencode({"dataset": source["dataset"], "config": source["config"], "split": split})
first, first_headers = request(base + "&offset=0&length=1")
revision = next((v for k, v in first_headers.items() if k.lower() == "x-revision"), None)
if revision != source["revision"]:
raise ValueError(f"Rows API revision mismatch: {domain}/{split}: {revision}")
total = json.loads(first)["num_rows_total"]
for page in range(pages):
offset = page * (total - 50) // (pages - 1)
jobs.append((domain, split, offset, total, base + f"&offset={offset}&length=50"))
def get_page(job):
domain, split, offset, total, url = job
raw, headers = request(url)
revision = next((v for k, v in headers.items() if k.lower() == "x-revision"), None)
payload = json.loads(raw)
if revision != SOURCES[domain]["revision"] or payload["num_rows_total"] != total or payload.get("partial"):
raise ValueError(f"Rows snapshot is incomplete or wrong revision: {url}")
name = f"{domain}-{split}-{offset:07d}.json"
(folder / name).write_bytes(raw)
return {"file": name, "url": url, "sha256": digest(raw), "domain": domain, "source_split": split, "offset": offset, "total_source_rows": total, "rows": len(payload["rows"]), "x_revision": revision, "kind": "rows"}
with concurrent.futures.ThreadPoolExecutor(max_workers=8) as pool:
files.extend(pool.map(get_page, jobs))
save_json(manifest_path, {"schema_version": 1, "sources": SOURCES, "selection": "50-row windows at evenly spaced integer offsets spanning each entire source split; source order preserved inside each window", "files": files})
print(json.dumps({"status": "fetched", "files": len(files), "rows": sum(x.get("rows", 0) for x in files)}))
def normalize_prompt(text):
return digest(" ".join(unicodedata.normalize("NFKC", text).casefold().split()).encode())
def tagged_json(text, tag):
pattern = re.compile(r"<" + tag + r">\s*(.*?)\s*</" + tag + r">", re.DOTALL)
matches = list(pattern.finditer(text))
if not matches or pattern.sub("", text).strip():
raise ValueError(f"Unparsed {tag} content")
return [json.loads(match.group(1)) for match in matches]
def convert(domain, row):
tools = None
if domain == "chat":
messages = row["messages"]
if any(set(m) != {"role", "content"} for m in messages):
raise ValueError("Unexpected chat schema")
elif domain in ("code", "multilingual"):
a, b = ("query", "answer") if domain == "code" else ("inputs", "targets")
messages = [{"role": "user", "content": row[a]}, {"role": "assistant", "content": row[b]}]
else:
tools = json.loads(row["tools"])
if not isinstance(tools, list) or not tools:
raise ValueError("No real tool schemas")
schemas = {}
for tool in tools:
if tool.get("type") != "function" or not isinstance(tool.get("function", {}).get("parameters"), dict):
raise ValueError("Malformed tool schema")
schemas[tool["function"]["name"]] = tool["function"]["parameters"]
messages, pending = [], []
calls_seen = 0
for turn in row["conversations"]:
role, content = turn["from"], turn["value"]
if role == "system":
blocks = re.findall(r"<tools>\s*(.*?)\s*</tools>", content, re.DOTALL)
# The source prose itself mentions <tools> </tools>; parse the final
# block containing actual JSON and remove source-format boilerplate.
if not blocks or json.loads(blocks[-1]) != tools:
raise ValueError("System tool signatures differ from schema")
continue
if role == "gpt" and "<tool_call>" in content:
if pending:
raise ValueError("Unanswered prior calls")
calls = tagged_json(content, "tool_call")
formatted = []
for call in calls:
name, arguments = call["name"], call["arguments"]
if name not in schemas or not isinstance(arguments, dict):
raise ValueError("Call missing schema or structured arguments")
if not set(schemas[name].get("required", [])) <= set(arguments):
raise ValueError("Required tool argument missing")
call_id = f"call_{calls_seen}"
calls_seen += 1
formatted.append({"id": call_id, "type": "function", "function": {"name": name, "arguments": arguments}})
pending.append((name, call_id))
messages.append({"role": "assistant", "content": "", "tool_calls": formatted})
elif role == "tool":
for response in tagged_json(content, "tool_response"):
if not pending or response["name"] != pending[0][0]:
raise ValueError("Tool response/call order mismatch")
name, call_id = pending.pop(0)
value = response["content"]
messages.append({"role": "tool", "name": name, "tool_call_id": call_id, "content": value if isinstance(value, str) else json.dumps(value, ensure_ascii=False)})
else:
if pending or role not in ("human", "gpt"):
raise ValueError("Unsupported or incomplete tool conversation")
messages.append({"role": "user" if role == "human" else "assistant", "content": content})
if pending or not calls_seen:
raise ValueError("No complete real tool roundtrip")
if not messages or messages[0]["role"] != "user" or messages[-1]["role"] != "assistant":
raise ValueError("Require a complete user-to-assistant conversation")
for message in messages:
if message["role"] not in ("user", "assistant", "tool") or not isinstance(message["content"], str):
raise ValueError("Malformed message")
if not message["content"].strip() and not message.get("tool_calls"):
raise ValueError("Empty message")
return messages, tools
def assigned_split(domain, source_split, row_idx):
bucket = int(digest(f"{domain}:{source_split}:{row_idx}:split-v1".encode())[:8], 16) % 10
if domain in ("chat", "multilingual"):
if source_split.startswith("train"):
return "calibration"
return "validation" if bucket < 5 else "heldout"
return "calibration" if bucket < 8 else "validation" if bucket == 8 else "heldout"
def prepare(root, weights):
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(weights, local_files_only=True)
names = ("tokenizer.json", "tokenizer_config.json", "chat_template.jinja")
tokenizer_hashes = {}
for name in names:
sidecar = weights / ".cache/huggingface/download" / (name + ".metadata")
if sidecar.read_text().splitlines()[0] != REVISION:
raise ValueError(f"Tokenizer revision mismatch: {name}")
tokenizer_hashes[name] = digest((weights / name).read_bytes())
template = (weights / "chat_template.jinja").read_text()
if tokenizer.chat_template != template:
raise ValueError("Loaded tokenizer does not use the exact pinned chat template")
folder = root / "instruction-sources"
manifest_path = folder / "snapshots.json"
manifest = json.loads(manifest_path.read_text())
if manifest["sources"] != SOURCES:
raise ValueError("Unrecognized snapshot sources")
candidates = collections.defaultdict(list)
rejected = collections.Counter()
for snapshot in manifest["files"]:
raw = (folder / snapshot["file"]).read_bytes()
if digest(raw) != snapshot["sha256"]:
raise ValueError(f"Snapshot changed: {snapshot['file']}")
if snapshot["kind"] != "rows":
continue
domain, source_split = snapshot["domain"], snapshot["source_split"]
for wrapped in json.loads(raw)["rows"]:
if wrapped.get("truncated_cells"):
rejected[f"{domain}:truncated_source_cell"] += 1
continue
row, index = wrapped["row"], wrapped["row_idx"]
split = assigned_split(domain, source_split, index)
try:
messages, tools = convert(domain, row)
ids = tokenizer.apply_chat_template(messages, tools=tools, tokenize=True, return_dict=False, add_generation_prompt=False, enable_thinking=False)
if len(ids) < 32 or len(ids) > 4096:
rejected[f"{domain}:outside_32_4096"] += 1
continue
prompt_hashes = sorted({normalize_prompt(m["content"]) for m in messages if m["role"] == "user"})
record = {"id": f"instruct-{domain}-{source_split}-{index}", "split": split, "domain": domain, "token_ids": ids, "messages": messages, "prompt_sha256": prompt_hashes, "source": {"dataset": SOURCES[domain]["dataset"], "revision": SOURCES[domain]["revision"], "split": source_split, "row_idx": index, "snapshot": snapshot["file"]}, "language": row.get("language_code", "en"), "code_language": row.get("lang") if domain == "code" else None}
if tools:
record["tools"] = tools
candidates[(split, domain)].append(record)
except (ValueError, TypeError, KeyError) as error:
rejected[f"{domain}:conversion:{str(error)[:100]}"] += 1
# Hash ordering spreads selected records over every downloaded window; eval
# reserves its normalized prompts before calibration is filled.
for (split, domain), values in candidates.items():
values.sort(key=lambda row: digest((row["id"] + ":order-v1").encode()))
if split != "calibration" and domain == "multilingual":
occurrences = collections.Counter()
ranks = {}
for row in values:
ranks[row["id"]] = occurrences[row["language"]]
occurrences[row["language"]] += 1
# Prefer different languages before adding another example in the
# same language; preserve hash order within each occurrence round.
values.sort(key=lambda row: ranks[row["id"]])
seen = set()
outputs = {"instruction-calibration.jsonl": [], "instruction-eval.jsonl": [], "instruction-long-context.jsonl": []}
for split in ("heldout", "validation", "calibration"):
for domain in SOURCES:
total = 0
target = TARGETS[domain] if split == "calibration" else 1024
destination = outputs["instruction-calibration.jsonl" if split == "calibration" else "instruction-eval.jsonl"]
for record in candidates[(split, domain)]:
size = len(record["token_ids"])
if size > (2048 if split == "calibration" else 1024):
continue
if split != "calibration" and total + size > target:
continue
if seen.intersection(record["prompt_sha256"]):
rejected[f"{domain}:normalized_prompt_duplicate"] += 1
continue
destination.append(record)
seen.update(record["prompt_sha256"])
total += size
if total >= target or (split != "calibration" and total >= 900):
break
if (split == "calibration" and total < target) or (split != "calibration" and total < 512):
raise ValueError(f"Insufficient {split}/{domain} tokens: {total}/{target}; rejections={dict(rejected)}")
for split in ("heldout", "validation"):
for domain in SOURCES:
for record in candidates[(split, domain)]:
if 2048 < len(record["token_ids"]) <= 4096 and not seen.intersection(record["prompt_sha256"]):
outputs["instruction-long-context.jsonl"].append(record)
seen.update(record["prompt_sha256"])
break
summaries = {}
for name, rows in outputs.items():
rows.sort(key=lambda row: digest((row["id"] + ":stream-v1").encode()))
path = root / name
path.write_text("".join(json.dumps(row, ensure_ascii=False, separators=(",", ":")) + "\n" for row in rows))
counts = collections.defaultdict(lambda: {"records": 0, "tokens": 0, "prediction_tokens": 0, "languages": {}})
for row in rows:
stats = counts[f"{row['split']}/{row['domain']}"]
stats["records"] += 1
stats["tokens"] += len(row["token_ids"])
stats["prediction_tokens"] += len(row["token_ids"]) - 1
stats["languages"][row["language"]] = stats["languages"].get(row["language"], 0) + 1
summaries[name] = {"sha256": digest(path.read_bytes()), "records": len(rows), "tokens": sum(len(row["token_ids"]) for row in rows), "min_record_tokens": min((len(row["token_ids"]) for row in rows), default=0), "max_record_tokens": max((len(row["token_ids"]) for row in rows), default=0), "counts": dict(counts)}
provenance = {"schema_version": 1, "purpose": "Diverse exact-template instruct PTQ calibration for TT sensitivity/precision ranking; not official Unsloth data or a llama.cpp imatrix; no training/QAT/QAD", "sources": SOURCES, "snapshot_manifest_sha256": digest(manifest_path.read_bytes()), "script_sha256": digest(Path(__file__).read_bytes()), "tokenizer": {"model": "Qwen/Qwen3.5-9B", "revision": REVISION, "files": tokenizer_hashes, "chat_template_applied": True, "add_generation_prompt": False, "enable_thinking": False, "tools": "Real Hermes function schemas, assistant.tool_calls argument objects and linked tool responses; source system tool boilerplate replaced by exact Qwen template tools rendering"}, "selection": {"source_rows": manifest["selection"], "candidate_order": "ascending SHA256(id + ':order-v1')", "split_assignment": "Official train/test for UltraChat and Aya; test hashed 50/50 validation/heldout. Code/Hermes train rows hashed 80/10/10. SHA256(domain:source_split:row_idx:split-v1) first8 hex mod10.", "dedup": "Global exact normalized user-prompt SHA256 (Unicode NFKC, casefold, whitespace collapse), across all domains and all splits including long contexts; reject entire conversation if any user prompt already used. Semantic/paraphrase dedup not claimed.", "context_policy": "No truncation or packing. Entire conversations only; 32..2048 calibration, 32..1024 small KL, 2049..4096 separate long contexts. Oversized or malformed/truncated source rows rejected.", "eval_independence": "No calibration prompt overlap; official heldout source splits where available, deterministic disjoint row partitions otherwise. Same dataset families, NOT independent-source generalization proof.", "calibration_target_tokens": TARGETS, "kl_total_token_cap": 8192}, "rejected_candidates": dict(rejected), "outputs": summaries}
provenance["packages"] = {name: importlib.metadata.version(name) for name in ("transformers", "tokenizers", "jinja2")}
provenance["selection"]["multilingual_eval_order"] = "Stable language-round-robin over hash order: first occurrence of every language before second occurrences, etc."
if summaries["instruction-calibration.jsonl"]["tokens"] < 300000 or summaries["instruction-eval.jsonl"]["tokens"] > 8192:
raise ValueError("Token budget contract violated")
save_json(root / "instruction-provenance.json", provenance)
print(json.dumps({name: {key: value for key, value in summary.items() if key != "counts"} for name, summary in summaries.items()}, indent=2))
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--root", type=Path, required=True)
parser.add_argument("--weights", type=Path)
modes = parser.add_mutually_exclusive_group(required=True)
modes.add_argument("--fetch", action="store_true")
modes.add_argument("--prepare", action="store_true")
args = parser.parse_args()
if args.fetch:
fetch(args.root)
else:
if args.weights is None:
parser.error("--prepare requires --weights")
prepare(args.root, args.weights)
if __name__ == "__main__":
main()