diff --git a/fp8_optimization.py b/fp8_optimization.py new file mode 100644 index 0000000..09f026d --- /dev/null +++ b/fp8_optimization.py @@ -0,0 +1,47 @@ +#based on ComfyUI's and MinusZoneAI's fp8_linear optimization + +import torch +import torch.nn as nn + +def fp8_linear_forward(cls, original_dtype, input): + weight_dtype = cls.weight.dtype + if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: + if len(input.shape) == 3: + if weight_dtype == torch.float8_e4m3fn: + inn = input.reshape(-1, input.shape[2]).to(torch.float8_e5m2) + else: + inn = input.reshape(-1, input.shape[2]).to(torch.float8_e4m3fn) + w = cls.weight.t() + + scale_weight = torch.ones((1), device=input.device, dtype=torch.float32) + scale_input = scale_weight + + bias = cls.bias.to(original_dtype) if cls.bias is not None else None + out_dtype = original_dtype + + if bias is not None: + o = torch._scaled_mm(inn, w, out_dtype=out_dtype, bias=bias, scale_a=scale_input, scale_b=scale_weight) + else: + o = torch._scaled_mm(inn, w, out_dtype=out_dtype, scale_a=scale_input, scale_b=scale_weight) + + if isinstance(o, tuple): + o = o[0] + + return o.reshape((-1, input.shape[1], cls.weight.shape[0])) + else: + cls.to(original_dtype) + out = cls.original_forward(input.to(original_dtype)) + cls.to(original_dtype) + return out + else: + return cls.original_forward(input) + +def convert_fp8_linear(module, original_dtype, params_to_keep={}): + setattr(module, "fp8_matmul_enabled", True) + + for name, module in module.named_modules(): + if not any(keyword in name for keyword in params_to_keep): + if isinstance(module, nn.Linear): + original_forward = module.forward + setattr(module, "original_forward", original_forward) + setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input)) diff --git a/nodes.py b/nodes.py index 68a4c4f..80ee7e4 100644 --- a/nodes.py +++ b/nodes.py @@ -101,7 +101,7 @@ class HyVideoModelLoader: "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), "base_precision": (["fp16", "fp32", "bf16"], {"default": "bf16"}), - "quantization": (['disabled', 'fp8_e4m3fn', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6"], {"default": 'disabled', "tooltip": "optional quantization method"}), + "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6"], {"default": 'disabled', "tooltip": "optional quantization method"}), "load_device": (["main_device", "offload_device"], {"default": "main_device"}), }, "optional": { @@ -121,7 +121,7 @@ class HyVideoModelLoader: CATEGORY = "HunyuanVideoWrapper" def loadmodel(self, model, base_precision, load_device, quantization, - compile_args=None, attention_mode="sdpa", enable_sequential_cpu_offload=False, block_swap_args=None): + compile_args=None, attention_mode="sdpa", block_swap_args=None): transformer = None manual_offloading = True if "sage" in attention_mode: @@ -164,7 +164,7 @@ class HyVideoModelLoader: ) log.info("Using accelerate to load and assign model weights to device...") - if quantization == "fp8_e4m3fn": + if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast": dtype = torch.float8_e4m3fn else: dtype = base_dtype @@ -221,6 +221,12 @@ class HyVideoModelLoader: manual_offloading = False # to disable manual .to(device) calls log.info(f"Quantized transformer blocks to {quantization}") + + elif quantization == "fp8_e4m3fn_fast": + from .fp8_optimization import convert_fp8_linear + if "1.5" in model: + params_to_keep.update({"ff"}) #otherwise NaNs + convert_fp8_linear(transformer, base_dtype, params_to_keep=params_to_keep) scheduler = FlowMatchDiscreteScheduler( shift=9.0,