#!/usr/bin/env python3 """ Script to delete experts from a MoE model layer by layer. Processes one safetensor file at a time. Can either: 1. Use a pre-computed deleted_experts.json file 2. Analyze original vs pruned model to find deleted experts first """ import argparse import gc import json import re import shutil from pathlib import Path from typing import Optional import numpy as np import torch from safetensors import safe_open from safetensors.torch import save_file from scipy.optimize import linear_sum_assignment def load_deleted_experts(deleted_file: Path) -> dict[int, list[int]]: """Load the deleted experts mapping from JSON file.""" with open(deleted_file, "r") as f: data = json.load(f) return {int(k): v for k, v in data["deleted_experts_per_layer"].items()} def find_deleted_experts( original_model_path: Path, pruned_model_path: Path, num_original_experts: int, num_pruned_experts: int, hidden_dim: int = 3072, ) -> dict[int, list[int]]: """ Find deleted experts by comparing router matrices between original and pruned models. Uses Hungarian algorithm for optimal matching based on cosine distance. Returns a dict mapping layer_num -> list of deleted expert indices. """ deleted_experts = {} num_layers = 62 # MiniMax-M2.x models have 62 layers for layer_num in range(num_layers): router_key = f"model.layers.{layer_num}.block_sparse_moe.gate.weight" # Load router from original model orig_router = None for sf_file in sorted(original_model_path.glob("model-*.safetensors")): with safe_open(sf_file, framework="pt") as f: if router_key in f.keys(): orig_router = f.get_tensor(router_key) break # Load router from pruned model prune_router = None for sf_file in sorted(pruned_model_path.glob("model-*.safetensors")): with safe_open(sf_file, framework="pt") as f: if router_key in f.keys(): prune_router = f.get_tensor(router_key) break if orig_router is None or prune_router is None: print(f" Layer {layer_num}: Router not found, skipping") continue # Convert bfloat16 if needed if orig_router.dtype == torch.bfloat16: orig_router = orig_router.to(torch.float32) if prune_router.dtype == torch.bfloat16: prune_router = prune_router.to(torch.float32) # L2 normalize rows for cosine distance orig_np = orig_router.numpy() prune_np = prune_router.numpy() orig_norm = orig_np / (np.linalg.norm(orig_np, axis=1, keepdims=True) + 1e-10) prune_norm = prune_np / (np.linalg.norm(prune_np, axis=1, keepdims=True) + 1e-10) # Cosine distance matrix distance_matrix = 1 - np.dot(orig_norm, prune_norm.T) # Hungarian algorithm row_ind, col_ind = linear_sum_assignment(distance_matrix) # Find unmatched (deleted) experts matched_original = set(row_ind) all_original = set(range(num_original_experts)) deleted = sorted(all_original - matched_original) deleted_experts[layer_num] = deleted if layer_num in [0, 1, 2, 10, 20, 30, 40, 50, 61]: print(f" Layer {layer_num}: {len(deleted)} deleted experts") return deleted_experts def get_retained_experts(num_original: int, deleted: list[int]) -> list[int]: """Get the list of retained expert indices (sorted).""" all_experts = set(range(num_original)) deleted_set = set(deleted) return sorted(all_experts - deleted_set) def get_layer_for_tensor(tensor_name: str) -> Optional[int]: """Get layer number for a tensor name.""" parts = tensor_name.split(".") for i, part in enumerate(parts): if part == "layers" and i + 1 < len(parts): try: return int(parts[i + 1]) except ValueError: pass return None def is_expert_tensor(key: str) -> bool: """ Check if a tensor key belongs to a specific numbered expert. Uses regex to match .experts.. pattern robustly, covering NVFP4 sub-tensors like w1.weight, w1.weight_scale, w1.weight_scale_2, w1.input_scale, etc. """ return bool(re.search(r'\.experts\.\d+\.', key)) # Keys to strip: asymmetric KV-cache quantization zero-points that vLLM # doesn't support (only scale factors are used). _STRIP_SUFFIXES = (".k_bias", ".v_bias") def should_strip_tensor(key: str) -> bool: """Return True for tensors that should be omitted from the output model.""" return any(key.endswith(sfx) for sfx in _STRIP_SUFFIXES) def process_tensor( key: str, tensor: torch.Tensor, deleted_experts: dict[int, list[int]], num_original_experts: int, ) -> tuple[str, torch.Tensor] | None: """ Process a single tensor. Returns (new_key, new_tensor) or None to skip. """ # Strip unsupported KV quantization zero-points. if should_strip_tensor(key): return None layer_num = get_layer_for_tensor(key) # Non-layer or non-MoE tensors pass through unchanged. if layer_num is None or "block_sparse_moe" not in key: return (key, tensor) if is_expert_tensor(key): # Expert weight/scale tensor – delete or renumber. parts = key.split(".") for i, part in enumerate(parts): if part == "experts" and i + 1 < len(parts): expert_idx = int(parts[i + 1]) deleted = deleted_experts.get(layer_num, []) retained = get_retained_experts(num_original_experts, deleted) if expert_idx not in retained: return None # Deleted expert – drop tensor. # Renumber to sequential 0-based index. new_expert_idx = retained.index(expert_idx) new_key_parts = parts.copy() new_key_parts[i + 1] = str(new_expert_idx) new_key = ".".join(new_key_parts) return (new_key, tensor) # Shouldn't happen, but fall through unchanged. return (key, tensor) elif "gate.weight" in key: deleted = deleted_experts.get(layer_num, []) retained = get_retained_experts(num_original_experts, deleted) indices = torch.tensor(retained, dtype=torch.long) new_tensor = tensor[indices].clone() print(f" gate.weight layer {layer_num}: {tuple(tensor.shape)} → {tuple(new_tensor.shape)}") return (key, new_tensor) elif "e_score_correction_bias" in key: deleted = deleted_experts.get(layer_num, []) retained = get_retained_experts(num_original_experts, deleted) indices = torch.tensor(retained, dtype=torch.long) return (key, tensor[indices].clone()) else: return (key, tensor) def process_file( input_path: Path, output_path: Path, deleted_experts: dict[int, list[int]], num_original_experts: int, ) -> tuple[int, int, int]: """ Process one safetensor file. Returns (kept, deleted, stripped) tensor counts. """ tensors = {} kept = deleted_count = stripped = 0 with safe_open(input_path, framework="pt") as f: all_keys = list(f.keys()) # Report stripped keys upfront for this file. stripped_keys = [k for k in all_keys if should_strip_tensor(k)] if stripped_keys: print(f" Stripping {len(stripped_keys)} KV-bias tensors: {stripped_keys[:4]}{'…' if len(stripped_keys) > 4 else ''}") with safe_open(input_path, framework="pt") as f: for key in all_keys: tensor = f.get_tensor(key) result = process_tensor(key, tensor, deleted_experts, num_original_experts) if result is not None: new_key, new_tensor = result tensors[new_key] = new_tensor kept += 1 else: if should_strip_tensor(key): stripped += 1 else: deleted_count += 1 save_file(tensors, output_path) del tensors gc.collect() return kept, deleted_count, stripped def get_file_to_tensors(model_path: Path) -> dict[str, list[str]]: """Get mapping from filename to list of tensor names.""" index_path = model_path / "model.safetensors.index.json" with open(index_path, "r") as f: index_data = json.load(f) weight_map = index_data["weight_map"] file_to_tensors: dict[str, list[str]] = {} for key, file_name in weight_map.items(): if file_name not in file_to_tensors: file_to_tensors[file_name] = [] file_to_tensors[file_name].append(key) return file_to_tensors def main(): parser = argparse.ArgumentParser( description="Delete experts from MoE model. " "Either provide --deleted-experts-file or use --compare-with to find deleted experts." ) parser.add_argument("input_model", type=Path, help="Input model path (256 experts)") parser.add_argument("output_model", type=Path, help="Output model path") parser.add_argument("--num-original-experts", type=int, default=256) parser.add_argument("--num-retained-experts", type=int, default=192) parser.add_argument("--deleted-experts-file", type=Path, help="JSON with pre-computed deleted experts (optional)") parser.add_argument("--compare-with", type=Path, help="Pruned model path (192 experts) to compare and find deleted experts") parser.add_argument("--save-deleted-experts", type=Path, help="Save found deleted experts to JSON file (for --compare-with mode)") args = parser.parse_args() # Either load from file or find via comparison if args.deleted_experts_file: print("Loading deleted experts from file...") deleted_experts = load_deleted_experts(args.deleted_experts_file) print(f" {len(deleted_experts)} layers") # Sanity-check: verify expected deletions per layer. counts = {l: len(v) for l, v in deleted_experts.items()} expected = args.num_original_experts - args.num_retained_experts bad = {l: c for l, c in counts.items() if c != expected} if bad: print(f" WARNING: {len(bad)} layers have unexpected deletion count " f"(expected {expected}): {dict(list(bad.items())[:5])}") elif args.compare_with: print("Finding deleted experts by comparing models...") print(f" Original: {args.input_model}") print(f" Pruned: {args.compare_with}") deleted_experts = find_deleted_experts( args.input_model, args.compare_with, args.num_original_experts, args.num_retained_experts, ) print(f" Found deleted experts for {len(deleted_experts)} layers") if args.save_deleted_experts: print(f"\nSaving deleted experts to {args.save_deleted_experts}...") data = {"deleted_experts_per_layer": {str(k): v for k, v in deleted_experts.items()}} with open(args.save_deleted_experts, "w") as f: json.dump(data, f, indent=2) else: parser.error("Either --deleted-experts-file or --compare-with is required") # Create output dir args.output_model.mkdir(parents=True, exist_ok=True) # Copy non-safetensor files print("\nCopying config/tokenizer files...") for fname in ["config.json", "configuration_minimax_m2.py", "tokenizer.json", "tokenizer_config.json", "vocab.json", "merges.txt", "special_tokens_map.json", "added_tokens.json", "generation_config.json", "chat_template.jinja", ".gitattributes", "hf_quant_config.json", "modeling_minimax_m2.py"]: src = args.input_model / fname if src.exists(): shutil.copy2(src, args.output_model / fname) print(f" {fname}") # Update config: set num_local_experts to the retained count. print("\nUpdating config...") with open(args.input_model / "config.json") as f: config = json.load(f) config["num_local_experts"] = args.num_retained_experts with open(args.output_model / "config.json", "w") as f: json.dump(config, f, indent=2) print(f" num_local_experts = {args.num_retained_experts}") # Process safetensors print("\nProcessing safetensor files...") file_to_tensors = get_file_to_tensors(args.input_model) total_kept = total_deleted = total_stripped = 0 for file_name in sorted(file_to_tensors.keys()): input_path = args.input_model / file_name output_path = args.output_model / file_name print(f" {file_name}...") kept, deleted_count, stripped = process_file( input_path, output_path, deleted_experts, args.num_original_experts ) print(f" kept={kept} deleted={deleted_count} stripped={stripped}") total_kept += kept total_deleted += deleted_count total_stripped += stripped gc.collect() print(f"\nTotal: kept={total_kept} deleted={total_deleted} stripped={total_stripped}") # Update index. # Bug fixes vs original: # 1. Use `layer_num is not None` instead of truthy `layer_num` (layer 0 == 0 is falsy). # 2. Use a `handled` flag so deleted experts are not re-added via the for-else clause. # 3. Use is_expert_tensor() regex instead of ".w" heuristic. print("\nUpdating model index...") with open(args.input_model / "model.safetensors.index.json") as f: index_data = json.load(f) weight_map = index_data["weight_map"] new_weight_map: dict[str, str] = {} idx_skipped = idx_stripped = idx_renamed = 0 for key, file_name in weight_map.items(): # Drop stripped tensors (KV bias). if should_strip_tensor(key): idx_stripped += 1 continue layer_num = get_layer_for_tensor(key) # FIX 1: use `is not None` so layer 0 (== 0, falsy) is handled correctly. if layer_num is not None and "block_sparse_moe" in key and is_expert_tensor(key): parts = key.split(".") handled = False for i, part in enumerate(parts): if part == "experts" and i + 1 < len(parts): expert_idx = int(parts[i + 1]) deleted = deleted_experts.get(layer_num, []) retained = get_retained_experts(args.num_original_experts, deleted) if expert_idx not in retained: # FIX 2: mark as handled and break so we don't fall # through to the for-else and re-add the old key. idx_skipped += 1 handled = True break new_expert_idx = retained.index(expert_idx) new_key_parts = parts.copy() new_key_parts[i + 1] = str(new_expert_idx) new_key = ".".join(new_key_parts) new_weight_map[new_key] = file_name idx_renamed += 1 handled = True break if not handled: # No expert index found in key (shouldn't happen), keep as-is. new_weight_map[key] = file_name else: new_weight_map[key] = file_name index_data["weight_map"] = new_weight_map with open(args.output_model / "model.safetensors.index.json", "w") as f: json.dump(index_data, f, indent=2) print(f" {len(new_weight_map)} index entries " f"(renamed={idx_renamed} skipped={idx_skipped} stripped={idx_stripped})") print("\n✓ Done!") if __name__ == "__main__": main()