[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' }}
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
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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)"
@@ -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"):
@@ -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"),
@@ -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),
@@ -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
@@ -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),
@@ -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
@@ -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)
@@ -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,
@@ -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
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:
@@ -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}: "
@@ -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."""
@@ -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))
@@ -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