Fix for some fp8_scaled models
This commit is contained in:
+1
-1
@@ -13,7 +13,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
||||
module_prefix = prefix + name + "."
|
||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights)
|
||||
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix and "face" not in module_prefix:
|
||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix:
|
||||
in_features = state_dict[module_prefix + "weight"].shape[1]
|
||||
out_features = state_dict[module_prefix + "weight"].shape[0]
|
||||
if scale_weights is not None:
|
||||
|
||||
@@ -861,14 +861,18 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
# GGUF: skip GGUFParameter params
|
||||
if gguf and isinstance(param, GGUFParameter):
|
||||
continue
|
||||
|
||||
value=sd[name.replace("_orig_mod.", "")]
|
||||
|
||||
key = name.replace("_orig_mod.", "")
|
||||
value=sd[key]
|
||||
|
||||
if gguf:
|
||||
dtype_to_use = torch.float32 if "patch_embedding" in name or "motion_encoder" in name else base_dtype
|
||||
else:
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else weight_dtype
|
||||
dtype_to_use = weight_dtype if value.dtype == weight_dtype else dtype_to_use
|
||||
scale_key = key.replace(".weight", ".scale_weight")
|
||||
if scale_key in sd:
|
||||
dtype_to_use = value.dtype
|
||||
if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name:
|
||||
dtype_to_use = base_dtype
|
||||
if "patch_embedding" in name or "motion_encoder" in name:
|
||||
|
||||
Reference in New Issue
Block a user