From ed70226ec18fdf077f097c97131253046c4232e3 Mon Sep 17 00:00:00 2001 From: Thomas Date: Mon, 11 May 2026 13:41:36 +0800 Subject: [PATCH] [Diffusion][NPU][GPU] Fix SANA model execution error (#24798) --- .../multimodal_gen/runtime/layers/activation.py | 9 +++++++++ .../runtime/models/encoders/gemma2.py | 17 ++++------------- .../scheduling_dpm_solver_multistep.py | 13 +++++++++++-- 3 files changed, 24 insertions(+), 15 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/layers/activation.py b/python/sglang/multimodal_gen/runtime/layers/activation.py index 2795bd9f0..9513420e5 100644 --- a/python/sglang/multimodal_gen/runtime/layers/activation.py +++ b/python/sglang/multimodal_gen/runtime/layers/activation.py @@ -82,6 +82,15 @@ class GeluAndMul(CustomOp): def forward_cuda(self, *args, **kwargs) -> Any: return self.forward_native(*args, **kwargs) + def forward_npu(self, x: torch.Tensor) -> torch.Tensor: + y_npu, _ = torch_npu.npu_geglu( + x, + dim=-1, + approximate=1 if self.approximate == "tanh" else 0, + activate_left=True, + ) + return y_npu + def forward_native(self, x: torch.Tensor) -> torch.Tensor: """PyTorch-native implementation equivalent to forward().""" d = x.shape[-1] // 2 diff --git a/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py b/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py index af18c4f2a..1832c709e 100644 --- a/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py +++ b/python/sglang/multimodal_gen/runtime/models/encoders/gemma2.py @@ -178,31 +178,22 @@ class Gemma2Attention(nn.Module): key = k.transpose(1, 2) value = v.transpose(1, 2) - attn_mask = torch.zeros( - (seq_len, seq_len), device=hidden_states.device, dtype=torch.float32 - ) - causal = torch.triu( + attn_mask = torch.tril( torch.ones( (seq_len, seq_len), device=hidden_states.device, dtype=torch.bool - ), - diagonal=1, + ) ) - attn_mask = attn_mask.masked_fill(causal, float("-inf")) if self.is_sliding and self.sliding_window is not None: idx = torch.arange(seq_len, device=hidden_states.device) dist = idx[None, :] - idx[:, None] too_far = dist > self.sliding_window - attn_mask = attn_mask.masked_fill(too_far, float("-inf")) + attn_mask = attn_mask.masked_fill(too_far, False) if attention_mask is not None: - key_pad = ~attention_mask.to(torch.bool) attn_mask = attn_mask[None, None, :, :].expand( batch_size, 1, seq_len, seq_len ) - attn_mask = attn_mask.masked_fill( - key_pad[:, None, None, :].expand(batch_size, 1, seq_len, seq_len), - float("-inf"), - ) + attn_mask = attn_mask & attention_mask.to(torch.bool)[:, None, None, :] attn_kwargs = { "attn_mask": attn_mask, diff --git a/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_dpm_solver_multistep.py b/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_dpm_solver_multistep.py index 8f72f22f0..03c695909 100644 --- a/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_dpm_solver_multistep.py +++ b/python/sglang/multimodal_gen/runtime/models/schedulers/scheduling_dpm_solver_multistep.py @@ -115,9 +115,18 @@ class DPMSolverMultistepScheduler(SchedulerMixin, ConfigMixin, BaseScheduler): model_output: torch.Tensor, timestep: int, sample: torch.Tensor, - **kwargs, + generator: torch.Generator | None = None, + variance_noise: torch.Tensor | None = None, + return_dict: bool = True, ): - return self._inner.step(model_output, timestep, sample, **kwargs) + return self._inner.step( + model_output, + timestep, + sample, + generator=generator, + variance_noise=variance_noise, + return_dict=return_dict, + ) @property def sigmas(self):