[Diffusion] Revert 18619 (#19510)

This commit is contained in:
Xiaoyu Zhang
2026-03-03 08:15:15 +08:00
committed by GitHub
parent 6822941514
commit 145ae518ac
@@ -30,9 +30,8 @@ from sglang.multimodal_gen.runtime.layers.layernorm import (
apply_qk_norm, apply_qk_norm,
) )
from sglang.multimodal_gen.runtime.layers.linear import ( from sglang.multimodal_gen.runtime.layers.linear import (
ColumnParallelLinear,
MergedColumnParallelLinear, MergedColumnParallelLinear,
RowParallelLinear, ReplicatedLinear,
) )
from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import ( from sglang.multimodal_gen.runtime.layers.quantization.configs.base_config import (
QuantizationConfig, QuantizationConfig,
@@ -89,109 +88,6 @@ def _get_qkv_projections(
return img_query, img_key, img_value, txt_query, txt_key, txt_value return img_query, img_key, img_value, txt_query, txt_key, txt_value
class GELU(nn.Module):
r"""
GELU activation function with tanh approximation support with `approximate="tanh"`.
Parameters:
dim_in (`int`): The number of channels in the input.
dim_out (`int`): The number of channels in the output.
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
quant_config: Quantization configure.
prefix: The name of the layer in the state dict.
"""
def __init__(
self,
dim_in: int,
dim_out: int,
approximate: str = "none",
bias: bool = True,
quant_config=None,
prefix: str = "",
):
super().__init__()
self.proj = ColumnParallelLinear(
dim_in,
dim_out,
bias=bias,
gather_output=False,
quant_config=quant_config,
prefix=f"{prefix}.proj" if prefix else "",
)
self.approximate = approximate
def forward(self, hidden_states):
hidden_states = self.proj(hidden_states)
return F.gelu(hidden_states[0], approximate=self.approximate)
class FeedForward(nn.Module):
r"""
A feed-forward layer.
Parameters:
dim (`int`): The number of channels in the input.
dim_out (`int`, *optional*): The number of channels in the output. If not given, defaults to `dim`.
mult (`int`, *optional*, defaults to 4): The multiplier to use for the hidden dimension.
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
quant_config: Quantization configure.
prefix: The name of the layer in the state dict.
"""
def __init__(
self,
dim: int,
dim_out: Optional[int] = None,
mult: int = 4,
activation_fn: str = "geglu",
inner_dim=None,
bias: bool = True,
quant_config=None,
prefix: str = "",
):
super().__init__()
if inner_dim is None:
inner_dim = int(dim * mult)
dim_out = dim_out if dim_out is not None else dim
if activation_fn == "gelu":
act_fn = GELU(dim, inner_dim, bias=bias, quant_config=None, prefix=prefix)
if activation_fn == "gelu-approximate":
act_fn = GELU(
dim,
inner_dim,
approximate="tanh",
bias=bias,
quant_config=None,
prefix=prefix,
)
else:
raise NotImplementedError(
f"activation_fn '{activation_fn}' is not supported."
)
self.net = nn.ModuleList([])
self.net.append(act_fn)
self.net.append(nn.Identity())
self.net.append(
RowParallelLinear(
inner_dim,
dim_out,
bias=True,
input_is_parallel=True,
quant_config=None,
)
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for module in self.net:
hidden_states = module(hidden_states)
return hidden_states
class QwenTimestepProjEmbeddings(nn.Module): class QwenTimestepProjEmbeddings(nn.Module):
def __init__(self, embedding_dim, use_additional_t_cond=False): def __init__(self, embedding_dim, use_additional_t_cond=False):
super().__init__() super().__init__()
@@ -624,27 +520,9 @@ class QwenImageCrossAttention(nn.Module):
) )
else: else:
# Use separate Q/K/V projections for non-quantized models # Use separate Q/K/V projections for non-quantized models
self.to_q = ColumnParallelLinear( self.to_q = ReplicatedLinear(dim, self.inner_dim, bias=True)
dim, self.to_k = ReplicatedLinear(dim, self.inner_dim, bias=True)
self.inner_dim, self.to_v = ReplicatedLinear(dim, self.inner_dim, bias=True)
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_q",
)
self.to_k = ColumnParallelLinear(
dim,
self.inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_k",
)
self.to_v = ColumnParallelLinear(
dim,
self.inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.to_v",
)
if self.qk_norm: if self.qk_norm:
self.norm_q = RMSNorm(head_dim, eps=eps) if qk_norm else nn.Identity() self.norm_q = RMSNorm(head_dim, eps=eps) if qk_norm else nn.Identity()
@@ -662,51 +540,25 @@ class QwenImageCrossAttention(nn.Module):
) )
else: else:
# Use separate Q/K/V projections for non-quantized models # Use separate Q/K/V projections for non-quantized models
self.add_q_proj = ColumnParallelLinear( self.add_q_proj = ReplicatedLinear(
added_kv_proj_dim, added_kv_proj_dim, self.inner_dim, bias=True
self.inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.add_q_proj",
) )
self.add_k_proj = ColumnParallelLinear( self.add_k_proj = ReplicatedLinear(
added_kv_proj_dim, added_kv_proj_dim, self.inner_dim, bias=True
self.inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.add_k_proj",
) )
self.add_v_proj = ColumnParallelLinear( self.add_v_proj = ReplicatedLinear(
added_kv_proj_dim, added_kv_proj_dim, self.inner_dim, bias=True
self.inner_dim,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.add_v_proj",
) )
if context_pre_only is not None and not context_pre_only: if context_pre_only is not None and not context_pre_only:
self.to_add_out = ColumnParallelLinear( self.to_add_out = ReplicatedLinear(self.inner_dim, self.dim, bias=out_bias)
self.inner_dim,
self.dim,
bias=out_bias,
gather_output=True,
quant_config=quant_config,
prefix=f"{prefix}.to_add_out",
)
else: else:
self.to_add_out = None self.to_add_out = None
if not pre_only: if not pre_only:
self.to_out = nn.ModuleList([]) self.to_out = nn.ModuleList([])
self.to_out.append( self.to_out.append(
ColumnParallelLinear( ReplicatedLinear(self.inner_dim, self.dim, bias=out_bias)
self.inner_dim,
self.dim,
bias=out_bias,
gather_output=True,
quant_config=quant_config,
prefix=f"{prefix}.to_out.0",
)
) )
else: else:
self.to_out = None self.to_out = None
@@ -848,13 +700,8 @@ class QwenImageTransformerBlock(nn.Module):
# Image processing modules # Image processing modules
self.img_mod = nn.Sequential( self.img_mod = nn.Sequential(
nn.SiLU(), nn.SiLU(),
ColumnParallelLinear( nn.Linear(
dim, dim, 6 * dim, bias=True
6 * dim,
bias=True,
gather_output=True,
quant_config=mod_quant_config,
prefix=f"{prefix}.img_mod",
), # For scale, shift, gate for norm1 and norm2 ), # For scale, shift, gate for norm1 and norm2
) )
self.img_norm1 = LayerNormScaleShift( self.img_norm1 = LayerNormScaleShift(
@@ -877,13 +724,8 @@ class QwenImageTransformerBlock(nn.Module):
# Text processing modules # Text processing modules
self.txt_mod = nn.Sequential( self.txt_mod = nn.Sequential(
nn.SiLU(), nn.SiLU(),
ColumnParallelLinear( nn.Linear(
dim, dim, 6 * dim, bias=True
6 * dim,
bias=True,
gather_output=True,
quant_config=mod_quant_config,
prefix=f"{prefix}.txt_mod",
), # For scale, shift, gate for norm1 and norm2 ), # For scale, shift, gate for norm1 and norm2
) )
self.txt_norm1 = LayerNormScaleShift( self.txt_norm1 = LayerNormScaleShift(
@@ -919,15 +761,11 @@ class QwenImageTransformerBlock(nn.Module):
dim=dim, dim=dim,
dim_out=dim, dim_out=dim,
activation_fn="gelu-approximate", activation_fn="gelu-approximate",
quant_config=quant_config,
prefix=f"{prefix}.img_mlp",
) )
self.txt_mlp = FeedForward( self.txt_mlp = FeedForward(
dim=dim, dim=dim,
dim_out=dim, dim_out=dim,
activation_fn="gelu-approximate", activation_fn="gelu-approximate",
quant_config=quant_config,
prefix=f"{prefix}.txt_mlp",
) )
if nunchaku_enabled: if nunchaku_enabled:
@@ -1043,8 +881,8 @@ class QwenImageTransformerBlock(nn.Module):
modulate_index: Optional[List[int]] = None, modulate_index: Optional[List[int]] = None,
) -> Tuple[torch.Tensor, torch.Tensor]: ) -> Tuple[torch.Tensor, torch.Tensor]:
# Get modulation parameters for both streams # Get modulation parameters for both streams
img_mod_params = self.img_mod[1](temb_img_silu)[0] # [B, 6*dim] img_mod_params = self.img_mod[1](temb_img_silu) # [B, 6*dim]
txt_mod_params = self.txt_mod[1](temb_txt_silu)[0] # [B, 6*dim] txt_mod_params = self.txt_mod[1](temb_txt_silu) # [B, 6*dim]
if ( if (
self.quant_config is not None self.quant_config is not None
@@ -1107,7 +945,7 @@ class QwenImageTransformerBlock(nn.Module):
gate_x=img_gate1, gate_x=img_gate1,
residual_x=hidden_states, residual_x=hidden_states,
) )
img_mlp_output = self.img_mlp(img_modulated2)[0] img_mlp_output = self.img_mlp(img_modulated2)
if img_mlp_output.dim() == 2: if img_mlp_output.dim() == 2:
img_mlp_output = img_mlp_output.unsqueeze(0) img_mlp_output = img_mlp_output.unsqueeze(0)
@@ -1123,7 +961,7 @@ class QwenImageTransformerBlock(nn.Module):
scale=txt_scale2, scale=txt_scale2,
) )
txt_gate2 = txt_gate2_raw.unsqueeze(1) txt_gate2 = txt_gate2_raw.unsqueeze(1)
txt_mlp_output = self.txt_mlp(txt_modulated2)[0] txt_mlp_output = self.txt_mlp(txt_modulated2)
if txt_mlp_output.dim() == 2: if txt_mlp_output.dim() == 2:
txt_mlp_output = txt_mlp_output.unsqueeze(0) txt_mlp_output = txt_mlp_output.unsqueeze(0)