| |
| """ |
| 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() |
| fp32_count += 1 |
| else: |
| converted[name] = tensor |
| |
| 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() |
| |
| |
| 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() |