[Diffusion][NPU][GPU] Fix SANA model execution error (#24798)

This commit is contained in:
Thomas
2026-05-11 08:41:36 +03:00
committed by GitHub
parent 9ec2880eca
commit ed70226ec1
3 changed files with 24 additions and 15 deletions
@@ -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
@@ -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,
@@ -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):