work now also with the big model

This commit is contained in:
Shmuel Ronen
2025-07-14 18:34:11 +03:00
committed by GitHub
parent 354e919a02
commit dc66416fc8
+58 -1
View File
@@ -254,6 +254,8 @@ class ThinkSoundModelLoader:
FUNCTION = "load_model"
CATEGORY = "ThinkSound"
# Replace the load_model method in ThinkSoundModelLoader class (around line 285)
def load_model(self, thinksound_model):
if not THINKSOUND_AVAILABLE:
raise ImportError("ThinkSound source code is not installed. Please download the ThinkSound repository from https://github.com/FunAudioLLM/ThinkSound and place it in the ComfyUI-ThinkSound folder.")
@@ -294,7 +296,62 @@ class ThinkSoundModelLoader:
# Load weights
thinksound_model_path = folder_paths.get_full_path_or_raise("thinksound", thinksound_model)
model_sd = load_torch_file(thinksound_model_path, device=offload_device)
model.load_state_dict(model_sd)
# 🔧 FIX: Handle different key formats in model checkpoints
def fix_state_dict_keys(state_dict):
"""Fix state dict keys for different ThinkSound model formats"""
new_state_dict = {}
# Check if we need to remove prefixes
sample_key = list(state_dict.keys())[0]
if sample_key.startswith('diffusion.'):
# Remove 'diffusion.' prefix from all keys
log.info("Removing 'diffusion.' prefix from model keys")
for key, value in state_dict.items():
if key.startswith('diffusion.'):
new_key = key[len('diffusion.'):]
new_state_dict[new_key] = value
else:
new_state_dict[key] = value
elif sample_key.startswith('model.'):
# Keys are already in correct format
new_state_dict = state_dict
else:
# Might need to add 'model.' prefix - check what the model expects
model_keys = set(model.state_dict().keys())
state_keys = set(state_dict.keys())
# If no keys match, try adding 'model.' prefix
if not model_keys.intersection(state_keys):
log.info("Adding 'model.' prefix to model keys")
for key, value in state_dict.items():
new_state_dict[f'model.{key}'] = value
else:
new_state_dict = state_dict
return new_state_dict
# Apply key fixing
model_sd = fix_state_dict_keys(model_sd)
# Load with strict=False to handle missing/extra keys gracefully
try:
model.load_state_dict(model_sd, strict=False)
log.info("✅ Model loaded successfully")
except RuntimeError as e:
log.error(f"❌ Model loading failed: {e}")
# Try loading only matching keys
model_keys = set(model.state_dict().keys())
checkpoint_keys = set(model_sd.keys())
log.info(f"Model expects {len(model_keys)} keys, checkpoint has {len(checkpoint_keys)} keys")
log.info(f"Matching keys: {len(model_keys.intersection(checkpoint_keys))}")
# Load only matching keys
filtered_sd = {k: v for k, v in model_sd.items() if k in model_keys}
model.load_state_dict(filtered_sd, strict=False)
log.warning("⚠️ Loaded model with partial weights")
model = model.eval().to(device=device, dtype=base_dtype)