diff --git a/nodes.py b/nodes.py index 765de2d..b2a7d0b 100644 --- a/nodes.py +++ b/nodes.py @@ -36,6 +36,18 @@ offload_device = mm.unet_offload_device() VAE_STRIDE = (4, 8, 8) PATCH_SIZE = (1, 2, 2) +try: + from .gguf.gguf import GGUFParameter +except: + pass + +class MetaParameter(torch.nn.Parameter): + def __new__(cls, dtype, quant_type=None): + data = torch.empty(0, dtype=dtype) + self = torch.nn.Parameter(data, requires_grad=False) + self.quant_type = quant_type + return self + def offload_transformer(transformer): for block in transformer.blocks: block.kv_cache = None @@ -55,6 +67,9 @@ def offload_transformer(transformer): if param.data.is_floating_point(): meta_param = torch.nn.Parameter(torch.empty_like(param.data, device='meta'), requires_grad=False) setattr(module, attr_name, meta_param) + elif isinstance(param.data, GGUFParameter): + quant_type = getattr(param, 'quant_type', None) + setattr(module, attr_name, MetaParameter(param.data.dtype, quant_type)) else: pass else: diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 57803b0..8024749 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -777,7 +777,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, weights = torch.from_numpy(tensor.data.copy()).to(load_device) sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights sd.update(unianimate_sd) - del unianimate_sd + del all_tensors, unianimate_sd if not getattr(transformer, "gguf_patched", False): transformer = _replace_with_gguf_linear(