#!/usr/bin/env python3 """Convert a local Nemotron-H NVFP4 Hugging Face checkpoint to MLX format (Apple Silicon). ModelOpt NVFP4 safetensors cannot be loaded directly by mlx_lm; this script dequantizes packed weights to bfloat16, loads them into mlx-lm's nemotron_h implementation, then re-quantizes to MLX NVFP4 for Metal inference. Recommended on Mac: slice to 12B or 23B first (see zero_shot_slicing.py) — the 30B variant needs ~60 GB RAM at conversion time and ~20 GB at inference. Setup (torch313-metal venv): source ~/python3-venv/torch313-metal/bin/activate uv pip install mlx-lm safetensors Example: # Optional: smaller variant for Mac memory budgets python zero_shot_slicing.py \\ --source-checkpoint . \\ --target-checkpoint ./nemotron-elastic-12b-nvfp4 \\ --size 12B --precision nvfp4 python convert_to_mlx.py \\ --hf-path ./nemotron-elastic-12b-nvfp4 \\ --mlx-path ./nemotron-12b-mlx python chat_mlx.py --model ./nemotron-12b-mlx """ from __future__ import annotations import argparse import gc import json import shutil import sys from pathlib import Path import mlx.core as mx import torch from safetensors.torch import load_file from mlx_lm.utils import _get_classes, quantize_model, save from mlx_lm.utils import load_tokenizer DEFAULT_HF_PATH = Path(__file__).resolve().parent SCALE_SUFFIXES = (".weight_scale", ".weight_scale_2", ".input_scale") E2M1_VALUES = torch.tensor( [0, 0.5, 1, 1.5, 2, 3, 4, 6, 0, -0.5, -1, -1.5, -2, -3, -4, -6], dtype=torch.float32, ) def load_exclude_modules(hf_path: Path) -> set[str]: quant_path = hf_path / "hf_quant_config.json" if not quant_path.is_file(): return set() with open(quant_path, encoding="utf-8") as f: data = json.load(f) return set(data.get("quantization", {}).get("exclude_modules", [])) def module_path(param_name: str) -> str: for suffix in (".weight", ".bias", *SCALE_SUFFIXES): if param_name.endswith(suffix): return param_name[: -len(suffix)] return param_name def is_excluded(param_name: str, exclude_modules: set[str]) -> bool: return module_path(param_name) in exclude_modules def dequant_nvfp4( weight: torch.Tensor, scale: torch.Tensor, scale2: torch.Tensor, *, group_size: int = 16, ) -> torch.Tensor: if weight.dtype != torch.uint8: raise TypeError(f"Expected uint8 NVFP4 weights, got {weight.dtype}") out_features, packed_in = weight.shape in_features = packed_in * 2 unpacked = torch.empty(out_features, in_features, dtype=torch.float32) unpacked[:, 1::2] = (weight >> 4).to(torch.long) unpacked[:, 0::2] = (weight & 0x0F).to(torch.long) unpacked = E2M1_VALUES[unpacked.long()] per_block = scale.to(torch.float32) * scale2.to(torch.float32) deq = unpacked.view(out_features, -1, group_size) * per_block.unsqueeze(-1) return deq.reshape(out_features, in_features).to(torch.bfloat16) def prepare_hf_weights( hf_path: Path, *, exclude_modules: set[str], group_size: int, ) -> dict[str, torch.Tensor]: shard_paths = sorted(hf_path.glob("model*.safetensors")) if not shard_paths: raise FileNotFoundError(f"No model*.safetensors files found in {hf_path}") # First pass: load ALL tensors from all shards (weights + scales) print("[convert] loading all shards into memory...") all_tensors: dict[str, torch.Tensor] = {} for shard_path in shard_paths: print(f"[convert] loading {shard_path.name} ...") raw = load_file(str(shard_path), device="cpu") all_tensors.update(raw) del raw gc.collect() print(f"[convert] loaded {len(all_tensors)} total tensors") # Second pass: dequantize weights that have scales prepared: dict[str, torch.Tensor] = {} for key, tensor in all_tensors.items(): if any(key.endswith(suffix) for suffix in SCALE_SUFFIXES): continue if key.endswith(".weight"): base = key[: -len(".weight")] scale_key = f"{base}.weight_scale" scale2_key = f"{base}.weight_scale_2" if scale_key in all_tensors and scale2_key in all_tensors and not is_excluded(key, exclude_modules): prepared[key] = dequant_nvfp4( tensor, all_tensors[scale_key], all_tensors[scale2_key], group_size=group_size, ) continue prepared[key] = tensor.to(torch.bfloat16) if tensor.dtype.is_floating_point else tensor del all_tensors gc.collect() return prepared def estimate_bf16_gb(num_params: int) -> float: return num_params * 2 / (1024**3) def convert_checkpoint( hf_path: Path, mlx_path: Path, *, quantize: bool, q_group_size: int, q_bits: int, q_mode: str, dry_run: bool, ) -> None: hf_path = hf_path.resolve() if not (hf_path / "config.json").is_file(): raise FileNotFoundError(f"Missing config.json in {hf_path}") if mlx_path.exists(): raise FileExistsError(f"Refusing to overwrite existing path: {mlx_path}") with open(hf_path / "config.json", encoding="utf-8") as f: config = json.load(f) exclude_modules = load_exclude_modules(hf_path) group_size = 16 quant_path = hf_path / "hf_quant_config.json" if quant_path.is_file(): with open(quant_path, encoding="utf-8") as f: group_size = json.load(f)["quantization"].get("group_size", group_size) hidden = config.get("hidden_size", 0) layers = config.get("num_hidden_layers", 0) print(f"[convert] model_type={config.get('model_type')} hidden={hidden} layers={layers}") print(f"[convert] excluded modules: {len(exclude_modules)}") prepared = prepare_hf_weights(hf_path, exclude_modules=exclude_modules, group_size=group_size) param_count = sum(t.numel() for t in prepared.values() if t.dtype.is_floating_point) print(f"[convert] prepared {len(prepared)} tensors (~{estimate_bf16_gb(param_count):.1f} GB bf16 peak)") if dry_run: print("[convert] dry run complete — weights prepared successfully") return model_class, model_args_class = _get_classes(config) model_args = model_args_class.from_dict(config) model = model_class(model_args) mx_weights = {k: mx.array(v) for k, v in prepared.items()} del prepared gc.collect() print("[convert] sanitizing weights...") try: mx_weights = model.sanitize(mx_weights) except Exception as e: print(f"[convert] sanitize failed: {e}") print("[convert] checking ALL MoE layers for shape mismatches...") # Check all layers for expert weight shape mismatches for layer_num in range(52): layer_experts = {} for k in sorted(mx_weights.keys()): if f'backbone.layers.{layer_num}.mixer.experts' in k and 'weight' in k and 'scale' not in k: parts = k.split('.') exp_idx = parts.index('experts') expert_num = int(parts[exp_idx + 1]) proj_type = parts[exp_idx + 2] if expert_num not in layer_experts: layer_experts[expert_num] = {} layer_experts[expert_num][proj_type] = mx_weights[k].shape if layer_experts: up_shapes = set() down_shapes = set() for exp_data in layer_experts.values(): if 'up_proj' in exp_data: up_shapes.add(exp_data['up_proj']) if 'down_proj' in exp_data: down_shapes.add(exp_data['down_proj']) if len(up_shapes) > 1 or len(down_shapes) > 1: print(f" Layer {layer_num}: SHAPE MISMATCH!") print(f" up_proj shapes: {up_shapes}") print(f" down_proj shapes: {down_shapes}") # Show which experts have which shapes for exp_num, shapes in sorted(layer_experts.items()): if 'up_proj' in shapes and shapes['up_proj'] not in up_shapes: print(f" Expert {exp_num}: up_proj={shapes.get('up_proj')}, down_proj={shapes.get('down_proj')}") raise print("[convert] loading weights into model...") try: model.load_weights(list(mx_weights.items()), strict=True) except Exception as e: print(f"[convert] load_weights failed: {e}") print("[convert] checking problematic weights...") for name, weight in mx_weights.items(): if hasattr(weight, 'shape'): print(f" {name}: {weight.shape}") raise del mx_weights gc.collect() if quantize: print(f"[convert] quantizing to MLX {q_mode} (group_size={q_group_size}, bits={q_bits})") model, config = quantize_model( model, config, q_group_size, q_bits, mode=q_mode, ) else: print("[convert] saving bfloat16 MLX checkpoint (no re-quantization)") tokenizer = load_tokenizer(hf_path, {"trust_remote_code": True}) mlx_path.mkdir(parents=True, exist_ok=False) save(mlx_path, hf_path, model, tokenizer, config) # Preserve Nemotron chat template if save() missed it for name in ("chat_template.jinja", "nano_v3_reasoning_parser.py"): src = hf_path / name if src.is_file(): shutil.copy2(src, mlx_path / name) print(f"[convert] wrote MLX model to {mlx_path}") def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Convert Nemotron-H ModelOpt NVFP4 checkpoint to MLX format", ) parser.add_argument( "--hf-path", type=Path, default=DEFAULT_HF_PATH, help="Path to the HF NVFP4 checkpoint (default: this directory)", ) parser.add_argument( "--mlx-path", type=Path, required=True, help="Output directory for the MLX model (must not exist)", ) parser.add_argument( "--no-quantize", action="store_true", help="Save bfloat16 MLX weights instead of re-quantizing to NVFP4", ) parser.add_argument( "--q-group-size", type=int, default=16, help="MLX quantization group size (default: 16)", ) parser.add_argument( "--q-bits", type=int, default=4, help="MLX quantization bits (default: 4)", ) parser.add_argument( "--q-mode", choices=["affine", "mxfp4", "nvfp4", "mxfp8"], default="nvfp4", help="MLX quantization mode (default: nvfp4)", ) parser.add_argument( "--dry-run", action="store_true", help="Prepare and validate weights without writing an MLX checkpoint", ) return parser.parse_args() def main() -> int: args = parse_args() try: convert_checkpoint( args.hf_path, args.mlx_path, quantize=not args.no_quantize, q_group_size=args.q_group_size, q_bits=args.q_bits, q_mode=args.q_mode, dry_run=args.dry_run, ) except Exception as exc: print(f"[convert] failed: {exc}", file=sys.stderr) return 1 return 0 if __name__ == "__main__": raise SystemExit(main())