Fix SageAttention naming, restore strict precision control, and fix crashes
- Re-implemented `precision` control (`fp16`, `bf16`, `bf32`, `auto`) in CLI, ComfyUI node, and backend logic to respect user choice. - Fixed `UnboundLocalError` in `apply_model_specific_config` by ensuring `compute_dtype` is always initialized before use. - Fixed `NameError` crash in `SeedVR2VideoUpscaler` by properly passing the restored `precision` argument. - Renamed `sd2`/`sd3` to `sa2`/`sa3` for clarity and fixed fallback logic to ensure SageAttention is correctly prioritized. - Added explicit logging of active attention backend and execution confirmation. - Updated `FP8CompatibleDiT` to exclude `FlashAttentionVarlen` modules from unnecessary wrapping.
This commit is contained in:
+5
-1
@@ -688,7 +688,8 @@ def _process_frames_core(
|
||||
dit_offload_device=dit_offload,
|
||||
vae_offload_device=vae_offload,
|
||||
tensor_offload_device=tensor_offload,
|
||||
debug=debug
|
||||
debug=debug,
|
||||
precision=args.precision
|
||||
)
|
||||
if runner_cache is not None:
|
||||
runner_cache['ctx'] = ctx
|
||||
@@ -1161,6 +1162,9 @@ Examples:
|
||||
perf_group.add_argument("--attention_mode", type=str, default="sdpa",
|
||||
choices=["sdpa", "flash_attn", "sa2", "sa3"],
|
||||
help="Attention backend: 'sdpa' (default), 'flash_attn' (faster), 'sa2' (SageAttention v2), 'sa3' (SageAttention v3)")
|
||||
perf_group.add_argument("--precision", type=str, default="auto",
|
||||
choices=["auto", "fp16", "bf16", "bf32"],
|
||||
help="Compute precision: 'auto' (default), 'fp16', 'bf16', or 'bf32' (TF32)")
|
||||
perf_group.add_argument("--compile_dit", action="store_true",
|
||||
help="Enable torch.compile for DiT model (20-40%% speedup, requires PyTorch 2.0+ and Triton)")
|
||||
perf_group.add_argument("--compile_vae", action="store_true",
|
||||
|
||||
@@ -318,7 +318,8 @@ def setup_generation_context(
|
||||
dit_offload_device: Optional[Union[str, torch.device]] = None,
|
||||
vae_offload_device: Optional[Union[str, torch.device]] = None,
|
||||
tensor_offload_device: Optional[Union[str, torch.device]] = None,
|
||||
debug: Optional['Debug'] = None
|
||||
debug: Optional['Debug'] = None,
|
||||
precision: str = 'auto'
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Initialize generation context with device configuration.
|
||||
@@ -333,6 +334,7 @@ def setup_generation_context(
|
||||
vae_offload_device: Device to offload VAE to when not in use (optional)
|
||||
tensor_offload_device: Device to offload intermediate tensors to (optional)
|
||||
debug: Debug instance for logging
|
||||
precision: Compute precision ('auto', 'fp16', 'bf16', 'bf32')
|
||||
|
||||
Returns:
|
||||
Dict[str, Any]: Generation context dictionary with torch.device objects
|
||||
@@ -365,8 +367,28 @@ def setup_generation_context(
|
||||
interrupt_fn = None
|
||||
comfyui_available = False
|
||||
|
||||
# Determine compute dtype (allow override in context setup if needed, but standard flow uses COMPUTE_DTYPE)
|
||||
compute_dtype = COMPUTE_DTYPE
|
||||
# Determine compute dtype based on precision request
|
||||
if precision == 'fp16':
|
||||
compute_dtype = torch.float16
|
||||
reason = "user requested fp16"
|
||||
elif precision == 'bf16':
|
||||
compute_dtype = torch.bfloat16
|
||||
reason = "user requested bf16"
|
||||
elif precision == 'bf32':
|
||||
# BF32 is usually implemented as float32 tensors with specific matmul settings (TF32)
|
||||
# For torch dtype context, we use float32
|
||||
compute_dtype = torch.float32
|
||||
reason = "user requested bf32 (TF32)"
|
||||
# Note: TF32 enablement should be handled globally or in model config
|
||||
else:
|
||||
# 'auto' - existing logic
|
||||
compute_dtype = COMPUTE_DTYPE
|
||||
if compute_dtype == torch.float32:
|
||||
reason = "quality"
|
||||
elif not BFLOAT16_SUPPORTED:
|
||||
reason = "compatibility (GPU lacks bfloat16 CUBLAS - 7B models unsupported, 3B may have artifacts)"
|
||||
else:
|
||||
reason = "performance"
|
||||
|
||||
# Create generation context
|
||||
ctx = {
|
||||
@@ -405,12 +427,6 @@ def setup_generation_context(
|
||||
f"LOCAL_RANK={os.environ['LOCAL_RANK']}",
|
||||
category="setup"
|
||||
)
|
||||
if ctx['compute_dtype'] == torch.float32:
|
||||
reason = "quality"
|
||||
elif not BFLOAT16_SUPPORTED:
|
||||
reason = "compatibility (GPU lacks bfloat16 CUBLAS - 7B models unsupported, 3B may have artifacts)"
|
||||
else:
|
||||
reason = "performance"
|
||||
debug.log(f"Unified compute dtype: {ctx['compute_dtype']} across entire pipeline for maximum {reason}", category="precision")
|
||||
|
||||
return ctx
|
||||
@@ -435,6 +451,7 @@ def prepare_runner(
|
||||
decode_tile_overlap: Optional[Tuple[int, int]] = None,
|
||||
tile_debug: str = "false",
|
||||
attention_mode: str = 'sdpa',
|
||||
precision: str = 'auto',
|
||||
torch_compile_args_dit: Optional[Dict[str, Any]] = None,
|
||||
torch_compile_args_vae: Optional[Dict[str, Any]] = None
|
||||
) -> Tuple['VideoDiffusionInfer', Dict[str, Any]]:
|
||||
@@ -505,6 +522,7 @@ def prepare_runner(
|
||||
decode_tile_overlap=decode_tile_overlap,
|
||||
tile_debug=tile_debug,
|
||||
attention_mode=attention_mode,
|
||||
precision=precision,
|
||||
torch_compile_args_dit=torch_compile_args_dit,
|
||||
torch_compile_args_vae=torch_compile_args_vae
|
||||
)
|
||||
|
||||
@@ -749,6 +749,7 @@ def configure_runner(
|
||||
decode_tile_overlap: Optional[Tuple[int, int]] = None,
|
||||
tile_debug: str = "false",
|
||||
attention_mode: str = 'sdpa',
|
||||
precision: str = 'auto',
|
||||
torch_compile_args_dit: Optional[Dict[str, Any]] = None,
|
||||
torch_compile_args_vae: Optional[Dict[str, Any]] = None
|
||||
) -> Tuple[VideoDiffusionInfer, Dict[str, Any]]:
|
||||
@@ -821,6 +822,9 @@ def configure_runner(
|
||||
block_swap_config, debug
|
||||
)
|
||||
|
||||
# Store precision setting
|
||||
runner._precision = precision
|
||||
|
||||
# Phase 4: Setup models (load from cache or create new)
|
||||
_setup_models(
|
||||
runner, cache_context, dit_model, vae_model,
|
||||
@@ -1182,10 +1186,20 @@ def apply_model_specific_config(model: torch.nn.Module, runner: VideoDiffusionIn
|
||||
"""
|
||||
if is_dit:
|
||||
# DiT-specific
|
||||
# Get compute_dtype from runner if available, fallback to bfloat16
|
||||
# Determine compute_dtype upfront (respect precision setting)
|
||||
compute_dtype = getattr(runner, '_compute_dtype', torch.bfloat16)
|
||||
|
||||
# Apply FP8 compatibility wrapper with compute_dtype
|
||||
# Apply precision override if set (redundant if ctx already handled it, but safe)
|
||||
precision_override = getattr(runner, '_precision', 'auto')
|
||||
if precision_override == 'fp16':
|
||||
compute_dtype = torch.float16
|
||||
elif precision_override == 'bf16':
|
||||
compute_dtype = torch.bfloat16
|
||||
elif precision_override == 'bf32':
|
||||
# TF32 context
|
||||
compute_dtype = torch.float32
|
||||
|
||||
# Apply FP8 compatibility wrapper with correct compute_dtype
|
||||
if not isinstance(model, FP8CompatibleDiT):
|
||||
debug.log("Applying FP8/RoPE compatibility wrapper to DiT model", category="setup")
|
||||
debug.start_timer("FP8CompatibleDiT")
|
||||
@@ -1193,6 +1207,10 @@ def apply_model_specific_config(model: torch.nn.Module, runner: VideoDiffusionIn
|
||||
debug.end_timer("FP8CompatibleDiT", "FP8/RoPE compatibility wrapper application")
|
||||
else:
|
||||
debug.log("Reusing existing FP8/RoPE compatibility wrapper", category="reuse")
|
||||
# Update compute_dtype if wrapper exists
|
||||
if model.compute_dtype != compute_dtype:
|
||||
debug.log(f"Updating FP8 wrapper compute_dtype to {compute_dtype}", category="setup")
|
||||
model.compute_dtype = compute_dtype
|
||||
|
||||
# Apply attention mode and compute_dtype to all FlashAttentionVarlen modules
|
||||
if hasattr(runner, '_dit_attention_mode'):
|
||||
|
||||
@@ -204,6 +204,18 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
|
||||
"• 'cuda:X': Offload to another GPU (good balance if available, faster than CPU)"
|
||||
)
|
||||
),
|
||||
io.Combo.Input("precision",
|
||||
options=["auto", "fp16", "bf16", "bf32"],
|
||||
default="auto",
|
||||
optional=True,
|
||||
tooltip=(
|
||||
"Precision for main generation process (default: auto).\n"
|
||||
"• auto: Automatically select based on device capabilities\n"
|
||||
"• fp16: Half precision (fastest, standard)\n"
|
||||
"• bf16: BFloat16 (better dynamic range, requires Ampere+ GPU)\n"
|
||||
"• bf32: Float32 with TF32 enabled (Ampere+ GPU)"
|
||||
)
|
||||
),
|
||||
io.Boolean.Input("enable_debug",
|
||||
default=False,
|
||||
optional=True,
|
||||
@@ -227,7 +239,7 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
|
||||
uniform_batch_size: bool = False, temporal_overlap: int = 0, prepend_frames: int = 0,
|
||||
color_correction: str = "wavelet", input_noise_scale: float = 0.0,
|
||||
latent_noise_scale: float = 0.0, offload_device: str = "none",
|
||||
enable_debug: bool = False) -> io.NodeOutput:
|
||||
precision: str = "auto", enable_debug: bool = False) -> io.NodeOutput:
|
||||
"""
|
||||
Execute SeedVR2 video upscaling with progress reporting
|
||||
|
||||
@@ -412,7 +424,8 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
|
||||
dit_offload_device=dit_offload_device,
|
||||
vae_offload_device=vae_offload_device,
|
||||
tensor_offload_device=tensor_offload_device,
|
||||
debug=debug
|
||||
debug=debug,
|
||||
precision=precision
|
||||
)
|
||||
|
||||
# Prepare runner with model state management and global cache
|
||||
@@ -435,6 +448,7 @@ class SeedVR2VideoUpscaler(io.ComfyNode):
|
||||
decode_tile_overlap=(decode_tile_overlap, decode_tile_overlap),
|
||||
tile_debug=tile_debug,
|
||||
attention_mode=attention_mode,
|
||||
precision=precision,
|
||||
torch_compile_args_dit=dit_torch_compile_args,
|
||||
torch_compile_args_vae=vae_torch_compile_args
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user