chore: clean diffusion dead code (#29789)

This commit is contained in:
Xiaoyu Zhang
2026-07-01 15:42:16 +08:00
committed by GitHub
parent 308d89e042
commit 8205aa3603
41 changed files with 37 additions and 108 deletions
@@ -164,7 +164,6 @@ class SGLDiffusionGenerator:
for k in diffusers_keys:
if k in sd:
new_sd[diffusers_keys[k]] = sd.pop(k)
offload_device = model_management.unet_offload_device()
if dtype is None:
unet_dtype = model_management.unet_dtype(
model_params=parameters,
@@ -26,8 +26,6 @@ class QwenImageExecutor(SGLDiffusionExecutor):
def _pack_latents(self, x):
"""Process hidden states for QwenImage model."""
bs, c, t, h, w = x.shape
patch_size = self.patch_size
latents = comfy.ldm.common_dit.pad_to_patch_size(
x, (1, self.patch_size, self.patch_size)
)
@@ -641,7 +641,7 @@ class SGLDiffusionServerSetLora:
# Call API
try:
response = sgld_client.set_lora(**request_params)
sgld_client.set_lora(**request_params)
return (sgld_client,)
except Exception as e:
raise RuntimeError(f"Failed to set LoRA adapter: {str(e)}")
@@ -685,7 +685,7 @@ class SGLDiffusionServerUnsetLora:
):
"""Unset LoRA adapter using SGLang Diffusion API."""
try:
response = sgld_client.unset_lora(target=target)
sgld_client.unset_lora(target=target)
return (sgld_client,)
except Exception as e:
raise RuntimeError(f"Failed to unset LoRA adapter: {str(e)}")
@@ -82,13 +82,6 @@ def test_comfyui_flux_pipeline_direct() -> None:
req.raw_latent_shape = torch.tensor(hidden_states.shape, dtype=torch.long)
clip_dim = 768
dummy_clip_embedding = torch.zeros(
batch_size,
77,
clip_dim,
device="cuda",
dtype=torch.bfloat16,
)
req.prompt_embeds = [pooled_projections, encoder_hidden_states]
if req.guidance_scale > 1.0:
@@ -149,7 +142,7 @@ def test_comfyui_flux_pipeline_direct() -> None:
noise_pred.dtype == torch.bfloat16
), f"noise_pred should be bfloat16, got {noise_pred.dtype}"
print(f"✓ Successfully retrieved noise_pred from OutputBatch!")
print("✓ Successfully retrieved noise_pred from OutputBatch!")
print(f" noise_pred shape: {noise_pred.shape}")
print(f" noise_pred dtype: {noise_pred.dtype}")
print(f" noise_pred device: {noise_pred.device}")
@@ -123,7 +123,7 @@ def test_comfyui_qwen_image_edit_pipeline_direct() -> None:
noise_pred.dtype == torch.bfloat16
), f"noise_pred should be bfloat16, got {noise_pred.dtype}"
print(f"✓ Successfully retrieved noise_pred from OutputBatch (Edit Mode)!")
print("✓ Successfully retrieved noise_pred from OutputBatch (Edit Mode)!")
print(f" noise_pred shape: {noise_pred.shape}")
print(f" noise_pred dtype: {noise_pred.dtype}")
print(f" noise_pred device: {noise_pred.device}")
@@ -107,7 +107,7 @@ def test_comfyui_qwen_image_pipeline_direct() -> None:
noise_pred.dtype == torch.bfloat16
), f"noise_pred should be bfloat16, got {noise_pred.dtype}"
print(f"✓ Successfully retrieved noise_pred from OutputBatch!")
print("✓ Successfully retrieved noise_pred from OutputBatch!")
print(f" noise_pred shape: {noise_pred.shape}")
print(f" noise_pred dtype: {noise_pred.dtype}")
print(f" noise_pred device: {noise_pred.device}")
@@ -109,7 +109,7 @@ def test_comfyui_zimage_pipeline_direct() -> None:
noise_pred.dtype == torch.bfloat16
), f"noise_pred should be bfloat16, got {noise_pred.dtype}"
print(f"✓ Successfully retrieved noise_pred from OutputBatch!")
print("✓ Successfully retrieved noise_pred from OutputBatch!")
print(f" noise_pred shape: {noise_pred.shape}")
print(f" noise_pred dtype: {noise_pred.dtype}")
print(f" noise_pred device: {noise_pred.device}")
@@ -174,8 +174,8 @@ def run_sgl_diffusion_webui(server_args: ServerArgs):
with gr.Blocks() as demo:
gr.Markdown("# 🚀 SGLang Diffusion Application")
with gr.Row():
launched_model_box = gr.Textbox(label="Model", value=server_args.model_path)
task_name_box = gr.Textbox(label="Task name", value=task_name)
gr.Textbox(label="Model", value=server_args.model_path)
gr.Textbox(label="Task name", value=task_name)
with gr.Row():
with gr.Column(scale=4):
@@ -567,7 +567,7 @@ async def benchmark(args):
else:
raise ValueError(f"Unknown dataset: {args.dataset}")
logger.info(f"Loading requests...")
logger.info("Loading requests...")
requests_list = dataset.get_requests()
logger.info(f"Prepared {len(requests_list)} requests from {args.dataset} dataset.")
@@ -93,7 +93,6 @@ class JoyImageEditPipelineConfig(ImagePipelineConfig):
while height >= step_height:
if max(height, width) / min(height, width) <= max_ratio:
ratio = height / width
buckets.append((1, 1, 1, height, width))
# Try to increase width or decrease height
if height * (width + step_width) <= target_pixels:
@@ -1,2 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
from .vmoba import (
moba_attn_varlen as moba_attn_varlen,
process_moba_input as process_moba_input,
process_moba_output as process_moba_output,
)
__all__ = ["moba_attn_varlen", "process_moba_input", "process_moba_output"]
@@ -239,7 +239,6 @@ def _select_threshold_overall(
Normalization and sorting happen across all valid entries.
"""
C, H, S = gate.shape
CHS = C * H * S
eps = 1e-6
# LSEstyle normalization globally across all valid entries
@@ -688,7 +687,6 @@ def moba_attn_varlen(
.expand(-1, num_head, -1)
)
amplification_factor = 1e9 # Example factor; adjust as needed.
origin_gate = gate.clone()
gate = gate.clone()
if select_mode == "topk":
gate[gate_self_chunk_mask] += amplification_factor
@@ -963,7 +961,6 @@ def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
random.seed(0)
torch.manual_seed(0)
torch.cuda.manual_seed(0)
device = torch.cuda.current_device()
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(
dtype=dtype, device="cuda"
@@ -9,7 +9,7 @@ Adapted from Hunyuan3D-2: https://github.com/Tencent/Hunyuan3D-2
from __future__ import annotations
import os
from typing import List, Tuple
from typing import Tuple
import torch
from sglang.multimodal_gen.csrc.render import load_extension_with_recovery
@@ -1262,11 +1262,11 @@ class SequenceParallelGroupCoordinator(GroupCoordinator):
ring_group = kwargs.get("ring_group", None)
if ulysses_group is None:
raise RuntimeError(
f"Please pass argument 'ulysses_group' when calling init func of SequenceParallelGroupCoordinator"
"Please pass argument 'ulysses_group' when calling init func of SequenceParallelGroupCoordinator"
)
if ring_group is None:
raise RuntimeError(
f"Please pass argument 'ring_group' when calling init func of SequenceParallelGroupCoordinator"
"Please pass argument 'ring_group' when calling init func of SequenceParallelGroupCoordinator"
)
self.ulysses_group = ulysses_group
self.ring_group = ring_group
@@ -54,7 +54,6 @@ from sglang.multimodal_gen.runtime.utils.trace_wrapper import (
logger = init_logger(__name__)
# TODO: move to somewhere appropriate
try:
# Set the start method to 'spawn' to avoid CUDA errors in forked processes.
# This must be done at the top level of the module, before any CUDA context
@@ -582,7 +582,7 @@ def save_materialized_output(
if not save_output:
return
if not save_file_path:
logger.info(f"No output path provided, output not saved")
logger.info("No output path provided, output not saved")
return
os.makedirs(os.path.dirname(save_file_path), exist_ok=True)
@@ -200,7 +200,7 @@ def launch_server(server_args: ServerArgs, launch_http_server: bool = True):
http_server_process = mp.Process(
target=launch_http_server_only,
args=(server_args,),
name=f"sglang-diffusion-webui",
name="sglang-diffusion-webui",
daemon=True,
)
http_server_process.start()
@@ -100,7 +100,7 @@ class Fp8Config(QuantizationConfig):
if weight_block_size is not None:
if not is_checkpoint_fp8_serialized:
raise ValueError(
f"The block-wise quantization only supports fp8-serialized checkpoint for now."
"The block-wise quantization only supports fp8-serialized checkpoint for now."
)
if len(weight_block_size) != 2:
raise ValueError(
@@ -536,7 +536,6 @@ class DualTowerConditionalBridge(
A tuple of ((cos_v, sin_v), (cos_a, sin_a)).
"""
f_v, h, w = grid_size
L_v = f_v * h * w
L_a = int(audio_steps)
device = device or next(self.parameters()).device
@@ -297,8 +297,6 @@ class ErnieImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin)
self.out_channels = arch.out_channels
self.inner_dim = self.hidden_size
tp_size = get_tp_world_size()
self.x_embedder = nn.ModuleDict(
{
"proj": nn.Conv2d(
@@ -901,7 +901,7 @@ class GlmImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
]
] = None,
###
guidance: torch.Tensor = None, # TODO: this should probably be removed
guidance: torch.Tensor = None,
) -> Tuple[torch.Tensor]:
if kv_caches is not None:
kv_caches.set_mode(kv_caches_mode)
@@ -655,7 +655,6 @@ class HeliosTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
latents_history_long=None,
**kwargs,
) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
@@ -421,7 +421,7 @@ class JoyTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
self.hidden_size,
self.out_channels * math.prod(self.patch_size),
quant_config=quant_config,
prefix=f"proj_out",
prefix="proj_out",
)
self.__post_init__()
@@ -1048,7 +1048,6 @@ class QwenImageTransformerBlock(nn.Module):
# Apply attention gates and add residual (like in Megatron)
# - residual_out = gate_x * x + residual_x
# - x = norm(residual_out) * (1 + scale) + shift
# TODO: clean code here
is_scale_residual = isinstance(norm_module, ScaleResidualLayerNormScaleShift)
shift, scale, gate = mod_params.chunk(3, dim=-1)
@@ -1389,7 +1388,7 @@ class QwenImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
txt_seq_lens: Optional[List[int]] = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] = None,
additional_t_cond: Optional[torch.Tensor] = None,
guidance: torch.Tensor = None, # TODO: this should probably be removed
guidance: torch.Tensor = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
controlnet_block_samples=None,
return_dict: bool = True,
@@ -2839,11 +2839,9 @@ class BidirectionalGDNUCPESinglePathLiteLA(nn.Module):
kv_proj = apply_kv(torch.cat([k_bhnd, v_bhnd], dim=1))
k_proj, v_proj = torch.chunk(kv_proj, chunks=2, dim=1)
q_pre_dn = q_bhnd.permute(0, 1, 3, 2)
q_dn = q_proj.permute(0, 1, 3, 2)
k_pre_dn = k_bhnd.permute(0, 1, 3, 2)
k_dn = k_proj.permute(0, 1, 3, 2)
v_pre_dn = v_bhnd.permute(0, 1, 3, 2)
v_dn = v_proj.permute(0, 1, 3, 2)
# No RMS downscale here: full post-UCPE q/k/v feed the scan; inflation
@@ -939,7 +939,7 @@ class WanTransformer3DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
config.out_channels * math.prod(config.patch_size),
bias=True,
gather_output=True,
prefix=f"proj_out",
prefix="proj_out",
quant_config=quant_config,
)
self.scale_shift_table = nn.Parameter(
@@ -954,7 +954,6 @@ class ZImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin):
cap_feats = self._as_caption_list(encoder_hidden_states)
timestep = 1000.0 - timestep
t = timestep
device = x[0].device
t = self.t_embedder(t)
adaln_input = t.to(dtype=x[0].dtype)
(
@@ -57,7 +57,7 @@ from sglang.multimodal_gen.runtime.utils.common import add_prefix
# limitations under the License.
"""Inference-only Qwen2-VL model compatible with HuggingFace weights."""
import logging
from typing import Callable, Iterable, Optional, Tuple, Union
from typing import Iterable, Optional, Tuple, Union
try:
from typing import Unpack # type: ignore[attr-defined]
@@ -76,7 +76,6 @@ from transformers.models.qwen2_5_vl.modeling_qwen2_5_vl import (
Qwen2_5_VLModelOutputWithPast,
Qwen2_5_VLRotaryEmbedding,
apply_multimodal_rotary_pos_emb,
eager_attention_forward,
)
logger = logging.getLogger(__name__)
@@ -260,26 +259,10 @@ class Qwen2_5_VLAttention(nn.Module):
key_states, value_states, self.layer_idx, cache_kwargs
)
attention_interface: Callable = eager_attention_forward
# if self.config._attn_implementation != "eager":
# attention_interface = ALL_ATTENTION_FUNCTIONS["sdpa"]
query_states = query_states.transpose(1, 2)
key_states = key_states.transpose(1, 2)
value_states = value_states.transpose(1, 2)
attn_output = self.attn(query_states, key_states, value_states)
#
# attn_output, attn_weights = attention_interface(
# self,
# query_states,
# key_states,
# value_states,
# attention_mask,
# dropout=0.0 if not self.training else self.attention_dropout,
# scaling=self.scaling,
# sliding_window=self.sliding_window,
# position_ids=position_ids, # pass positions for FA2
# **kwargs,
# )
attn_output = attn_output.reshape(bsz, q_len, -1).contiguous()
attn_output = _linear_output(self.o_proj, attn_output)
@@ -455,13 +455,6 @@ class AutoencoderKL(nn.Module, LayerwiseOffloadableModuleMixin):
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
deprecation_message = (
"The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the "
"implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able "
"to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value."
)
# deprecate("tiled_encode", "1.0.0", deprecation_message, standard_warn=False)
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
@@ -64,7 +64,6 @@ class AutoencoderKLFlux2(ParallelTiledVAE):
latent_channels: int = arch_config.latent_channels
norm_num_groups: int = arch_config.norm_num_groups
sample_size: int = arch_config.sample_size
force_upcast: bool = arch_config.force_upcast
use_quant_conv: bool = arch_config.use_quant_conv
use_post_quant_conv: bool = arch_config.use_post_quant_conv
mid_block_add_attention: bool = arch_config.mid_block_add_attention
@@ -411,12 +410,6 @@ class AutoencoderKLFlux2(ParallelTiledVAE):
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
`tuple` is returned.
"""
deprecation_message = (
"The tiled_encode implementation supporting the `return_dict` parameter is deprecated. In the future, the "
"implementation of this method will be replaced with that of `_tiled_encode` and you will no longer be able "
"to pass `return_dict`. You will also have to create a `DiagonalGaussianDistribution()` from the returned value."
)
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
row_limit = self.tile_latent_min_size - blend_extent
@@ -556,8 +556,6 @@ def generate_dense_grid_points(
def extract_near_surface_volume_fn(input_tensor: torch.Tensor, alpha: float):
"""Extract near-surface voxels for hierarchical decoding."""
device = input_tensor.device
val = input_tensor + alpha
valid_mask = val > -9000
@@ -833,7 +831,6 @@ class FlashVDMVolumeDecoding:
dtype = latents.dtype
resolutions = []
orig_resolution = octree_resolution
if octree_resolution < min_resolution:
resolutions.append(octree_resolution)
while octree_resolution >= min_resolution:
@@ -17,8 +17,10 @@ from sglang.multimodal_gen.runtime.pipelines_core.lora_pipeline import LoRAPipel
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import Req
from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
maybe_download_model,
verify_model_config_and_directory,
maybe_download_model as maybe_download_model,
)
from sglang.multimodal_gen.runtime.utils.hf_diffusers_utils import (
verify_model_config_and_directory as verify_model_config_and_directory,
)
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -888,7 +888,6 @@ class MOVADenoisingStage(PipelineStage):
"""
min_layers = min(len(visual_dit.blocks), len(self.audio_dit.blocks))
visual_layers = len(visual_dit.blocks)
sp_size = get_sp_world_size()
# Build RoPE frequencies for cross-attention if needed (only used when SP == 1)
# When SP > 1, we rebuild freqs inside the loop after gathering full sequences
@@ -579,7 +579,6 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
# 5. Prepare timesteps
scheduler = self.scheduler
sigmas = np.linspace(1.0, 0, num_inference_steps + 1)[:-1]
image_seq_len = latents.shape[1]
base_seqlen = 256 * 256 / 16 / 16
mu = (image_latents.shape[1] / base_seqlen) ** 0.5
timesteps, num_inference_steps = retrieve_timesteps(
@@ -594,8 +593,6 @@ the image\n<|vision_start|><|image_pad|><|vision_end|><|im_end|>\n<|im_start|>as
negative_txt_seq_lens = _seq_lens_from_optional_mask(
negative_prompt_embeds, negative_prompt_embeds_mask
)
is_rgb = torch.tensor([0]).to(device=device, dtype=torch.long)
batch.prompt_embeds = [prompt_embeds]
batch.prompt_embeds_mask = [prompt_embeds_mask]
batch.prompt_seq_lens = [txt_seq_lens]
@@ -194,7 +194,6 @@ def get_quant_config(
return quant_cls.from_config(hf_quant_config)
model_name_or_path = model_config["model_path"]
is_local = os.path.isdir(model_name_or_path)
hf_folder = model_name_or_path
possible_config_filenames = quant_cls.get_config_filenames()
@@ -105,7 +105,7 @@ def main():
"1" # Skip consistency checks in GT gen mode
)
logger.info(f"GT generation mode enabled")
logger.info("GT generation mode enabled")
logger.info(f"Output directory: {out_dir}")
# Resolve test files path (same as run_suite.py)
@@ -112,7 +112,6 @@ def _run_case(case: DiffusionTestCase) -> dict:
ctx = mgr.start()
try:
sp = case.sampling_params
output_size = os.environ.get("SGLANG_TEST_OUTPUT_SIZE", sp.output_size)
client = _openai_client(ctx.port)
gen = get_generate_fn(
model_path=case.server_args.model_path,
@@ -110,11 +110,11 @@ def upload_file_to_slack(
token = os.environ.get("SGLANG_DIFFUSION_SLACK_TOKEN")
if not token:
logger.info(f"Slack upload failed: no token")
logger.info("Slack upload failed: no token")
return False
if not file_path or not os.path.exists(file_path):
logger.info(f"Slack upload failed: no file path")
logger.info("Slack upload failed: no file path")
return False
origin_paths = []
@@ -101,7 +101,6 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
default_port = get_dynamic_server_port()
port = int(os.environ.get("SGLANG_TEST_SERVER_PORT", default_port))
sampling_params = case.sampling_params
extra_args = os.environ.get("SGLANG_TEST_SERVE_ARGS", "")
extra_args = f"--model-type diffusion {extra_args}".strip()
@@ -114,7 +113,7 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
extra_args += f" --ulysses-degree {server_args.ulysses_degree}"
if server_args.dit_layerwise_offload:
extra_args += f" --dit-layerwise-offload true"
extra_args += " --dit-layerwise-offload true"
if server_args.dit_offload_prefetch_size:
extra_args += (
@@ -122,7 +121,7 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
)
if server_args.text_encoder_cpu_offload:
extra_args += f" --text-encoder-cpu-offload"
extra_args += " --text-encoder-cpu-offload"
if server_args.ring_degree is not None:
extra_args += f" --ring-degree {server_args.ring_degree}"
@@ -204,17 +203,6 @@ def diffusion_server(case: DiffusionTestCase) -> ServerContext:
)
raise
try:
# Reconstruct output size for OpenAI API
# Allow override via environment variable (useful for AMD where large resolutions can cause GPU hang)
output_size = os.environ.get(
"SGLANG_TEST_OUTPUT_SIZE", sampling_params.output_size
)
except Exception as exc:
logger.error("Warm-up failed for %s: %s", case.id, exc)
ctx.cleanup()
raise
try:
yield ctx
finally:
@@ -425,14 +413,11 @@ class DiffusionServerBase:
if not is_baseline_generation_mode:
missing_scenario = True
# Check for missing estimated_full_test_time_s
missing_estimated_time = False
if (
not missing_scenario
and not is_baseline_generation_mode
and scenario.estimated_full_test_time_s is None
):
missing_estimated_time = True
_MISSING_ESTIMATED_TIME_CASES.add(case.id)
validator_name = case.server_args.custom_validator or "default"
@@ -1065,7 +1065,7 @@ class TestIdeogram4(unittest.TestCase):
self.assertTrue(layer.mlp.down_proj.input_is_parallel)
self.assertTrue(layer.mlp.down_proj.reduce_results)
def test_denoise_and_decode_shape_smoke(self):
def test_denoise_and_decode_shape_check(self):
import sglang.multimodal_gen.runtime.server_args as server_args_module
cfg = Ideogram4PipelineConfig()