add progress status
This commit is contained in:
+16
-38
@@ -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
|
||||
|
||||
@@ -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"""
|
||||
|
||||
Reference in New Issue
Block a user