From ad783a1dfe531cf7237df7d993d3a08a8f39855e Mon Sep 17 00:00:00 2001 From: laksjdjf Date: Tue, 3 Jun 2025 22:38:47 +0900 Subject: [PATCH] auau --- __init__.py | 12 ++++ imatrix_data/a.txt | 1 + nodes.py | 94 ++++++++++++++++++++++++++++++ utils.py | 140 +++++++++++++++++++++++++++++++++++++++++++++ 4 files changed, 247 insertions(+) create mode 100644 __init__.py create mode 100644 imatrix_data/a.txt create mode 100644 nodes.py create mode 100644 utils.py diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..17673c5 --- /dev/null +++ b/__init__.py @@ -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] \ No newline at end of file diff --git a/imatrix_data/a.txt b/imatrix_data/a.txt new file mode 100644 index 0000000..c747b19 --- /dev/null +++ b/imatrix_data/a.txt @@ -0,0 +1 @@ +>q< \ No newline at end of file diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..0f05ec3 --- /dev/null +++ b/nodes.py @@ -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 {} \ No newline at end of file diff --git a/utils.py b/utils.py new file mode 100644 index 0000000..f81bd0d --- /dev/null +++ b/utils.py @@ -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(" 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(" 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("