Mixed Precision POC

This commit is contained in:
lspindler 2025-08-25 08:20:51 +02:00
parent ff57793659
commit 803fa58fc8

View File

@ -23,8 +23,60 @@ from comfy.cli_args import args, PerformanceFeature
import comfy.float import comfy.float
import comfy.rmsnorm import comfy.rmsnorm
import contextlib 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): def scaled_dot_product_attention(q, k, v, *args, **kwargs):
return torch.nn.functional.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 = [] bias_function = []
class disable_weight_init: 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): def reset_parameters(self):
return None return None
def forward_comfy_cast_weights(self, input): def _init_parameters_from_sd(self, state_dict, prefix):
weight, bias = cast_bias_weight(self, input) 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) 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): class Conv1d(torch.nn.Conv1d, CastWeightBiasOp):
def reset_parameters(self): 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): 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) fp8_compute = comfy.model_management.supports_fp8_compute(load_device)
if scaled_fp8 is not None: 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 ( if (
fp8_compute and fp8_compute and