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
This commit is contained in:
Adrien Toupet
2025-10-22 23:56:14 -04:00
parent c9dce827c0
commit 01cbdf8bc3
16 changed files with 224 additions and 183 deletions
+10 -9
View File
@@ -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)")
+26 -48
View File
@@ -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:
+3 -11
View File
@@ -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"):
+28 -13
View File
@@ -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
+10 -10
View File
@@ -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),
+2 -3
View File
@@ -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,
+9 -1
View File
@@ -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,
+7 -5
View File
@@ -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)
@@ -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()
),
+3 -2
View File
@@ -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
+9 -1
View File
@@ -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,
+10 -8
View File
@@ -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)
+4 -3
View File
@@ -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()
),
+3 -2
View File
@@ -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
@@ -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)
+64 -39
View File
@@ -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