mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-15 01:36:40 +08:00
Minor compatibility fixes
This commit is contained in:
parent
803fa58fc8
commit
70cb59eed6
12
comfy/ops.py
12
comfy/ops.py
@ -23,7 +23,6 @@ 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
|
import types
|
||||||
|
|
||||||
|
|
||||||
@ -39,7 +38,7 @@ def quantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype):
|
|||||||
return x
|
return x
|
||||||
|
|
||||||
def dequantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype):
|
def dequantizer(x: torch.Tensor, scale: torch.Tensor, dtype: torch.dtype):
|
||||||
x = (x * scale).to(dtype=dtype)
|
x = x.to(dtype=dtype) * scale
|
||||||
return x
|
return x
|
||||||
|
|
||||||
def woq_fwd(self, x):
|
def woq_fwd(self, x):
|
||||||
@ -72,6 +71,7 @@ def get_quantized_forward(scale_weight, scale_input):
|
|||||||
return quantized_fwd
|
return quantized_fwd
|
||||||
|
|
||||||
def get_quantizer_fn(scale_weight, scale_input):
|
def get_quantizer_fn(scale_weight, scale_input):
|
||||||
|
# TODO Block Scaling, MX Scaling, Double-Q-NVFP4
|
||||||
return quantizer
|
return quantizer
|
||||||
|
|
||||||
def get_dequantizer_fn(scale_weight, scale_input):
|
def get_dequantizer_fn(scale_weight, scale_input):
|
||||||
@ -164,7 +164,7 @@ class disable_weight_init:
|
|||||||
self.out_features = out_features
|
self.out_features = out_features
|
||||||
|
|
||||||
self.device = device
|
self.device = device
|
||||||
self.dtype = dtype
|
self.compute_dtype = 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))
|
||||||
@ -178,7 +178,7 @@ class disable_weight_init:
|
|||||||
if not state_dict:
|
if not state_dict:
|
||||||
logging.warning("No state dict provided.")
|
logging.warning("No state dict provided.")
|
||||||
weight = torch.nn.Parameter(
|
weight = torch.nn.Parameter(
|
||||||
torch.empty((self.out_features, self.in_features))
|
torch.empty((self.out_features, self.in_features), dtype=self.compute_dtype, device=self.device)
|
||||||
)
|
)
|
||||||
self.register_buffer('weight', weight)
|
self.register_buffer('weight', weight)
|
||||||
return
|
return
|
||||||
@ -186,7 +186,7 @@ class disable_weight_init:
|
|||||||
device = state_dict[f"{prefix}weight"].device
|
device = state_dict[f"{prefix}weight"].device
|
||||||
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=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)
|
||||||
@ -204,7 +204,7 @@ class disable_weight_init:
|
|||||||
|
|
||||||
self.forward = types.MethodType(get_quantized_forward(scale_weight, scale_input), self)
|
self.forward = types.MethodType(get_quantized_forward(scale_weight, scale_input), self)
|
||||||
setattr(self, "quantizer", get_quantizer_fn(scale_weight, scale_input))
|
setattr(self, "quantizer", get_quantizer_fn(scale_weight, scale_input))
|
||||||
setattr(self, "dequanizer", get_dequantizer_fn(scale_weight, scale_input))
|
setattr(self, "dequantizer", get_dequantizer_fn(scale_weight, scale_input))
|
||||||
|
|
||||||
def _load_from_state_dict(
|
def _load_from_state_dict(
|
||||||
self,
|
self,
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user