From c54dc4582fddc11151567809c53e8cb421c83b7d Mon Sep 17 00:00:00 2001 From: Elizaveta Martirosian Date: Fri, 7 Aug 2026 06:19:58 +0300 Subject: [PATCH] [Diffusion]Skipping tensor copying for non-BCG GLM-Image workflows (#33688) Co-authored-by: Elizaveta Martirosian Co-authored-by: ronnie_zheng --- .../multimodal_gen/runtime/models/dits/glm_image.py | 12 +++++++++++- .../test/scripts/gen_perf_baselines.py | 6 +++++- .../test/server/ascend/perf_baselines_npu.json | 2 +- 3 files changed, 17 insertions(+), 3 deletions(-) diff --git a/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py b/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py index d6948d477..9ce0fcddd 100644 --- a/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py +++ b/python/sglang/multimodal_gen/runtime/models/dits/glm_image.py @@ -56,6 +56,9 @@ from sglang.multimodal_gen.runtime.platforms import ( current_platform, ) from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph import ( + is_in_breakable_cuda_graph, +) logger = init_logger(__name__) @@ -977,8 +980,15 @@ class GlmImageTransformer2DModel(CachableDiT, LayerwiseOffloadableModuleMixin): hidden_states = self.image_projector(hidden_states) encoder_hidden_states = self.glyph_projector(encoder_hidden_states) + prior_embedding = self.prior_token_embedding(prior_token_id) - prior_embedding = prior_embedding.masked_fill(prior_token_drop.unsqueeze(-1), 0) + if is_in_breakable_cuda_graph(): + prior_embedding = prior_embedding.masked_fill( + prior_token_drop.unsqueeze(-1), 0 + ) + else: + prior_embedding[prior_token_drop] *= 0.0 + prior_hidden_states = self.prior_projector(prior_embedding) # SP: when latents are H-sharded, hidden_states has fewer patches than prior_hidden_states. # Shard prior_hidden_states along seq dim to match (prior is row-major, same as latent patches). diff --git a/python/sglang/multimodal_gen/test/scripts/gen_perf_baselines.py b/python/sglang/multimodal_gen/test/scripts/gen_perf_baselines.py index c0a4ba2ce..80b7e36bb 100644 --- a/python/sglang/multimodal_gen/test/scripts/gen_perf_baselines.py +++ b/python/sglang/multimodal_gen/test/scripts/gen_perf_baselines.py @@ -8,6 +8,7 @@ from pathlib import Path from openai import OpenAI +from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.test.server.test_server_utils import ( ServerManager, get_generate_fn, @@ -23,7 +24,10 @@ from sglang.multimodal_gen.test.test_utils import ( def _all_cases() -> list[DiffusionTestCase]: - import sglang.multimodal_gen.test.server.testcase_configs as cfg + if current_platform.is_npu(): + import sglang.multimodal_gen.test.server.ascend.testcase_configs_npu as cfg + else: + import sglang.multimodal_gen.test.server.testcase_configs as cfg cases: list[DiffusionTestCase] = [] for _, v in inspect.getmembers(cfg): diff --git a/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json b/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json index 75a072685..1708db6bb 100644 --- a/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json +++ b/python/sglang/multimodal_gen/test/server/ascend/perf_baselines_npu.json @@ -395,7 +395,7 @@ "GlmImageAR": 69249.94, "GlmImageBeforeDenoisingStage": 61.78, "DenoisingStage": 17750.92, - "DecodingStage": 795.68 + "DecodingStage": 277.36 }, "denoise_step_ms": { "0": 272.46,