mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-15 03:50:02 +08:00
Mixed Precision POC
This commit is contained in:
parent
ff57793659
commit
803fa58fc8
144
comfy/ops.py
144
comfy/ops.py
@ -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
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user