[diffusion] chore: change default seed to 42 (#23836)
This commit is contained in:
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user