Fix quant_state None on AMD GPUs by caching quant_state_dict at load time

This commit is contained in:
0xDELUXA
2026-03-27 01:12:45 +02:00
parent 7013f8630b
commit 2af2e9d920
2 changed files with 11 additions and 8 deletions
+1 -1
View File
@@ -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
+10 -7
View File
@@ -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