catplusplus's picture
Upload folder using huggingface_hub
2143e89 verified
Raw History Blame Contribute Delete
16 kB
#!/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.<N>. 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()