mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-15 05:10:02 +08:00
Introduce class factory for creating Opset with settings
This commit is contained in:
parent
4e20714b94
commit
e266831037
@ -143,6 +143,7 @@ class PerformanceFeature(enum.Enum):
|
|||||||
Fp16Accumulation = "fp16_accumulation"
|
Fp16Accumulation = "fp16_accumulation"
|
||||||
Fp8MatrixMultiplication = "fp8_matrix_mult"
|
Fp8MatrixMultiplication = "fp8_matrix_mult"
|
||||||
CublasOps = "cublas_ops"
|
CublasOps = "cublas_ops"
|
||||||
|
DynamicQuantizer = "dynamic_quantizer"
|
||||||
|
|
||||||
parser.add_argument("--fast", nargs="*", type=PerformanceFeature, help="Enable some untested and potentially quality deteriorating optimizations. --fast with no arguments enables everything. You can pass a list specific optimizations if you only want to enable specific ones. Current valid optimizations: fp16_accumulation fp8_matrix_mult cublas_ops")
|
parser.add_argument("--fast", nargs="*", type=PerformanceFeature, help="Enable some untested and potentially quality deteriorating optimizations. --fast with no arguments enables everything. You can pass a list specific optimizations if you only want to enable specific ones. Current valid optimizations: fp16_accumulation fp8_matrix_mult cublas_ops")
|
||||||
|
|
||||||
|
|||||||
@ -3,7 +3,7 @@ import torch
|
|||||||
from torch import nn
|
from torch import nn
|
||||||
from .ldm.modules.attention import CrossAttention, FeedForward
|
from .ldm.modules.attention import CrossAttention, FeedForward
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
ops = comfy.ops.manual_cast
|
ops = comfy.ops.operator_factory(comfy_cast_weights=True)
|
||||||
|
|
||||||
|
|
||||||
class GatedCrossAttentionDense(nn.Module):
|
class GatedCrossAttentionDense(nn.Module):
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
import torch
|
import torch
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
|
|
||||||
ops = comfy.ops.manual_cast
|
ops = comfy.ops.operator_factory(comfy_cast_weights=True)
|
||||||
|
|
||||||
class ReduxImageEncoder(torch.nn.Module):
|
class ReduxImageEncoder(torch.nn.Module):
|
||||||
def __init__(
|
def __init__(
|
||||||
|
|||||||
309
comfy/ops.py
309
comfy/ops.py
@ -23,62 +23,9 @@ from comfy.cli_args import args, PerformanceFeature
|
|||||||
import comfy.float
|
import comfy.float
|
||||||
import comfy.rmsnorm
|
import comfy.rmsnorm
|
||||||
import contextlib
|
import contextlib
|
||||||
|
from comfy.quant_tensor import Q_TYPES, tensor_quantizer, tensor_dequantizer, dynamic_tensor_quantizer, woq_fwd, quantized_fwd
|
||||||
import types
|
import types
|
||||||
|
import inspect
|
||||||
|
|
||||||
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.float()
|
|
||||||
|
|
||||||
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.to(dtype=scale.dtype) * scale).to(dtype=dtype)
|
|
||||||
return x
|
|
||||||
|
|
||||||
def woq_fwd(self, x):
|
|
||||||
dq_weight = self.dequantizer(self.weight, self.scale_weight, x.dtype)
|
|
||||||
bias = self.bias
|
|
||||||
if bias is not None and bias.dtype != x.dtype:
|
|
||||||
bias = self.dequantizer(bias, self.scale_weight, x.dtype)
|
|
||||||
return torch.nn.functional.linear(x, dq_weight, 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):
|
|
||||||
# TODO Block Scaling, MX Scaling, Double-Q-NVFP4
|
|
||||||
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)
|
||||||
@ -149,6 +96,8 @@ class CastWeightBiasOp:
|
|||||||
comfy_cast_weights = False
|
comfy_cast_weights = False
|
||||||
weight_function = []
|
weight_function = []
|
||||||
bias_function = []
|
bias_function = []
|
||||||
|
fp8_compute: bool = False
|
||||||
|
use_dynamic_quantizer: bool = False
|
||||||
|
|
||||||
class disable_weight_init:
|
class disable_weight_init:
|
||||||
class Linear(torch.nn.Module, CastWeightBiasOp):
|
class Linear(torch.nn.Module, CastWeightBiasOp):
|
||||||
@ -165,8 +114,8 @@ class disable_weight_init:
|
|||||||
self.in_features = in_features
|
self.in_features = in_features
|
||||||
self.out_features = out_features
|
self.out_features = out_features
|
||||||
|
|
||||||
self.device = device
|
self.device = torch.device("cpu") if device is None else device
|
||||||
self.compute_dtype = dtype
|
self.compute_dtype = torch.float32 if dtype is None else dtype
|
||||||
|
|
||||||
if bias:
|
if bias:
|
||||||
self.bias = torch.nn.Parameter(torch.empty(out_features, **factory_kwargs))
|
self.bias = torch.nn.Parameter(torch.empty(out_features, **factory_kwargs))
|
||||||
@ -176,6 +125,26 @@ class disable_weight_init:
|
|||||||
def reset_parameters(self):
|
def reset_parameters(self):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
def _set_quantizer_fn(self, scale_weight, scale_input):
|
||||||
|
if scale_weight.ndim != 0 and scale_weight.shape[0] != 1:
|
||||||
|
raise ValueError("Blockwise quantization is not supported")
|
||||||
|
if self.use_dynamic_quantizer:
|
||||||
|
setattr(self, "quantizer", dynamic_tensor_quantizer)
|
||||||
|
else:
|
||||||
|
setattr(self, "quantizer", tensor_quantizer)
|
||||||
|
|
||||||
|
def _set_dequantizer_fn(self, scale_weight):
|
||||||
|
if scale_weight.ndim != 0 and scale_weight.shape[0] != 1:
|
||||||
|
raise ValueError("Blockwise quantization is not supported")
|
||||||
|
setattr(self, "dequantizer", tensor_dequantizer)
|
||||||
|
|
||||||
|
def _set_quantized_forward(self):
|
||||||
|
if not self.fp8_compute:
|
||||||
|
q_fwd = woq_fwd
|
||||||
|
else:
|
||||||
|
q_fwd = quantized_fwd
|
||||||
|
self.forward = types.MethodType(q_fwd, self)
|
||||||
|
|
||||||
def _init_parameters_from_sd(self, state_dict, prefix):
|
def _init_parameters_from_sd(self, state_dict, prefix):
|
||||||
if not state_dict:
|
if not state_dict:
|
||||||
logging.warning("No state dict provided.")
|
logging.warning("No state dict provided.")
|
||||||
@ -187,8 +156,7 @@ class disable_weight_init:
|
|||||||
|
|
||||||
weight_dtype = state_dict[f"{prefix}weight"].dtype
|
weight_dtype = state_dict[f"{prefix}weight"].dtype
|
||||||
weight = torch.nn.Parameter(
|
weight = torch.nn.Parameter(
|
||||||
torch.empty((self.out_features, self.in_features), device=self.device, dtype=weight_dtype)
|
torch.empty((self.out_features, self.in_features), device=self.device, dtype=weight_dtype))
|
||||||
)
|
|
||||||
|
|
||||||
self.register_buffer('weight', weight)
|
self.register_buffer('weight', weight)
|
||||||
if weight_dtype not in Q_TYPES:
|
if weight_dtype not in Q_TYPES:
|
||||||
@ -199,17 +167,21 @@ class disable_weight_init:
|
|||||||
logging.warning("Using quantized Weights requires a scale to be present! Falling back to 1.0")
|
logging.warning("Using quantized Weights requires a scale to be present! Falling back to 1.0")
|
||||||
scale_weight = torch.ones(1)
|
scale_weight = torch.ones(1)
|
||||||
state_dict[f"{prefix}scale_weight"] = scale_weight
|
state_dict[f"{prefix}scale_weight"] = scale_weight
|
||||||
self.register_buffer('scale_weight', scale_weight.to(device=self.device, dtype=torch.float32))
|
self.register_buffer('scale_weight', scale_weight.to(device=self.device))
|
||||||
|
|
||||||
scale_input = state_dict.get(f"{prefix}scale_input", None)
|
scale_input = state_dict.get(f"{prefix}scale_input", None)
|
||||||
|
if scale_input is None and self.use_dynamic_quantizer:
|
||||||
|
scale_input = torch.ones(1) # Placeholder for API
|
||||||
if scale_input is not None:
|
if scale_input is not None:
|
||||||
if scale_input != 1: # TODO not really nice but e.g. Qwen VL has an input scale but does not use it?
|
self.register_buffer('scale_input', scale_input.to(device=self.device))
|
||||||
self.register_buffer('scale_input', scale_input.to(device=self.device, dtype=torch.float32))
|
|
||||||
else:
|
if self.bias is not None:
|
||||||
scale_input = None
|
# WAR not really nice, but Qwen VL has an input scale but uses f32 intermediates and quantized bias
|
||||||
self.forward = types.MethodType(get_quantized_forward(scale_weight, scale_input), self)
|
self.fp8_compute = not self.bias.dtype in Q_TYPES
|
||||||
setattr(self, "quantizer", get_quantizer_fn(scale_weight, scale_input))
|
|
||||||
setattr(self, "dequantizer", get_dequantizer_fn(scale_weight, scale_input))
|
self._set_quantizer_fn(scale_weight, scale_input)
|
||||||
|
self._set_dequantizer_fn(scale_weight)
|
||||||
|
self._set_quantized_forward()
|
||||||
|
|
||||||
def _load_from_state_dict(
|
def _load_from_state_dict(
|
||||||
self,
|
self,
|
||||||
@ -240,7 +212,6 @@ class disable_weight_init:
|
|||||||
weight, bias = cast_bias_weight(self, input)
|
weight, bias = cast_bias_weight(self, input)
|
||||||
return torch.nn.functional.linear(input, weight, bias)
|
return torch.nn.functional.linear(input, weight, bias)
|
||||||
|
|
||||||
|
|
||||||
class Conv1d(torch.nn.Conv1d, CastWeightBiasOp):
|
class Conv1d(torch.nn.Conv1d, CastWeightBiasOp):
|
||||||
def reset_parameters(self):
|
def reset_parameters(self):
|
||||||
return None
|
return None
|
||||||
@ -405,152 +376,6 @@ class disable_weight_init:
|
|||||||
else:
|
else:
|
||||||
raise ValueError(f"unsupported dimensions: {dims}")
|
raise ValueError(f"unsupported dimensions: {dims}")
|
||||||
|
|
||||||
|
|
||||||
class manual_cast(disable_weight_init):
|
|
||||||
class Linear(disable_weight_init.Linear):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class Conv1d(disable_weight_init.Conv1d):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class Conv2d(disable_weight_init.Conv2d):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class Conv3d(disable_weight_init.Conv3d):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class GroupNorm(disable_weight_init.GroupNorm):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class LayerNorm(disable_weight_init.LayerNorm):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class ConvTranspose2d(disable_weight_init.ConvTranspose2d):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class ConvTranspose1d(disable_weight_init.ConvTranspose1d):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class RMSNorm(disable_weight_init.RMSNorm):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
class Embedding(disable_weight_init.Embedding):
|
|
||||||
comfy_cast_weights = True
|
|
||||||
|
|
||||||
|
|
||||||
def fp8_linear(self, input):
|
|
||||||
dtype = self.weight.dtype
|
|
||||||
if dtype not in [torch.float8_e4m3fn]:
|
|
||||||
return None
|
|
||||||
|
|
||||||
tensor_2d = False
|
|
||||||
if len(input.shape) == 2:
|
|
||||||
tensor_2d = True
|
|
||||||
input = input.unsqueeze(1)
|
|
||||||
|
|
||||||
input_shape = input.shape
|
|
||||||
input_dtype = input.dtype
|
|
||||||
if len(input.shape) == 3:
|
|
||||||
w, bias = cast_bias_weight(self, input, dtype=dtype, bias_dtype=input_dtype)
|
|
||||||
w = w.t()
|
|
||||||
|
|
||||||
scale_weight = self.scale_weight
|
|
||||||
scale_input = self.scale_input
|
|
||||||
if scale_weight is None:
|
|
||||||
scale_weight = torch.ones((), device=input.device, dtype=torch.float32)
|
|
||||||
else:
|
|
||||||
scale_weight = scale_weight.to(input.device)
|
|
||||||
|
|
||||||
if scale_input is None:
|
|
||||||
scale_input = torch.ones((), device=input.device, dtype=torch.float32)
|
|
||||||
input = torch.clamp(input, min=-448, max=448, out=input)
|
|
||||||
input = input.reshape(-1, input_shape[2]).to(dtype).contiguous()
|
|
||||||
else:
|
|
||||||
scale_input = scale_input.to(input.device)
|
|
||||||
input = (input * (1.0 / scale_input).to(input_dtype)).reshape(-1, input_shape[2]).to(dtype).contiguous()
|
|
||||||
|
|
||||||
if bias is not None:
|
|
||||||
o = torch._scaled_mm(input, w, out_dtype=input_dtype, bias=bias, scale_a=scale_input, scale_b=scale_weight)
|
|
||||||
else:
|
|
||||||
o = torch._scaled_mm(input, w, out_dtype=input_dtype, scale_a=scale_input, scale_b=scale_weight)
|
|
||||||
|
|
||||||
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]))
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
class fp8_ops(manual_cast):
|
|
||||||
class Linear(manual_cast.Linear):
|
|
||||||
def reset_parameters(self):
|
|
||||||
self.scale_weight = None
|
|
||||||
self.scale_input = None
|
|
||||||
return None
|
|
||||||
|
|
||||||
def forward_comfy_cast_weights(self, input):
|
|
||||||
try:
|
|
||||||
out = fp8_linear(self, input)
|
|
||||||
if out is not None:
|
|
||||||
return out
|
|
||||||
except Exception as e:
|
|
||||||
logging.info("Exception during fp8 op: {}".format(e))
|
|
||||||
|
|
||||||
weight, bias = cast_bias_weight(self, input)
|
|
||||||
return torch.nn.functional.linear(input, weight, bias)
|
|
||||||
|
|
||||||
def scaled_fp8_ops(fp8_matrix_mult=False, scale_input=False, override_dtype=None):
|
|
||||||
logging.info("Using scaled fp8: fp8 matrix mult: {}, scale input: {}".format(fp8_matrix_mult, scale_input))
|
|
||||||
class scaled_fp8_op(manual_cast):
|
|
||||||
class Linear(manual_cast.Linear):
|
|
||||||
def __init__(self, *args, **kwargs):
|
|
||||||
if override_dtype is not None:
|
|
||||||
kwargs['dtype'] = override_dtype
|
|
||||||
super().__init__(*args, **kwargs)
|
|
||||||
|
|
||||||
def reset_parameters(self):
|
|
||||||
if not hasattr(self, 'scale_weight'):
|
|
||||||
self.scale_weight = torch.nn.parameter.Parameter(data=torch.ones((), device=self.weight.device, dtype=torch.float32), requires_grad=False)
|
|
||||||
|
|
||||||
if not scale_input:
|
|
||||||
self.scale_input = None
|
|
||||||
|
|
||||||
if not hasattr(self, 'scale_input'):
|
|
||||||
self.scale_input = torch.nn.parameter.Parameter(data=torch.ones((), device=self.weight.device, dtype=torch.float32), requires_grad=False)
|
|
||||||
return None
|
|
||||||
|
|
||||||
def forward_comfy_cast_weights(self, input):
|
|
||||||
if fp8_matrix_mult:
|
|
||||||
out = fp8_linear(self, input)
|
|
||||||
if out is not None:
|
|
||||||
return out
|
|
||||||
|
|
||||||
weight, bias = cast_bias_weight(self, input)
|
|
||||||
|
|
||||||
if weight.numel() < input.numel(): #TODO: optimize
|
|
||||||
return torch.nn.functional.linear(input, weight * self.scale_weight.to(device=weight.device, dtype=weight.dtype), bias)
|
|
||||||
else:
|
|
||||||
return torch.nn.functional.linear(input * self.scale_weight.to(device=weight.device, dtype=weight.dtype), weight, bias)
|
|
||||||
|
|
||||||
def convert_weight(self, weight, inplace=False, **kwargs):
|
|
||||||
if inplace:
|
|
||||||
weight *= self.scale_weight.to(device=weight.device, dtype=weight.dtype)
|
|
||||||
return weight
|
|
||||||
else:
|
|
||||||
return weight * self.scale_weight.to(device=weight.device, dtype=weight.dtype)
|
|
||||||
|
|
||||||
def set_weight(self, weight, inplace_update=False, seed=None, **kwargs):
|
|
||||||
weight = comfy.float.stochastic_rounding(weight / self.scale_weight.to(device=weight.device, dtype=weight.dtype), self.weight.dtype, seed=seed)
|
|
||||||
if inplace_update:
|
|
||||||
self.weight.data.copy_(weight)
|
|
||||||
else:
|
|
||||||
self.weight = torch.nn.Parameter(weight, requires_grad=False)
|
|
||||||
|
|
||||||
return scaled_fp8_op
|
|
||||||
|
|
||||||
CUBLAS_IS_AVAILABLE = False
|
CUBLAS_IS_AVAILABLE = False
|
||||||
try:
|
try:
|
||||||
from cublas_ops import CublasLinear
|
from cublas_ops import CublasLinear
|
||||||
@ -570,32 +395,30 @@ if CUBLAS_IS_AVAILABLE: # TODO check if this is actually faster I call BS
|
|||||||
def forward(self, *args, **kwargs):
|
def forward(self, *args, **kwargs):
|
||||||
return super().forward(*args, **kwargs)
|
return super().forward(*args, **kwargs)
|
||||||
|
|
||||||
def pick_operations(weight_dtype, compute_dtype, load_device=None, disable_fast_fp8=False, fp8_optimizations=False, scaled_fp8=None):
|
op_class_list = [
|
||||||
fp8_compute = comfy.model_management.supports_fp8_compute(load_device)
|
cls for name, cls in inspect.getmembers(disable_weight_init, inspect.isclass)
|
||||||
# TODO consider the different support cases
|
if cls.__module__ == disable_weight_init.__module__ and name != "__class__"
|
||||||
# Potentially also allow for auto-quant with dynamic input quantizers
|
]
|
||||||
return disable_weight_init
|
|
||||||
# if scaled_fp8 is not None:
|
def operator_factory(**factory_kwargs):
|
||||||
# return disable_weight_init
|
class OpSet:
|
||||||
# # return scaled_fp8_ops(fp8_matrix_mult=fp8_compute and fp8_optimizations, scale_input=fp8_optimizations, override_dtype=scaled_fp8)
|
pass
|
||||||
#
|
op_set = OpSet()
|
||||||
# if (
|
|
||||||
# fp8_compute and
|
for k, v in factory_kwargs.items():
|
||||||
# (fp8_optimizations or PerformanceFeature.Fp8MatrixMultiplication in args.fast) and
|
assert hasattr(CastWeightBiasOp, k)
|
||||||
# not disable_fast_fp8
|
|
||||||
# ):
|
for base_class in op_class_list:
|
||||||
# return fp8_ops
|
new_class = type(base_class.__name__, (base_class,), factory_kwargs)
|
||||||
#
|
setattr(op_set, base_class.__name__, new_class)
|
||||||
# if (
|
|
||||||
# PerformanceFeature.CublasOps in args.fast and
|
return op_set
|
||||||
# CUBLAS_IS_AVAILABLE and
|
|
||||||
# weight_dtype == torch.float16 and
|
# TODO might be nicer to have a unified interface to the factory
|
||||||
# (compute_dtype == torch.float16 or compute_dtype is None)
|
# TODO logic might not be 1-1 match to original implementation
|
||||||
# ):
|
def pick_operations(weight_dtype=None, compute_dtype=None, load_device=None, disable_fast_fp8=False):
|
||||||
# logging.info("Using cublas ops")
|
fp8_compute = (comfy.model_management.supports_fp8_compute(load_device) and not disable_fast_fp8)
|
||||||
# return cublas_ops
|
use_dynamic_quantizer = PerformanceFeature.DynamicQuantizer in args.fast
|
||||||
#
|
manual_cast = not((weight_dtype == compute_dtype) or use_dynamic_quantizer or fp8_compute)
|
||||||
# if compute_dtype is None or weight_dtype == compute_dtype:
|
return operator_factory(comfy_cast_weights=manual_cast, use_dynamic_quantizer=use_dynamic_quantizer, fp8_compute=fp8_compute)
|
||||||
# return disable_weight_init
|
|
||||||
#
|
|
||||||
# return manual_cast
|
|
||||||
|
|||||||
91
comfy/quant_tensor.py
Normal file
91
comfy/quant_tensor.py
Normal file
@ -0,0 +1,91 @@
|
|||||||
|
import torch
|
||||||
|
|
||||||
|
"""
|
||||||
|
x: 512x1024 w:1024x1024
|
||||||
|
- For TensorWise scaling, a and b should be float8, scales should be float and singletons.
|
||||||
|
- For RowWise scaling, a and b should be float8, scales should be float, scale_a should be (512, 1) and scale_b should be (1, 1024), and both should be contiguous.
|
||||||
|
- For BlockWise 1x128 scaling, a and b should be float8, scales should be float, scale_a should be (512, 8) and scale_b should be (8, 1024), and both should be outer-dim-major.
|
||||||
|
- For BlockWise 128x128 scaling, a and b should be float8, scales should be float, scale_a should be (4, 8) and scale_b should be (8, 8), and both should be near-inner-dim-major (with 16-byte aligned strides).
|
||||||
|
- For Blockwise 1x32 scaling, a and b should be float8, scales should be float8_e8m0fnu, scale_a should have 16384 elements and scale_b should have 32768 elements, and both should be contiguous.
|
||||||
|
- For Blockwise 1x16 scaling, a and b should be float4 (packed 2x), scales should be float8_e4m3fn, scale_a should have 65536 elements and scale_b should have 131072 elements, and both should be contiguous.
|
||||||
|
"""
|
||||||
|
Q_TYPES = [torch.float8_e4m3fn, torch.float4_e2m1fn_x2]
|
||||||
|
|
||||||
|
def dynamic_tensor_quantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype):
|
||||||
|
input_scale = torch.abs(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.float()
|
||||||
|
|
||||||
|
def mxfp8_quantizer(x: torch.Tensor, dtype: torch.dtype):
|
||||||
|
block_size = 32
|
||||||
|
orig_shape = x.shape
|
||||||
|
x = x.reshape(-1, block_size)
|
||||||
|
scale = (torch.amax(torch.abs(x), dim=-1) / torch.finfo(dtype).max)
|
||||||
|
x = (x / scale.unsqueeze(-1)).clamp(torch.finfo(dtype).min, torch.finfo(dtype).max).to(dtype=dtype).contiguous()
|
||||||
|
x = x.view(orig_shape)
|
||||||
|
|
||||||
|
return x, scale.to(dtype=torch.float8_e8m0fnu).contiguous()
|
||||||
|
|
||||||
|
def tensor_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, scale.float()
|
||||||
|
|
||||||
|
def nvfp4_quantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype):
|
||||||
|
"""
|
||||||
|
orig_shape = x.shape
|
||||||
|
x = x.reshape(orig_shape[0], -1, block_size)
|
||||||
|
max_abs = torch.amax(torch.abs(x), dim=-1)
|
||||||
|
block_scale = (max_abs / torch.finfo(torch.float4_e2m1fn_x2.max))-float()
|
||||||
|
scaled_block_scales = block_scale / scale
|
||||||
|
scaled_block_scales_fp8 = torch.clamp(
|
||||||
|
scaled_block_scales, min=E4M3_EPS, max=F8E4M3_MAX
|
||||||
|
).to(torch.float8_e4m3fn)
|
||||||
|
scaled_block_scales_fp32 = scaled_block_scales_fp8.to(torch.float32)
|
||||||
|
# We "temporarily" dequant the scaled_block_scales_fp32 to get the per_tensor_scale
|
||||||
|
# To apply to data
|
||||||
|
total_scale = scale * scaled_block_scales_fp32
|
||||||
|
data_scaled = x / total_scale.unsqueeze(-1)
|
||||||
|
out_scales = scaled_block_scales_fp8
|
||||||
|
|
||||||
|
data_scaled = torch.clamp(data_scaled, -F4_E2M1_MAX, F4_E2M1_MAX)
|
||||||
|
data_scaled = data_scaled.view(orig_shape)
|
||||||
|
data_lp = f32_to_f4_unpacked(data_scaled)
|
||||||
|
# TODO: NotImplementedError: "copy_kernel" not implemented for 'Float4_e2m1fn_x2'
|
||||||
|
# data_lp = pack_uint4(data_lp).view(torch.float4_e2m1fn_x2)
|
||||||
|
data_lp = pack_uint4(data_lp)
|
||||||
|
return out_scales, data_lp
|
||||||
|
"""
|
||||||
|
block_size: int = 16
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
|
def tensor_dequantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype):
|
||||||
|
x = (x.to(dtype=scale.dtype) * scale).to(dtype=dtype)
|
||||||
|
return x
|
||||||
|
|
||||||
|
def woq_fwd(self, x):
|
||||||
|
dq_weight = self.dequantizer(self.weight, self.scale_weight, x.dtype)
|
||||||
|
bias = self.bias
|
||||||
|
if bias is not None and bias.dtype == self.weight.dtype:
|
||||||
|
bias = self.dequantizer(bias, self.scale_weight, x.dtype)
|
||||||
|
return torch.nn.functional.linear(x, dq_weight, 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, input_scale = 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=input_scale, scale_b=self.scale_weight.float(),
|
||||||
|
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]))
|
||||||
@ -113,9 +113,11 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
if operations is None:
|
if operations is None:
|
||||||
scaled_fp8 = model_options.get("scaled_fp8", None)
|
scaled_fp8 = model_options.get("scaled_fp8", None)
|
||||||
if scaled_fp8 is not None:
|
if scaled_fp8 is not None:
|
||||||
operations = comfy.ops.scaled_fp8_ops(fp8_matrix_mult=False, override_dtype=scaled_fp8)
|
# operations = comfy.ops.scaled_fp8_ops(fp8_matrix_mult=False, override_dtype=scaled_fp8) # TODO
|
||||||
|
operations = comfy.ops.pick_operations(compute_dtype=torch.float8_e4m3fn, weight_dtype=torch.float8_e4m3fn)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
operations = comfy.ops.manual_cast
|
operations = comfy.ops.operator_factory(comfy_cast_weights=True)
|
||||||
|
|
||||||
self.operations = operations
|
self.operations = operations
|
||||||
self.transformer = model_class(config, dtype, device, self.operations)
|
self.transformer = model_class(config, dtype, device, self.operations)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user