#!/usr/bin/env python3 """ Simple FP32 to FP16 converter for safetensors files. Converts ALL FP32 tensors to FP16 without discrimination. USAGE: python convert_fp16.py input.safetensors [output.safetensors] """ import argparse from pathlib import Path from typing import Dict import torch from safetensors.torch import load_file, save_file def convert_fp32_to_fp16(state_dict: Dict[str, torch.Tensor]) -> Dict[str, torch.Tensor]: """Convert all FP32 tensors to FP16.""" converted = {} fp32_count = 0 total_tensors = len(state_dict) for name, tensor in state_dict.items(): if tensor.dtype == torch.float32: converted[name] = tensor.half() # Convert to FP16 fp32_count += 1 else: converted[name] = tensor # Keep as-is print(f"Converted {fp32_count}/{total_tensors} tensors from FP32 to FP16") return converted def calculate_size_reduction(original: Dict[str, torch.Tensor], converted: Dict[str, torch.Tensor]) -> None: """Calculate and print size reduction.""" original_bytes = sum(t.numel() * t.element_size() for t in original.values()) converted_bytes = sum(t.numel() * t.element_size() for t in converted.values()) saved_bytes = original_bytes - converted_bytes original_gb = original_bytes / (1024**3) converted_gb = converted_bytes / (1024**3) saved_gb = saved_bytes / (1024**3) print(f"Size: {original_gb:.3f} GB → {converted_gb:.3f} GB (saved {saved_gb:.3f} GB)") def main(): parser = argparse.ArgumentParser(description="Convert all FP32 tensors to FP16 in a safetensors file") parser.add_argument("input", help="Input safetensors file") parser.add_argument("output", nargs="?", help="Output safetensors file (optional)") args = parser.parse_args() # Generate output filename if not provided if args.output is None: input_path = Path(args.input) stem = input_path.stem if not stem.endswith("_fp16"): stem += "_fp16" args.output = str(input_path.with_name(stem + input_path.suffix)) print(f"Loading: {args.input}") state_dict = load_file(args.input) print("Converting FP32 → FP16...") converted_state_dict = convert_fp32_to_fp16(state_dict) calculate_size_reduction(state_dict, converted_state_dict) print(f"Saving: {args.output}") save_file(converted_state_dict, args.output) print("Done!") if __name__ == "__main__": main()