#!/usr/bin/env python3 """Build the pinned Mellum2.1 export from local source weights; never serve it.""" from __future__ import annotations import argparse import hashlib import importlib.metadata import json import random import shutil from contextlib import contextmanager from pathlib import Path from audit_export import SOURCE_REVISION, audit, tensor_files MODEL_ID = "JetBrains/Mellum2.1-12B-A2.5B-Thinking" CONVERSATION_DATASET = "HuggingFaceH4/ultrachat_200k" CONVERSATION_REVISION = "8049631c405ae6576f93f445c6b8166f76f5505a" TOOL_DATASET = "NousResearch/hermes-function-calling-v1" TOOL_REVISION = "dae3e1d28cfbcf4b915c04ea1e072030529b4bda" @contextmanager def calibration_attention_outputs(model): """Discard unused attention probabilities through temporary PyTorch hooks. MellumDecoderLayer consumes only attention tuple item zero, as does AWQ's loss. AWQ 0.14 first collects every tuple before extracting that item, which otherwise retains all eager attention matrices on the calibration GPU. The primary output and eager attention calculations stay unchanged. """ from transformers.models.mellum.modeling_mellum import MellumAttention import torch def discard_probabilities(_module, _inputs, output): if not isinstance(output, tuple) or len(output) != 2 or not isinstance(output[0], torch.Tensor): raise RuntimeError("Unexpected MellumAttention output contract") return output[0], None names, handles = [], [] try: for name, module in model.named_modules(): if isinstance(module, MellumAttention): names.append(name) handles.append(module.register_forward_hook(discard_probabilities)) yield names finally: for handle in handles: handle.remove() def normalize_messages(row: dict) -> list[dict]: if "messages" in row: return row["messages"] roles = {"human": "user", "gpt": "assistant", "system": "system", "tool": "tool", "function": "tool", "user": "user", "assistant": "assistant"} messages = [] for turn in row.get("conversations", []): role = roles.get(turn.get("from")) if role is None: raise ValueError(f"Unsupported public dataset role: {turn.get('from')!r}") messages.append({"role": role, "content": turn.get("value", "")}) return messages def build_calibration(args, tokenizer): from datasets import Dataset, load_dataset # Keep conversation content, tool results and schema examples together. # Hermes rows already include their tool/schema declarations in system text; # adding the separate tools field again would duplicate those declarations. tool_count = args.num_samples // 4 sources = [ (args.conversation_dataset, args.conversation_revision, "default", "train_sft", args.num_samples - 2 * tool_count), (args.tool_dataset, args.tool_revision, "func_calling", "train", tool_count), (args.tool_dataset, args.tool_revision, "json_mode_agentic", "train", tool_count), ] rows, manifest = [], [] for repo, revision, config, split, count in sources: stream = load_dataset(repo, name=config, split=split, revision=revision, streaming=True).shuffle(seed=args.seed, buffer_size=1024) selected = 0 for ordinal, row in enumerate(stream): if ordinal >= max(1000, count * 20): break messages = normalize_messages(row) if not messages: continue text = tokenizer.apply_chat_template( messages, tokenize=False, add_generation_prompt=False, enable_thinking=False) if not text.strip(): continue # Materialize exact bounded public text used by the compressor. rows.append({"text": text}) manifest.append({"dataset": repo, "revision": revision, "config": config, "split": split, "shuffled_ordinal": ordinal, "row_id": row.get("id", row.get("prompt_id")), "text_sha256": hashlib.sha256(text.encode()).hexdigest()}) selected += 1 if selected == count: break if selected != count: raise ValueError(f"Needed {count} usable rows from {repo}/{config}; got {selected}") random.Random(args.seed).shuffle(rows) return Dataset.from_list(rows), manifest, rows def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--source-dir", type=Path, required=True, help="Local HF snapshot downloaded at the pinned source revision") parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--num-samples", type=int, default=256) parser.add_argument("--max-seq-length", type=int, default=2048) parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--torch-threads", type=int, default=4) parser.add_argument("--recipe", type=Path, default=Path(__file__).with_name("recipe.yaml")) parser.add_argument("--conversation-dataset", default=CONVERSATION_DATASET) parser.add_argument("--conversation-revision", default=CONVERSATION_REVISION) parser.add_argument("--tool-dataset", default=TOOL_DATASET) parser.add_argument("--tool-revision", default=TOOL_REVISION) args = parser.parse_args() if args.num_samples < 4 or args.max_seq_length < 128 or args.torch_threads < 1: parser.error("Use at least 4 calibration samples and a sequence length of at least 128") for revision in (args.conversation_revision, args.tool_revision): if len(revision) != 40 or any(c not in "0123456789abcdef" for c in revision): parser.error("Dataset revisions must be immutable 40-character commit hashes") if args.output_dir.exists(): parser.error("Output directory already exists; choose a new directory") source = args.source_dir.resolve() if not source.is_dir(): parser.error("Source directory does not exist") index = json.loads((source / "model.safetensors.index.json").read_text()) if index.get("metadata", {}).get("total_parameters") != 12149923072: parser.error("Source index does not match the pinned Mellum2.1 parameter count") if any("mtp" in key.lower() for key in index["weight_map"]): parser.error("This workflow expects the published checkpoint without an MTP head") # The receipt binds a local directory to the exact source download. It is a # provenance statement, not a substitute for HF's download integrity checks. receipt = json.loads((source / "source-provenance.json").read_text()) if receipt != {"model_id": MODEL_ID, "revision": SOURCE_REVISION}: parser.error("Source download receipt does not match the pinned model and revision") tensor_files(source) # Reject truncated/missing safetensors shards before calibration. import torch from transformers import AutoModelForCausalLM, AutoTokenizer from transformers.core_model_loading import WeightRenaming from llmcompressor import oneshot from llmcompressor.modeling import patch_moe_mappings from llmcompressor.utils import load_context versions = {p: importlib.metadata.version(p) for p in ("llmcompressor", "compressed-tensors", "transformers", "torch", "datasets")} expected = {"llmcompressor": "0.14.0", "compressed-tensors": "0.19.0", "transformers": "5.17.0"} if any(versions[p] != v for p, v in expected.items()): raise RuntimeError(f"Unexpected quantizer versions: {versions}") torch.manual_seed(args.seed) torch.set_num_threads(args.torch_threads) torch.set_num_interop_threads(1) tokenizer = AutoTokenizer.from_pretrained(source, local_files_only=True) dataset, calibration_manifest, rows = build_calibration(args, tokenizer) print(f"Prepared {len(rows)} public calibration samples; loading local source", flush=True) # Published Mellum tensors are already per-expert 2D weights. Register that # checkpoint layout through the quantizer's supported interface to avoid # loading them as fused 3D weights and immediately splitting them again. mappings = [WeightRenaming( source_patterns=rf"\.experts\.(\d+)\.{projection}\.", target_patterns=rf".experts.\1.{projection}.") for projection in ("gate_proj", "up_proj", "down_proj")] with patch_moe_mappings("mellum", mappings, remove_targets=[ "mlp.experts.gate_up_proj", "mlp.experts.down_proj"]), load_context(): model = AutoModelForCausalLM.from_pretrained( source, dtype=torch.bfloat16, local_files_only=True) expert_names = [name for name, module in model.named_modules() if ".mlp.experts." in name and isinstance(module, torch.nn.Linear)] if len(expert_names) != 5376: raise RuntimeError(f"Expected 5376 linearized expert projections; got {len(expert_names)}") with calibration_attention_outputs(model) as attention_hook_names: if len(attention_hook_names) != 28: raise RuntimeError(f"Expected 28 Mellum attention hooks; got {len(attention_hook_names)}") oneshot(model=model, processor=tokenizer, dataset=dataset, recipe=str(args.recipe), pipeline="sequential", sequential_targets=["MellumDecoderLayer"], sequential_offload_device="cpu", sequential_prefetch=False, batch_size=1, dataloader_num_workers=0, preprocessing_num_workers=1, max_seq_length=args.max_seq_length, num_calibration_samples=args.num_samples, output_dir=str(args.output_dir), save_compressed=True) tokenizer.save_pretrained(args.output_dir) for filename in ("chat_template.jinja", "generation_config.json"): if (source / filename).exists(): shutil.copy2(source / filename, args.output_dir / filename) shutil.copy2(args.recipe, args.output_dir / "quantization-recipe.yaml") evidence = args.output_dir / "quantization-evidence" evidence.mkdir() (evidence / "calibration.jsonl").write_text( "".join(json.dumps(row, ensure_ascii=False) + "\n" for row in rows)) provenance = {"model_id": MODEL_ID, "source_revision": SOURCE_REVISION, "versions": versions, "torch_cuda": torch.version.cuda, "num_samples": args.num_samples, "max_seq_length": args.max_seq_length, "seed": args.seed, "torch_threads": args.torch_threads, "calibration_attention_hook": { "interface": "torch.nn.Module.register_forward_hook", "purpose": "Discard unused eager attention tuple item 1 before AWQ retains batches", "primary_output": "unchanged", "modules": attention_hook_names, "removed_after_calibration": True}, "calibration": calibration_manifest, "recipe_sha256": hashlib.sha256(args.recipe.read_bytes()).hexdigest(), "status": "unverified until export audit and semantic serving checks pass"} (evidence / "provenance.json").write_text(json.dumps(provenance, indent=2) + "\n") freeze = Path("/opt/mellum/quantizer-freeze.txt") if freeze.exists(): shutil.copy2(freeze, evidence / freeze.name) report = audit(args.output_dir, source) (evidence / "export-audit.json").write_text(json.dumps(report, indent=2) + "\n") print(json.dumps(report, indent=2)) if __name__ == "__main__": main()