whisper / convert.py
vpermilp's picture
Add scaled FP8 conversion support
7ca597c verified
Raw
History Blame Contribute Delete
5.04 kB
#!/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()