From 01cbdf8bc319d0050fc0bce1e70dcc52baed744b Mon Sep 17 00:00:00 2001 From: Adrien Toupet Date: Wed, 22 Oct 2025 23:56:14 -0400 Subject: [PATCH] Optimize VAE defaults and standardize dtype pipeline for quality/performance VAE Changes: - Enable encode tiling by default (prevents noise artifacts at high resolution) - Increase tile size to 1024px (down from 512px) for optimal quality - Increase tile overlap to 128px for better blending Dtype Pipeline: - Hardcode compute_dtype to bfloat16 for consistent quality/performance/VRAM balance - Ensure all pipeline steps are using compute_dtype when relevant - Refactor code for improved performance and memory management --- inference_cli.py | 19 ++-- src/core/generation.py | 74 +++++-------- src/core/infer.py | 14 +-- src/core/model_manager.py | 41 ++++--- src/interfaces/vae_model_loader.py | 20 ++-- src/interfaces/video_upscaler.py | 5 +- src/models/dit_3b/attention.py | 10 +- src/models/dit_3b/modulation.py | 12 +- .../dit_3b/nablocks/attention/mmattn.py | 14 ++- src/models/dit_3b/normalization.py | 5 +- src/models/dit_7b/attention.py | 10 +- src/models/dit_7b/modulation.py | 18 +-- src/models/dit_7b/nablocks/mmsr_block.py | 7 +- src/models/dit_7b/normalization.py | 5 +- .../video_vae_v3/modules/attn_video_vae.py | 50 +++++---- src/optimization/compatibility.py | 103 +++++++++++------- 16 files changed, 224 insertions(+), 183 deletions(-) diff --git a/inference_cli.py b/inference_cli.py index d699d95..7788038 100644 --- a/inference_cli.py +++ b/inference_cli.py @@ -614,18 +614,19 @@ def parse_arguments() -> argparse.Namespace: help="Device to offload intermediate tensors between phases (default: cpu). " "Options: 'cpu', 'none'. Use 'cpu' to prevent VRAM accumulation for long videos (recommended), " "'none' to keep all tensors on GPU (faster but uses more VRAM)") - parser.add_argument("--vae_encode_tiling_enabled", action="store_true", - help="Enable VAE encode tiling for improved VRAM usage") - parser.add_argument("--vae_encode_tile_size", action=OneOrTwoValues, nargs='+', default=(512, 512), - help="VAE encode tile size (default: 512). Use single integer or two integers 'h w'. Only used if --vae_encode_tiling_enabled is set") + parser.add_argument("--disable_vae_encode_tiling", action="store_false", + dest="vae_encode_tiling_enabled", default=True, + help="Disable VAE encode tiling. By default, tiling is enabled to prevent noise artifacts at high resolution.") + parser.add_argument("--vae_encode_tile_size", action=OneOrTwoValues, nargs='+', default=(1024, 1024), + help="VAE encode tile size in pixels (default: 1024). Only used when encode tiling is enabled. Can be reduced for memory, but increasing above 1024 causes noise artifacts. Use single integer or two integers 'h w'.") parser.add_argument("--vae_encode_tile_overlap", action=OneOrTwoValues, nargs='+', default=(128, 128), - help="VAE encode tile overlap (default: 128). Use single integer or two integers 'h w'. Only used if --vae_encode_tiling_enabled is set") + help="VAE encode tile overlap in pixels (default: 128). Only used when encode tiling is enabled. Higher values improve blending at the cost of slower processing. Use single integer or two integers 'h w'.") parser.add_argument("--vae_decode_tiling_enabled", action="store_true", - help="Enable VAE decode tiling for improved VRAM usage") - parser.add_argument("--vae_decode_tile_size", action=OneOrTwoValues, nargs='+', default=(512, 512), - help="VAE decode tile size (default: 512). Use single integer or two integers 'h w'. Only used if --vae_decode_tiling_enabled is set") + help="Enable VAE decode tiling for VRAM reduction during decoding. Disabled by default.") + parser.add_argument("--vae_decode_tile_size", action=OneOrTwoValues, nargs='+', default=(1024, 1024), + help="VAE decode tile size in pixels (default: 1024). Only used when decode tiling is enabled. Adjust based on available VRAM. Use single integer or two integers 'h w'.") parser.add_argument("--vae_decode_tile_overlap", action=OneOrTwoValues, nargs='+', default=(128, 128), - help="VAE decode tile overlap (default: 128). Use single integer or two integers 'h w'. Only used if --vae_decode_tiling_enabled is set") + help="VAE decode tile overlap in pixels (default: 128). Only used when decode tiling is enabled. Higher values improve blending at the cost of slower processing. Use single integer or two integers 'h w'.") parser.add_argument("--attention_mode", type=str, default="sdpa", choices=["sdpa", "flash_attn"], help="Attention computation backend: 'sdpa' (default, always available) or 'flash_attn' (requires flash-attn package, faster)") diff --git a/src/core/generation.py b/src/core/generation.py index 10e0437..0e93482 100644 --- a/src/core/generation.py +++ b/src/core/generation.py @@ -15,7 +15,7 @@ Key Features: - Four-phase pipeline (encode-all → upscale-all → decode-all → postprocess-all) for efficiency - Native FP8 pipeline support for 2x speedup and 50% VRAM reduction - Temporal overlap support for smooth transitions between batches -- Adaptive dtype detection and optimal autocast configuration +- Adaptive dtype detection and configuration - Memory-efficient pre-allocated batch processing - Stream-based assembly eliminates memory spikes for long videos - Advanced video format handling (4n+1 constraint) @@ -256,19 +256,21 @@ def _ensure_precision_initialized( debug: Optional['Debug'] = None ) -> None: """ - Initialize compute_dtype and autocast_dtype based on actual model dtypes. + Log model dtypes for debugging. Compute dtype is hardcoded in context. - Lazily initializes compute_dtype and autocast_dtype by inspecting actual - model weights to determine optimal precision settings. Only checks models - that are materialized (not on meta device). Safe to call multiple times. + Since compute_dtype is hardcoded to bfloat16 in setup_generation_context(), + this function only logs model dtypes for informational purposes. Args: - ctx: Generation context dictionary to update with precision settings + ctx: Generation context dictionary (compute_dtype already set) runner: VideoDiffusionInfer instance with loaded models debug: Optional Debug instance for logging """ + if not debug: + return + try: - # Check which models are materialized (not on meta device) + # Get model dtypes for informational logging dit_dtype = None vae_dtype = None @@ -288,34 +290,19 @@ def _ensure_precision_initialized( except StopIteration: pass - # Need at least one materialized model - if dit_dtype is None and vae_dtype is None: - return + # Build precision info string + parts = [] + if dit_dtype is not None: + parts.append(f"DiT={dit_dtype}") + if vae_dtype is not None: + parts.append(f"VAE={vae_dtype}") + parts.append(f"compute={ctx['compute_dtype']}") - # Initialize compute dtype once - if ctx.get('compute_dtype') is None: - ctx['compute_dtype'] = torch.bfloat16 - ctx['autocast_dtype'] = torch.bfloat16 - - # Always log current state (what's materialized) - if debug: - parts = [] - if dit_dtype is not None: - parts.append(f"DiT={dit_dtype}") - if vae_dtype is not None: - parts.append(f"VAE={vae_dtype}") - parts.append(f"compute={ctx['compute_dtype']}") - parts.append(f"autocast={ctx['autocast_dtype']}") - - debug.log(f"Initialized precision: {', '.join(parts)}", category="precision") + if parts: + debug.log(f"Model precision: {', '.join(parts)}", category="precision") except Exception as e: - # Fallback to safe defaults - ctx['compute_dtype'] = torch.bfloat16 - ctx['autocast_dtype'] = torch.bfloat16 - - if debug: - debug.log(f"Could not detect model dtypes: {e}, falling back to BFloat16", level="WARNING", category="model", force=True) + debug.log(f"Could not log model dtypes: {e}", level="WARNING", category="precision", force=True) def setup_generation_context( @@ -347,11 +334,6 @@ def setup_generation_context( - offload device configurations - state containers (latents, samples, etc.) - ComfyUI integration hooks - - Note: - Precision settings (compute_dtype, autocast_dtype) are lazily initialized - when first needed via _ensure_precision_initialized() to detect actual - model dtypes and configure optimal computation settings. """ # Apply default devices if not specified default_device = "cpu" @@ -380,8 +362,8 @@ def setup_generation_context( 'dit_offload_device': dit_offload_device, 'vae_offload_device': vae_offload_device, 'tensor_offload_device': tensor_offload_device, - 'compute_dtype': None, - 'autocast_dtype': None, + 'compute_dtype': torch.bfloat16, # Hardcoded - gives the best compromise between memory & quality without artifacts + 'interrupt_fn': interrupt_fn, 'video_transform': None, 'text_embeds': None, 'all_transformed_videos': [], @@ -412,6 +394,7 @@ def setup_generation_context( f"LOCAL_RANK={os.environ['LOCAL_RANK']}", category="setup" ) + debug.log(f"Unified compute dtype: torch.bfloat16 across entire pipeline for maximum compatibility", category="precision") return ctx @@ -489,6 +472,7 @@ def prepare_runner( vae_model=vae_model, base_cache_dir=model_dir, debug=debug, + ctx=ctx, dit_cache=dit_cache, vae_cache=vae_cache, dit_id=dit_id, @@ -504,12 +488,6 @@ def prepare_runner( torch_compile_args_dit=torch_compile_args_dit, torch_compile_args_vae=torch_compile_args_vae ) - - # Store device configuration on runner for submodule access (e.g., BlockSwap, Cleanup) - runner._dit_device = ctx['dit_device'] - runner._vae_device = ctx['vae_device'] - runner._dit_offload_device = ctx['dit_offload_device'] - runner._vae_offload_device = ctx['vae_offload_device'] return runner, cache_context @@ -828,7 +806,7 @@ def encode_all_batches( tensor_name=f"latent_{encode_idx+1}", dtype=ctx['compute_dtype'], debug=debug, - reason="storing encoded latents for upscaling (VAE dtype → compute dtype)", + reason="storing encoded latents for upscaling", indent_level=1 ) else: @@ -1036,7 +1014,7 @@ def upscale_all_batches( # Run inference debug.start_timer(f"dit_inference_{upscale_idx+1}") with torch.no_grad(): - with torch.autocast(str(ctx['dit_device']), ctx['autocast_dtype'], enabled=True): + with torch.autocast(str(ctx['dit_device']), ctx['compute_dtype'], enabled=True): upscaled_latents = runner.inference( noises=noises, conditions=conditions, @@ -1226,7 +1204,7 @@ def decode_all_batches( tensor_name=f"sample_{decode_idx+1}", dtype=ctx['compute_dtype'], debug=debug, - reason="storing decoded samples for post-processing (VAE dtype → compute dtype)", + reason="storing decoded samples for post-processing", indent_level=1 ) else: diff --git a/src/core/infer.py b/src/core/infer.py index 589be32..93a6384 100644 --- a/src/core/infer.py +++ b/src/core/infer.py @@ -168,10 +168,6 @@ class VideoDiffusionInfer(): else: batches = [sample.unsqueeze(0) for sample in samples] - use_encode_tiling = self.encode_tiled - if use_encode_tiling: - self.debug.log(f"Using VAE tiled encoding (Tile: {self.encode_tile_size}, Overlap: {self.encode_tile_overlap})", category="vae", force=True, indent_level=1) - # VAE process by each group. for sample in batches: sample = sample.to(device, dtype) @@ -180,11 +176,11 @@ class VideoDiffusionInfer(): sample = self.vae.preprocess(sample) if use_sample: - latent = self.vae.encode(sample, tiled=use_encode_tiling, tile_size=self.encode_tile_size, + latent = self.vae.encode(sample, tiled=self.encode_tiled, tile_size=self.encode_tile_size, tile_overlap=self.encode_tile_overlap).latent else: # Deterministic vae encode, only used for i2v inference (optionally) - latent = self.vae.encode(sample, tiled=use_encode_tiling, tile_size=self.encode_tile_size, + latent = self.vae.encode(sample, tiled=self.encode_tiled, tile_size=self.encode_tile_size, tile_overlap=self.encode_tile_overlap).posterior.mode().squeeze(2) latent = latent.unsqueeze(2) if latent.ndim == 4 else latent @@ -231,10 +227,6 @@ class VideoDiffusionInfer(): else: latents = [latent.unsqueeze(0) for latent in latents] - use_decode_tiling = self.decode_tiled - if use_decode_tiling: - self.debug.log(f"Using VAE tiled decoding (Tile: {self.decode_tile_size}, Overlap: {self.decode_tile_overlap})", category="vae", force=True, indent_level=1) - self.debug.log(f"Latents shape: {latents[0].shape}", category="info", indent_level=1) for i, latent in enumerate(latents): @@ -245,7 +237,7 @@ class VideoDiffusionInfer(): sample = self.vae.decode( latent, - tiled=use_decode_tiling, tile_size=self.decode_tile_size, + tiled=self.decode_tiled, tile_size=self.decode_tile_size, tile_overlap=self.decode_tile_overlap).sample if hasattr(self.vae, "postprocess"): diff --git a/src/core/model_manager.py b/src/core/model_manager.py index 11a46ef..d64b849 100644 --- a/src/core/model_manager.py +++ b/src/core/model_manager.py @@ -509,6 +509,7 @@ def configure_runner( vae_model: str, base_cache_dir: str, debug: 'Debug', + ctx: Dict[str, Any], dit_cache: bool = False, vae_cache: bool = False, dit_id: Optional[int] = None, @@ -534,6 +535,7 @@ def configure_runner( vae_model: VAE model filename (e.g., "ema_vae_fp16.safetensors") base_cache_dir: Base directory containing model files debug: Debug instance for logging (required) + ctx: Generation context from setup_generation_context dit_cache: Whether to cache DiT model between runs vae_cache: Whether to cache VAE model between runs dit_id: Node instance ID for DiT model caching (required if dit_cache=True) @@ -580,7 +582,7 @@ def configure_runner( # Phase 3: Configure runner settings _configure_runner_settings( - runner, + runner, ctx, encode_tiled, encode_tile_size, encode_tile_overlap, decode_tiled, decode_tile_size, decode_tile_overlap, attention_mode, @@ -799,6 +801,7 @@ def _create_new_runner( def _configure_runner_settings( runner: VideoDiffusionInfer, + ctx: Dict[str, Any], encode_tiled: bool, encode_tile_size: Optional[Tuple[int, int]], encode_tile_overlap: Optional[Tuple[int, int]], @@ -821,6 +824,7 @@ def _configure_runner_settings( Args: runner: VideoDiffusionInfer instance to configure + ctx: Generation context from setup_generation_context encode_tiled: Enable tiled VAE encoding to reduce VRAM during encoding encode_tile_size: Tile dimensions (height, width) for encoding in pixels encode_tile_overlap: Overlap dimensions (height, width) between encoding tiles @@ -856,6 +860,13 @@ def _configure_runner_settings( 'decode_tile_overlap': decode_tile_overlap } + # Store device configuration on runner for submodule access (e.g., BlockSwap, Cleanup) + runner._dit_device = ctx['dit_device'] + runner._vae_device = ctx['vae_device'] + runner._dit_offload_device = ctx['dit_offload_device'] + runner._vae_offload_device = ctx['vae_offload_device'] + runner._compute_dtype = ctx['compute_dtype'] + runner.debug = debug @@ -1069,11 +1080,11 @@ def _setup_vae_model( runner.config.vae.model = OmegaConf.merge(runner.config.vae.model, vae_config) - if torch.mps.is_available(): - original_vae_dtype = runner.config.vae.dtype - runner.config.vae.dtype = "bfloat16" - debug.log(f"MPS detected: Setting VAE dtype from {original_vae_dtype} to {runner.config.vae.dtype} for compatibility", - category="precision", force=True) + # Set VAE dtype from runner's compute_dtype + compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16) + vae_dtype_str = str(compute_dtype).split('.')[-1] + runner.config.vae.dtype = vae_dtype_str + runner._vae_dtype_override = compute_dtype vae_checkpoint_path = find_model_file(vae_model, base_cache_dir) runner = prepare_model_structure(runner, "vae", vae_checkpoint_path, @@ -1439,8 +1450,6 @@ def prepare_model_structure( else: runner.vae = model runner._vae_checkpoint = checkpoint_path - # Store VAE dtype override if needed - runner._vae_dtype_override = getattr(torch, config.vae.dtype) if torch.mps.is_available() else None return runner @@ -1978,23 +1987,28 @@ def apply_model_specific_config(model: torch.nn.Module, runner: VideoDiffusionIn """ if is_dit: # DiT-specific - # Apply FP8 compatibility wrapper + # Apply FP8 compatibility wrapper with compute_dtype if not isinstance(model, FP8CompatibleDiT): debug.log("Applying FP8/RoPE compatibility wrapper to DiT model", category="setup") debug.start_timer("FP8CompatibleDiT") - model = FP8CompatibleDiT(model, debug, skip_conversion=False) + # Get compute_dtype from runner if available, fallback to bfloat16 + compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16) + model = FP8CompatibleDiT(model, debug, compute_dtype=compute_dtype, skip_conversion=False) debug.end_timer("FP8CompatibleDiT", "FP8/RoPE compatibility wrapper application") else: debug.log("Reusing existing FP8/RoPE compatibility wrapper", category="reuse") - # Apply attention mode to all FlashAttentionVarlen modules + # Apply attention mode and compute_dtype to all FlashAttentionVarlen modules if hasattr(runner, '_dit_attention_mode'): requested_attention_mode = runner._dit_attention_mode or 'sdpa' # Validate and get final attention_mode (with warning if fallback needed) attention_mode = validate_flash_attention_availability(requested_attention_mode, debug) - debug.log(f"Applying {attention_mode} attention mode", category="setup") + # Get compute_dtype from runner + compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16) + debug.log(f"Applying {attention_mode} attention mode and {compute_dtype} compute dtype to model", category="setup") + # Get the actual model (unwrap if needed) actual_model = model.dit_model if hasattr(model, 'dit_model') else model @@ -2003,10 +2017,11 @@ def apply_model_specific_config(model: torch.nn.Module, runner: VideoDiffusionIn for module in actual_model.modules(): if type(module).__name__ == 'FlashAttentionVarlen': module.attention_mode = attention_mode + module.compute_dtype = compute_dtype updated_count += 1 if updated_count > 0: - debug.log(f"Applied {attention_mode} to {updated_count} modules", category="success") + debug.log(f"Applied {attention_mode} and compute_dtype={compute_dtype} to {updated_count} modules", category="success") # Apply BlockSwap before torch.compile (only if not already active) # BlockSwap wraps forward methods, and torch.compile needs to capture the wrapped version diff --git a/src/interfaces/vae_model_loader.py b/src/interfaces/vae_model_loader.py index fc9e4bd..80f1670 100644 --- a/src/interfaces/vae_model_loader.py +++ b/src/interfaces/vae_model_loader.py @@ -50,23 +50,23 @@ class SeedVR2LoadVAEModel(io.ComfyNode): tooltip="Device for VAE inference (encoding/decoding)" ), io.Boolean.Input("encode_tiled", - default=False, + default=True, optional=True, - tooltip="Enable tiled encoding to reduce VRAM during encoding" + tooltip="Enable tiled encoding (ON by default to prevent noise artifacts at high resolution). Disable only for low-resolution inputs to improve speed." ), io.Int.Input("encode_tile_size", - default=512, + default=1024, min=64, step=32, optional=True, - tooltip="Size of encoding tiles in pixels. Smaller = less VRAM but more seams/artifacts and slower. Larger = more VRAM but better quality and faster." + tooltip="Encoding tile size in pixels (default: 1024). Can be reduced if running out of memory, but increasing above 1024 is not recommended as it causes noise artifacts in the latent." ), io.Int.Input("encode_tile_overlap", - default=64, + default=128, min=0, step=32, optional=True, - tooltip="Pixel overlap between encoding tiles to reduce visible seams. Higher = better blending but slower processing." + tooltip="Pixel overlap between encoding tiles to reduce visible seams (default: 128). Higher values improve blending at the cost of slower processing." ), io.Boolean.Input("decode_tiled", default=False, @@ -74,18 +74,18 @@ class SeedVR2LoadVAEModel(io.ComfyNode): tooltip="Enable tiled decoding to reduce VRAM during decoding" ), io.Int.Input("decode_tile_size", - default=512, + default=1024, min=64, step=32, optional=True, - tooltip="Size of decoding tiles in pixels. Smaller = less VRAM but more seams/artifacts and slower. Larger = more VRAM but better quality and faster." + tooltip="Decoding tile size in pixels (default: 1024). Adjust based on available VRAM." ), io.Int.Input("decode_tile_overlap", - default=64, + default=128, min=0, step=32, optional=True, - tooltip="Pixel overlap between decoding tiles to reduce visible seams. Higher = better blending but slower processing." + tooltip="Pixel overlap between decoding tiles to reduce visible seams (default: 128). Higher values improve blending at the cost of slower processing." ), io.Combo.Input("offload_device", options=get_device_list(include_none=True, include_cpu=True), diff --git a/src/interfaces/video_upscaler.py b/src/interfaces/video_upscaler.py index 12f0ad9..acc7bc1 100644 --- a/src/interfaces/video_upscaler.py +++ b/src/interfaces/video_upscaler.py @@ -96,9 +96,8 @@ class SeedVR2VideoUpscaler(io.ComfyNode): ), io.Combo.Input("color_correction", options=["lab", "wavelet", "wavelet_adaptive", "hsv", "adain", "none"], - default="wavelet", - optional=True, - tooltip="Color correction method to match upscaled output to input. 'wavelet' and 'wavelet_adaptive' provide best results." + default="lab", + tooltip="Color correction method: 'lab' (full perceptual color matching with detail preservation, recommended), 'wavelet' (frequency-based natural colors, preserves details), 'wavelet_adaptive' (wavelet base + targeted saturation correction), 'hsv' (hue-conditional saturation matching), 'adain' (statistical style transfer), 'none' (no correction)" ), io.Float.Input("input_noise_scale", default=0.0, diff --git a/src/models/dit_3b/attention.py b/src/models/dit_3b/attention.py index b717a50..75df3d6 100644 --- a/src/models/dit_3b/attention.py +++ b/src/models/dit_3b/attention.py @@ -120,15 +120,17 @@ class FlashAttentionVarlen(nn.Module): - Flash Attention: Uses @torch._dynamo.disable wrapper (C++ extension) """ - def __init__(self, attention_mode: str = 'sdpa'): + def __init__(self, attention_mode: str = 'sdpa', compute_dtype: torch.dtype = None): """ Initialize with specified attention backend. Args: attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_flash_attention_availability) + compute_dtype: Compute dtype for attention (set by pipeline, defaults to None for auto-detection) """ super().__init__() self.attention_mode = attention_mode + self.compute_dtype = compute_dtype def tflops(self, args, kwargs, output) -> float: cu_seqlens_q = kwargs["cu_seqlens_q"] @@ -141,6 +143,12 @@ class FlashAttentionVarlen(nn.Module): def forward(self, q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): kwargs["deterministic"] = torch.are_deterministic_algorithms_enabled() + # Convert to pipeline compute_dtype if configured (handles FP8 → fp16/bf16) + if self.compute_dtype is not None and q.dtype != self.compute_dtype: + q = q.to(self.compute_dtype) + k = k.to(self.compute_dtype) + v = v.to(self.compute_dtype) + if self.attention_mode == 'flash_attn': return _call_flash_attn_varlen_func( q, k, v, cu_seqlens_q, cu_seqlens_k, diff --git a/src/models/dit_3b/modulation.py b/src/models/dit_3b/modulation.py index 6a0590d..854ae09 100644 --- a/src/models/dit_3b/modulation.py +++ b/src/models/dit_3b/modulation.py @@ -92,17 +92,19 @@ class AdaSingle(nn.Module): getattr(self, f"{layer}_gate", None), ) - # Handle potential FP8 parameters - convert to computation dtype + # Handle potential FP8 parameters - convert to input computation dtype if hasattr(torch, 'float8_e4m3fn'): fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) + # Use input tensor's dtype as target (respects pipeline precision) + target_dtype = hid.dtype - # Convert FP8 parameters to BFloat16 for arithmetic operations + # Convert FP8 parameters to match input dtype for arithmetic operations if shiftB is not None and shiftB.dtype in fp8_types: - shiftB = shiftB.to(torch.bfloat16) + shiftB = shiftB.to(target_dtype) if scaleB is not None and scaleB.dtype in fp8_types: - scaleB = scaleB.to(torch.bfloat16) + scaleB = scaleB.to(target_dtype) if gateB is not None and gateB.dtype in fp8_types: - gateB = gateB.to(torch.bfloat16) + gateB = gateB.to(target_dtype) if mode == "in": return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB) diff --git a/src/models/dit_3b/nablocks/attention/mmattn.py b/src/models/dit_3b/nablocks/attention/mmattn.py index 3a6dcf6..a311449 100644 --- a/src/models/dit_3b/nablocks/attention/mmattn.py +++ b/src/models/dit_3b/nablocks/attention/mmattn.py @@ -122,10 +122,11 @@ class NaMMAttention(nn.Module): concat, unconcat = cache("mm_pnp", lambda: na.concat_idx(vid_len, txt_len)) + # Attention handles dtype conversion internally using pipeline compute_dtype attn = self.attn( - q=concat(vid_q, txt_q).bfloat16(), - k=concat(vid_k, txt_k).bfloat16(), - v=concat(vid_v, txt_v).bfloat16(), + q=concat(vid_q, txt_q), + k=concat(vid_k, txt_k), + v=concat(vid_v, txt_v), cu_seqlens_q=cache("mm_seqlens", lambda: safe_pad_operation(all_len.cumsum(0), (1, 0)).int()), cu_seqlens_k=cache("mm_seqlens", lambda: safe_pad_operation(all_len.cumsum(0), (1, 0)).int()), max_seqlen_q=cache("mm_maxlen", lambda: all_len.max()), @@ -240,10 +241,11 @@ class NaSwinAttention(NaMMAttention): else: vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win) + # Attention handles dtype conversion internally using pipeline compute_dtype out = self.attn( - q=concat_win(vid_q, txt_q).bfloat16(), - k=concat_win(vid_k, txt_k).bfloat16(), - v=concat_win(vid_v, txt_v).bfloat16(), + q=concat_win(vid_q, txt_q), + k=concat_win(vid_k, txt_k), + v=concat_win(vid_v, txt_v), cu_seqlens_q=cache_win( "vid_seqlens_q", lambda: safe_pad_operation(all_len_win.cumsum(0), (1, 0)).int() ), diff --git a/src/models/dit_3b/normalization.py b/src/models/dit_3b/normalization.py index 612d971..155c998 100644 --- a/src/models/dit_3b/normalization.py +++ b/src/models/dit_3b/normalization.py @@ -97,11 +97,12 @@ class CustomRMSNorm(nn.Module): normalized = input / rms if self.elementwise_affine: - # Convert FP8 weight to BFloat16 for arithmetic operations + # Convert FP8 weight to match input dtype for arithmetic operations if hasattr(torch, 'float8_e4m3fn'): fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) if self.weight.dtype in fp8_types: - weight = self.weight.to(torch.bfloat16) + # Use input dtype as target (respects pipeline precision) + weight = self.weight.to(input.dtype) return normalized * weight return normalized * self.weight diff --git a/src/models/dit_7b/attention.py b/src/models/dit_7b/attention.py index b717a50..75df3d6 100644 --- a/src/models/dit_7b/attention.py +++ b/src/models/dit_7b/attention.py @@ -120,15 +120,17 @@ class FlashAttentionVarlen(nn.Module): - Flash Attention: Uses @torch._dynamo.disable wrapper (C++ extension) """ - def __init__(self, attention_mode: str = 'sdpa'): + def __init__(self, attention_mode: str = 'sdpa', compute_dtype: torch.dtype = None): """ Initialize with specified attention backend. Args: attention_mode: 'flash_attn' or 'sdpa' (validated externally by validate_flash_attention_availability) + compute_dtype: Compute dtype for attention (set by pipeline, defaults to None for auto-detection) """ super().__init__() self.attention_mode = attention_mode + self.compute_dtype = compute_dtype def tflops(self, args, kwargs, output) -> float: cu_seqlens_q = kwargs["cu_seqlens_q"] @@ -141,6 +143,12 @@ class FlashAttentionVarlen(nn.Module): def forward(self, q, k, v, cu_seqlens_q, cu_seqlens_k, max_seqlen_q, max_seqlen_k, **kwargs): kwargs["deterministic"] = torch.are_deterministic_algorithms_enabled() + # Convert to pipeline compute_dtype if configured (handles FP8 → fp16/bf16) + if self.compute_dtype is not None and q.dtype != self.compute_dtype: + q = q.to(self.compute_dtype) + k = k.to(self.compute_dtype) + v = v.to(self.compute_dtype) + if self.attention_mode == 'flash_attn': return _call_flash_attn_varlen_func( q, k, v, cu_seqlens_q, cu_seqlens_k, diff --git a/src/models/dit_7b/modulation.py b/src/models/dit_7b/modulation.py index 1c348dd..38fff07 100644 --- a/src/models/dit_7b/modulation.py +++ b/src/models/dit_7b/modulation.py @@ -87,17 +87,19 @@ class AdaSingle(nn.Module): getattr(self, f"{layer}_gate"), ) - # Handle potential FP8 parameters - convert to computation dtype + # Handle potential FP8 parameters - convert to input computation dtype if hasattr(torch, 'float8_e4m3fn'): fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) + # Use input tensor's dtype as target (respects pipeline precision) + target_dtype = hid.dtype - # Convert FP8 parameters to BFloat16 for arithmetic operations - if shiftB.dtype in fp8_types: - shiftB = shiftB.to(torch.bfloat16) - if scaleB.dtype in fp8_types: - scaleB = scaleB.to(torch.bfloat16) - if gateB.dtype in fp8_types: - gateB = gateB.to(torch.bfloat16) + # Convert FP8 parameters to match input dtype for arithmetic operations + if shiftB is not None and shiftB.dtype in fp8_types: + shiftB = shiftB.to(target_dtype) + if scaleB is not None and scaleB.dtype in fp8_types: + scaleB = scaleB.to(target_dtype) + if gateB is not None and gateB.dtype in fp8_types: + gateB = gateB.to(target_dtype) if mode == "in": return hid.mul_(scaleA + scaleB).add_(shiftA + shiftB) diff --git a/src/models/dit_7b/nablocks/mmsr_block.py b/src/models/dit_7b/nablocks/mmsr_block.py index 3714bea..fe9010a 100644 --- a/src/models/dit_7b/nablocks/mmsr_block.py +++ b/src/models/dit_7b/nablocks/mmsr_block.py @@ -127,10 +127,11 @@ class NaSwinAttention(MMWindowAttention): if self.rope: vid_q, vid_k = self.rope(vid_q, vid_k, window_shape, cache_win) + # Attention handles dtype conversion internally using pipeline compute_dtype out = self.attn( - q=concat_win(vid_q, txt_q).bfloat16(), - k=concat_win(vid_k, txt_k).bfloat16(), - v=concat_win(vid_v, txt_v).bfloat16(), + q=concat_win(vid_q, txt_q), + k=concat_win(vid_k, txt_k), + v=concat_win(vid_v, txt_v), cu_seqlens_q=cache_win( "vid_seqlens_q", lambda: safe_pad_operation(all_len_win.cumsum(0), (1, 0)).int() ), diff --git a/src/models/dit_7b/normalization.py b/src/models/dit_7b/normalization.py index 549e13f..c34513a 100644 --- a/src/models/dit_7b/normalization.py +++ b/src/models/dit_7b/normalization.py @@ -86,11 +86,12 @@ class CustomRMSNorm(nn.Module): normalized = input / rms if self.elementwise_affine: - # Convert FP8 weight to BFloat16 for arithmetic operations + # Convert FP8 weight to match input dtype for arithmetic operations if hasattr(torch, 'float8_e4m3fn'): fp8_types = (torch.float8_e4m3fn, torch.float8_e5m2) if self.weight.dtype in fp8_types: - weight = self.weight.to(torch.bfloat16) + # Use input dtype as target (respects pipeline precision) + weight = self.weight.to(input.dtype) return normalized * weight return normalized * self.weight diff --git a/src/models/video_vae_v3/modules/attn_video_vae.py b/src/models/video_vae_v3/modules/attn_video_vae.py index 794c9ed..1ac8971 100644 --- a/src/models/video_vae_v3/modules/attn_video_vae.py +++ b/src/models/video_vae_v3/modules/attn_video_vae.py @@ -1307,6 +1307,13 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): x = x.unsqueeze(2) b, c, f, H, W = x.shape + tile_h, tile_w = tile_size + + # Only tile if input resolution requires multiple tiles + if H <= tile_h and W <= tile_w: + return self.slicing_encode(x) + else: + self.debug.log(f"Using VAE tiled encoding (Tile: {tile_size}, Overlap: {tile_overlap})", category="vae", force=True, indent_level=1) # Spatial scale factor (output/latent) scale_factor = self.spatial_downsample_factor @@ -1368,19 +1375,15 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): tile_id += 1 tile_sample = x[:, :, :, y_out:y_out_end, x_out:x_out_end] - # Log progress periodically instead of every tile - if self.debug and (tile_id == 1 or tile_id % 5 == 0 or tile_id == num_tiles): - end_tile = min(tile_id + 4, num_tiles) + # Log progress periodically instead of every tile (at 1, 6, 11, 16, ...) + if self.debug and (tile_id % 5 == 1 or tile_id == num_tiles): if tile_id == num_tiles: - self.debug.log( - f"Encoding tile {tile_id} / {num_tiles}", - category="vae", - ) + # Only log final tile if not covered by previous range + if (tile_id - 1) % 5 == 0: + self.debug.log(f"Encoding tile {tile_id} / {num_tiles}", category="vae") else: - self.debug.log( - f"Encoding tiles {tile_id}-{end_tile} / {num_tiles}", - category="vae", - ) + end_tile = min(tile_id + 4, num_tiles) + self.debug.log(f"Encoding tiles {tile_id}-{end_tile} / {num_tiles}", category="vae") encoded_tile = self.slicing_encode(tile_sample) @@ -1451,6 +1454,13 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): latent_tile_h = max(1, tile_h // scale_factor) latent_tile_w = max(1, tile_w // scale_factor) + + # Only tile if latent resolution requires multiple tiles + if H <= latent_tile_h and W <= latent_tile_w: + return self.slicing_decode(z) + else: + self.debug.log(f"Using VAE tiled decoding (Tile: {tile_size}, Overlap: {tile_overlap})", category="vae", force=True, indent_level=1) + latent_overlap_h = max(0, min((overlap_h // scale_factor), latent_tile_h - 1)) latent_overlap_w = max(0, min((overlap_w // scale_factor), latent_tile_w - 1)) @@ -1494,19 +1504,15 @@ class VideoAutoencoderKL(diffusers.AutoencoderKL): tile_id += 1 tile_latent = z[:, :, :, y_lat:y_lat_end, x_lat:x_lat_end] - # Log progress periodically instead of every tile - if self.debug and (tile_id == 1 or tile_id % 5 == 0 or tile_id == num_tiles): - end_tile = min(tile_id + 4, num_tiles) + # Log progress periodically instead of every tile (at 1, 6, 11, 16, ...) + if self.debug and (tile_id % 5 == 1 or tile_id == num_tiles): if tile_id == num_tiles: - self.debug.log( - f"Decoding tile {tile_id} / {num_tiles}", - category="vae", - ) + # Only log final tile if not covered by previous range + if (tile_id - 1) % 5 == 0: + self.debug.log(f"Decoding tile {tile_id} / {num_tiles}", category="vae") else: - self.debug.log( - f"Decoding tiles {tile_id}-{end_tile} / {num_tiles}", - category="vae", - ) + end_tile = min(tile_id + 4, num_tiles) + self.debug.log(f"Decoding tiles {tile_id}-{end_tile} / {num_tiles}", category="vae") decoded_tile = self.slicing_decode(tile_latent) diff --git a/src/optimization/compatibility.py b/src/optimization/compatibility.py index 2ce5706..c75ec85 100644 --- a/src/optimization/compatibility.py +++ b/src/optimization/compatibility.py @@ -174,35 +174,44 @@ def call_rope_with_stability(method, *args, **kwargs): class FP8CompatibleDiT(torch.nn.Module): """ Wrapper for DiT models with automatic compatibility management + advanced optimizations - - FP8: Keeps native FP8 parameters, converts inputs/outputs - - FP16: Uses native FP16 - - Mixed Precision: Stabilizes RoPE for models with FP16 blocks - - RoPE: Converted from FP8 to BFloat16 only when detected as FP8 + + Precision Handling: + - FP8: Keeps native FP8 parameters (memory efficient), converts inputs/outputs to compute_dtype for arithmetic + - FP16: Uses native FP16 precision throughout + - BFloat16: Uses native BFloat16 precision throughout + - Float32: Uses full precision for maximum quality + - RoPE: Converted from FP8 to compute_dtype for numerical consistency + + Optimizations: - Flash Attention: Automatic optimization of attention layers + - RoPE Stabilization: Error handling for numerical stability in mixed precision + - MPS Compatibility: Unified dtype conversion for Apple Silicon backends """ - def __init__(self, dit_model, debug: 'Debug', skip_conversion: bool = False): + def __init__(self, dit_model, debug: 'Debug', compute_dtype: torch.dtype = torch.bfloat16, skip_conversion: bool = False): super().__init__() self.dit_model = dit_model self.debug = debug + self.compute_dtype = compute_dtype self.model_dtype = self._detect_model_dtype() self.is_fp8_model = self.model_dtype in (torch.float8_e4m3fn, torch.float8_e5m2) self.is_fp16_model = self.model_dtype == torch.float16 # Only convert if not already done (e.g., when reusing cached weights) if not skip_conversion and self.is_fp8_model: - # Only FP8 models need RoPE frequency conversion + # FP8 models need RoPE frequency conversion to compute dtype model_variant = self._get_model_variant() self.debug.log(f"Detected NaDiT {model_variant} FP8 - Converting RoPE freqs for FP8 compatibility", - category="precision") + category="precision") self.debug.start_timer("_convert_rope_freqs") - self._convert_rope_freqs() + self._convert_rope_freqs(target_dtype=self.compute_dtype) self.debug.end_timer("_convert_rope_freqs", "RoPE freqs conversion") + if torch.mps.is_available(): self.debug.log(f"Also converting NaDiT parameters/buffers for MPS backend", category="setup", force=True) - self.debug.start_timer("_force_nadit_bfloat16") - self._force_nadit_bfloat16() - self.debug.end_timer("_force_nadit_bfloat16", "NaDiT parameters/buffers conversion") + self.debug.start_timer("_force_nadit_precision") + self._force_nadit_precision(target_dtype=self.compute_dtype) + self.debug.end_timer("_force_nadit_precision", "NaDiT parameters/buffers conversion") # Apply RoPE stabilization for numerical stability self.debug.log(f"Stabilizing RoPE computations for numerical stability", category="setup") @@ -232,52 +241,66 @@ class FP8CompatibleDiT(torch.nn.Module): else: return "Unknown" - def _convert_rope_freqs(self) -> None: - """Convert RoPE frequency buffers from FP8 to BFloat16 for compatibility""" + def _convert_rope_freqs(self, target_dtype: torch.dtype = torch.bfloat16) -> None: + """ + Convert RoPE frequency buffers from FP8 to target dtype for compatibility. + + Args: + target_dtype: Target dtype for RoPE freqs (default: bfloat16 for stability) + """ converted = 0 for module in self.dit_model.modules(): if 'RotaryEmbedding' in type(module).__name__: if hasattr(module, 'rope') and hasattr(module.rope, 'freqs'): if module.rope.freqs.dtype in (torch.float8_e4m3fn, torch.float8_e5m2): - if module.rope.freqs.device.type == "mps": - module.rope.freqs.data = module.rope.freqs.to("cpu").to(torch.bfloat16).to("mps") + if module.rope.freqs.device.type == "mps": + module.rope.freqs.data = module.rope.freqs.to("cpu").to(target_dtype).to("mps") else: - module.rope.freqs.data = module.rope.freqs.to(torch.bfloat16) + module.rope.freqs.data = module.rope.freqs.to(target_dtype) converted += 1 - self.debug.log(f"Converted {converted} RoPE frequency buffers from FP8 to BFloat16 for compatibility", category="success") + self.debug.log(f"Converted {converted} RoPE frequency buffers from FP8 to {target_dtype} for compatibility", category="success") - def _force_nadit_bfloat16(self) -> None: - """🎯 Force ALL NaDiT parameters to BFloat16 to avoid promotion errors""" + def _force_nadit_precision(self, target_dtype: torch.dtype = torch.bfloat16) -> None: + """ + Force ALL NaDiT parameters to target dtype to avoid promotion errors (MPS requirement). + + Args: + target_dtype: Target dtype for all parameters (default: bfloat16 for MPS compatibility) + """ converted_count = 0 original_dtype = None - # Convert ALL parameters to BFloat16 (FP8, FP16, etc.) + # Convert ALL parameters to target dtype for name, param in self.dit_model.named_parameters(): if original_dtype is None: original_dtype = param.dtype - if param.dtype != torch.bfloat16: + if param.dtype != target_dtype: if param.device.type == "mps": - param_data = param.data.to("cpu").to(torch.bfloat16).to("mps") + temp_cpu = param.data.to("cpu") + temp_converted = temp_cpu.to(target_dtype) + param.data = temp_converted.to("mps") + del temp_cpu, temp_converted else: - param_data = param.data.to(torch.bfloat16) - param.data = param_data + param.data = param.data.to(target_dtype) converted_count += 1 # Also convert buffers for name, buffer in self.dit_model.named_buffers(): - if buffer.dtype != torch.bfloat16: - if param.device.type == "mps": - buffer_data = buffer.data.to("cpu").to(torch.bfloat16).to("mps") + if buffer.dtype != target_dtype: + if buffer.device.type == "mps": + temp_cpu = buffer.data.to("cpu") + temp_converted = temp_cpu.to(target_dtype) + buffer.data = temp_converted.to("mps") + del temp_cpu, temp_converted else: - buffer_data = buffer.data.to(torch.bfloat16) - buffer.data = buffer_data + buffer.data = buffer.data.to(target_dtype) converted_count += 1 - self.debug.log(f"Converted {converted_count} NaDiT parameters/buffers for MPS", category="success") + self.debug.log(f"Converted {converted_count} NaDiT parameters/buffers to {target_dtype} for MPS", category="success") # Update detected dtype - self.model_dtype = torch.bfloat16 - self.is_fp8_model = False # Model is no longer FP8 after conversion + self.model_dtype = target_dtype + self.is_fp8_model = (target_dtype in (torch.float8_e4m3fn, torch.float8_e5m2)) def _stabilize_rope_computations(self): """ @@ -523,23 +546,25 @@ class FP8CompatibleDiT(torch.nn.Module): return module._original_forward(x, *args, **kwargs) def forward(self, *args, **kwargs): - """Forward pass with minimal dtype conversion overhead + """ + Forward pass with minimal dtype conversion overhead Conversion strategy: - - FP16 models: Keep everything in FP16 (no conversion needed) - - FP8 models: Convert FP8 tensors to BFloat16 (required for arithmetic) - - BFloat16 models: No conversion needed + - FP16/BFloat16/Float32 models: Use native precision (no conversion needed) + - FP8 models: Convert FP8 tensors to compute_dtype for arithmetic operations + (FP8 parameters stay in FP8 for memory efficiency, only converted for computation) """ # Only convert if we have an FP8 model for arithmetic operations if self.is_fp8_model: fp8_dtypes = (torch.float8_e4m3fn, torch.float8_e5m2) + target_dtype = self.compute_dtype # Convert args converted_args = [] for arg in args: if isinstance(arg, torch.Tensor) and arg.dtype in fp8_dtypes: - converted_args.append(arg.to(torch.bfloat16)) + converted_args.append(arg.to(target_dtype)) else: converted_args.append(arg) @@ -547,7 +572,7 @@ class FP8CompatibleDiT(torch.nn.Module): converted_kwargs = {} for key, value in kwargs.items(): if isinstance(value, torch.Tensor) and value.dtype in fp8_dtypes: - converted_kwargs[key] = value.to(torch.bfloat16) + converted_kwargs[key] = value.to(target_dtype) else: converted_kwargs[key] = value @@ -560,7 +585,7 @@ class FP8CompatibleDiT(torch.nn.Module): except Exception as e: self.debug.log(f"Forward pass error: {e}", level="ERROR", category="generation", force=True) if self.is_fp8_model: - self.debug.log(f"FP8 model - converted FP8 tensors to BFloat16", category="info", force=True) + self.debug.log(f"FP8 model - converted FP8 tensors to {self.compute_dtype}", category="info", force=True) else: self.debug.log(f"{self.model_dtype} model - no conversion applied", category="info", force=True) raise