ComfyUI_extracted_lora / convert_fp8_simple.py
morisoba's picture
add fp8 version
bde8c2c
Raw
History Blame Contribute Delete
2.28 kB
import torch
from safetensors import safe_open
from safetensors.torch import save_file
from tqdm import tqdm
import argparse
import os
def convert_to_fp8_simple(input_path, output_path=None):
"""
LoRAファイルを通常のFP8(非スケール版)に変換
"""
if output_path is None:
# 出力ファイル名を自動生成
base_name = os.path.splitext(input_path)[0]
output_path = f"{base_name}_fp8.safetensors"
print(f"入力ファイル: {input_path}")
print(f"出力ファイル: {output_path}")
# メタデータを読み込む
metadata = {}
with safe_open(input_path, framework="pt", device="cpu") as f:
if f.metadata() is not None:
metadata = f.metadata()
# テンソルを変換
converted_tensors = {}
print("テンソルを変換中...")
with safe_open(input_path, framework="pt", device="cpu") as f:
for key in tqdm(f.keys(), desc="変換中"):
tensor = f.get_tensor(key)
# シンプルにFP8に変換
converted_tensors[key] = tensor.to(torch.float8_e4m3fn)
# メタデータに形式情報を追加
metadata["format"] = "pt"
metadata["fp8_type"] = "simple" # fp8-scaledではないことを明示
# 保存
print(f"変換したモデルを保存中: {output_path}")
save_file(converted_tensors, output_path, metadata=metadata)
# ファイルサイズの比較
original_size = os.path.getsize(input_path) / (1024**3) # GB
converted_size = os.path.getsize(output_path) / (1024**3) # GB
print(f"\n✅ 変換完了!")
print(f"元のサイズ: {original_size:.2f} GB")
print(f"変換後: {converted_size:.2f} GB")
print(f"削減率: {(1 - converted_size/original_size)*100:.1f}%")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="LoRAファイルを通常のFP8に変換")
parser.add_argument("input", type=str, help="入力LoRAファイルのパス")
parser.add_argument("--output", "-o", type=str, default=None,
help="出力ファイルのパス(省略時は自動生成)")
args = parser.parse_args()
convert_to_fp8_simple(args.input, args.output)