mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-08-04 03:06:58 +08:00
add weight_scale_inv impl
This commit is contained in:
parent
eaf68c9b5b
commit
701e4999dc
28
comfy/ops.py
28
comfy/ops.py
@ -464,6 +464,22 @@ class fp8_ops(manual_cast):
|
|||||||
uncast_bias_weight(self, weight, bias, offload_stream)
|
uncast_bias_weight(self, weight, bias, offload_stream)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def scale_hadamard(larger, smaller, k_value=None, divide=False):
|
||||||
|
if smaller.shape == torch.Size([]):
|
||||||
|
# unset
|
||||||
|
return larger
|
||||||
|
h, w = smaller.shape
|
||||||
|
if k_value is None:
|
||||||
|
#calculate from larger compared to smaller
|
||||||
|
k_value = larger.shape[-1] // smaller.shape[-1]
|
||||||
|
expected_shape = (h * k_value, w * k_value)
|
||||||
|
assert larger.shape == expected_shape, "weight_scale_inv mismatch, skipping"
|
||||||
|
if divide:
|
||||||
|
smaller = 1.0 / smaller
|
||||||
|
result = larger.view(h, k_value, w, k_value) * smaller.view(h, 1, w, 1)
|
||||||
|
return result.reshape(h * k_value, w * k_value)
|
||||||
|
|
||||||
def scaled_fp8_ops(fp8_matrix_mult=False, scale_input=False, override_dtype=None):
|
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))
|
logging.info("Using scaled fp8: fp8 matrix mult: {}, scale input: {}".format(fp8_matrix_mult, scale_input))
|
||||||
class scaled_fp8_op(manual_cast):
|
class scaled_fp8_op(manual_cast):
|
||||||
@ -477,6 +493,9 @@ def scaled_fp8_ops(fp8_matrix_mult=False, scale_input=False, override_dtype=None
|
|||||||
if not hasattr(self, 'scale_weight'):
|
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)
|
self.scale_weight = torch.nn.parameter.Parameter(data=torch.ones((), device=self.weight.device, dtype=torch.float32), requires_grad=False)
|
||||||
|
|
||||||
|
if not hasattr(self, 'weight_scale_inv'):
|
||||||
|
self.weight_scale_inv = torch.nn.parameter.Parameter(data=torch.ones((), device=self.weight.device, dtype=torch.float32), requires_grad=False)
|
||||||
|
|
||||||
if not scale_input:
|
if not scale_input:
|
||||||
self.scale_input = None
|
self.scale_input = None
|
||||||
|
|
||||||
@ -502,12 +521,17 @@ def scaled_fp8_ops(fp8_matrix_mult=False, scale_input=False, override_dtype=None
|
|||||||
def convert_weight(self, weight, inplace=False, **kwargs):
|
def convert_weight(self, weight, inplace=False, **kwargs):
|
||||||
if inplace:
|
if inplace:
|
||||||
weight *= self.scale_weight.to(device=weight.device, dtype=weight.dtype)
|
weight *= self.scale_weight.to(device=weight.device, dtype=weight.dtype)
|
||||||
|
if self.weight_scale_inv.shape == torch.Size([]):
|
||||||
|
return weight
|
||||||
|
weight = scale_hadamard(weight, self.weight_scale_inv.to(device=weight.device, dtype=weight.dtype), k_value=None, divide=False)
|
||||||
return weight
|
return weight
|
||||||
else:
|
else:
|
||||||
return weight.to(dtype=torch.float32) * self.scale_weight.to(device=weight.device, dtype=torch.float32)
|
return scale_hadamard(weight.to(dtype=torch.float32) * self.scale_weight.to(device=weight.device, dtype=torch.float32),
|
||||||
|
self.weight_scale_inv.to(device=weight.device, dtype=torch.float32), k_value=None, divide=False)
|
||||||
|
|
||||||
def set_weight(self, weight, inplace_update=False, seed=None, return_weight=False, **kwargs):
|
def set_weight(self, weight, inplace_update=False, seed=None, return_weight=False, **kwargs):
|
||||||
weight = comfy.float.stochastic_rounding(weight / self.scale_weight.to(device=weight.device, dtype=weight.dtype), self.weight.dtype, seed=seed)
|
weight = comfy.float.stochastic_rounding(scale_hadamard(weight / self.scale_weight.to(device=weight.device, dtype=weight.dtype),
|
||||||
|
self.weight_scale_inv.to(device=weight.device, dtype=weight.dtype), k_value=None), self.weight.dtype, seed=seed)
|
||||||
if return_weight:
|
if return_weight:
|
||||||
return weight
|
return weight
|
||||||
if inplace_update:
|
if inplace_update:
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user