[AMD] Replace triton rotary_emb with aiter rotary_emb for Wan2.2 denoise (#22422)

This commit is contained in:
jacky.cheng
2026-04-09 18:21:02 -07:00
committed by GitHub
parent 5638d40f3a
commit d283808457
2 changed files with 55 additions and 0 deletions
+44
View File
@@ -50,6 +50,9 @@ from sglang.multimodal_gen.runtime.layers.visual_embedding import (
)
from sglang.multimodal_gen.runtime.managers.forward_context import get_forward_context
from sglang.multimodal_gen.runtime.models.dits.base import CachableDiT
from sglang.multimodal_gen.runtime.models.utils import (
_use_aiter,
)
from sglang.multimodal_gen.runtime.platforms import (
AttentionBackendEnum,
current_platform,
@@ -62,6 +65,9 @@ from sglang.srt.utils import add_prefix
logger = init_logger(__name__)
_is_cuda = current_platform.is_cuda()
if _use_aiter:
from aiter.ops.rope import rope_cached_2c_fwd_inplace
class WanImageEmbedding(torch.nn.Module):
@@ -547,6 +553,25 @@ class WanTransformerBlock(nn.Module):
query, key = apply_flashinfer_rope_qk_inplace(
query, key, cos_sin_cache, is_neox=False
)
elif _use_aiter:
query_shape = query.shape
key_shape = key.shape
num_tokens = query.shape[:-2].numel()
q_sbhd = query.view(num_tokens, 1, query_shape[-2], query_shape[-1])
k_sbhd = key.view(num_tokens, 1, key_shape[-2], key_shape[-1])
cos_sbhd = cos.contiguous().view(num_tokens, 1, 1, -1)
sin_sbhd = sin.contiguous().view(num_tokens, 1, 1, -1)
rope_cached_2c_fwd_inplace(
q_sbhd,
k_sbhd,
cos_sbhd,
sin_sbhd,
1, # GPTJ rotate style
True, # reuse_freqs_front_part
False, # nope_first
)
query = q_sbhd.view(query_shape)
key = k_sbhd.view(key_shape)
else:
query, key = _apply_rotary_emb(
query, cos, sin, is_neox_style=False
@@ -774,6 +799,25 @@ class WanTransformerBlock_VSA(nn.Module):
query, key = apply_flashinfer_rope_qk_inplace(
query, key, cos_sin_cache, is_neox=False
)
elif _use_aiter:
query_shape = query.shape
key_shape = key.shape
num_tokens = query.shape[:-2].numel()
q_sbhd = query.view(num_tokens, 1, query_shape[-2], query_shape[-1])
k_sbhd = key.view(num_tokens, 1, key_shape[-2], key_shape[-1])
cos_sbhd = cos.contiguous().view(num_tokens, 1, 1, -1)
sin_sbhd = sin.contiguous().view(num_tokens, 1, 1, -1)
rope_cached_2c_fwd_inplace(
q_sbhd,
k_sbhd,
cos_sbhd,
sin_sbhd,
1, # GPTJ rotate style
True, # reuse_freqs_front_part
False, # nope_first
)
query = q_sbhd.view(query_shape)
key = k_sbhd.view(key_shape)
else:
query, key = _apply_rotary_emb(
query, cos, sin, is_neox_style=False
@@ -8,6 +8,17 @@ from typing import Any
import torch
from sglang.srt.utils import (
get_bool_env_var,
is_gfx95_supported,
is_hip,
)
_is_hip = is_hip()
_is_gfx95_supported = is_gfx95_supported()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_use_aiter_gfx95 = _use_aiter and _is_gfx95_supported
def set_weight_attrs(
weight: torch.Tensor,