[diffusion] fix: fix diffusion output stability on mps (#30017)
This commit is contained in:
@@ -157,9 +157,24 @@ class RMSNorm(CustomOp):
|
||||
|
||||
x_var = x[..., : self.variance_size_override]
|
||||
|
||||
if x.device.type == "mps" and self.variance_size_override is None:
|
||||
weight = self.weight.to(dtype=torch.float32)
|
||||
x = F.rms_norm(
|
||||
x,
|
||||
(self.hidden_size,),
|
||||
weight,
|
||||
self.variance_epsilon,
|
||||
).to(orig_dtype)
|
||||
if residual is None:
|
||||
return x
|
||||
return x, residual
|
||||
|
||||
variance = x_var.pow(2).mean(dim=-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + self.variance_epsilon)
|
||||
x = (x * self.weight).to(orig_dtype)
|
||||
weight = self.weight
|
||||
if x.device.type == "mps" and weight.dtype != x.dtype:
|
||||
weight = weight.to(dtype=x.dtype)
|
||||
x = (x * weight).to(orig_dtype)
|
||||
if residual is None:
|
||||
return x
|
||||
else:
|
||||
|
||||
@@ -154,6 +154,17 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
||||
def apply(
|
||||
self, layer: torch.nn.Module, x: torch.Tensor, bias: torch.Tensor | None = None
|
||||
) -> torch.Tensor:
|
||||
if x.device.type == "mps" and (
|
||||
x.dtype != torch.float32
|
||||
or layer.weight.dtype != torch.float32
|
||||
or (bias is not None and bias.dtype != torch.float32)
|
||||
):
|
||||
return F.linear(
|
||||
x.to(torch.float32),
|
||||
layer.weight.to(torch.float32),
|
||||
None if bias is None else bias.to(torch.float32),
|
||||
).to(x.dtype)
|
||||
|
||||
output = (
|
||||
F.linear(x, layer.weight, bias)
|
||||
if IS_AMP_SUPPORTED or bias is None
|
||||
|
||||
@@ -23,6 +23,7 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.layerwise_offload_co
|
||||
is_text_encoder_component_name,
|
||||
is_vae_component_name,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import DiffusionNvtxHooks
|
||||
@@ -94,6 +95,8 @@ class ComponentResidencyPipeline(Protocol):
|
||||
def should_cpu_offload_component(
|
||||
component_name: str, module: nn.Module, server_args: ServerArgs
|
||||
) -> bool:
|
||||
if current_platform.is_mps():
|
||||
return False
|
||||
if server_args.use_fsdp_inference or is_fsdp_managed_module(module):
|
||||
return False
|
||||
if is_dit_component_name(component_name):
|
||||
|
||||
@@ -18,6 +18,38 @@ from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _mps_safe_linear(linear: nn.Linear, x: torch.Tensor) -> torch.Tensor:
|
||||
if x.device.type != "mps":
|
||||
return linear(x)
|
||||
|
||||
return F.linear(
|
||||
x.to(torch.float32),
|
||||
linear.weight.to(torch.float32),
|
||||
None if linear.bias is None else linear.bias.to(torch.float32),
|
||||
).to(x.dtype)
|
||||
|
||||
|
||||
def _mps_safe_conv2d(conv: nn.Conv2d, x: torch.Tensor) -> torch.Tensor:
|
||||
if x.device.type != "mps":
|
||||
return conv(x)
|
||||
|
||||
return F.conv2d(
|
||||
x.to(torch.float32),
|
||||
conv.weight.to(torch.float32),
|
||||
None if conv.bias is None else conv.bias.to(torch.float32),
|
||||
conv.stride,
|
||||
conv.padding,
|
||||
conv.dilation,
|
||||
conv.groups,
|
||||
).to(x.dtype)
|
||||
|
||||
|
||||
def _mps_match_dtype(tensor: torch.Tensor, ref: torch.Tensor) -> torch.Tensor:
|
||||
if ref.device.type == "mps" and tensor.dtype != ref.dtype:
|
||||
return tensor.to(dtype=ref.dtype)
|
||||
return tensor
|
||||
|
||||
|
||||
class SanaCombinedTimestepSizeEmbeddings(nn.Module):
|
||||
def __init__(self, embedding_dim):
|
||||
super().__init__()
|
||||
@@ -32,7 +64,16 @@ class SanaCombinedTimestepSizeEmbeddings(nn.Module):
|
||||
timesteps_proj = self.time_proj(timestep)
|
||||
if hidden_dtype is not None:
|
||||
timesteps_proj = timesteps_proj.to(dtype=hidden_dtype)
|
||||
timesteps_emb = self.timestep_embedder(timesteps_proj)
|
||||
if timesteps_proj.device.type == "mps":
|
||||
embedder = self.timestep_embedder
|
||||
timesteps_emb = _mps_safe_linear(embedder.linear_1, timesteps_proj)
|
||||
if embedder.act is not None:
|
||||
timesteps_emb = embedder.act(timesteps_emb)
|
||||
timesteps_emb = _mps_safe_linear(embedder.linear_2, timesteps_emb)
|
||||
if embedder.post_act is not None:
|
||||
timesteps_emb = embedder.post_act(timesteps_emb)
|
||||
else:
|
||||
timesteps_emb = self.timestep_embedder(timesteps_proj)
|
||||
return timesteps_emb
|
||||
|
||||
|
||||
@@ -45,7 +86,7 @@ class SanaAdaLayerNormSingle(nn.Module):
|
||||
|
||||
def forward(self, timestep, hidden_dtype=None):
|
||||
embedded_timestep = self.emb(timestep, hidden_dtype=hidden_dtype)
|
||||
out = self.linear(self.silu(embedded_timestep))
|
||||
out = _mps_safe_linear(self.linear, self.silu(embedded_timestep))
|
||||
return out, embedded_timestep
|
||||
|
||||
|
||||
@@ -56,6 +97,7 @@ class SanaModulatedNorm(nn.Module):
|
||||
|
||||
def forward(self, x, temb, scale_shift_table):
|
||||
x = self.norm(x)
|
||||
scale_shift_table = _mps_match_dtype(scale_shift_table, temb)
|
||||
shift, scale = (scale_shift_table[None] + temb[:, None]).chunk(2, dim=1)
|
||||
x = x * (1 + scale) + shift
|
||||
return x
|
||||
@@ -80,12 +122,12 @@ class GLUMBConv(nn.Module):
|
||||
self.conv_point = nn.Conv2d(hidden_channels, out_channels, 1, 1, 0, bias=False)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states = self.conv_inverted(hidden_states)
|
||||
hidden_states = _mps_safe_conv2d(self.conv_inverted, hidden_states)
|
||||
hidden_states = self.nonlinearity(hidden_states)
|
||||
hidden_states = self.conv_depth(hidden_states)
|
||||
hidden_states = _mps_safe_conv2d(self.conv_depth, hidden_states)
|
||||
hidden_states, gate = torch.chunk(hidden_states, 2, dim=1)
|
||||
hidden_states = hidden_states * self.nonlinearity(gate)
|
||||
hidden_states = self.conv_point(hidden_states)
|
||||
hidden_states = _mps_safe_conv2d(self.conv_point, hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -129,7 +171,7 @@ class SanaLinearAttention(nn.Module):
|
||||
hidden_states = qkv / normalizer
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(B, S, -1)
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
hidden_states = _mps_safe_linear(self.to_out[0], hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -156,7 +198,7 @@ class SanaCrossAttention(nn.Module):
|
||||
B, S, _ = hidden_states.shape
|
||||
T = encoder_hidden_states.shape[1]
|
||||
|
||||
query = self.to_q(hidden_states)
|
||||
query = _mps_safe_linear(self.to_q, hidden_states)
|
||||
kv, _ = self.to_kv(encoder_hidden_states)
|
||||
key, value = kv.split([self.inner_dim, self.inner_dim], dim=-1)
|
||||
|
||||
@@ -173,7 +215,7 @@ class SanaCrossAttention(nn.Module):
|
||||
query, key, value, attn_mask=attn_mask
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).reshape(B, S, -1)
|
||||
hidden_states = self.to_out[0](hidden_states)
|
||||
hidden_states = _mps_safe_linear(self.to_out[0], hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -224,8 +266,9 @@ class SanaTransformerBlock(nn.Module):
|
||||
):
|
||||
batch_size = hidden_states.shape[0]
|
||||
|
||||
scale_shift_table = _mps_match_dtype(self.scale_shift_table, timestep)
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||
self.scale_shift_table[None] + timestep.reshape(batch_size, 6, -1)
|
||||
scale_shift_table[None] + timestep.reshape(batch_size, 6, -1)
|
||||
).chunk(6, dim=1)
|
||||
|
||||
norm_hidden = self.norm1(hidden_states)
|
||||
@@ -339,7 +382,7 @@ class SanaTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
|
||||
hidden_states = self.patch_embed["proj"](hidden_states)
|
||||
hidden_states = _mps_safe_conv2d(self.patch_embed["proj"], hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
timestep_emb, embedded_timestep = self.time_embed(
|
||||
@@ -349,7 +392,16 @@ class SanaTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
||||
if isinstance(encoder_attention_mask, (list, tuple)):
|
||||
encoder_attention_mask = encoder_attention_mask[0]
|
||||
|
||||
encoder_hidden_states = self.caption_projection(encoder_hidden_states)
|
||||
if encoder_hidden_states.device.type == "mps":
|
||||
encoder_hidden_states = _mps_safe_linear(
|
||||
self.caption_projection.linear_1, encoder_hidden_states
|
||||
)
|
||||
encoder_hidden_states = self.caption_projection.act_1(encoder_hidden_states)
|
||||
encoder_hidden_states = _mps_safe_linear(
|
||||
self.caption_projection.linear_2, encoder_hidden_states
|
||||
)
|
||||
else:
|
||||
encoder_hidden_states = self.caption_projection(encoder_hidden_states)
|
||||
if encoder_hidden_states.shape[0] != batch_size:
|
||||
encoder_hidden_states = encoder_hidden_states.expand(
|
||||
batch_size, -1, -1
|
||||
@@ -379,7 +431,7 @@ class SanaTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
|
||||
hidden_states = self.norm_out(
|
||||
hidden_states, embedded_timestep, self.scale_shift_table
|
||||
)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
hidden_states = _mps_safe_linear(self.proj_out, hidden_states)
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, post_patch_height, post_patch_width, p, p, self.out_channels
|
||||
)
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
from collections.abc import Iterable
|
||||
|
||||
import torch
|
||||
from diffusers.models.autoencoders.vae import DecoderOutput
|
||||
from torch import nn
|
||||
|
||||
from sglang.multimodal_gen.configs.models.vaes.sana import SanaVAEConfig
|
||||
@@ -41,6 +42,31 @@ class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin):
|
||||
self._loaded_state_dict: dict[str, torch.Tensor] = {}
|
||||
self._spatial_parallel_decode_enabled = False
|
||||
|
||||
@staticmethod
|
||||
def _target_device_type(args, kwargs) -> str | None:
|
||||
device = kwargs.get("device")
|
||||
if device is None and args:
|
||||
first_arg = args[0]
|
||||
if isinstance(first_arg, torch.Tensor):
|
||||
device = first_arg.device
|
||||
elif isinstance(first_arg, (str, torch.device)):
|
||||
device = first_arg
|
||||
if device is None:
|
||||
return None
|
||||
return torch.device(device).type
|
||||
|
||||
@staticmethod
|
||||
def _target_dtype(args, kwargs) -> torch.dtype | None:
|
||||
dtype = kwargs.get("dtype")
|
||||
if dtype is not None:
|
||||
return dtype
|
||||
for arg in args:
|
||||
if isinstance(arg, torch.dtype):
|
||||
return arg
|
||||
if isinstance(arg, torch.Tensor) and arg.is_floating_point():
|
||||
return arg.dtype
|
||||
return None
|
||||
|
||||
def _ensure_inner_model(self, state_dict: dict[str, torch.Tensor] | None = None):
|
||||
if self._inner_model is not None:
|
||||
return
|
||||
@@ -115,6 +141,22 @@ class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin):
|
||||
|
||||
def decode(self, z: torch.Tensor, **kwargs):
|
||||
self._ensure_inner_model()
|
||||
if z.device.type == "mps":
|
||||
orig_device = z.device
|
||||
torch.mps.synchronize()
|
||||
self._inner_model = self._inner_model.to("cpu", dtype=torch.float32)
|
||||
torch.mps.empty_cache()
|
||||
z = z.to(device="cpu", dtype=torch.float32)
|
||||
decoded = self._inner_model.decode(z, **kwargs)
|
||||
if isinstance(decoded, DecoderOutput):
|
||||
return DecoderOutput(sample=decoded.sample.to(device=orig_device))
|
||||
if isinstance(decoded, tuple):
|
||||
sample = decoded[0].to(device=orig_device)
|
||||
return (sample, *decoded[1:])
|
||||
if isinstance(decoded, torch.Tensor):
|
||||
return decoded.to(device=orig_device)
|
||||
return decoded
|
||||
|
||||
z = z.to(dtype=self.dtype)
|
||||
if not self._spatial_parallel_decode_enabled:
|
||||
return self._inner_model.decode(z, **kwargs)
|
||||
@@ -161,6 +203,15 @@ class AutoencoderDC(nn.Module, LayerwiseOffloadableModuleMixin):
|
||||
return loaded_params
|
||||
|
||||
def to(self, *args, **kwargs):
|
||||
if self._target_device_type(args, kwargs) == "mps":
|
||||
# AutoencoderDC decode is unstable on MPS in the full Sana pipeline.
|
||||
dtype = self._target_dtype(args, kwargs)
|
||||
if dtype is not None:
|
||||
if self._inner_model is not None:
|
||||
self._inner_model = self._inner_model.to(dtype=dtype)
|
||||
return super().to(dtype=dtype)
|
||||
return self
|
||||
|
||||
if self._inner_model is not None:
|
||||
self._inner_model = self._inner_model.to(*args, **kwargs)
|
||||
return super().to(*args, **kwargs)
|
||||
|
||||
@@ -27,7 +27,6 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import (
|
||||
AttentionBackendEnum,
|
||||
current_platform,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
RealtimeCausalDiTState,
|
||||
@@ -35,6 +34,9 @@ from sglang.multimodal_gen.runtime.realtime.states import (
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_enabled as precision_autocast_enabled,
|
||||
)
|
||||
@@ -640,9 +642,9 @@ class CausalDMDDenoisingStage(DenoisingStage):
|
||||
autocast_enabled: bool,
|
||||
) -> torch.Tensor:
|
||||
with (
|
||||
torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=target_dtype,
|
||||
precision_autocast_context(
|
||||
target_dtype,
|
||||
disable_autocast=False,
|
||||
enabled=autocast_enabled,
|
||||
),
|
||||
set_forward_context(
|
||||
|
||||
@@ -32,6 +32,7 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context,
|
||||
autocast_enabled,
|
||||
resolve_precision,
|
||||
temporary_module_dtype,
|
||||
@@ -223,13 +224,12 @@ class DecodingStage(PipelineStage):
|
||||
latents = server_args.pipeline_config.preprocess_decoding(
|
||||
latents, server_args, vae=self.vae
|
||||
)
|
||||
if latents.device.type == "mps":
|
||||
torch.mps.synchronize()
|
||||
torch.mps.empty_cache()
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
with autocast_context(vae_dtype, server_args.disable_autocast):
|
||||
try:
|
||||
# TODO: make it more specific
|
||||
if server_args.pipeline_config.vae_tiling:
|
||||
|
||||
@@ -5,6 +5,7 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import gc
|
||||
import inspect
|
||||
import math
|
||||
import time
|
||||
@@ -108,6 +109,9 @@ from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import maybe_nvtx_range
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_enabled as precision_autocast_enabled,
|
||||
)
|
||||
@@ -1206,13 +1210,17 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
||||
|
||||
# 5. Advance the scheduler state with the predicted noise.
|
||||
with maybe_nvtx_range("scheduler_step", use_nvtx):
|
||||
ctx.latents = ctx.scheduler.step(
|
||||
latents_dtype = ctx.latents.dtype
|
||||
latents = ctx.scheduler.step(
|
||||
model_output=noise_pred,
|
||||
timestep=step.t_device,
|
||||
sample=ctx.latents,
|
||||
**ctx.extra_step_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if latents.dtype != latents_dtype and latents.device.type == "mps":
|
||||
latents = latents.to(latents_dtype)
|
||||
ctx.latents = latents
|
||||
|
||||
# 6. Re-apply any model-specific latent constraints after the update.
|
||||
ctx.latents = self.post_forward_for_ti2v_task(
|
||||
@@ -1337,10 +1345,13 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
||||
self._component_residency_manager.remove_nvtx_hooks_for_module(
|
||||
self.transformer
|
||||
)
|
||||
self._component_residency_manager.strategy_for.cache_clear()
|
||||
del self.transformer
|
||||
if pipeline is not None and "transformer" in pipeline.modules:
|
||||
del pipeline.modules["transformer"]
|
||||
server_args.model_loaded["transformer"] = False
|
||||
gc.collect()
|
||||
torch.mps.empty_cache()
|
||||
logger.info(
|
||||
"Memory after deallocating transformer: %s",
|
||||
torch.mps.current_allocated_memory(),
|
||||
@@ -1576,9 +1587,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
||||
use_nvtx = self._apply_nvtx_gate(ctx.is_warmup)
|
||||
|
||||
with (
|
||||
torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=ctx.target_dtype,
|
||||
precision_autocast_context(
|
||||
ctx.target_dtype,
|
||||
server_args.disable_autocast,
|
||||
enabled=ctx.autocast_enabled,
|
||||
),
|
||||
maybe_nvtx_range("denoising_loop", use_nvtx),
|
||||
@@ -1642,6 +1653,8 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
|
||||
(denoising_end_time - denoising_start_time) / len(ctx.timesteps),
|
||||
)
|
||||
|
||||
if "step" in locals():
|
||||
del step
|
||||
self._finish_active_component_use()
|
||||
|
||||
# Rollout postprocessing must run BEFORE _finalize_denoising_loop so
|
||||
|
||||
@@ -14,10 +14,12 @@ from sglang.multimodal_gen.runtime.pipelines_core.diffusion_scheduler_utils impo
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages import DenoisingStage
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import dict_to_3d_list
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -150,9 +152,9 @@ class DmdDenoisingStage(DenoisingStage):
|
||||
)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=target_dtype,
|
||||
with precision_autocast_context(
|
||||
target_dtype,
|
||||
server_args.disable_autocast,
|
||||
enabled=autocast_enabled,
|
||||
):
|
||||
attn_metadata = self._build_attn_metadata(i, batch, server_args)
|
||||
|
||||
@@ -20,10 +20,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
VerificationResult,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context,
|
||||
autocast_enabled,
|
||||
resolve_precision,
|
||||
temporary_module_dtype,
|
||||
@@ -106,11 +106,7 @@ class EncodingStage(PipelineStage):
|
||||
self.vae = vae
|
||||
|
||||
# Encode image to latents
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
with autocast_context(vae_dtype, server_args.disable_autocast):
|
||||
if server_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if server_args.vae_sp:
|
||||
|
||||
@@ -36,11 +36,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
|
||||
VerificationResult,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
align_tensor_to_module_dtype,
|
||||
autocast_context,
|
||||
autocast_enabled,
|
||||
resolve_precision,
|
||||
temporary_module_dtype,
|
||||
@@ -609,11 +609,7 @@ class LTX2ImageEncodingStage(PipelineStage):
|
||||
)
|
||||
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
with autocast_context(vae_dtype, server_args.disable_autocast):
|
||||
try:
|
||||
if server_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
@@ -650,13 +646,8 @@ class LTX2ImageEncodingStage(PipelineStage):
|
||||
vae_dtype = resolve_precision(
|
||||
server_args, "vae", precision_attr="vae_precision"
|
||||
)
|
||||
vae_autocast_enabled = autocast_enabled(vae_dtype, server_args.disable_autocast)
|
||||
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
with autocast_context(vae_dtype, server_args.disable_autocast):
|
||||
return self._condition_image_encoder(video_condition)
|
||||
|
||||
@staticmethod
|
||||
@@ -932,11 +923,7 @@ class ImageVAEEncodingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
with autocast_context(vae_dtype, server_args.disable_autocast):
|
||||
if server_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if server_args.vae_sp:
|
||||
|
||||
+5
-5
@@ -6,11 +6,11 @@ from sglang.multimodal_gen.runtime.managers.memory_managers.component_manager im
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import DecodingStage
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
align_tensor_to_module_dtype,
|
||||
autocast_context,
|
||||
autocast_enabled,
|
||||
resolve_precision,
|
||||
temporary_module_dtype,
|
||||
@@ -77,9 +77,9 @@ class LTX2AVDecodingStage(DecodingStage):
|
||||
latents, server_args, vae=self.vae
|
||||
)
|
||||
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
with autocast_context(
|
||||
dtype=vae_dtype,
|
||||
disable_autocast=server_args.disable_autocast,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
try:
|
||||
@@ -167,9 +167,9 @@ class LTX2AVDecodingStage(DecodingStage):
|
||||
should_cast_audio_vae = not audio_vae_autocast_enabled
|
||||
with (
|
||||
torch.no_grad(),
|
||||
torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
autocast_context(
|
||||
dtype=audio_vae_dtype,
|
||||
disable_autocast=server_args.disable_autocast,
|
||||
enabled=audio_vae_autocast_enabled,
|
||||
),
|
||||
):
|
||||
|
||||
+5
-2
@@ -68,6 +68,9 @@ from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs, get_global_server_args
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.perf_logger import StageProfiler
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.utils.profiler import SGLDiffusionProfiler
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
from sglang.srt.utils.common import get_compiler_backend
|
||||
@@ -971,9 +974,9 @@ class MOVADecodingStage(PipelineStage):
|
||||
batch.latents, self.video_vae
|
||||
)
|
||||
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
with precision_autocast_context(
|
||||
dtype=vae_dtype,
|
||||
disable_autocast=server_args.disable_autocast,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
if server_args.pipeline_config.vae_tiling:
|
||||
|
||||
+6
-4
@@ -50,9 +50,11 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.progressive_resolution.upsample import (
|
||||
apply_upsample,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -537,9 +539,9 @@ class ProgressiveDenoisingStage(DenoisingStage):
|
||||
# ── Stage loop ────────────────────────────────────────────────────────
|
||||
# DenoisingStage.forward() wraps its denoising loop in torch.autocast;
|
||||
# we bypass that path, so we must apply the same context here.
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=ctx.target_dtype,
|
||||
with precision_autocast_context(
|
||||
ctx.target_dtype,
|
||||
server_args.disable_autocast,
|
||||
enabled=ctx.autocast_enabled,
|
||||
):
|
||||
for stage in range(1, num_stages + 1):
|
||||
|
||||
@@ -14,11 +14,13 @@ from sglang.multimodal_gen.runtime.pipelines_core.stages.decoding import (
|
||||
from sglang.multimodal_gen.runtime.pipelines_core.stages.image_encoding import (
|
||||
ImageVAEEncodingStage,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.runtime.realtime.session import (
|
||||
BaseRealtimeState,
|
||||
)
|
||||
from sglang.multimodal_gen.runtime.server_args import ServerArgs
|
||||
from sglang.multimodal_gen.runtime.utils.precision import (
|
||||
autocast_context as precision_autocast_context,
|
||||
)
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
@@ -140,9 +142,9 @@ class CausalVaeDecodingStage(DecodingStage):
|
||||
latents, server_args, vae=self.vae
|
||||
)
|
||||
|
||||
with torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
with precision_autocast_context(
|
||||
dtype=vae_dtype,
|
||||
disable_autocast=server_args.disable_autocast,
|
||||
enabled=vae_autocast_enabled,
|
||||
):
|
||||
try:
|
||||
|
||||
@@ -1133,8 +1133,8 @@ class ServerArgs(DisaggServerArgsMixin):
|
||||
or self.vae_cpu_offload
|
||||
):
|
||||
logger.warning(
|
||||
"Disabling component CPU offload on MPS because CPU-to-MPS "
|
||||
"module relocation can produce invalid diffusion outputs."
|
||||
"Disabling component CPU offload on MPS because the component "
|
||||
"residency offload strategy is only validated on CUDA."
|
||||
)
|
||||
self.dit_cpu_offload = False
|
||||
self.text_encoder_cpu_offload = False
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from contextlib import contextmanager
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from typing import Iterator, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.multimodal_gen.runtime.platforms import current_platform
|
||||
from sglang.multimodal_gen.utils import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
@@ -70,7 +71,29 @@ def resolve_component_precision(server_args, module_name: str) -> Optional[torch
|
||||
|
||||
|
||||
def autocast_enabled(dtype: torch.dtype, disable_autocast: bool) -> bool:
|
||||
return dtype != torch.float32 and not disable_autocast
|
||||
return (
|
||||
dtype != torch.float32
|
||||
and not disable_autocast
|
||||
and current_platform.is_amp_supported()
|
||||
)
|
||||
|
||||
|
||||
def autocast_context(
|
||||
dtype: torch.dtype,
|
||||
disable_autocast: bool,
|
||||
*,
|
||||
enabled: Optional[bool] = None,
|
||||
):
|
||||
autocast_is_enabled = (
|
||||
autocast_enabled(dtype, disable_autocast) if enabled is None else enabled
|
||||
)
|
||||
if not autocast_is_enabled and current_platform.is_mps():
|
||||
return nullcontext()
|
||||
return torch.autocast(
|
||||
device_type=current_platform.device_type,
|
||||
dtype=dtype,
|
||||
enabled=autocast_is_enabled,
|
||||
)
|
||||
|
||||
|
||||
def get_module_dtype(module, default: torch.dtype = torch.float32) -> torch.dtype:
|
||||
|
||||
@@ -75,9 +75,9 @@
|
||||
},
|
||||
"zimage_image_t2i_fp8": {
|
||||
"clip_threshold": 0.97,
|
||||
"ssim_threshold": 0.95,
|
||||
"psnr_threshold": 30.0,
|
||||
"mean_abs_diff_threshold": 4.0
|
||||
"ssim_threshold": 0.84,
|
||||
"psnr_threshold": 18.0,
|
||||
"mean_abs_diff_threshold": 13.0
|
||||
},
|
||||
"qwen_image_edit_2509_ti2i": {
|
||||
"clip_threshold": 0.91,
|
||||
|
||||
@@ -2,6 +2,7 @@ import importlib.util
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from contextlib import nullcontext
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
@@ -9,11 +10,15 @@ import torch
|
||||
|
||||
|
||||
def _load_precision_module():
|
||||
stub_names = (
|
||||
package_names = (
|
||||
"sglang",
|
||||
"sglang.multimodal_gen",
|
||||
"sglang.multimodal_gen.runtime",
|
||||
"sglang.multimodal_gen.runtime.utils",
|
||||
)
|
||||
stub_names = (
|
||||
*package_names,
|
||||
"sglang.multimodal_gen.runtime.platforms",
|
||||
"sglang.multimodal_gen.utils",
|
||||
)
|
||||
missing = object()
|
||||
@@ -26,10 +31,17 @@ def _load_precision_module():
|
||||
"bf16": torch.bfloat16,
|
||||
"fp32": torch.float32,
|
||||
}
|
||||
for package_name in stub_names[:-1]:
|
||||
platforms_module = types.ModuleType("sglang.multimodal_gen.runtime.platforms")
|
||||
platforms_module.current_platform = SimpleNamespace(
|
||||
device_type="cpu",
|
||||
is_mps=lambda: False,
|
||||
is_amp_supported=lambda: True,
|
||||
)
|
||||
for package_name in package_names:
|
||||
package = types.ModuleType(package_name)
|
||||
package.__path__ = []
|
||||
sys.modules[package_name] = package
|
||||
sys.modules["sglang.multimodal_gen.runtime.platforms"] = platforms_module
|
||||
sys.modules["sglang.multimodal_gen.utils"] = utils_module
|
||||
|
||||
precision_path = (
|
||||
@@ -53,6 +65,7 @@ def _load_precision_module():
|
||||
|
||||
precision = _load_precision_module()
|
||||
align_tensor_to_module_dtype = precision.align_tensor_to_module_dtype
|
||||
autocast_context = precision.autocast_context
|
||||
autocast_enabled = precision.autocast_enabled
|
||||
get_module_dtype = precision.get_module_dtype
|
||||
precision_to_dtype = precision.precision_to_dtype
|
||||
@@ -74,6 +87,19 @@ class _ParameterDtypeWinsModule(torch.nn.Module):
|
||||
self.weight = torch.nn.Parameter(torch.ones(1, dtype=torch.float16))
|
||||
|
||||
|
||||
class _FakePlatform:
|
||||
def __init__(self, device_type: str, *, is_mps: bool, amp_supported: bool):
|
||||
self.device_type = device_type
|
||||
self._is_mps = is_mps
|
||||
self._amp_supported = amp_supported
|
||||
|
||||
def is_mps(self):
|
||||
return self._is_mps
|
||||
|
||||
def is_amp_supported(self):
|
||||
return self._amp_supported
|
||||
|
||||
|
||||
class TestDiffusionPrecisionConsistency(unittest.TestCase):
|
||||
def _server_args(self, **overrides):
|
||||
config = {
|
||||
@@ -156,6 +182,27 @@ class TestDiffusionPrecisionConsistency(unittest.TestCase):
|
||||
aligned_tokens = align_tensor_to_module_dtype(tokens, module_without_parameters)
|
||||
self.assertEqual(aligned_tokens.dtype, torch.long)
|
||||
|
||||
def test_autocast_context_honors_explicit_override(self):
|
||||
original_platform = precision.current_platform
|
||||
try:
|
||||
precision.current_platform = _FakePlatform(
|
||||
"cpu", is_mps=False, amp_supported=True
|
||||
)
|
||||
disabled_context = autocast_context(
|
||||
torch.bfloat16, disable_autocast=False, enabled=False
|
||||
)
|
||||
self.assertNotIsInstance(disabled_context, nullcontext)
|
||||
|
||||
precision.current_platform = _FakePlatform(
|
||||
"mps", is_mps=True, amp_supported=False
|
||||
)
|
||||
mps_disabled_context = autocast_context(
|
||||
torch.bfloat16, disable_autocast=False, enabled=False
|
||||
)
|
||||
self.assertIsInstance(mps_disabled_context, nullcontext)
|
||||
finally:
|
||||
precision.current_platform = original_platform
|
||||
|
||||
def test_temporary_module_dtype(self):
|
||||
module = torch.nn.Linear(2, 2).to(dtype=torch.float32)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user