Fix quant_state None on AMD GPUs by caching quant_state_dict at load time
This commit is contained in:
@@ -42,7 +42,7 @@ git clone https://github.com/mengqin/ComfyUI-UnetBnbModelLoader ComfyUI/custom_n
|
||||
.\python_embeded\python.exe -s -m pip install -r .\ComfyUI\custom_nodes\ComfyUI-UnetBnbModelLoader\requirements.txt
|
||||
```
|
||||
|
||||
Because this plugin relies on bitsandbytes, we are unable to support macOS and AMD GPUs.
|
||||
Because this plugin relies on bitsandbytes, we are unable to support macOS.
|
||||
|
||||
## Usage
|
||||
|
||||
|
||||
@@ -55,7 +55,8 @@ class LazyLayer(torch.nn.Module):
|
||||
data=weight_data, quantized_stats=quant_state_dict, device=device
|
||||
)
|
||||
self.weight = bnb_param
|
||||
|
||||
self._bnb_quant_state_dict = quant_state_dict
|
||||
|
||||
for k in bnb_state_dict.keys():
|
||||
state_dict.pop(k)
|
||||
if k in unexpected_keys: unexpected_keys.remove(k)
|
||||
@@ -94,9 +95,10 @@ class LazyOps(comfy.ops.manual_cast):
|
||||
if getattr(self, "is_bnb_quantized", lambda : False)():
|
||||
if not patches_for_this_layer:
|
||||
bias = self.bias.to(device=x.device, dtype=x.dtype) if self.bias is not None else None
|
||||
return bnb.matmul_4bit(
|
||||
x, self.weight.t(), bias=bias, quant_state=getattr(self.weight, "quant_state", None)
|
||||
).to(x.dtype)
|
||||
qs = getattr(self.weight, "quant_state", None)
|
||||
if qs is None and hasattr(self, "_bnb_quant_state_dict"):
|
||||
qs = bnb.functional.QuantState.from_dict(self._bnb_quant_state_dict, device=x.device)
|
||||
return bnb.matmul_4bit(x, self.weight.t(), bias=bias, quant_state=qs).to(x.dtype)
|
||||
|
||||
try:
|
||||
base_w = self.weight.to(x.device)
|
||||
@@ -113,9 +115,10 @@ class LazyOps(comfy.ops.manual_cast):
|
||||
|
||||
if weight_final_fp32 is None:
|
||||
bias = self.bias.to(device=x.device, dtype=x.dtype) if self.bias is not None else None
|
||||
return bnb.matmul_4bit(
|
||||
x, self.weight.t(), bias=bias, quant_state=getattr(self.weight, "quant_state", None)
|
||||
).to(x.dtype)
|
||||
qs = getattr(self.weight, "quant_state", None)
|
||||
if qs is None and hasattr(self, "_bnb_quant_state_dict"):
|
||||
qs = bnb.functional.QuantState.from_dict(self._bnb_quant_state_dict, device=x.device)
|
||||
return bnb.matmul_4bit(x, self.weight.t(), bias=bias, quant_state=qs).to(x.dtype)
|
||||
|
||||
weight_final = comfy.float.stochastic_rounding(weight_final_fp32, x.dtype)
|
||||
bias = self.bias.to(device=x.device, dtype=x.dtype) if self.bias is not None else None
|
||||
|
||||
Reference in New Issue
Block a user