add progress status

This commit is contained in:
NumZ
2025-06-30 21:54:25 +02:00
parent ad45b9821d
commit e803398800
2 changed files with 43 additions and 42 deletions
+16 -38
View File
@@ -169,7 +169,7 @@ def cut_videos(videos):
return result
def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_size=90, preserve_vram=False, temporal_overlap=0, debug=False):
def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_size=90, preserve_vram=False, temporal_overlap=0, debug=False, progress_callback=None):
"""
Main generation loop with context-aware temporal processing
@@ -182,6 +182,7 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
batch_size (int): Batch size for processing
preserve_vram (str/bool): VRAM preservation mode
temporal_overlap (int): Frames for temporal continuity
progress_callback (callable): Optional callback for progress reporting
Returns:
torch.Tensor: Generated video frames
@@ -192,9 +193,11 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
- Memory-optimized batch processing
- Advanced video transformation pipeline
- Intelligent VRAM management throughout process
- Real-time progress reporting
"""
device = "cuda" if torch.cuda.is_available() else "cpu"
# Adaptive model dtype detection for maximum performance
model_dtype = None
try:
@@ -277,6 +280,9 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
step = batch_size
temporal_overlap = 0
# Calculate total batches for progress reporting
total_batches = len(range(0, len(images), step))
# Move images to CPU for memory efficiency
#t = time.time()
#images = images.to("cpu")
@@ -284,7 +290,7 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
try:
# Main processing loop with context awareness
for batch_idx in range(0, len(images), step):
for batch_count, batch_idx in enumerate(range(0, len(images), step)):
# Calculate batch indices with overlap
comfy.model_management.throw_exception_if_processing_interrupted()
if batch_idx == 0:
@@ -304,8 +310,13 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
tps_loop = time.time()
batch_number = (batch_idx // step + 1) if step > 0 else 1
current_frames = end_idx - start_idx
print(f"\n🎬 Batch {batch_number}: frames {start_idx}-{end_idx-1}")
# Progress callback - batch start
if progress_callback:
progress_callback(batch_count, total_batches, current_frames, "Processing batch...")
# Process current batch
video = images[start_idx:end_idx]
if debug:
@@ -360,18 +371,6 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
# Post-process samples
sample = samples[0]
del samples
# 🔧 DIAGNOSTIC: Vérifier les valeurs après generation_step
'''
if debug:
print(f"🔍 DIAGNOSTIC - Après generation_step:")
print(f" 📊 Sample shape: {sample.shape}")
print(f" 📊 Sample dtype: {sample.dtype}")
print(f" 📊 Sample device: {sample.device}")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
print(f" 📊 Sample std: {sample.std():.6f}")
print(f" 📊 Non-zero count: {(sample != 0).sum()}/{sample.numel()}")
'''
#del samples
if ori_lengths[0] < sample.shape[0]:
sample = sample[:ori_lengths[0]]
@@ -390,33 +389,10 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
#del transformed_video
sample = wavelet_reconstruction(sample, input_video[0][:sample.size(0)])
del input_video
'''
if debug:
# 🔧 DIAGNOSTIC: Vérifier les valeurs après wavelet_reconstruction
print(f"🔍 DIAGNOSTIC - Après wavelet_reconstruction:")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
print(f" 📊 Non-zero count: {(sample != 0).sum()}/{sample.numel()}")
'''
#del input_video
# Convert to final image format
sample = optimized_sample_to_image_format(sample)
'''
if debug:
# 🔧 DIAGNOSTIC: Vérifier les valeurs après optimized_sample_to_image_format
print(f"🔍 DIAGNOSTIC - Après optimized_sample_to_image_format:")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
'''
sample = sample.clip(-1, 1).mul_(0.5).add_(0.5)
'''
if debug:
# 🔧 DIAGNOSTIC: Vérifier les valeurs finales
print(f"🔍 DIAGNOSTIC - Valeurs finales:")
print(f" 📊 Sample range: [{sample.min():.6f}, {sample.max():.6f}]")
print(f" 📊 Sample mean: {sample.mean():.6f}")
print(f" 🎯 Est-ce que l'image est noire? {sample.max() < 0.01}")
'''
sample_cpu = sample.to("cpu")
del sample
batch_samples.append(sample_cpu)
@@ -441,6 +417,8 @@ def generation_loop(runner, images, cfg_scale=1.0, seed=666, res_w=720, batch_si
#del text_pos_embeds, text_neg_embeds
#clear_vram_cache()
for i in range(len(batch_samples)):
batch_samples[i] = batch_samples[i].to(device)
# Concatenate all batch results
+27 -4
View File
@@ -13,6 +13,10 @@ from src.utils.downloads import download_weight, get_base_cache_dir
from src.core.model_manager import configure_runner
from src.core.generation import generation_loop
from src.optimization.memory_manager import clear_rope_lru_caches, fast_model_cleanup
# Import ComfyUI progress reporting
from server import PromptServer
script_directory = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
class SeedVR2:
@@ -24,6 +28,7 @@ class SeedVR2:
- Adaptive VRAM management
- Advanced dtype compatibility
- Optimized inference pipeline
- Real-time progress reporting
"""
def __init__(self):
@@ -86,10 +91,11 @@ class SeedVR2:
def execute(self, images: torch.Tensor, model: str, seed: int, new_resolution: int,
batch_size: int, preserve_vram: bool) -> Tuple[torch.Tensor]:
"""Execute SeedVR2 video upscaling"""
"""Execute SeedVR2 video upscaling with progress reporting"""
temporal_overlap = 0
print(f"🔄 Preparing model: {model}")
download_weight(model)
debug = False
cfg_scale = 1.0
@@ -160,8 +166,10 @@ class SeedVR2:
def _internal_execute(self, images, model, seed, new_resolution, cfg_scale, batch_size, preserve_vram, temporal_overlap, debug):
"""Internal execution logic"""
"""Internal execution logic with progress tracking"""
total_start_time = time.time()
# Configure runner
if debug:
print("🔄 Configuring inference runner...")
@@ -170,20 +178,35 @@ class SeedVR2:
if debug:
print(f"🔄 Runner configuration time: {time.time() - runner_start:.2f}s")
if debug:
print("🚀 Starting video upscaling generation...")
# Execute generation
# Execute generation with progress callback
sample = generation_loop(
self.runner, images, cfg_scale, seed, new_resolution,
batch_size, preserve_vram, temporal_overlap, debug
batch_size, preserve_vram, temporal_overlap, debug,
progress_callback=self._progress_callback
)
print(f"✅ Video upscaling completed successfully!")
# Cleanup
print(f"🔄 Total execution time: {time.time() - total_start_time:.2f}s")
self.cleanup(force_ram_cleanup=True)
return (sample,)
def _progress_callback(self, batch_idx, total_batches, current_batch_frames, message=""):
"""Progress callback for generation loop"""
# Send numerical progress
progress_value = int((batch_idx / total_batches) * 100)
progress_data = {
"value": progress_value,
"max": 100,
"node": "seedvr2_node"
}
PromptServer.instance.send_sync("progress", progress_data, None)
def __del__(self):
"""Destructor"""