diff --git a/comfy/ops.py b/comfy/ops.py index 18e7db705..bd3628ac6 100644 --- a/comfy/ops.py +++ b/comfy/ops.py @@ -23,8 +23,60 @@ from comfy.cli_args import args, PerformanceFeature import comfy.float import comfy.rmsnorm import contextlib +from torch.nn.attention import SDPBackend, sdpa_kernel +import types +Q_TYPES = [torch.float8_e4m3fn, torch.float4_e2m1fn_x2] + +def dynamic_quantizer(x: torch.Tensor, dtype: torch.dtype): + input_scale = x.max() / torch.finfo(dtype).max + x = (x / input_scale).clamp(torch.finfo(dtype).min, torch.finfo(dtype).max).to(dtype=dtype) + return x, input_scale + +def quantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype): + x = (x / scale).clamp(torch.finfo(dtype).min, torch.finfo(dtype).max).to(dtype=dtype).contiguous() + return x + +def dequantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype): + x = (x * scale).to(dtype=dtype) + return x + +def woq_fwd(self, x): + dq_weight = self.dequantizer(self.weight, self.scale_weight, x.dtype) + return torch.nn.functional.linear(x, dq_weight, self.bias) + +def quantized_fwd(self, input): + tensor_2d = False + if len(input.shape) == 2: + tensor_2d = True + input = input.unsqueeze(1) + + input_shape = input.shape + input_dtype = input.dtype + assert len(input_shape) == 3, "input must be 3D" + + q_input = self.quantizer(input, self.scale_input, self.weight.dtype) + q_input = q_input.reshape(-1, input_shape[2]) + o = torch._scaled_mm(q_input, self.weight.T, scale_a=self.scale_input, scale_b=self.scale_weight, bias=self.bias, out_dtype=input_dtype) + if isinstance(o, tuple): + o = o[0] + if tensor_2d: + return o.reshape(input_shape[0], -1) + return o.reshape((-1, input_shape[1], self.weight.shape[0])) + +def get_quantized_forward(scale_weight, scale_input): + if scale_input is None: + return woq_fwd + else: + return quantized_fwd + +def get_quantizer_fn(scale_weight, scale_input): + return quantizer + +def get_dequantizer_fn(scale_weight, scale_input): + return dequantizer + def scaled_dot_product_attention(q, k, v, *args, **kwargs): return torch.nn.functional.scaled_dot_product_attention(q, k, v, *args, **kwargs) @@ -97,19 +149,92 @@ class CastWeightBiasOp: bias_function = [] class disable_weight_init: - class Linear(torch.nn.Linear, CastWeightBiasOp): + class Linear(torch.nn.Module, CastWeightBiasOp): + def __init__( + self, + in_features: int, + out_features: int, + bias: bool = True, + device=None, + dtype=None, + ) -> None: + factory_kwargs = {"device": device, "dtype": dtype} + super().__init__() + self.in_features = in_features + self.out_features = out_features + + self.device = device + self.dtype = dtype + + if bias: + self.bias = torch.nn.Parameter(torch.empty(out_features, **factory_kwargs)) + else: + self.register_parameter("bias", None) + def reset_parameters(self): return None - def forward_comfy_cast_weights(self, input): - weight, bias = cast_bias_weight(self, input) + def _init_parameters_from_sd(self, state_dict, prefix): + if not state_dict: + logging.warning("No state dict provided.") + weight = torch.nn.Parameter( + torch.empty((self.out_features, self.in_features)) + ) + self.register_buffer('weight', weight) + return + + device = state_dict[f"{prefix}weight"].device + weight_dtype = state_dict[f"{prefix}weight"].dtype + weight = torch.nn.Parameter( + torch.empty((self.out_features, self.in_features), device=device, dtype=weight_dtype) + ) + + self.register_buffer('weight', weight) + if weight_dtype not in Q_TYPES: + return + + scale_weight = state_dict.get(f"{prefix}scale_weight", None) + if scale_weight is None: + raise Exception("Using quantized Weights requires a scale to be present!") + self.register_buffer('scale_weight', scale_weight.float()) + + scale_input = state_dict.get(f"{prefix}scale_input", None) + if scale_input is not None: + self.register_buffer('scale_input', scale_input.float()) + + self.forward = types.MethodType(get_quantized_forward(scale_weight, scale_input), self) + setattr(self, "quantizer", get_quantizer_fn(scale_weight, scale_input)) + setattr(self, "dequanizer", get_dequantizer_fn(scale_weight, scale_input)) + + def _load_from_state_dict( + self, + state_dict, + prefix, + local_metadata, + strict, + missing_keys, + unexpected_keys, + error_msgs, + ): + self._init_parameters_from_sd(state_dict, prefix) + + return super()._load_from_state_dict( + state_dict, + prefix, + local_metadata, + strict, + missing_keys, + unexpected_keys, + error_msgs, + ) + + def forward(self, input): + weight = self.weight + bias = self.bias + if self.comfy_cast_weights or len(self.weight_function) > 0 or len(self.bias_function) > 0: + weight, bias = cast_bias_weight(self, input) return torch.nn.functional.linear(input, weight, bias) - def forward(self, *args, **kwargs): - if self.comfy_cast_weights or len(self.weight_function) > 0 or len(self.bias_function) > 0: - return self.forward_comfy_cast_weights(*args, **kwargs) - else: - return super().forward(*args, **kwargs) class Conv1d(torch.nn.Conv1d, CastWeightBiasOp): def reset_parameters(self): @@ -443,7 +568,8 @@ if CUBLAS_IS_AVAILABLE: def pick_operations(weight_dtype, compute_dtype, load_device=None, disable_fast_fp8=False, fp8_optimizations=False, scaled_fp8=None): fp8_compute = comfy.model_management.supports_fp8_compute(load_device) if scaled_fp8 is not None: - return scaled_fp8_ops(fp8_matrix_mult=fp8_compute and fp8_optimizations, scale_input=fp8_optimizations, override_dtype=scaled_fp8) + return disable_weight_init + # return scaled_fp8_ops(fp8_matrix_mult=fp8_compute and fp8_optimizations, scale_input=fp8_optimizations, override_dtype=scaled_fp8) if ( fp8_compute and