This commit is contained in:
laksjdjf
2025-06-03 22:38:47 +09:00
committed by GitHub
parent 227f8755b9
commit ad783a1dfe
4 changed files with 247 additions and 0 deletions
+12
View File
@@ -0,0 +1,12 @@
from .nodes import ImatrixUNETLoader, SaveImatrix
NODE_CLASS_MAPPINGS = {
"ImatrixUNETLoader": ImatrixUNETLoader,
"SaveImatrix": SaveImatrix,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ImatrixUNETLoader": "Imatrix UNet Loader",
"SaveImatrix": "Save Imatrix",
}
__all__ = [NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS]
+1
View File
@@ -0,0 +1 @@
>q<
+94
View File
@@ -0,0 +1,94 @@
import comfy
import folder_paths
import torch
import comfy.ops
from .utils import save_imatrix
import os
CURRENT_DIR = os.path.dirname(os.path.realpath(__file__))
DATE_DIR = os.path.join(CURRENT_DIR, "imatrix_data")
class ImatrixOps(comfy.ops.manual_cast):
class Linear(comfy.ops.manual_cast.Linear):
def __init__(self, in_features, out_features, *args, **kwargs):
super().__init__(in_features, out_features, *args, **kwargs)
self.imatrix = torch.ones(in_features)
self.num_counts = 0
def forward(self, x, *args, **kwargs):
self.num_counts += 1
imatrix = x.detach().clone().float().pow(2).mean(dim=list(range(len(x.shape)))[:-1]).cpu()
self.imatrix = imatrix / self.num_counts + self.imatrix * (self.num_counts - 1) / self.num_counts
return super().forward(x, *args, **kwargs)
class Conv2d(comfy.ops.manual_cast.Conv2d):
def __init__(self, in_channels, out_channels, kernel_size, *args, **kwargs):
super().__init__(in_channels, out_channels, kernel_size, *args, **kwargs)
self.imatrix = torch.ones(in_channels * self.kernel_size[0] * self.kernel_size[1])
self.num_counts = 0
def forward(self, x, *args, **kwargs):
self.num_counts += 1
imatrix = x.detach().clone().float().pow(2).mean(dim=(0, 2, 3)).cpu().repeat(self.kernel_size[0] * self.kernel_size[1])
self.imatrix = imatrix / self.num_counts + self.imatrix * (self.num_counts - 1) / self.num_counts
return super().forward(x, *args, **kwargs)
class ImatrixUNETLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"unet_name": (folder_paths.get_filename_list("diffusion_models"), ),
"weight_dtype": (["default", "fp8_e4m3fn", "fp8_e4m3fn_fast", "fp8_e5m2"],)
}
}
RETURN_TYPES = ("MODEL",)
FUNCTION = "load_unet"
CATEGORY = "imatrix"
def load_unet(self, unet_name, weight_dtype="default"):
model_options = {}
if weight_dtype == "fp8_e4m3fn":
model_options["dtype"] = torch.float8_e4m3fn
elif weight_dtype == "fp8_e4m3fn_fast":
model_options["dtype"] = torch.float8_e4m3fn
model_options["fp8_optimizations"] = True
elif weight_dtype == "fp8_e5m2":
model_options["dtype"] = torch.float8_e5m2
model_options["custom_operations"] = ImatrixOps()
unet_path = folder_paths.get_full_path_or_raise("diffusion_models", unet_name)
model = comfy.sd.load_diffusion_model(unet_path, model_options=model_options)
return (model,)
class SaveImatrix:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("MODEL",),
"file_name": ("STRING", {"default": "imatrix"}),
},
"optional": {
"image_not_used": ("IMAGE", ),
}
}
RETURN_TYPES = ()
FUNCTION = "save_imatrix"
CATEGORY = "imatrix"
OUTPUT_NODE = True
def save_imatrix(self, model, file_name, image_not_used=None):
imatrix_data = {}
for name, module in model.model.diffusion_model.named_modules():
if hasattr(module, "imatrix") and module.imatrix is not None:
imatrix_data[name + ".weight"] = module.imatrix.float().cpu().numpy().tolist()
imatrix_file = os.path.join(DATE_DIR, f"{file_name}.dat")
save_imatrix(imatrix_file, imatrix_data)
print(f"Saved importance matrix to {imatrix_file}")
return {}
+140
View File
@@ -0,0 +1,140 @@
import os
import struct
from typing import Dict, List, Tuple, Sequence, Mapping, Union
def load_imatrix(
imatrix_file: str,
trace_env: str = "LLAMA_TRACE",
) -> Tuple[Dict[str, List[float]], str, int]:
"""
Parameters
----------
imatrix_file : str
読み込む .imatrix バイナリへのパス
trace_env : str, default "LLAMA_TRACE"
この環境変数がセットされていればデバッグ出力を行う
Returns
-------
imatrix_data : dict[str, list[float]]
エントリ名 → 重要度ベクトル
imatrix_dataset : str
ファイル末尾に埋め込まれたデータセット名(なければ "")
m_last_call : int
行列計算時のチャンク数(なければ 0)
"""
imatrix_data: Dict[str, List[float]] = {}
imatrix_dataset = ""
m_last_call = 0
with open(imatrix_file, "rb") as f:
# --- 1. 先頭: エントリ総数 ------------------------------------------
n_entries_bytes = f.read(4)
if len(n_entries_bytes) < 4:
raise ValueError(f"{imatrix_file}: no data")
n_entries = struct.unpack("<i", n_entries_bytes)[0]
if n_entries < 1:
raise ValueError(f"{imatrix_file}: n_entries < 1")
# --- 2. 各エントリを読み取る ----------------------------------------
for i in range(n_entries):
# 2-a. 名前長と名前文字列
name_len = struct.unpack("<i", f.read(4))[0]
name = f.read(name_len).decode("utf-8", errors="replace")
# 2-b. 呼び出し回数 ncall と 値の個数 nval
ncall, nval = struct.unpack("<ii", f.read(8))
if nval < 1:
raise ValueError(f"entry {i}: nval < 1")
# 2-c. nval 個の float32
buf = f.read(4 * nval)
if len(buf) < 4 * nval:
raise ValueError(f"entry {i}: data truncated")
values = list(struct.unpack(f"<{nval}f", buf))
# 2-d. 平均化
if ncall > 0:
values = [v / ncall for v in values]
imatrix_data[name] = values
if os.getenv(trace_env):
print(
f"load_imatrix: loaded data (size = {nval:6d}, "
f"ncall = {ncall:6d}) for '{name}'"
)
# --- 3. 末尾に追加メタ情報があるか確認 ------------------------------
tail = f.read(4)
if tail: # まだバイトが残っている場合のみ
m_last_call = struct.unpack("<i", tail)[0]
dataset_len = struct.unpack("<i", f.read(4))[0]
imatrix_dataset = f.read(dataset_len).decode("utf-8", errors="replace")
if os.getenv(trace_env):
print(f"load_imatrix: imatrix dataset = '{imatrix_dataset}'")
print(
f"load_imatrix: loaded {len(imatrix_data)} importance matrix entries "
f"from {imatrix_file} computed on {m_last_call} chunks"
)
return imatrix_data, imatrix_dataset, m_last_call
def save_imatrix(
imatrix_file: str,
imatrix_data: Mapping[str, Sequence[float]],
*,
call_counts: Union[int, Mapping[str, int]] = 1,
imatrix_dataset: str = "",
m_last_call: int = 0,
) -> None:
"""
Parameters
----------
imatrix_file : str
出力する .imatrix バイナリのパス
imatrix_data : dict[str, Sequence[float]]
エントリ名 → 値のベクトル(平均済み/未平均どちらでも OK)
call_counts : int | dict[str, int], default 1
・各エントリ共通の呼び出し回数 (= 平均化係数) を 1 つの int で指定
・あるいはエントリごとに dict で指定
※「すでに平均済みの値」を保存したいときは 0 を渡してください
imatrix_dataset : str, default ""
ファイル末尾に埋め込むデータセット名(空文字なら書き込まない)
m_last_call : int, default 0
全体のチャンク数。dataset 名を書くときはセットで入れる
"""
# 入力バリデーション ---------------------------------------------------
if isinstance(call_counts, int):
call_counts = {k: call_counts for k in imatrix_data.keys()}
else:
# dict で来た場合、全キーがそろっているか確認
missing = set(imatrix_data) - set(call_counts)
if missing:
raise KeyError(f"call_counts is missing keys: {missing}")
with open(imatrix_file, "wb") as f:
# 1. 先頭: エントリ総数 -------------------------------------------
n_entries = len(imatrix_data)
f.write(struct.pack("<i", n_entries))
# 2. 各エントリを順番に書き込む -----------------------------------
for name, values in imatrix_data.items():
name_b = name.encode("utf-8")
f.write(struct.pack("<i", len(name_b))) # name length
f.write(name_b) # name bytes
ncall = call_counts[name]
nval = len(values)
f.write(struct.pack("<ii", ncall, nval)) # ncall, nval
# nval 個の float32
fmt = f"<{nval}f"
f.write(struct.pack(fmt, *values))
# 3. オプションのメタ情報 -----------------------------------------
if imatrix_dataset:
f.write(struct.pack("<i", m_last_call))
dataset_b = imatrix_dataset.encode("utf-8")
f.write(struct.pack("<i", len(dataset_b)))
f.write(dataset_b)