From 21e608353699e8e0108c0d64a24e18b2dfa5d8b5 Mon Sep 17 00:00:00 2001 From: Fill Date: Thu, 1 Jan 2026 13:23:04 -0800 Subject: [PATCH] Fix low-memory mode shape mismatch for embedding layers MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Add _resize_embeddings_for_checkpoint() helper method - Apply embedding resize logic to low-mem mode (fixes #4) - Handles tokenizer version differences between checkpoint and model - Bump version to 1.1.2 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 --- fl_utils/songgen_wrapper.py | 63 +++++++++++++++++++++++++++++++++++++ pyproject.toml | 2 +- 2 files changed, 64 insertions(+), 1 deletion(-) diff --git a/fl_utils/songgen_wrapper.py b/fl_utils/songgen_wrapper.py index f74ba5c..22b05aa 100644 --- a/fl_utils/songgen_wrapper.py +++ b/fl_utils/songgen_wrapper.py @@ -337,6 +337,10 @@ class SongGenWrapper: for k, v in checkpoint.items() if k.startswith('audiolm') } + + # Resize embedding layers to match checkpoint (fixes tokenizer version mismatch) + audiolm = self._resize_embeddings_for_checkpoint(audiolm, audiolm_state_dict) + audiolm.load_state_dict(audiolm_state_dict, strict=False) audiolm = audiolm.eval().cuda().to(torch.float16) @@ -510,6 +514,65 @@ class SongGenWrapper: # No prompt return None, None, None, True + def _resize_embeddings_for_checkpoint(self, model: torch.nn.Module, state_dict: dict) -> torch.nn.Module: + """ + Resize embedding layers in the model to match checkpoint dimensions. + This fixes tokenizer version mismatches where vocab sizes differ. + + Args: + model: The model to resize embeddings in + state_dict: The checkpoint state dict to match + + Returns: + The model with resized embeddings + """ + import torch.nn as nn + + def get_nested_attr(obj, attr_path): + """Get nested attribute from object using dot-separated path.""" + parts = attr_path.split('.') + for part in parts: + if hasattr(obj, part): + obj = getattr(obj, part) + elif hasattr(obj, '_modules') and part in obj._modules: + obj = obj._modules[part] + else: + return None + return obj + + def set_nested_attr(obj, attr_path, value): + """Set nested attribute on object using dot-separated path.""" + parts = attr_path.split('.') + for part in parts[:-1]: + if hasattr(obj, part): + obj = getattr(obj, part) + elif hasattr(obj, '_modules') and part in obj._modules: + obj = obj._modules[part] + else: + return False + setattr(obj, parts[-1], value) + return True + + for key, ckpt_tensor in state_dict.items(): + if 'output_proj.weight' in key: + # Get the path to the parent module (remove .weight) + module_path = key.rsplit('.', 1)[0] + current_module = get_nested_attr(model, module_path) + + if current_module is not None and hasattr(current_module, 'weight'): + current_size = current_module.weight.shape[0] + checkpoint_size = ckpt_tensor.shape[0] + + if current_size != checkpoint_size: + print(f"[FL SongGen LowMem] Resizing embedding {module_path}: {current_size} -> {checkpoint_size}") + # Create new embedding with checkpoint size + embed_dim = ckpt_tensor.shape[1] + padding_idx = getattr(current_module, 'padding_idx', None) + new_embedding = nn.Embedding(checkpoint_size, embed_dim, padding_idx=padding_idx) + set_nested_attr(model, module_path, new_embedding) + + return model + def _separate_audio(self, audio: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """ Separate audio into full, vocal, and bgm using Demucs. diff --git a/pyproject.toml b/pyproject.toml index 8a46e41..93b452b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,7 +1,7 @@ [project] name = "comfyui_fl-songgen" description = "FL Song Gen - AI-powered song generation nodes for ComfyUI. Generate complete songs with vocals and instrumentals from lyrics using Tencent's SongGeneration (LeVo) model. Features style transfer, auto style presets, dual-track output, and up to 4m30s song generation." -version = "1.1.1" +version = "1.1.2" license = "Apache-2.0" dependencies = [ "torch>=2.0.0",