#!/usr/bin/env python3 """Convert Hugging Face Whisper weights to compact inference safetensors. Usage: uv run --with numpy --with safetensors --with ml-dtypes python convert.py \ model.safetensors model-f16.safetensors """ import argparse import json from pathlib import Path import numpy as np from safetensors import safe_open from safetensors.numpy import save_file def remap_key(key: str) -> str: key = key.removeprefix("model.") key = { "encoder.embed_positions.weight": "encoder.positional_embedding", "decoder.embed_positions.weight": "decoder.positional_embedding", }.get(key, key) key = key.replace("decoder.embed_tokens", "decoder.token_embedding", 1) key = key.replace("encoder.layer_norm", "encoder.ln_post", 1) key = key.replace("encoder.layers.", "encoder.blocks.", 1) key = key.replace("decoder.layer_norm", "decoder.ln", 1) key = key.replace("decoder.layers.", "decoder.blocks.", 1) key = key.replace("self_attn_layer_norm", "attn_ln") key = key.replace("encoder_attn_layer_norm", "cross_attn_ln") key = key.replace("encoder_attn", "cross_attn") key = key.replace("self_attn", "attn") key = key.replace("q_proj", "query") key = key.replace("k_proj", "key") key = key.replace("v_proj", "value") key = key.replace("out_proj", "out") key = key.replace("fc1", "mlp.0") key = key.replace("fc2", "mlp.2") return key.replace("final_layer_norm", "mlp_ln") def keeps_float32(key: str) -> bool: return ( key in {"encoder.positional_embedding", "decoder.positional_embedding"} or key.startswith("encoder.ln_post.") or key.startswith("decoder.ln.") or any(part in key for part in (".attn_ln.", ".cross_attn_ln.", ".mlp_ln.")) ) def quantizes_fp8(key: str, tensor: np.ndarray, compute_dtype: str) -> bool: return ( compute_dtype == "float8_e4m3fn" and tensor.ndim == 2 and key != "decoder.token_embedding.weight" and not keeps_float32(key) ) def target_dtype(key: str, tensor: np.ndarray, compute_dtype: str): if keeps_float32(key): return np.float32 if quantizes_fp8(key, tensor, compute_dtype): import ml_dtypes return ml_dtypes.float8_e4m3fn return np.float16 def quantize_fp8(tensor: np.ndarray): import ml_dtypes value = tensor.astype(np.float32) axes = tuple(range(1, value.ndim)) scale = np.max(np.abs(value), axis=axes, keepdims=True) / 448.0 scale = np.where(scale == 0, 1.0, scale).astype(np.float16) quantized = (value / scale.astype(np.float32)).astype(ml_dtypes.float8_e4m3fn) return quantized, scale def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("input", type=Path, help="Source model.safetensors") parser.add_argument("output", type=Path, help="Converted model.safetensors") parser.add_argument("--source", help="Source repository and revision for metadata") parser.add_argument("--compute-dtype", choices=("float16", "float8_e4m3fn"), default="float16") args = parser.parse_args() tensors = {} with safe_open(args.input, framework="numpy") as checkpoint: source_metadata = checkpoint.metadata() or {} for source_key in checkpoint.keys(): key = remap_key(source_key) if key in tensors: raise ValueError(f"duplicate normalized key: {key}") tensor = checkpoint.get_tensor(source_key) if np.issubdtype(tensor.dtype, np.floating): dtype = target_dtype(key, tensor, args.compute_dtype) if quantizes_fp8(key, tensor, args.compute_dtype): tensor, scale = quantize_fp8(tensor) tensors[f"{key}.weight_scale"] = np.ascontiguousarray(scale) else: tensor = tensor.astype(dtype) tensors[key] = np.ascontiguousarray(tensor) transform = "HF keys normalized; compute weights FP16; positional embeddings and LayerNorm FP32" if args.compute_dtype == "float8_e4m3fn": transform = ( "HF keys normalized; linear weights float8_e4m3fn with per-output-channel scales; " "token/positional embeddings, convolutions, biases, and scales FP16 except LayerNorm/positional FP32" ) metadata = { **source_metadata, "format": "pt", "precision": f"mixed-{args.compute_dtype}-f16-f32", "transform": transform, } if args.source: metadata["source"] = args.source args.output.parent.mkdir(parents=True, exist_ok=True) save_file(tensors, args.output, metadata=metadata) counts = {str(dtype): sum(t.dtype == dtype for t in tensors.values()) for dtype in {t.dtype for t in tensors.values()}} size = args.output.stat().st_size / 2**30 print(json.dumps({"output": str(args.output), "size_gib": round(size, 3), "tensors": len(tensors), "dtypes": counts})) if __name__ == "__main__": main()