added loading bars
This commit is contained in:
@@ -261,10 +261,11 @@ class CodecLM:
|
||||
if self.duration <= self.max_duration:
|
||||
# generate by sampling from LM, simple case.
|
||||
with self.autocast:
|
||||
gen_tokens = self.lm.generate(texts=texts,
|
||||
descriptions=descriptions,
|
||||
audio_qt_embs=audio_qt_embs,
|
||||
max_gen_len=total_gen_len,
|
||||
gen_tokens = self.lm.generate(texts=texts,
|
||||
descriptions=descriptions,
|
||||
audio_qt_embs=audio_qt_embs,
|
||||
max_gen_len=total_gen_len,
|
||||
callback=_progress_callback,
|
||||
**self.generation_params)
|
||||
else:
|
||||
raise NotImplementedError(f"duration {self.duration} < max duration {self.max_duration}")
|
||||
|
||||
@@ -348,9 +348,10 @@ class LmModel(StreamingModule):
|
||||
top_k: int = 250,
|
||||
top_p: float = 0.0,
|
||||
cfg_coef: tp.Optional[float] = None,
|
||||
check: bool = False,
|
||||
check: bool = False,
|
||||
record_tokens: bool = True,
|
||||
record_window: int = 150
|
||||
record_window: int = 150,
|
||||
callback: tp.Optional[tp.Callable[[int, int], None]] = None
|
||||
) -> torch.Tensor:
|
||||
"""Generate tokens sampling from the model given a prompt or unconditionally. Generation can
|
||||
be perform in a greedy fashion or using sampling with top K and top P strategies.
|
||||
@@ -452,6 +453,12 @@ class LmModel(StreamingModule):
|
||||
gen_sequence = gen_sequence[..., :offset+1]
|
||||
break
|
||||
prev_offset = offset
|
||||
|
||||
# Call progress callback if provided
|
||||
if callback is not None:
|
||||
current_step = offset - start_offset_sequence + 1
|
||||
total_steps = gen_sequence_len - start_offset_sequence
|
||||
callback(current_step, total_steps)
|
||||
|
||||
# ensure sequence has been entirely filled
|
||||
assert not (gen_sequence == unknown_token).any()
|
||||
|
||||
@@ -8,6 +8,8 @@ import os
|
||||
from typing import Tuple
|
||||
import importlib.util
|
||||
|
||||
from comfy.utils import ProgressBar
|
||||
|
||||
# Get the package root directory
|
||||
_PACKAGE_ROOT = os.path.dirname(os.path.dirname(__file__))
|
||||
|
||||
@@ -105,11 +107,18 @@ class FL_SongGen_ModelLoader:
|
||||
print(f"{'='*60}\n")
|
||||
|
||||
try:
|
||||
# 4 steps for full model loading
|
||||
pbar = ProgressBar(4)
|
||||
|
||||
def progress_callback(current, total):
|
||||
pbar.update_absolute(current)
|
||||
|
||||
model_info = load_model(
|
||||
variant=model_variant,
|
||||
low_mem=low_mem,
|
||||
use_flash_attn=False,
|
||||
force_reload=force_reload
|
||||
force_reload=force_reload,
|
||||
progress_callback=progress_callback
|
||||
)
|
||||
print(f"[FL SongGen] Model loaded successfully!")
|
||||
return (model_info,)
|
||||
|
||||
@@ -307,7 +307,8 @@ def load_model(
|
||||
low_mem: bool = False,
|
||||
use_flash_attn: bool = False,
|
||||
force_reload: bool = False,
|
||||
device: Optional[str] = None
|
||||
device: Optional[str] = None,
|
||||
progress_callback: Optional[callable] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Load SongGeneration model.
|
||||
@@ -318,6 +319,7 @@ def load_model(
|
||||
use_flash_attn: Use Flash Attention 2
|
||||
force_reload: Force reload even if cached
|
||||
device: Device to load model on (default: auto-detect)
|
||||
progress_callback: Optional callback(current, total) for progress updates
|
||||
|
||||
Returns:
|
||||
Dict containing model components and configuration
|
||||
@@ -431,9 +433,11 @@ def load_model(
|
||||
model_info["ckpt_path"] = str(ckpt_path)
|
||||
model_info["loaded"] = False
|
||||
print("[FL SongGen] Low memory mode: model will be loaded on-demand")
|
||||
if progress_callback:
|
||||
progress_callback(1, 1) # Complete immediately for low_mem mode
|
||||
else:
|
||||
# Normal mode: load everything now
|
||||
model_info = _load_full_model(model_info, cfg, ckpt_path, device)
|
||||
model_info = _load_full_model(model_info, cfg, ckpt_path, device, progress_callback)
|
||||
|
||||
# Load auto prompts if available
|
||||
auto_prompts_path = get_auto_prompts_path()
|
||||
@@ -455,11 +459,22 @@ def _load_full_model(
|
||||
model_info: dict,
|
||||
cfg: OmegaConf,
|
||||
ckpt_path: Path,
|
||||
device: str
|
||||
device: str,
|
||||
progress_callback: Optional[callable] = None
|
||||
) -> dict:
|
||||
"""Load full model (non-low-memory mode)."""
|
||||
from codeclm.models import builders, CodecLM
|
||||
|
||||
# Total steps: audio_tokenizer, separate_tokenizer, language_model, create_wrapper
|
||||
total_steps = 4
|
||||
current_step = 0
|
||||
|
||||
def update_progress():
|
||||
nonlocal current_step
|
||||
current_step += 1
|
||||
if progress_callback:
|
||||
progress_callback(current_step, total_steps)
|
||||
|
||||
# Load audio tokenizer for prompt encoding
|
||||
print("[FL SongGen] Loading audio tokenizer...")
|
||||
audio_tokenizer = builders.get_audio_tokenizer_model(cfg.audio_tokenizer_checkpoint, cfg)
|
||||
@@ -468,6 +483,7 @@ def _load_full_model(
|
||||
if device == "cuda":
|
||||
audio_tokenizer = audio_tokenizer.cuda()
|
||||
model_info["audio_tokenizer"] = audio_tokenizer
|
||||
update_progress()
|
||||
|
||||
# Load separate tokenizer for vocal/bgm encoding
|
||||
print("[FL SongGen] Loading separate tokenizer...")
|
||||
@@ -480,6 +496,7 @@ def _load_full_model(
|
||||
else:
|
||||
separate_tokenizer = None
|
||||
model_info["separate_tokenizer"] = separate_tokenizer
|
||||
update_progress()
|
||||
|
||||
# Load LM
|
||||
print("[FL SongGen] Loading language model...")
|
||||
@@ -540,8 +557,10 @@ def _load_full_model(
|
||||
|
||||
if device == "cuda":
|
||||
audiolm = audiolm.cuda().to(torch.float16)
|
||||
update_progress()
|
||||
|
||||
# Create CodecLM wrapper
|
||||
print("[FL SongGen] Creating model wrapper...")
|
||||
model = CodecLM(
|
||||
name=model_info["variant"],
|
||||
lm=audiolm,
|
||||
@@ -553,6 +572,7 @@ def _load_full_model(
|
||||
model_info["model"] = model
|
||||
model_info["audiolm"] = audiolm
|
||||
model_info["loaded"] = True
|
||||
update_progress()
|
||||
|
||||
# Cleanup checkpoint to save memory
|
||||
del checkpoint
|
||||
|
||||
+1
-1
@@ -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.0.4"
|
||||
version = "1.0.5"
|
||||
license = "Apache-2.0"
|
||||
dependencies = [
|
||||
"torch>=2.0.0",
|
||||
|
||||
Reference in New Issue
Block a user