From 144038fbae146e48d3a3c873898a76418348b8f8 Mon Sep 17 00:00:00 2001 From: Mick Date: Tue, 28 Apr 2026 20:39:23 +0800 Subject: [PATCH] [diffusion] chore: change default seed to 42 (#23836) --- .github/workflows/diffusion-ci-gt-gen.yml | 74 ++++++++++++--- .github/workflows/pr-test-multimodal-gen.yml | 2 +- .github/workflows/pr-test.yml | 2 +- .../runtime/disaggregation/scheduler_mixin.py | 15 +++- .../runtime/entrypoints/http_server.py | 2 - .../runtime/entrypoints/openai/image_api.py | 2 +- .../runtime/entrypoints/openai/protocol.py | 4 +- .../runtime/entrypoints/openai/video_api.py | 2 +- .../entrypoints/post_training/io_struct.py | 2 +- .../pipelines_core/stages/denoising.py | 5 +- .../pipelines_core/stages/hunyuan3d_shape.py | 44 ++++----- .../multimodal_gen/test/partitioning.py | 32 +++++++ .../sglang/multimodal_gen/test/run_suite.py | 37 +++++--- .../test/scripts/gen_diffusion_ci_outputs.py | 89 ++++++++++++++++--- .../test/unit/test_disagg_trace.py | 39 ++++++++ .../diffusion/compute_diffusion_partitions.py | 89 ++++++++++--------- .../utils/diffusion/diffusion_case_parser.py | 14 ++- 17 files changed, 339 insertions(+), 115 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/partitioning.py diff --git a/.github/workflows/diffusion-ci-gt-gen.yml b/.github/workflows/diffusion-ci-gt-gen.yml index 4c5c4f3f7..b517a7a1a 100644 --- a/.github/workflows/diffusion-ci-gt-gen.yml +++ b/.github/workflows/diffusion-ci-gt-gen.yml @@ -38,12 +38,49 @@ env: OUTPUT_NAME: ${{ inputs.output_name || 'diffusion-ci-outputs' }} jobs: - multimodal-diffusion-gen-1gpu: + compute-diffusion-partitions: if: github.repository == 'sgl-project/sglang' + runs-on: ubuntu-latest + outputs: + matrix-1gpu: ${{ steps.compute.outputs.matrix-1gpu }} + matrix-2gpu: ${{ steps.compute.outputs.matrix-2gpu }} + matrix-b200: ${{ steps.compute.outputs.matrix-b200 }} + partition-count-1gpu: ${{ steps.compute.outputs['partition-count-1gpu'] }} + partition-count-2gpu: ${{ steps.compute.outputs['partition-count-2gpu'] }} + partition-count-b200: ${{ steps.compute.outputs['partition-count-b200'] }} + plan-1gpu: ${{ steps.compute.outputs.plan-1gpu }} + plan-2gpu: ${{ steps.compute.outputs.plan-2gpu }} + plan-b200: ${{ steps.compute.outputs.plan-b200 }} + steps: + - name: Checkout code + uses: actions/checkout@v4 + with: + ref: ${{ inputs.ref || github.ref }} + + - name: Set up Python + uses: actions/setup-python@v5 + with: + python-version: '3.10' + + - name: Compute partitions + id: compute + run: | + python scripts/ci/utils/diffusion/compute_diffusion_partitions.py \ + --min-time 1200 \ + --target-time 1800 \ + --max-time 2400 \ + --max-partitions 10 \ + --parametrized-only + + multimodal-diffusion-gen-1gpu: + needs: compute-diffusion-partitions + if: | + needs.compute-diffusion-partitions.result == 'success' && + needs.compute-diffusion-partitions.outputs.matrix-1gpu != '{"include":[]}' runs-on: 1-gpu-h100 strategy: - matrix: - part: [0, 1] + fail-fast: false + matrix: ${{ fromJson(needs.compute-diffusion-partitions.outputs.matrix-1gpu) }} timeout-minutes: 150 steps: - name: Checkout code @@ -75,12 +112,14 @@ jobs: - name: Generate outputs env: RUNAI_STREAMER_MEMORY_LIMIT: 0 + PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-1gpu }} run: | cd python python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \ --suite 1-gpu \ --partition-id ${{ matrix.part }} \ - --total-partitions 2 \ + --total-partitions ${{ needs.compute-diffusion-partitions.outputs['partition-count-1gpu'] }} \ + --partition-plan-json "$PARTITION_PLAN_JSON" \ --out-dir ./${{ env.OUTPUT_NAME }} \ ${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }} @@ -100,11 +139,14 @@ jobs: ${{ inputs.output_name != '' && format('--target-dir diffusion-ci/{0}', inputs.output_name) || '' }} multimodal-diffusion-gen-2gpu: - if: github.repository == 'sgl-project/sglang' + needs: compute-diffusion-partitions + if: | + needs.compute-diffusion-partitions.result == 'success' && + needs.compute-diffusion-partitions.outputs.matrix-2gpu != '{"include":[]}' runs-on: 2-gpu-h100 strategy: - matrix: - part: [0, 1] + fail-fast: false + matrix: ${{ fromJson(needs.compute-diffusion-partitions.outputs.matrix-2gpu) }} timeout-minutes: 150 steps: - name: Checkout code @@ -136,12 +178,14 @@ jobs: - name: Generate outputs env: RUNAI_STREAMER_MEMORY_LIMIT: 0 + PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-2gpu }} run: | cd python python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \ --suite 2-gpu \ --partition-id ${{ matrix.part }} \ - --total-partitions 2 \ + --total-partitions ${{ needs.compute-diffusion-partitions.outputs['partition-count-2gpu'] }} \ + --partition-plan-json "$PARTITION_PLAN_JSON" \ --out-dir ./${{ env.OUTPUT_NAME }} \ ${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }} @@ -161,8 +205,14 @@ jobs: ${{ inputs.output_name != '' && format('--target-dir diffusion-ci/{0}', inputs.output_name) || '' }} multimodal-diffusion-gen-b200: - if: github.repository == 'sgl-project/sglang' + needs: compute-diffusion-partitions + if: | + needs.compute-diffusion-partitions.result == 'success' && + needs.compute-diffusion-partitions.outputs.matrix-b200 != '{"include":[]}' runs-on: 4-gpu-b200 + strategy: + fail-fast: false + matrix: ${{ fromJson(needs.compute-diffusion-partitions.outputs.matrix-b200) }} timeout-minutes: 240 steps: - name: Checkout code @@ -194,17 +244,21 @@ jobs: - name: Generate outputs env: RUNAI_STREAMER_MEMORY_LIMIT: 0 + PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-b200 }} run: | cd python python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \ --suite 1-gpu-b200 \ + --partition-id ${{ matrix.part }} \ + --total-partitions ${{ needs.compute-diffusion-partitions.outputs['partition-count-b200'] }} \ + --partition-plan-json "$PARTITION_PLAN_JSON" \ --out-dir ./${{ env.OUTPUT_NAME }} \ ${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }} - name: Upload artifact uses: actions/upload-artifact@v4 with: - name: ${{ env.OUTPUT_NAME }}-b200 + name: ${{ env.OUTPUT_NAME }}-b200-part${{ matrix.part }} path: python/${{ env.OUTPUT_NAME }} retention-days: 7 diff --git a/.github/workflows/pr-test-multimodal-gen.yml b/.github/workflows/pr-test-multimodal-gen.yml index 110c812c0..938e2dae5 100644 --- a/.github/workflows/pr-test-multimodal-gen.yml +++ b/.github/workflows/pr-test-multimodal-gen.yml @@ -80,7 +80,7 @@ jobs: - name: Compute partitions id: compute run: | - python scripts/ci/utils/diffusion/compute_diffusion_partitions.py --min-time 1200 --target-time 1800 --max-time 2400 --max-partitions 10 + python scripts/ci/utils/diffusion/compute_diffusion_partitions.py --min-time 1200 --target-time 1800 --max-time 2400 --max-partitions 10 --parametrized-only multimodal-gen-test-1-gpu: needs: compute-diffusion-partitions diff --git a/.github/workflows/pr-test.yml b/.github/workflows/pr-test.yml index 649eb3c5b..2287eea8f 100644 --- a/.github/workflows/pr-test.yml +++ b/.github/workflows/pr-test.yml @@ -138,7 +138,7 @@ jobs: - ".github/workflows/pr-gate.yml" - ".github/actions/**" - "python/pyproject.toml" - - "python/sglang/!(multimodal_gen|jit_kernel/diffusion|jit_kernel/tests/diffusion|jit_kernel/benchmark/diffusion|cli)/**/!(*.md)" + - "python/sglang/!(multimodal_gen)/**/!(*.md)" - "scripts/ci/cuda/*" - "scripts/ci/utils/*" - "test/**/!(*.md)" diff --git a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py index 615dfca3c..0cb1b0aac 100644 --- a/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py +++ b/python/sglang/multimodal_gen/runtime/disaggregation/scheduler_mixin.py @@ -264,6 +264,11 @@ def extract_transfer_fields(req) -> tuple[dict, dict]: except (TypeError, ValueError): pass + if getattr(req, "generator", None) is not None: + seed = getattr(req, "seed", None) + if seed is not None: + scalar_fields["seed"] = _to_json_serializable(seed) + if _debug_transfer: import torch as _torch @@ -1316,9 +1321,13 @@ class SchedulerDisaggMixin: # Recreate torch.Generator from seed (not serializable over transfer) seed = scalar_fields.get("seed") if seed is not None: - gen = torch.Generator(device="cpu") - gen.manual_seed(int(seed)) - req.generator = gen + if isinstance(seed, list): + req.generator = [ + torch.Generator(device="cpu").manual_seed(int(item)) + for item in seed + ] + else: + req.generator = torch.Generator(device="cpu").manual_seed(int(seed)) # Rebuild trace_ctx from the propagated __getstate__ dict so this role's # spans nest under the sender's trace (same mechanism SRT uses via pickle). if trace_state and trace_state.get("tracing_enable"): diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py index 63162cb66..b4249c804 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/http_server.py @@ -35,7 +35,6 @@ if TYPE_CHECKING: logger = init_logger(__name__) -DEFAULT_SEED = 1024 VERTEX_ROUTE = os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate") @@ -278,7 +277,6 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput): rid, prompt=inst.get("prompt") or inst.get("text"), image_path=inst.get("image") or inst.get("image_url"), - seed=params.get("seed", DEFAULT_SEED), num_frames=params.get("num_frames"), fps=params.get("fps"), width=params.get("width"), diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py index 43cb01d75..d19e9648c 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/image_api.py @@ -225,7 +225,7 @@ async def edits( size: Optional[str] = Form(None), output_format: Optional[str] = Form(None), background: Optional[str] = Form("auto"), - seed: Optional[int] = Form(1024), + seed: Optional[int] = Form(None), generator_device: Optional[str] = Form("cuda"), user: Optional[str] = Form(None), negative_prompt: Optional[str] = Form(None), diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py index d860fc153..ed627777d 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/protocol.py @@ -44,7 +44,7 @@ class ImageGenerationsRequest(BaseModel): true_cfg_scale: Optional[float] = ( None # for CFG vs guidance distillation (e.g., QwenImage) ) - seed: Optional[Union[int, List[int]]] = 1024 + seed: Optional[Union[int, List[int]]] = None generator_device: Optional[str] = "cuda" negative_prompt: Optional[str] = None output_quality: Optional[str] = "default" @@ -93,7 +93,7 @@ class VideoGenerationsRequest(BaseModel): size: Optional[str] = "" fps: Optional[int] = None num_frames: Optional[int] = None - seed: Optional[Union[int, List[int]]] = 1024 + seed: Optional[Union[int, List[int]]] = None generator_device: Optional[str] = "cuda" # SGLang extensions width: Optional[int] = None diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py index e20b1c700..ccf07918d 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/openai/video_api.py @@ -196,7 +196,7 @@ async def create_video( size: Optional[str] = Form(None), fps: Optional[int] = Form(None), num_frames: Optional[int] = Form(None), - seed: Optional[int] = Form(1024), + seed: Optional[int] = Form(None), generator_device: Optional[str] = Form("cuda"), negative_prompt: Optional[str] = Form(None), guidance_scale: Optional[float] = Form(None), diff --git a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py index 341fb87fa..4e3f540ad 100644 --- a/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py +++ b/python/sglang/multimodal_gen/runtime/entrypoints/post_training/io_struct.py @@ -27,7 +27,7 @@ class GetWeightsChecksumReqInput: class RolloutRequest(BaseModel): prompt: str negative_prompt: Optional[str] = None - seed: int = 1024 + seed: Optional[int] = None generator_device: str = "cuda" width: Optional[int] = None diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py index 2dc23b4c2..3cea29d1d 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/denoising.py @@ -563,12 +563,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin): # Get latents and embeddings latents = batch.latents - prompt_embeds = batch.prompt_embeds # Removed Tensor truthiness assert to avoid GPU sync - neg_prompt_embeds = None if batch.do_classifier_free_guidance: - neg_prompt_embeds = batch.negative_prompt_embeds - assert neg_prompt_embeds is not None + assert batch.negative_prompt_embeds is not None # Removed Tensor truthiness assert to avoid GPU sync should_preprocess_for_wan_ti2v = should_apply_wan_ti2v(batch, server_args) diff --git a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/hunyuan3d_shape.py b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/hunyuan3d_shape.py index 75fd02472..eed87c056 100644 --- a/python/sglang/multimodal_gen/runtime/pipelines_core/stages/hunyuan3d_shape.py +++ b/python/sglang/multimodal_gen/runtime/pipelines_core/stages/hunyuan3d_shape.py @@ -22,7 +22,10 @@ from sglang.multimodal_gen.runtime.loader.component_loaders.transformer_loader i from sglang.multimodal_gen.runtime.managers.forward_context import set_forward_context from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch, Req from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage -from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import DenoisingStage +from sglang.multimodal_gen.runtime.pipelines_core.stages.denoising import ( + DenoisingContext, + DenoisingStage, +) from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import ( StageValidators as V, ) @@ -302,26 +305,25 @@ class Hunyuan3DShapeDenoisingStage(DenoisingStage): pos_cond_kwargs = {"encoder_hidden_states": cond} neg_cond_kwargs = {} - return { - "extra_step_kwargs": extra_step_kwargs, - "scheduler": scheduler, - "target_dtype": target_dtype, - "autocast_enabled": autocast_enabled, - "timesteps": timesteps, - "num_inference_steps": num_inference_steps, - "num_warmup_steps": num_warmup_steps, - "image_kwargs": {}, - "pos_cond_kwargs": pos_cond_kwargs, - "neg_cond_kwargs": neg_cond_kwargs, - "latents": latents, - "prompt_embeds": batch.prompt_embeds, - "neg_prompt_embeds": None, - "boundary_timestep": None, - "z": None, - "reserved_frames_mask": None, - "seq_len": None, - "guidance": guidance, - } + return DenoisingContext( + scheduler=scheduler, + extra_step_kwargs=extra_step_kwargs, + target_dtype=target_dtype, + autocast_enabled=autocast_enabled, + timesteps=timesteps, + num_inference_steps=num_inference_steps, + num_warmup_steps=num_warmup_steps, + image_kwargs={}, + pos_cond_kwargs=pos_cond_kwargs, + neg_cond_kwargs=neg_cond_kwargs, + latents=latents, + boundary_timestep=None, + z=None, + reserved_frames_mask=None, + seq_len=None, + guidance=guidance, + is_warmup=batch.is_warmup, + ) def _predict_noise( self, diff --git a/python/sglang/multimodal_gen/test/partitioning.py b/python/sglang/multimodal_gen/test/partitioning.py new file mode 100644 index 000000000..7bf189f73 --- /dev/null +++ b/python/sglang/multimodal_gen/test/partitioning.py @@ -0,0 +1,32 @@ +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class PartitionItem: + kind: str + item_id: str + est_time: float + used_fallback_estimate: bool = False + + +def partition_items_by_lpt( + items: list[PartitionItem], num_partitions: int +) -> list[list[PartitionItem]]: + if not items or num_partitions <= 0: + return [] + + sorted_items = sorted( + items, + key=lambda item: (-item.est_time, item.kind, item.item_id), + ) + partitions: list[list[PartitionItem]] = [[] for _ in range(num_partitions)] + partition_sums = [0.0] * num_partitions + + for item in sorted_items: + min_idx = partition_sums.index(min(partition_sums)) + partitions[min_idx].append(item) + partition_sums[min_idx] += item.est_time + + return partitions diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index b6c4c76dc..2ddfa9975 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -20,6 +20,10 @@ from pathlib import Path import tabulate from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.test.partitioning import ( + PartitionItem, + partition_items_by_lpt, +) from sglang.multimodal_gen.test.server.gpu_cases import ( ONE_GPU_CASES, TWO_GPU_CASES, @@ -177,16 +181,21 @@ def auto_partition( if not cases or size <= 0: return [] - sorted_cases = sorted(cases, key=lambda c: get_case_est_time(c.id), reverse=True) - partitions: list[list[DiffusionTestCase]] = [[] for _ in range(size)] - partition_sums = [0.0] * size + case_by_id = {case.id: case for case in cases} + items = [ + PartitionItem(kind="case", item_id=case.id, est_time=get_case_est_time(case.id)) + for case in cases + ] + partitions = partition_items_by_lpt(items, size) + if rank >= len(partitions): + return [] + return [case_by_id[item.item_id] for item in partitions[rank]] - for case in sorted_cases: - min_idx = partition_sums.index(min(partition_sums)) - partitions[min_idx].append(case) - partition_sums[min_idx] += get_case_est_time(case.id) - return partitions[rank] if rank < size else [] +def get_suite_files_rel(suite: str, parametrized_only: bool = False) -> list[str]: + if parametrized_only and suite in PARAMETRIZED_CASE_GROUPS: + return [filename for filename, _ in PARAMETRIZED_CASE_GROUPS[suite]] + return SUITES[suite] def _normalize_standalone_key(standalone_file: str) -> str: @@ -746,14 +755,18 @@ def run_pytest( ) -def partition_test_files(files, partition_id, total_partitions): +def partition_items_by_index( + items: list[str], partition_id: int, total_partitions: int +) -> list[str]: return [ - file_path - for i, file_path in enumerate(files) - if i % total_partitions == partition_id + item for i, item in enumerate(items) if i % total_partitions == partition_id ] +def partition_test_files(files, partition_id, total_partitions): + return partition_items_by_index(files, partition_id, total_partitions) + + def run_component_accuracy_files(files, filter_expr=None, continue_on_error=False): exit_code = 0 for file_path in files: diff --git a/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py b/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py index b604074fa..c03651a49 100755 --- a/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py +++ b/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py @@ -19,8 +19,13 @@ from pathlib import Path from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.test.run_suite import ( SUITES, + PartitionItem, _maybe_pin_update_weights_model_pair, collect_test_items, + get_case_est_time, + get_suite_files_rel, + parse_partition_plan, + partition_items_by_lpt, run_pytest, ) @@ -67,6 +72,12 @@ def main(): required=False, help="Specific case IDs to run (space-separated). If provided, only these cases will be run.", ) + parser.add_argument( + "--partition-plan-json", + type=str, + required=False, + help="Full partition plan JSON for the current suite.", + ) args = parser.parse_args() @@ -78,6 +89,10 @@ def main(): parser.error( "Both --partition-id and --total-partitions must be provided together" ) + if args.partition_plan_json and ( + args.partition_id is None or args.total_partitions is None + ): + parser.error("--partition-plan-json requires partition-id and total-partitions") # Create output directory out_dir = Path(args.out_dir) @@ -98,8 +113,10 @@ def main(): test_root_dir = current_file_path.parent.parent # scripts -> test target_dir = test_root_dir / "server" - # Get files from suite (same as run_suite.py) - suite_files_rel = SUITES[args.suite] + # GT generation only runs DiffusionTestCase parametrized cases. Standalone + # server tests such as disagg validate behavior but do not produce GT images. + suite_files_rel = get_suite_files_rel(args.suite, parametrized_only=True) + _maybe_pin_update_weights_model_pair(suite_files_rel) suite_files_abs = [] for f_rel in suite_files_rel: @@ -113,30 +130,76 @@ def main(): logger.error(f"No valid test files found for suite '{args.suite}'.") sys.exit(1) - # Build pytest filter for case_ids if provided + partition_id = args.partition_id if args.partition_id is not None else 0 + total_partitions = args.total_partitions if args.total_partitions is not None else 1 + + selected_plan_case_ids = None + if args.partition_plan_json: + assignment = parse_partition_plan( + suite=args.suite, + partition_id=partition_id, + total_partitions=total_partitions, + plan_json=args.partition_plan_json, + ) + selected_plan_case_ids = assignment.case_ids + if args.case_ids: + requested_case_ids = set(args.case_ids) + selected_plan_case_ids = [ + case_id + for case_id in selected_plan_case_ids + if case_id in requested_case_ids + ] + if not selected_plan_case_ids: + logger.warning("No testcase cases assigned to this partition.") + sys.exit(0) + + # Build pytest filter for case_ids if provided. filter_expr = None - if args.case_ids: + if selected_plan_case_ids is not None: + filters = [ + f"test_diffusion_generation[{case_id}]" + for case_id in selected_plan_case_ids + ] + filter_expr = " or ".join(filters) + logger.info(f"Filtering by partition plan case IDs: {selected_plan_case_ids}") + elif args.case_ids: # pytest parametrized test format: test_diffusion_generation[case_id] filters = [f"test_diffusion_generation[{case_id}]" for case_id in args.case_ids] filter_expr = " or ".join(filters) logger.info(f"Filtering by case IDs: {args.case_ids}") - # Collect all test items (same as run_suite.py) + # Collect all test items and keep only testcase-based GT generators. all_test_items = collect_test_items(suite_files_abs, filter_expr=filter_expr) + all_test_items = [ + item for item in all_test_items if "test_diffusion_generation[" in item + ] if not all_test_items: logger.warning(f"No test items found for suite '{args.suite}'.") sys.exit(0) - # Partition by test items (same as run_suite.py) - partition_id = args.partition_id if args.partition_id is not None else 0 - total_partitions = args.total_partitions if args.total_partitions is not None else 1 + if selected_plan_case_ids is not None: + selected_case_id_set = set(selected_plan_case_ids) + my_items = [ + item + for item in all_test_items + if item[item.index("[") + 1 : item.rindex("]")] in selected_case_id_set + ] + else: + # Partition by test items with the same LPT strategy used by CI partitioning. + partition_items = [] + for item in all_test_items: + case_id = item[item.index("[") + 1 : item.rindex("]")] + partition_items.append( + PartitionItem( + kind="case", + item_id=item, + est_time=get_case_est_time(case_id), + ) + ) - my_items = [ - item - for i, item in enumerate(all_test_items) - if i % total_partitions == partition_id - ] + partitions = partition_items_by_lpt(partition_items, total_partitions) + my_items = [item.item_id for item in partitions[partition_id]] logger.info( f"Partition {partition_id}/{total_partitions}: " diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py b/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py index c8d3861ab..0cb314234 100644 --- a/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py +++ b/python/sglang/multimodal_gen/test/unit/test_disagg_trace.py @@ -20,6 +20,8 @@ from __future__ import annotations import json import unittest +import torch + from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import ( SchedulerDisaggMixin, extract_transfer_fields, @@ -74,6 +76,43 @@ def _roundtrip_scalar_fields(scalar_fields: dict) -> dict: class TestDisaggTracePropagation(unittest.TestCase): + def test_transfer_keeps_seed_needed_to_rebuild_generator(self): + req = Req(request_id="test-seed", prompt="x") + req.generator = torch.Generator(device="cpu").manual_seed(req.seed) + + _, scalar_fields = extract_transfer_fields(req) + + self.assertEqual(scalar_fields["seed"], 42) + + rebuilt = SchedulerDisaggMixin._build_disagg_req(None, dict(scalar_fields), {}) + self.assertIsInstance(rebuilt.generator, torch.Generator) + self.assertEqual(rebuilt.seed, 42) + + expected = torch.rand( + (), generator=torch.Generator(device="cpu").manual_seed(42) + ) + actual = torch.rand((), generator=rebuilt.generator) + self.assertEqual(actual.item(), expected.item()) + + def test_build_disagg_req_rebuilds_generator_list(self): + scalar_fields = { + "request_id": "test-seed-list", + "prompt": "x", + "num_outputs_per_prompt": 2, + "seed": [11, 12], + } + + rebuilt = SchedulerDisaggMixin._build_disagg_req(None, dict(scalar_fields), {}) + + self.assertEqual(rebuilt.seed, [11, 12]) + self.assertEqual(len(rebuilt.generator), 2) + for seed, generator in zip(rebuilt.seed, rebuilt.generator): + expected = torch.rand( + (), generator=torch.Generator(device="cpu").manual_seed(seed) + ) + actual = torch.rand((), generator=generator) + self.assertEqual(actual.item(), expected.item()) + def test_tracing_disabled_omits_trace_state(self): """With a default TraceNullContext Req, no _trace_state is emitted and the JSON codec does not encounter any live OTel objects.""" diff --git a/scripts/ci/utils/diffusion/compute_diffusion_partitions.py b/scripts/ci/utils/diffusion/compute_diffusion_partitions.py index 50569dfe9..c4d304ab5 100755 --- a/scripts/ci/utils/diffusion/compute_diffusion_partitions.py +++ b/scripts/ci/utils/diffusion/compute_diffusion_partitions.py @@ -7,11 +7,11 @@ AST parsing to extract parametrized cases plus standalone files from source. """ import argparse +import importlib.util import json import math import os import sys -from dataclasses import dataclass from pathlib import Path from diffusion_case_parser import ( @@ -22,21 +22,25 @@ from diffusion_case_parser import ( resolve_case_config_path, ) -SUITE_OUTPUT_NAMES = { - "1-gpu": "1gpu", - "2-gpu": "2gpu", -} + +def _load_partitioning_helpers(): + repo_root = Path(__file__).resolve().parents[4] + helper_path = repo_root / "python/sglang/multimodal_gen/test/partitioning.py" + spec = importlib.util.spec_from_file_location( + "diffusion_test_partitioning", helper_path + ) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module.PartitionItem, module.partition_items_by_lpt + + +PartitionItem, partition_items_by_lpt = _load_partitioning_helpers() + +SUITE_OUTPUT_NAMES = {"1-gpu": "1gpu", "2-gpu": "2gpu", "1-gpu-b200": "b200"} DEFAULT_STANDALONE_EST_TIME_SECONDS = 300.0 -@dataclass(frozen=True) -class PartitionItem: - kind: str - item_id: str - est_time: float - used_fallback_estimate: bool = False - - def validate_suite_case_coverage(suites: dict[str, DiffusionSuiteInfo]) -> None: """ Guardrail: dynamic diffusion suites must contain parametrized cases. @@ -85,11 +89,16 @@ def compute_partition_count( return max(min_partition_count, min(preferred_count, max_partition_count)) -def build_partition_items(suite_info: DiffusionSuiteInfo) -> list[PartitionItem]: +def build_partition_items( + suite_info: DiffusionSuiteInfo, include_standalone: bool = True +) -> list[PartitionItem]: items = [ PartitionItem(kind="case", item_id=case.case_id, est_time=case.est_time) for case in suite_info.cases ] + if not include_standalone: + return items + items.extend( PartitionItem( kind="standalone", @@ -106,27 +115,6 @@ def build_partition_items(suite_info: DiffusionSuiteInfo) -> list[PartitionItem] return items -def lpt_partition( - items: list[PartitionItem], num_partitions: int -) -> list[list[PartitionItem]]: - if not items or num_partitions <= 0: - return [] - - sorted_items = sorted( - items, - key=lambda item: (-item.est_time, item.kind, item.item_id), - ) - partitions: list[list[PartitionItem]] = [[] for _ in range(num_partitions)] - partition_sums = [0.0] * num_partitions - - for item in sorted_items: - min_idx = partition_sums.index(min(partition_sums)) - partitions[min_idx].append(item) - partition_sums[min_idx] += item.est_time - - return partitions - - def build_matrix(partition_count: int) -> dict: if partition_count <= 0: return {"include": []} @@ -180,11 +168,20 @@ def print_suite_summary( suite_name: str, suite_info: DiffusionSuiteInfo, partitions: list[list[PartitionItem]], + include_standalone: bool = True, ) -> None: - total_time = sum(item.est_time for item in build_partition_items(suite_info)) + total_time = sum( + item.est_time + for item in build_partition_items( + suite_info, include_standalone=include_standalone + ) + ) print(f"{suite_name.upper()} suite:") print(f" Cases: {len(suite_info.cases)}") - print(f" Standalone files: {len(suite_info.standalone_files)}") + standalone_label = "Standalone files" + if not include_standalone: + standalone_label = "Standalone files ignored" + print(f" {standalone_label}: {len(suite_info.standalone_files)}") print( f" Missing standalone estimates: {len(suite_info.missing_standalone_estimates)}" ) @@ -247,6 +244,11 @@ def main(): default=10, help="Maximum number of partitions (default: 10)", ) + parser.add_argument( + "--parametrized-only", + action="store_true", + help="Only partition DiffusionTestCase parametrized cases.", + ) args = parser.parse_args() script_dir = Path(__file__).resolve().parent @@ -281,7 +283,9 @@ def main(): if suite_name not in SUITE_OUTPUT_NAMES: continue - items = build_partition_items(suite_info) + items = build_partition_items( + suite_info, include_standalone=not args.parametrized_only + ) total_time = sum(item.est_time for item in items) partition_count = compute_partition_count( total_time_seconds=total_time, @@ -290,9 +294,14 @@ def main(): max_time_seconds=args.max_time, max_partitions=args.max_partitions, ) - partitions = lpt_partition(items, partition_count) + partitions = partition_items_by_lpt(items, partition_count) - print_suite_summary(suite_name, suite_info, partitions) + print_suite_summary( + suite_name, + suite_info, + partitions, + include_standalone=not args.parametrized_only, + ) output_name = SUITE_OUTPUT_NAMES[suite_name] output_github_value(f"matrix-{output_name}", build_matrix(partition_count)) diff --git a/scripts/ci/utils/diffusion/diffusion_case_parser.py b/scripts/ci/utils/diffusion/diffusion_case_parser.py index 290ed2d2c..1393e7daa 100755 --- a/scripts/ci/utils/diffusion/diffusion_case_parser.py +++ b/scripts/ci/utils/diffusion/diffusion_case_parser.py @@ -106,6 +106,12 @@ class DiffusionTestCaseVisitor(ast.NodeVisitor): if not isinstance(target, ast.Name) or not isinstance(op, ast.Add): return + if isinstance(value, ast.Name): + target_suite = CASE_LIST_TO_SUITE.get(target.id) + value_suite = CASE_LIST_TO_SUITE.get(value.id) + if target_suite and value_suite and target_suite != value_suite: + return + rhs_case_ids = self._extract_case_ids(value) if rhs_case_ids is None: return @@ -152,9 +158,11 @@ class DiffusionTestCaseVisitor(ast.NodeVisitor): if not isinstance(node, ast.Call): return None - # Check if it's a DiffusionTestCase call - if isinstance(node.func, ast.Name) and node.func.id == "DiffusionTestCase": - # First positional argument is the case_id + # First positional argument is the case_id. + if isinstance(node.func, ast.Name) and node.func.id in { + "DiffusionTestCase", + "_make_modelopt_ci_case", + }: if node.args and isinstance(node.args[0], ast.Constant): return node.args[0].value