[Diffusion] Qwen image edit support qknorm optimization (#16062)
This commit is contained in:
@@ -10,6 +10,7 @@ import torch.nn as nn
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from sgl_kernel import fused_add_rmsnorm, rmsnorm
|
from sgl_kernel import fused_add_rmsnorm, rmsnorm
|
||||||
|
|
||||||
|
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm, fused_inplace_qknorm
|
||||||
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
from sglang.multimodal_gen.runtime.layers.custom_op import CustomOp
|
||||||
from sglang.multimodal_gen.runtime.layers.triton_ops import (
|
from sglang.multimodal_gen.runtime.layers.triton_ops import (
|
||||||
fuse_scale_shift_kernel,
|
fuse_scale_shift_kernel,
|
||||||
@@ -421,3 +422,43 @@ class LayerNormScaleShift(nn.Module):
|
|||||||
output = output.to(x.dtype)
|
output = output.to(x.dtype)
|
||||||
|
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def apply_qk_norm(
|
||||||
|
q: torch.Tensor,
|
||||||
|
k: torch.Tensor,
|
||||||
|
q_norm: "RMSNorm",
|
||||||
|
k_norm: "RMSNorm",
|
||||||
|
head_dim: int,
|
||||||
|
allow_inplace: bool = True,
|
||||||
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
|
"""Apply QK normalization for query and key tensors.
|
||||||
|
|
||||||
|
Minimal multimodal_gen-only implementation: only the JIT fused inplace
|
||||||
|
QK-norm kernel path is supported (no fallback).
|
||||||
|
"""
|
||||||
|
|
||||||
|
batch_size = q.size(0)
|
||||||
|
q_eps = q_norm.variance_epsilon
|
||||||
|
k_eps = k_norm.variance_epsilon
|
||||||
|
# Only try fused path on CUDA and when it won't introduce implicit copies.
|
||||||
|
if (
|
||||||
|
q.is_cuda
|
||||||
|
and allow_inplace
|
||||||
|
and (q_eps == k_eps)
|
||||||
|
and can_use_fused_inplace_qknorm(head_dim)
|
||||||
|
):
|
||||||
|
fused_inplace_qknorm(
|
||||||
|
q=q.view(batch_size, -1, head_dim),
|
||||||
|
k=k.view(batch_size, -1, head_dim),
|
||||||
|
q_weight=q_norm.weight,
|
||||||
|
k_weight=k_norm.weight,
|
||||||
|
head_dim=head_dim,
|
||||||
|
eps=q_eps,
|
||||||
|
)
|
||||||
|
return q, k
|
||||||
|
|
||||||
|
raise RuntimeError(
|
||||||
|
"apply_qk_norm: fused inplace QK-norm is not applicable "
|
||||||
|
"(expected CUDA, contiguous q/k, matching eps, and supported head_dim)"
|
||||||
|
)
|
||||||
|
|||||||
@@ -17,7 +17,11 @@ from diffusers.models.normalization import AdaLayerNormContinuous
|
|||||||
|
|
||||||
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
|
from sglang.multimodal_gen.configs.models.dits.qwenimage import QwenImageDitConfig
|
||||||
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
from sglang.multimodal_gen.runtime.layers.attention import USPAttention
|
||||||
from sglang.multimodal_gen.runtime.layers.layernorm import LayerNorm, RMSNorm
|
from sglang.multimodal_gen.runtime.layers.layernorm import (
|
||||||
|
LayerNorm,
|
||||||
|
RMSNorm,
|
||||||
|
apply_qk_norm,
|
||||||
|
)
|
||||||
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
from sglang.multimodal_gen.runtime.layers.linear import ReplicatedLinear
|
||||||
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
from sglang.multimodal_gen.runtime.layers.rotary_embedding import (
|
||||||
apply_flashinfer_rope_qk_inplace,
|
apply_flashinfer_rope_qk_inplace,
|
||||||
@@ -547,14 +551,23 @@ class QwenImageCrossAttention(nn.Module):
|
|||||||
txt_value = txt_value.unflatten(-1, (self.num_heads, -1))
|
txt_value = txt_value.unflatten(-1, (self.num_heads, -1))
|
||||||
|
|
||||||
# Apply QK normalization
|
# Apply QK normalization
|
||||||
if self.norm_q is not None:
|
if self.qk_norm:
|
||||||
img_query = self.norm_q(img_query)
|
img_query, img_key = apply_qk_norm(
|
||||||
if self.norm_k is not None:
|
q=img_query,
|
||||||
img_key = self.norm_k(img_key)
|
k=img_key,
|
||||||
if self.norm_added_q is not None:
|
q_norm=self.norm_q,
|
||||||
txt_query = self.norm_added_q(txt_query)
|
k_norm=self.norm_k,
|
||||||
if self.norm_added_k is not None:
|
head_dim=img_query.shape[-1],
|
||||||
txt_key = self.norm_added_k(txt_key)
|
allow_inplace=True,
|
||||||
|
)
|
||||||
|
txt_query, txt_key = apply_qk_norm(
|
||||||
|
q=txt_query,
|
||||||
|
k=txt_key,
|
||||||
|
q_norm=self.norm_added_q,
|
||||||
|
k_norm=self.norm_added_k,
|
||||||
|
head_dim=txt_query.shape[-1],
|
||||||
|
allow_inplace=True,
|
||||||
|
)
|
||||||
|
|
||||||
# Apply RoPE
|
# Apply RoPE
|
||||||
if image_rotary_emb is not None:
|
if image_rotary_emb is not None:
|
||||||
|
|||||||
Reference in New Issue
Block a user