GGUF fix
This commit is contained in:
+1
-1
@@ -36,6 +36,7 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
|
||||
if (
|
||||
isinstance(module, nn.Linear)
|
||||
and not isinstance(module, GGUFLinear)
|
||||
and _should_convert_to_gguf(state_dict, module_prefix)
|
||||
and name not in modules_to_not_convert
|
||||
):
|
||||
@@ -54,7 +55,6 @@ def _replace_with_gguf_linear(model, compute_dtype, state_dict, prefix="", modul
|
||||
model._modules[name].source_cls = type(module)
|
||||
# Force requires_grad to False to avoid unexpected errors
|
||||
model._modules[name].requires_grad_(False)
|
||||
|
||||
return model
|
||||
|
||||
def set_lora_params_gguf(module, patches, module_prefix=""):
|
||||
|
||||
@@ -755,6 +755,8 @@ def load_weights(transformer, sd, weight_dtype, base_dtype, transformer_load_dev
|
||||
pbar.update(100)
|
||||
#for name, param in transformer.named_parameters():
|
||||
# print(name, param.dtype, param.device, param.shape)
|
||||
#for name, param in transformer.blocks[0].motion_attn.named_parameters():
|
||||
# print(name, param.data)
|
||||
pbar.update_absolute(param_count)
|
||||
pbar.update_absolute(0)
|
||||
|
||||
@@ -782,12 +784,12 @@ def load_weights_gguf(transformer, reader, sd, base_dtype, transformer_load_devi
|
||||
continue
|
||||
#print(name, param.dtype, param.device, param.shape)
|
||||
if isinstance(param, GGUFParameter):
|
||||
dtype_to_use = torch.uint8
|
||||
continue
|
||||
elif "patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
else:
|
||||
dtype_to_use = base_dtype
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
set_module_tensor_to_device(patcher.model.diffusion_model, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name.replace("_orig_mod.", "")])
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
@@ -1142,6 +1144,7 @@ class WanVideoModelLoader:
|
||||
"add_ref_conv": True if "ref_conv.weight" in sd else False,
|
||||
"in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None,
|
||||
"add_control_adapter": True if "control_adapter.conv.weight" in sd else False,
|
||||
"use_motion_attn": True if "blocks.0.motion_attn.k.weight" in sd else False
|
||||
}
|
||||
|
||||
with init_empty_weights():
|
||||
|
||||
Reference in New Issue
Block a user