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:
+10
-9
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
|
||||
@@ -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),
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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()
|
||||
),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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()
|
||||
),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user