NSFW_MMaudio / convert_fp32_to_fp16.py
phazei's picture
Initial commit
59e9071
Raw
History Blame Contribute Delete
2.52 kB
#!/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()