feat: Add robust model management, auto-batching, and critical bug fixes

This commit is contained in:
if-ai
2025-09-05 21:46:37 +01:00
parent da0353e230
commit df3029ca21
5 changed files with 263 additions and 75 deletions
+4
View File
@@ -90,3 +90,7 @@ venv.bak/
*.njsproj
*.sln
*.swp
# AI tool-specific files
.claude/
.serena/
+30
View File
@@ -88,6 +88,13 @@ This node takes the main model from the loader and then loads all the smaller, a
This is an optional but highly recommended performance-enhancing node. It uses `torch.compile` to optimize the model's code for your specific hardware.
- **Note**: The very first time you run a workflow with this node, it will take a minute or two to perform the compilation. However, every subsequent run will be significantly faster (often 20-30%).
- **`compile_mode`**: This controls the trade-off between compilation time and the amount of performance gain.
- `default`: The best balance. It provides a good speedup with a reasonable initial compile time.
- `reduce-overhead`: Compiles more slowly but can reduce the overhead of running the model, which might be faster for very small audio generations.
- `max-autotune`: Takes the longest to compile initially, but it tries many different optimizations to find the absolute fastest option for your specific hardware.
- **`backend`**: This is an advanced setting that changes the underlying compiler used by PyTorch. For most users, the default `inductor` is the best choice.
#### 4. HunyuanVideo-Foley Generator (Advanced)
This is the main workhorse node where the audio generation happens.
@@ -111,6 +118,29 @@ The two memory-related checkboxes on the Generator node are crucial for managing
- **Important Distinction:** This process is smart. It will **only** unload the model if it was loaded by the generator node itself (the simple workflow). If the model was passed in from the `HunyuanVideoFoleyModelLoader` (the advanced workflow), it will **not** unload it, respecting the fact that you may want to reuse the pre-loaded model for another generation.
- **Use this when:** You are finished with audio generation and want to free up as much memory as possible for completely different tasks.
### Performance Tuning & VRAM Usage
The most memory-intensive part of the process is visual feature extraction. We've implemented batched processing to prevent out-of-memory errors with longer videos or on GPUs with less VRAM. You can control this with two settings on the **Generator (Advanced)** node:
- **`feature_extraction_batch_size`**: This determines how many video frames are processed by the feature extractor models at once.
- **Lower values** significantly reduce peak VRAM usage at the cost of slightly slower processing.
- **Higher values** speed up processing but require more VRAM.
- **`enable_profiling`**: If you check this box, the node will print detailed performance timings and peak VRAM usage for the feature extraction step to the console. This is highly recommended for finding the optimal batch size for your specific hardware.
#### Recommended Batch Sizes
These are general starting points. The optimal value can vary based on your exact GPU, driver version, and other running processes.
| VRAM Tier | Video Resolution | Recommended Batch Size | Notes |
| :--- | :--- | :--- | :--- |
| **≤ 8 GB** | 480p | 4 - 8 | Start with 4. If successful, you can try increasing it. |
| | 720p | 2 - 4 | Start with 2. 720p videos are demanding on low VRAM cards. |
| **12-16 GB** | 480p | 16 - 32 | The default of 16 should work well. Can be increased for more speed. |
| | 720p | 8 - 16 | Start with 8 or 16. |
| **≥ 24 GB**| 480p | 32 - 64 | You can safely increase the batch size for maximum performance. |
| | 720p | 16 - 32 | A batch size of 32 should be easily achievable. |
## Usage
### Node Types
+76
View File
@@ -0,0 +1,76 @@
import os
import torch
import comfy.utils
from loguru import logger
import folder_paths
from huggingface_hub import hf_hub_download
# --- Constants ---
FOLEY_MODEL_NAMES = ["hunyuanvideo_foley.pth", "vae_128d_48k.pth", "synchformer_state_dict.pth"]
SIGLIP_MODEL_REPO = "google/siglip-base-patch16-512"
CLAP_MODEL_REPO = "laion/clap-htsat-unfused"
# --- Path Management ---
def get_model_dir(subfolder=""):
"""Returns the primary Foley models directory."""
return os.path.join(folder_paths.get_folder_paths("foley")[0], subfolder)
def get_full_model_path(model_name, subfolder=""):
"""Returns the full path for a given model name."""
return os.path.join(get_model_dir(subfolder), model_name)
# --- Core Functionality ---
def find_or_download(model_name, repo_id, subfolder="", subfolder_in_repo=""):
"""
Finds a model file, downloading it if it's not found in standard locations.
- Checks the main ComfyUI foley models directory first.
- Falls back to downloading from Hugging Face.
"""
local_path = get_full_model_path(model_name, subfolder)
if os.path.exists(local_path):
logger.info(f"Found local model: {local_path}")
return local_path
logger.warning(f"Could not find {model_name} locally. Attempting to download from {repo_id}...")
try:
downloaded_path = hf_hub_download(
repo_id=repo_id,
filename=model_name,
subfolder=subfolder_in_repo,
local_dir=get_model_dir(subfolder),
local_dir_use_symlinks=False
)
logger.info(f"Successfully downloaded model to: {downloaded_path}")
return downloaded_path
except Exception as e:
logger.error(f"Failed to download {model_name} from {repo_id}: {e}")
raise FileNotFoundError(f"Could not find or download {model_name}. Please check your connection or download it manually.")
def get_siglip_path():
"""Special handling for the SigLIP model which is a directory."""
return find_or_download_directory(repo_id=SIGLIP_MODEL_REPO, local_dir_name="siglip-base-patch16-512")
def get_clap_path():
"""Special handling for the CLAP model which is a directory."""
return find_or_download_directory(repo_id=CLAP_MODEL_REPO, local_dir_name="clap-htsat-unfused")
def find_or_download_directory(repo_id, local_dir_name):
"""
Finds a model directory, downloading it if it's not found.
This is for models like SigLIP that are not single files.
"""
local_path = get_model_dir(local_dir_name)
if os.path.exists(local_path) and os.listdir(local_path):
logger.info(f"Found local model directory: {local_path}")
return local_path
logger.warning(f"Could not find {local_dir_name} directory locally. Attempting to download from {repo_id}...")
# We can't use hf_hub_download for a whole directory in the same way,
# but the transformers library will handle this caching for us automatically
# when `from_pretrained` is called. We just need to return the repo_id.
# The actual "download" is implicit.
return repo_id
+60 -65
View File
@@ -54,6 +54,7 @@ try:
from .utils import feature_process_unified, extract_video_path, create_node_exit_values
from hunyuanvideo_foley.utils.media_utils import merge_audio_video
from .model_urls import get_model_url
from .model_management import find_or_download, get_siglip_path, get_clap_path, get_model_dir
except ImportError as e:
logger.error(f"Failed to import HunyuanVideo-Foley modules: {e}")
logger.error("Make sure the HunyuanVideo-Foley package is installed and accessible")
@@ -144,6 +145,19 @@ class HunyuanVideoFoleyNode:
"placeholder": "Prefix for output filename"
}),
# Memory optimization options
"feature_extraction_batch_size": ("INT", {
"default": 0, "min": 0, "max": 128, "step": 2,
"tooltip": "Frames to process at once. Set to 0 for auto-detection based on VRAM. Default is 16 for >16GB VRAM."
}),
"syncformer_batch_size": ("INT", {
"default": 8, "min": 1, "max": 64, "step": 1,
"tooltip": "(Advanced) Internal batch size for the Syncformer model. Lower this if you still get OOM errors with long videos."
}),
"enable_profiling": ("BOOLEAN", {
"default": False,
"display": "checkbox",
"tooltip": "Enable detailed performance and memory logging in the console."
}),
"memory_efficient": ("BOOLEAN", {
"default": False,
"display": "checkbox",
@@ -201,89 +215,54 @@ class HunyuanVideoFoleyNode:
return device
@classmethod
def download_models(cls, model_name: str, destination_dir: str) -> Tuple[bool, str]:
"""Download models from URLs if they don't exist."""
model_info = get_model_url(model_name)
if not model_info:
return False, f"Model '{model_name}' not found in configuration."
os.makedirs(destination_dir, exist_ok=True)
for model_file in model_info.get("models", []):
filename = model_file.get("filename")
url = model_file.get("url")
if not filename or not url:
logger.warning("Skipping a model file due to missing filename or URL.")
continue
destination_path = os.path.join(destination_dir, filename)
if os.path.exists(destination_path):
logger.info(f"'{filename}' already exists. Skipping download.")
continue
logger.info(f"Downloading '{filename}' from '{url}'...")
try:
# Use comfy.utils for resumable downloads with progress bar
comfy.utils.pb_download(url, destination_path)
except Exception as e:
return False, f"Failed to download '{filename}': {str(e)}"
return True, "All models downloaded or already exist."
@classmethod
def load_models(cls, model_path: str = "", config_path: str = "",
def load_models(cls, model_path: str = "", config_path: str = "",
memory_efficient: bool = False, cpu_offload: bool = False) -> Tuple[bool, str]:
"""Load models if not already loaded or if path changed"""
try:
# --- New Robust Model Pathing ---
# Use our model manager as a pre-flight check to ensure all models
# are downloaded and present before we proceed.
logger.info("Verifying local model integrity...")
find_or_download("hunyuanvideo_foley.pth", "Tencent-Hunyuan/HunyuanVideo-Foley", "hunyuanvideo-foley-xxl")
find_or_download("vae_128d_48k.pth", "Tencent-Hunyuan/HunyuanVideo-Foley", "hunyuanvideo-foley-xxl")
find_or_download("synchformer_state_dict.pth", "Tencent-Hunyuan/HunyuanVideo-Foley", "hunyuanvideo-foley-xxl")
get_siglip_path() # This ensures the Hugging Face cache is populated if needed
get_clap_path()
logger.info("All models found or downloaded successfully.")
# Set default paths if empty
if not model_path.strip():
foley_models_dir = folder_paths.get_folder_paths("foley")[0]
model_path = os.path.join(foley_models_dir, "hunyuanvideo-foley-xxl")
# --- Auto-Downloader Integration ---
if not os.path.exists(model_path):
logger.info(f"Model directory not found at '{model_path}'. Attempting to download...")
success, message = cls.download_models("hunyuanvideo-foley-xxl", model_path)
if not success:
return False, f"Auto-download failed: {message}"
# --- End of Integration ---
# Now, get the single directory path that the original library expects.
model_dir = get_model_dir("hunyuanvideo-foley-xxl")
if not config_path.strip():
current_dir = os.path.dirname(os.path.abspath(__file__))
config_path = os.path.join(current_dir, "configs", "hunyuanvideo-foley-xxl.yaml")
# Check if models are already loaded with the same path and memory mode
# Also check for "preloaded" which means models came from pipeline
if (cls._model_dict is not None and
cls._cfg is not None and
(cls._model_path == model_path or cls._model_path == "preloaded") and
# --- Caching Check (checks against the directory path string) ---
if (cls._model_dict is not None and
cls._cfg is not None and
(cls._model_path == model_dir or cls._model_path == "preloaded") and
cls._memory_efficient == memory_efficient):
return True, "Models already loaded"
# Setup device
cls._device = cls.setup_device("auto", 0)
logger.info(f"Loading models from: {model_path}")
logger.info(f"Loading models from directory: {model_dir}")
logger.info(f"Config: {config_path}")
# Load models
cls._model_dict, cls._cfg = load_model(model_path, config_path, cls._device)
# Do not set _model_path here if preloaded, to prevent stale state.
# Call the original library loader with the directory path it expects.
cls._model_dict, cls._cfg = load_model(model_dir, config_path, cls._device)
if cls._model_path != "preloaded":
cls._model_path = model_path
cls._model_path = model_dir
cls._memory_efficient = memory_efficient
logger.info("Models loaded successfully!")
return True, "Models loaded successfully!"
except Exception as e:
import traceback
error_msg = f"Failed to load models: {str(e)}"
logger.error(error_msg)
logger.error(traceback.format_exc())
cls._model_dict = None
cls._cfg = None
cls._device = None
@@ -379,9 +358,22 @@ class HunyuanVideoFoleyNode:
memory_efficient: bool = False,
cpu_offload: bool = False,
enabled: bool = True,
silent_audio: bool = True):
silent_audio: bool = True,
feature_extraction_batch_size: int = 16,
syncformer_batch_size: int = 8,
enable_profiling: bool = False):
"""Generate audio for the input video/images with the given text prompt"""
# --- Input Validation & Auto Batching ---
effective_batch_size = feature_extraction_batch_size
if effective_batch_size == 0:
from .utils import get_auto_batch_size
effective_batch_size = get_auto_batch_size()
if effective_batch_size < 1:
logger.warning(f"Batch size cannot be less than 1. Clamping value from {feature_extraction_batch_size} to 1.")
effective_batch_size = 1
if not enabled:
logger.info("HunyuanVideo-Foley node is disabled. Passing through inputs.")
return create_node_exit_values(
@@ -424,7 +416,10 @@ class HunyuanVideoFoleyNode:
negative_prompt=negative_prompt,
model_dict=self._model_dict,
cfg=self._cfg,
fps_hint=fps
fps_hint=fps,
batch_size=effective_batch_size,
sync_batch_size=syncformer_batch_size,
enable_profiling=enable_profiling
)
# --- State Correction and VRAM Management for Denoising ---
+93 -10
View File
@@ -12,12 +12,81 @@ import decord
from tqdm import tqdm
from PIL import Image
from einops import rearrange
import time
# We need to import the original library functions that our safe wrappers will call.
from hunyuanvideo_foley.utils.feature_utils import encode_text_feat, encode_video_with_siglip2, encode_video_with_sync
from hunyuanvideo_foley.utils.config_utils import AttributeDict
def _encode_video_with_siglip2_safely(pixel_values, model_dict):
"""
A wrapper to handle different versions of the transformers library for SigLIP2.
Expects a 4D tensor of shape (B, C, H, W).
"""
if hasattr(model_dict.siglip2_model, 'get_image_features'):
# Older transformers versions
return model_dict.siglip2_model.get_image_features(pixel_values=pixel_values)
else:
# Newer transformers versions
return model_dict.siglip2_model(pixel_values=pixel_values).image_embeds
def get_auto_batch_size():
"""
Automatically determines a safe batch size based on available GPU VRAM.
"""
if not torch.cuda.is_available():
logger.info("CUDA not available, returning default batch size of 4 for CPU.")
return 4 # A safe default for CPU
try:
total_vram_gb = torch.cuda.get_device_properties(0).total_memory / (1024**3)
logger.info(f"Detected {total_vram_gb:.2f} GB of total VRAM.")
if total_vram_gb <= 8:
batch_size = 4
elif total_vram_gb <= 12:
batch_size = 8
elif total_vram_gb <= 16:
batch_size = 16
else: # > 16GB
batch_size = 32
logger.info(f"Setting automatic batch size to {batch_size}.")
return batch_size
except Exception as e:
logger.warning(f"Could not determine VRAM, falling back to default batch size of 8. Error: {e}")
return 8 # Fallback
class SimpleProfiler:
"""A simple profiler for timing and CUDA memory logging."""
def __init__(self, name, enabled=True):
self.name = name
self.enabled = enabled
self.start_time = None
self.device = get_optimal_device()
def __enter__(self):
if not self.enabled: return self
self.start_time = time.time()
if self.device.type == 'cuda':
torch.cuda.reset_peak_memory_stats(self.device)
start_mem = torch.cuda.memory_allocated(self.device) / 1024**2
logger.info(f"[Profiler:{self.name}] Entering block. Start VRAM: {start_mem:.2f} MB")
return self
def __exit__(self, exc_type, exc_val, exc_tb):
if not self.enabled: return
elapsed_time = time.time() - self.start_time
if self.device.type == 'cuda':
peak_mem = torch.cuda.max_memory_allocated(self.device) / 1024**2
end_mem = torch.cuda.memory_allocated(self.device) / 1024**2
logger.info(f"[Profiler:{self.name}] Exiting block. Time: {elapsed_time:.4f}s. VRAM Usage (End/Peak): {end_mem:.2f} / {peak_mem:.2f} MB")
else:
logger.info(f"[Profiler:{self.name}] Exiting block. Time: {elapsed_time:.4f}s. (Running on CPU)")
def tensor_to_video(video_tensor: torch.Tensor, output_path: str, fps: int = 30) -> str:
"""
Convert a video tensor to a video file
@@ -146,7 +215,7 @@ def extract_video_path(video):
return None
@torch.inference_mode()
def _encode_visual_features_safely(frames_uint8, model_dict, fps_hint):
def _encode_visual_features_safely(frames_uint8, model_dict, fps_hint, batch_size=16, sync_batch_size=8, enable_profiling=False):
""" Memory-safe visual feature extraction from a list of pre-loaded frames. """
dev = model_dict.device
@@ -159,23 +228,37 @@ def _encode_visual_features_safely(frames_uint8, model_dict, fps_hint):
visual_features = {}
pil_list = [Image.fromarray(f).convert("RGB") for f in frames_uint8]
# --- BATCHED SigLIP2 ---
all_siglip_feats = []
try:
logger.info("Moving SigLIP2 to device for feature extraction...")
model_dict.siglip2_model.to(dev)
siglip_list = [model_dict.siglip2_preprocess(im) for im in pil_list]
clip_frames = torch.stack(siglip_list, dim=0).unsqueeze(0).to(dev)
visual_features['siglip2_feat'] = encode_video_with_siglip2(clip_frames, model_dict).to(dev)
with SimpleProfiler("SigLIP2 Batching", enabled=enable_profiling):
for i in tqdm(range(0, len(pil_list), batch_size), desc="Processing SigLIP2 batches"):
batch_pils = pil_list[i:i+batch_size]
siglip_list = [model_dict.siglip2_preprocess(im) for im in batch_pils]
clip_frames = torch.stack(siglip_list, dim=0).to(dev)
batch_feat = _encode_video_with_siglip2_safely(clip_frames, model_dict)
all_siglip_feats.append(batch_feat.cpu())
if all_siglip_feats:
# Concatenate along the batch dimension (dim=0)
final_feats = torch.cat(all_siglip_feats, dim=0)
# Reshape to what the pipeline expects: (1, total_frames, feature_dim)
visual_features['siglip2_feat'] = final_feats.unsqueeze(0).to(dev)
finally:
logger.info("Offloading SigLIP2 from device."); model_dict.siglip2_model.to("cpu")
if dev.type == 'cuda': torch.cuda.empty_cache()
# --- BATCHED Syncformer ---
try:
logger.info("Moving Syncformer to device for feature extraction...")
model_dict.syncformer_model.to(dev)
# Correctly preprocess frames for syncformer one by one
sync_list = [model_dict.syncformer_preprocess(im) for im in pil_list]
sync_frames = torch.stack(sync_list, dim=0).unsqueeze(0).to(dev)
visual_features['syncformer_feat'] = encode_video_with_sync(sync_frames, model_dict)
with SimpleProfiler("Syncformer Processing", enabled=enable_profiling):
# Syncformer needs all frames at once, but we can batch its internal processing.
sync_list = [model_dict.syncformer_preprocess(im) for im in pil_list]
sync_frames = torch.stack(sync_list, dim=0).unsqueeze(0).to(dev)
visual_features['syncformer_feat'] = encode_video_with_sync(sync_frames, model_dict, batch_size=sync_batch_size)
finally:
logger.info("Offloading Syncformer from device."); model_dict.syncformer_model.to("cpu")
if dev.type == 'cuda': torch.cuda.empty_cache()
@@ -183,7 +266,7 @@ def _encode_visual_features_safely(frames_uint8, model_dict, fps_hint):
audio_len_in_s = len(frames_uint8) / fps_hint
return AttributeDict(visual_features), audio_len_in_s
def feature_process_unified(video_input, image_input, prompt, model_dict, cfg, negative_prompt="", fps_hint=24.0, max_frames=450):
def feature_process_unified(video_input, image_input, prompt, model_dict, cfg, negative_prompt="", fps_hint=24.0, max_frames=450, batch_size=16, sync_batch_size=8, enable_profiling=False):
""" Unified, memory-safe feature processing for either a video path or an image tensor. """
frames_uint8 = None; fps = fps_hint
@@ -204,7 +287,7 @@ def feature_process_unified(video_input, image_input, prompt, model_dict, cfg, n
if frames_uint8 is None: raise ValueError("No valid video or image frames to process.")
visual_feats, audio_len_in_s = _encode_visual_features_safely(frames_uint8, model_dict, fps)
visual_feats, audio_len_in_s = _encode_visual_features_safely(frames_uint8, model_dict, fps, batch_size, sync_batch_size, enable_profiling)
# Use the provided negative prompt, or fall back to a default.
neg_prompt = negative_prompt if negative_prompt and negative_prompt.strip() else "noisy, harsh"