[diffusion] chore: change default seed to 42 (#23836)

This commit is contained in:
Mick
2026-04-28 20:39:23 +08:00
committed by GitHub
parent 69a71219cb
commit 144038fbae
17 changed files with 339 additions and 115 deletions
+64 -10
View File
@@ -38,12 +38,49 @@ env:
OUTPUT_NAME: ${{ inputs.output_name || 'diffusion-ci-outputs' }} OUTPUT_NAME: ${{ inputs.output_name || 'diffusion-ci-outputs' }}
jobs: jobs:
multimodal-diffusion-gen-1gpu: compute-diffusion-partitions:
if: github.repository == 'sgl-project/sglang' 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 runs-on: 1-gpu-h100
strategy: strategy:
matrix: fail-fast: false
part: [0, 1] matrix: ${{ fromJson(needs.compute-diffusion-partitions.outputs.matrix-1gpu) }}
timeout-minutes: 150 timeout-minutes: 150
steps: steps:
- name: Checkout code - name: Checkout code
@@ -75,12 +112,14 @@ jobs:
- name: Generate outputs - name: Generate outputs
env: env:
RUNAI_STREAMER_MEMORY_LIMIT: 0 RUNAI_STREAMER_MEMORY_LIMIT: 0
PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-1gpu }}
run: | run: |
cd python cd python
python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \ python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \
--suite 1-gpu \ --suite 1-gpu \
--partition-id ${{ matrix.part }} \ --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 }} \ --out-dir ./${{ env.OUTPUT_NAME }} \
${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }} ${{ 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) || '' }} ${{ inputs.output_name != '' && format('--target-dir diffusion-ci/{0}', inputs.output_name) || '' }}
multimodal-diffusion-gen-2gpu: 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 runs-on: 2-gpu-h100
strategy: strategy:
matrix: fail-fast: false
part: [0, 1] matrix: ${{ fromJson(needs.compute-diffusion-partitions.outputs.matrix-2gpu) }}
timeout-minutes: 150 timeout-minutes: 150
steps: steps:
- name: Checkout code - name: Checkout code
@@ -136,12 +178,14 @@ jobs:
- name: Generate outputs - name: Generate outputs
env: env:
RUNAI_STREAMER_MEMORY_LIMIT: 0 RUNAI_STREAMER_MEMORY_LIMIT: 0
PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-2gpu }}
run: | run: |
cd python cd python
python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \ python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \
--suite 2-gpu \ --suite 2-gpu \
--partition-id ${{ matrix.part }} \ --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 }} \ --out-dir ./${{ env.OUTPUT_NAME }} \
${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }} ${{ 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) || '' }} ${{ inputs.output_name != '' && format('--target-dir diffusion-ci/{0}', inputs.output_name) || '' }}
multimodal-diffusion-gen-b200: 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 runs-on: 4-gpu-b200
strategy:
fail-fast: false
matrix: ${{ fromJson(needs.compute-diffusion-partitions.outputs.matrix-b200) }}
timeout-minutes: 240 timeout-minutes: 240
steps: steps:
- name: Checkout code - name: Checkout code
@@ -194,17 +244,21 @@ jobs:
- name: Generate outputs - name: Generate outputs
env: env:
RUNAI_STREAMER_MEMORY_LIMIT: 0 RUNAI_STREAMER_MEMORY_LIMIT: 0
PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-b200 }}
run: | run: |
cd python cd python
python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \ python -m sglang.multimodal_gen.test.scripts.gen_diffusion_ci_outputs \
--suite 1-gpu-b200 \ --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 }} \ --out-dir ./${{ env.OUTPUT_NAME }} \
${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }} ${{ inputs.case_ids != '' && format('--case-ids {0}', inputs.case_ids) || '' }}
- name: Upload artifact - name: Upload artifact
uses: actions/upload-artifact@v4 uses: actions/upload-artifact@v4
with: with:
name: ${{ env.OUTPUT_NAME }}-b200 name: ${{ env.OUTPUT_NAME }}-b200-part${{ matrix.part }}
path: python/${{ env.OUTPUT_NAME }} path: python/${{ env.OUTPUT_NAME }}
retention-days: 7 retention-days: 7
+1 -1
View File
@@ -80,7 +80,7 @@ jobs:
- name: Compute partitions - name: Compute partitions
id: compute id: compute
run: | 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: multimodal-gen-test-1-gpu:
needs: compute-diffusion-partitions needs: compute-diffusion-partitions
+1 -1
View File
@@ -138,7 +138,7 @@ jobs:
- ".github/workflows/pr-gate.yml" - ".github/workflows/pr-gate.yml"
- ".github/actions/**" - ".github/actions/**"
- "python/pyproject.toml" - "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/cuda/*"
- "scripts/ci/utils/*" - "scripts/ci/utils/*"
- "test/**/!(*.md)" - "test/**/!(*.md)"
@@ -264,6 +264,11 @@ def extract_transfer_fields(req) -> tuple[dict, dict]:
except (TypeError, ValueError): except (TypeError, ValueError):
pass 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: if _debug_transfer:
import torch as _torch import torch as _torch
@@ -1316,9 +1321,13 @@ class SchedulerDisaggMixin:
# Recreate torch.Generator from seed (not serializable over transfer) # Recreate torch.Generator from seed (not serializable over transfer)
seed = scalar_fields.get("seed") seed = scalar_fields.get("seed")
if seed is not None: if seed is not None:
gen = torch.Generator(device="cpu") if isinstance(seed, list):
gen.manual_seed(int(seed)) req.generator = [
req.generator = gen 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 # 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). # spans nest under the sender's trace (same mechanism SRT uses via pickle).
if trace_state and trace_state.get("tracing_enable"): if trace_state and trace_state.get("tracing_enable"):
@@ -35,7 +35,6 @@ if TYPE_CHECKING:
logger = init_logger(__name__) logger = init_logger(__name__)
DEFAULT_SEED = 1024
VERTEX_ROUTE = os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate") VERTEX_ROUTE = os.environ.get("AIP_PREDICT_ROUTE", "/vertex_generate")
@@ -278,7 +277,6 @@ async def vertex_generate(vertex_req: VertexGenerateReqInput):
rid, rid,
prompt=inst.get("prompt") or inst.get("text"), prompt=inst.get("prompt") or inst.get("text"),
image_path=inst.get("image") or inst.get("image_url"), image_path=inst.get("image") or inst.get("image_url"),
seed=params.get("seed", DEFAULT_SEED),
num_frames=params.get("num_frames"), num_frames=params.get("num_frames"),
fps=params.get("fps"), fps=params.get("fps"),
width=params.get("width"), width=params.get("width"),
@@ -225,7 +225,7 @@ async def edits(
size: Optional[str] = Form(None), size: Optional[str] = Form(None),
output_format: Optional[str] = Form(None), output_format: Optional[str] = Form(None),
background: Optional[str] = Form("auto"), background: Optional[str] = Form("auto"),
seed: Optional[int] = Form(1024), seed: Optional[int] = Form(None),
generator_device: Optional[str] = Form("cuda"), generator_device: Optional[str] = Form("cuda"),
user: Optional[str] = Form(None), user: Optional[str] = Form(None),
negative_prompt: Optional[str] = Form(None), negative_prompt: Optional[str] = Form(None),
@@ -44,7 +44,7 @@ class ImageGenerationsRequest(BaseModel):
true_cfg_scale: Optional[float] = ( true_cfg_scale: Optional[float] = (
None # for CFG vs guidance distillation (e.g., QwenImage) 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" generator_device: Optional[str] = "cuda"
negative_prompt: Optional[str] = None negative_prompt: Optional[str] = None
output_quality: Optional[str] = "default" output_quality: Optional[str] = "default"
@@ -93,7 +93,7 @@ class VideoGenerationsRequest(BaseModel):
size: Optional[str] = "" size: Optional[str] = ""
fps: Optional[int] = None fps: Optional[int] = None
num_frames: 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" generator_device: Optional[str] = "cuda"
# SGLang extensions # SGLang extensions
width: Optional[int] = None width: Optional[int] = None
@@ -196,7 +196,7 @@ async def create_video(
size: Optional[str] = Form(None), size: Optional[str] = Form(None),
fps: Optional[int] = Form(None), fps: Optional[int] = Form(None),
num_frames: 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"), generator_device: Optional[str] = Form("cuda"),
negative_prompt: Optional[str] = Form(None), negative_prompt: Optional[str] = Form(None),
guidance_scale: Optional[float] = Form(None), guidance_scale: Optional[float] = Form(None),
@@ -27,7 +27,7 @@ class GetWeightsChecksumReqInput:
class RolloutRequest(BaseModel): class RolloutRequest(BaseModel):
prompt: str prompt: str
negative_prompt: Optional[str] = None negative_prompt: Optional[str] = None
seed: int = 1024 seed: Optional[int] = None
generator_device: str = "cuda" generator_device: str = "cuda"
width: Optional[int] = None width: Optional[int] = None
@@ -563,12 +563,9 @@ class DenoisingStage(PipelineStage, RolloutDenoisingMixin):
# Get latents and embeddings # Get latents and embeddings
latents = batch.latents latents = batch.latents
prompt_embeds = batch.prompt_embeds
# Removed Tensor truthiness assert to avoid GPU sync # Removed Tensor truthiness assert to avoid GPU sync
neg_prompt_embeds = None
if batch.do_classifier_free_guidance: if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds assert batch.negative_prompt_embeds is not None
assert neg_prompt_embeds is not None
# Removed Tensor truthiness assert to avoid GPU sync # Removed Tensor truthiness assert to avoid GPU sync
should_preprocess_for_wan_ti2v = should_apply_wan_ti2v(batch, server_args) should_preprocess_for_wan_ti2v = should_apply_wan_ti2v(batch, server_args)
@@ -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.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.schedule_batch import OutputBatch, Req
from sglang.multimodal_gen.runtime.pipelines_core.stages.base import PipelineStage 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 ( from sglang.multimodal_gen.runtime.pipelines_core.stages.validators import (
StageValidators as V, StageValidators as V,
) )
@@ -302,26 +305,25 @@ class Hunyuan3DShapeDenoisingStage(DenoisingStage):
pos_cond_kwargs = {"encoder_hidden_states": cond} pos_cond_kwargs = {"encoder_hidden_states": cond}
neg_cond_kwargs = {} neg_cond_kwargs = {}
return { return DenoisingContext(
"extra_step_kwargs": extra_step_kwargs, scheduler=scheduler,
"scheduler": scheduler, extra_step_kwargs=extra_step_kwargs,
"target_dtype": target_dtype, target_dtype=target_dtype,
"autocast_enabled": autocast_enabled, autocast_enabled=autocast_enabled,
"timesteps": timesteps, timesteps=timesteps,
"num_inference_steps": num_inference_steps, num_inference_steps=num_inference_steps,
"num_warmup_steps": num_warmup_steps, num_warmup_steps=num_warmup_steps,
"image_kwargs": {}, image_kwargs={},
"pos_cond_kwargs": pos_cond_kwargs, pos_cond_kwargs=pos_cond_kwargs,
"neg_cond_kwargs": neg_cond_kwargs, neg_cond_kwargs=neg_cond_kwargs,
"latents": latents, latents=latents,
"prompt_embeds": batch.prompt_embeds, boundary_timestep=None,
"neg_prompt_embeds": None, z=None,
"boundary_timestep": None, reserved_frames_mask=None,
"z": None, seq_len=None,
"reserved_frames_mask": None, guidance=guidance,
"seq_len": None, is_warmup=batch.is_warmup,
"guidance": guidance, )
}
def _predict_noise( def _predict_noise(
self, self,
@@ -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
+25 -12
View File
@@ -20,6 +20,10 @@ from pathlib import Path
import tabulate import tabulate
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger 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 ( from sglang.multimodal_gen.test.server.gpu_cases import (
ONE_GPU_CASES, ONE_GPU_CASES,
TWO_GPU_CASES, TWO_GPU_CASES,
@@ -177,16 +181,21 @@ def auto_partition(
if not cases or size <= 0: if not cases or size <= 0:
return [] return []
sorted_cases = sorted(cases, key=lambda c: get_case_est_time(c.id), reverse=True) case_by_id = {case.id: case for case in cases}
partitions: list[list[DiffusionTestCase]] = [[] for _ in range(size)] items = [
partition_sums = [0.0] * size 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: 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 [ return [
file_path item for i, item in enumerate(items) if i % total_partitions == partition_id
for i, file_path in enumerate(files)
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): def run_component_accuracy_files(files, filter_expr=None, continue_on_error=False):
exit_code = 0 exit_code = 0
for file_path in files: for file_path in files:
@@ -19,8 +19,13 @@ from pathlib import Path
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
from sglang.multimodal_gen.test.run_suite import ( from sglang.multimodal_gen.test.run_suite import (
SUITES, SUITES,
PartitionItem,
_maybe_pin_update_weights_model_pair, _maybe_pin_update_weights_model_pair,
collect_test_items, collect_test_items,
get_case_est_time,
get_suite_files_rel,
parse_partition_plan,
partition_items_by_lpt,
run_pytest, run_pytest,
) )
@@ -67,6 +72,12 @@ def main():
required=False, required=False,
help="Specific case IDs to run (space-separated). If provided, only these cases will be run.", 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() args = parser.parse_args()
@@ -78,6 +89,10 @@ def main():
parser.error( parser.error(
"Both --partition-id and --total-partitions must be provided together" "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 # Create output directory
out_dir = Path(args.out_dir) out_dir = Path(args.out_dir)
@@ -98,8 +113,10 @@ def main():
test_root_dir = current_file_path.parent.parent # scripts -> test test_root_dir = current_file_path.parent.parent # scripts -> test
target_dir = test_root_dir / "server" target_dir = test_root_dir / "server"
# Get files from suite (same as run_suite.py) # GT generation only runs DiffusionTestCase parametrized cases. Standalone
suite_files_rel = SUITES[args.suite] # 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) _maybe_pin_update_weights_model_pair(suite_files_rel)
suite_files_abs = [] suite_files_abs = []
for f_rel in suite_files_rel: 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}'.") logger.error(f"No valid test files found for suite '{args.suite}'.")
sys.exit(1) 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 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] # pytest parametrized test format: test_diffusion_generation[case_id]
filters = [f"test_diffusion_generation[{case_id}]" for case_id in args.case_ids] filters = [f"test_diffusion_generation[{case_id}]" for case_id in args.case_ids]
filter_expr = " or ".join(filters) filter_expr = " or ".join(filters)
logger.info(f"Filtering by case IDs: {args.case_ids}") 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 = 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: if not all_test_items:
logger.warning(f"No test items found for suite '{args.suite}'.") logger.warning(f"No test items found for suite '{args.suite}'.")
sys.exit(0) sys.exit(0)
# Partition by test items (same as run_suite.py) if selected_plan_case_ids is not None:
partition_id = args.partition_id if args.partition_id is not None else 0 selected_case_id_set = set(selected_plan_case_ids)
total_partitions = args.total_partitions if args.total_partitions is not None else 1 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 = [ partitions = partition_items_by_lpt(partition_items, total_partitions)
item my_items = [item.item_id for item in partitions[partition_id]]
for i, item in enumerate(all_test_items)
if i % total_partitions == partition_id
]
logger.info( logger.info(
f"Partition {partition_id}/{total_partitions}: " f"Partition {partition_id}/{total_partitions}: "
@@ -20,6 +20,8 @@ from __future__ import annotations
import json import json
import unittest import unittest
import torch
from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import ( from sglang.multimodal_gen.runtime.disaggregation.scheduler_mixin import (
SchedulerDisaggMixin, SchedulerDisaggMixin,
extract_transfer_fields, extract_transfer_fields,
@@ -74,6 +76,43 @@ def _roundtrip_scalar_fields(scalar_fields: dict) -> dict:
class TestDisaggTracePropagation(unittest.TestCase): 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): def test_tracing_disabled_omits_trace_state(self):
"""With a default TraceNullContext Req, no _trace_state is emitted and """With a default TraceNullContext Req, no _trace_state is emitted and
the JSON codec does not encounter any live OTel objects.""" the JSON codec does not encounter any live OTel objects."""
@@ -7,11 +7,11 @@ AST parsing to extract parametrized cases plus standalone files from source.
""" """
import argparse import argparse
import importlib.util
import json import json
import math import math
import os import os
import sys import sys
from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from diffusion_case_parser import ( from diffusion_case_parser import (
@@ -22,21 +22,25 @@ from diffusion_case_parser import (
resolve_case_config_path, resolve_case_config_path,
) )
SUITE_OUTPUT_NAMES = {
"1-gpu": "1gpu", def _load_partitioning_helpers():
"2-gpu": "2gpu", 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 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: def validate_suite_case_coverage(suites: dict[str, DiffusionSuiteInfo]) -> None:
""" """
Guardrail: dynamic diffusion suites must contain parametrized cases. 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)) 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 = [ items = [
PartitionItem(kind="case", item_id=case.case_id, est_time=case.est_time) PartitionItem(kind="case", item_id=case.case_id, est_time=case.est_time)
for case in suite_info.cases for case in suite_info.cases
] ]
if not include_standalone:
return items
items.extend( items.extend(
PartitionItem( PartitionItem(
kind="standalone", kind="standalone",
@@ -106,27 +115,6 @@ def build_partition_items(suite_info: DiffusionSuiteInfo) -> list[PartitionItem]
return items 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: def build_matrix(partition_count: int) -> dict:
if partition_count <= 0: if partition_count <= 0:
return {"include": []} return {"include": []}
@@ -180,11 +168,20 @@ def print_suite_summary(
suite_name: str, suite_name: str,
suite_info: DiffusionSuiteInfo, suite_info: DiffusionSuiteInfo,
partitions: list[list[PartitionItem]], partitions: list[list[PartitionItem]],
include_standalone: bool = True,
) -> None: ) -> 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"{suite_name.upper()} suite:")
print(f" Cases: {len(suite_info.cases)}") 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( print(
f" Missing standalone estimates: {len(suite_info.missing_standalone_estimates)}" f" Missing standalone estimates: {len(suite_info.missing_standalone_estimates)}"
) )
@@ -247,6 +244,11 @@ def main():
default=10, default=10,
help="Maximum number of partitions (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() args = parser.parse_args()
script_dir = Path(__file__).resolve().parent script_dir = Path(__file__).resolve().parent
@@ -281,7 +283,9 @@ def main():
if suite_name not in SUITE_OUTPUT_NAMES: if suite_name not in SUITE_OUTPUT_NAMES:
continue 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) total_time = sum(item.est_time for item in items)
partition_count = compute_partition_count( partition_count = compute_partition_count(
total_time_seconds=total_time, total_time_seconds=total_time,
@@ -290,9 +294,14 @@ def main():
max_time_seconds=args.max_time, max_time_seconds=args.max_time,
max_partitions=args.max_partitions, 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_name = SUITE_OUTPUT_NAMES[suite_name]
output_github_value(f"matrix-{output_name}", build_matrix(partition_count)) output_github_value(f"matrix-{output_name}", build_matrix(partition_count))
@@ -106,6 +106,12 @@ class DiffusionTestCaseVisitor(ast.NodeVisitor):
if not isinstance(target, ast.Name) or not isinstance(op, ast.Add): if not isinstance(target, ast.Name) or not isinstance(op, ast.Add):
return 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) rhs_case_ids = self._extract_case_ids(value)
if rhs_case_ids is None: if rhs_case_ids is None:
return return
@@ -152,9 +158,11 @@ class DiffusionTestCaseVisitor(ast.NodeVisitor):
if not isinstance(node, ast.Call): if not isinstance(node, ast.Call):
return None return None
# Check if it's a DiffusionTestCase call # First positional argument is the case_id.
if isinstance(node.func, ast.Name) and node.func.id == "DiffusionTestCase": if isinstance(node.func, ast.Name) and node.func.id in {
# First positional argument is the case_id "DiffusionTestCase",
"_make_modelopt_ci_case",
}:
if node.args and isinstance(node.args[0], ast.Constant): if node.args and isinstance(node.args[0], ast.Constant):
return node.args[0].value return node.args[0].value