From 4d23a4fa6d062a9a7ccc3008cb32f8a5454198e2 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Mon, 7 Sep 2026 15:13:59 +0800 Subject: [PATCH] [Test] Consolidate test cleanup and CI taxonomy (net -11.4K lines) (#37436) Co-authored-by: Mick Qian --- .claude/rules/unit-test-admission.md | 9 + .claude/skills/write-sglang-test/SKILL.md | 54 +- .github/workflows/pr-test-multimodal-gen.yml | 80 +- python/sglang/kernels/ops/diffusion/README.md | 4 +- .../sglang/kernels/ops/diffusion/__init__.py | 4 +- .../sglang/multimodal_gen/test/run_suite.py | 689 +------------ .../test/runner/diffusion_suite_runner.py | 690 +++++++++++++ .../test/scripts/gen_diffusion_ci_outputs.py | 4 +- .../test/unit/realtime/test_realtime_webui.py | 112 --- .../test/unit/test_consistency_metrics.py | 12 - .../multimodal_gen/test/unit/test_cosmos3.py | 12 - .../unit/test_diffusion_benchmark_skill.py | 658 ------------ .../test/unit/test_disagg_roles.py | 7 - .../test/unit/test_nvtx_pytorch_hooks.py | 11 - .../test/unit/test_qwen2_5vl_generation.py | 52 - .../test/unit/test_qwen3vl_vision.py | 9 - .../test/unit/test_server_args.py | 5 - .../test/unit/test_server_warmup_progress.py | 36 - .../unit/test_single_rank_device_group.py | 5 - .../test/unit/test_sol_attn_backend.py | 48 - .../multimodal_gen/test/unit/test_sp_shard.py | 6 - .../test/unit/test_srt_clip_reuse.py | 6 - .../unit/test_subblock_sparse_attention.py | 12 - .../test/unit/test_suite_partitioning.py | 4 +- .../test/unit/test_vae_loader.py | 9 - .../unit/test_vae_spatial_parallel_decode.py | 6 - .../test/unit/test_wan_attention_backend.py | 36 - .../arg_groups/model_overrides/__init__.py | 3 +- .../srt/arg_groups/model_overrides/minicpm.py | 4 +- python/sglang/srt/server_args.py | 3 +- python/sglang/test/ci/ci_register.py | 13 - .../sglang/test/ci/diffusion_suite_bridge.py | 35 + .../utils/diffusion/diffusion_case_parser.py | 4 +- .../diffusion/verify_diffusion_coverage.py | 3 +- scripts/lint/check_registered_tests.py | 106 ++ sgl-model-gateway/bindings/python/README.md | 3 +- .../python/tests/test_router_config.py | 423 -------- .../python/tests/test_startup_sequence.py | 845 +--------------- .../bindings/python/tests/test_validation.py | 509 ---------- test/README.md | 41 +- .../models}/test_text_models_gsm8k_eval.py | 0 .../models}/test_vlms_mmmu_eval.py | 0 .../amd/accuracy/mi30x/test_grok_eval_amd.py | 290 ------ .../mi35x/test_glm52_fp8_eval_mi35x.py | 2 +- .../registered/core/test_engine_child_pids.py | 5 +- .../core/test_request_queue_validation.py | 5 +- .../cp/test_deepseek_v4_flash_fp4_b200_cp.py | 2 +- test/registered/cpu/test_spec_eagle_cpu.py | 71 -- .../cpu/test_spec_eagle_parity_cpu.py | 29 - .../cpu/test_spec_eagle_topk_cpu.py | 68 -- .../aligner/unsharder/test_executor.py | 263 ----- .../tensor_comparator/test_comparator.py | 73 -- test/registered/debug_utils/test_dumper.py | 41 - .../test_tensor_dump_forward_hook.py | 58 +- .../e2e/diffusion/test_diffusion_1_gpu.py | 13 + .../diffusion/test_diffusion_1_gpu_5090.py | 13 + .../diffusion/test_diffusion_1_gpu_b200.py | 13 + .../e2e/diffusion/test_diffusion_2_gpu.py | 13 + .../e2e/diffusion/test_diffusion_bcg.py | 13 + .../test_diffusion_component_accuracy.py | 13 + .../e2e/diffusion/test_diffusion_unit.py | 13 + .../models}/test_compressed_tensors_models.py | 0 .../models}/test_deepseek_v3_fp4.py | 0 .../models}/test_deepseek_v3_mtp.py | 0 .../test_deepseek_v4_flash_fp4_b200.py | 0 .../test_deepseek_v4_flash_fp4_h200.py | 0 ...test_deepseek_v4_flash_fp4_megamoe_b200.py | 0 .../test_deepseek_v4_flash_fp8_h200.py | 0 .../models}/test_dsa_glm52_dp_mtp.py | 0 .../models}/test_dsa_glm52_hisparse.py | 0 .../models}/test_dsa_glm52_nvfp4_dp_mtp.py | 0 .../models}/test_dsa_glm52_nvfp4_tp_mtp.py | 0 .../test_dsa_glm52_pd_mtp_cp_layersplit.py | 0 .../models}/test_dsa_glm52_tp_mtp.py | 0 .../test_gemma4_fp8_per_expert_loading.py | 0 .../models}/test_generation_models.py | 0 .../models}/test_glm53_flash_b200.py | 0 .../models}/test_glm53_flash_h200.py | 0 .../models}/test_gpt_oss_4gpu_mxfp4.py | 0 .../models}/test_gpt_oss_sm120.py | 0 .../models}/test_inkling.py | 0 .../models}/test_inkling_small_nvfp4.py | 0 .../models}/test_inkling_unified.py | 0 .../models}/test_kimi_k3_b300.py | 0 .../models}/test_kimi_k3_b300_low_latency.py | 0 .../models}/test_kimi_linear_models.py | 0 .../test_kimi_linear_unified_memory.py | 0 ...imi_linear_unified_memory_dcp_blackwell.py | 0 .../models}/test_layernorm_sp.py | 0 .../models}/test_mimo_v2.py | 0 .../models}/test_mimo_v2_flash.py | 0 .../models}/test_minimax_m25_basic.py | 0 .../models}/test_ministral4_models.py | 0 .../models}/test_nvidia_nemotron_3_nano.py | 0 .../test_nvidia_nemotron_3_super_bf16.py | 0 .../test_nvidia_nemotron_3_super_bf16_mtp.py | 0 .../models}/test_qwen35_fp4_mtp.py | 0 .../models}/test_qwen3_next_models.py | 0 .../models}/test_qwen3_next_models_extra.py | 0 .../models}/test_qwen3_next_models_mtp.py | 0 .../models}/test_step3p5_flash_chain_mtp.py | 0 .../models}/test_transformers_backend_eval.py | 0 .../models}/test_transformers_models.py | 0 .../models}/test_vlm_models.py | 0 .../test_deepseek_v32_indexcache.py | 0 .../test_deepseek_v3_cutedsl_4gpu.py | 0 .../models_large}/test_glm52_fp8.py | 0 .../models_large}/test_glm_46.py | 0 .../models_large}/test_gpt_oss_120b.py | 0 .../test_inkling_nvfp4_nightly.py | 2 +- .../models_large}/test_kimi_k25.py | 0 .../test_laguna_nvfp4_nightly.py | 0 .../models_large}/test_ling_2_6_flash.py | 0 .../test_longcat_flash_lite_fp8.py | 0 .../models_large}/test_minimax_m25.py | 0 .../models_large}/test_mistral_large3.py | 0 .../test_nvidia_nemotron_3_super_nightly.py | 0 .../test_nvidia_nemotron_3_super_nvfp4.py | 0 .../models_large}/test_qwen35.py | 0 .../models_large}/test_ring_2_5_1t.py | 0 .../ops/diffusion/test_import_surface.py | 228 ----- .../ops/layernorm/test_kernels_namespace.py | 134 +-- .../kernels/test_kernel_inventory.py | 239 ----- .../test_gpt_oss_mlx_correctness.py | 3 +- .../test_qwen2_moe_mlx_correctness.py | 3 +- .../test_qwen3_moe_mlx_correctness.py | 3 +- .../models_e2e/test_dummy_grok_models.py | 41 - .../models_e2e/test_ministral3_models.py | 34 - .../test_npu_openai_function_calling.py | 940 ------------------ .../observability/test_priority_metrics.py | 2 - test/registered/observability/test_tracing.py | 5 +- .../openai_server/basic/test_http2_server.py | 5 +- .../function_call/test_anthropic_tool_use.py | 4 +- .../test_openai_function_calling.py | 53 +- .../validation/test_large_max_new_tokens.py | 4 +- .../validation/test_matched_stop.py | 4 +- .../test_request_length_validation.py | 5 +- .../{ => models}/test_text_models_perf.py | 0 .../perf/{ => models}/test_vlms_perf.py | 0 test/registered/profiling/test_profile_v2.py | 109 -- ..._unified_radix_cache_kl_hybrid_bitexact.py | 2 +- test/registered/reasoning/test_reasoning.py | 5 +- .../rl/test_update_weights_from_disk.py | 447 --------- .../spec/eagle/test_spec_eagle_stress.py | 2 +- .../{ => models}/test_stress_deepseek_v3.py | 0 .../{ => models}/test_stress_glm_4_6.py | 0 .../{ => models}/test_stress_kimi_k2.py | 0 .../{ => models}/test_stress_qwen3_235b.py | 0 test/registered/unit/README.md | 13 +- .../test_pynccl_allocator_import.py | 62 -- .../test_effective_state_surfaces.py | 401 -------- .../mlx/test_attention_patching.py | 3 +- .../mlx/test_attn_dp_request_capacity.py | 3 +- .../mlx/test_max_running_requests.py | 5 +- .../mlx/test_metal_profiler.py | 3 +- .../mlx/test_mlx_pool_dtype.py | 3 +- .../mlx/test_mlx_reference_correctness.py | 3 +- .../mlx/test_mlx_runner_pool_contract.py | 3 +- .../hardware_backend/mlx/test_mlx_sampling.py | 3 +- .../mlx/test_muse_glimmer_mlx_model.py | 3 +- .../hardware_backend/mlx/test_quantization.py | 3 +- .../mlx/test_runner_init_contract.py | 85 -- .../mlx/test_scheduler_mixin.py | 3 +- .../mlx/test_sliding_window_attention.py | 3 +- .../mlx/test_swa_radix_pool.py | 3 +- .../mlx/test_tp_worker_routing.py | 3 +- .../mlx/test_windowed_kv_cache.py | 3 +- .../moe/test_flashinfer_cutedsl_dispatch.py | 64 -- .../unit/lora/test_deepseek_mla_correction.py | 24 - .../test_session_unified_radix_cache.py | 73 -- .../runner/test_flashinfer_autotune.py | 97 -- .../models/test_draft_entry_hook_parity.py | 7 - .../unit/models/test_fusion_gate_coverage.py | 181 ---- .../models/test_kimi_k25_mm_projection.py | 41 - .../test_model_config_reads_resolved_input.py | 765 -------------- .../test_no_public_non_field_slot.py | 86 -- .../test_record_member_calls_resolve.py | 144 --- .../test_resolution_declarations.py | 225 +---- .../test_resolution_is_reproducible.py | 564 ----------- .../test_resolution_reads_no_bag.py | 333 ------- .../unit/test_bench_long_context.py | 128 --- .../unit/test_context_accessor_shadowing.py | 191 ---- ...test_dead_server_args_parameter_ratchet.py | 86 -- .../unit/test_legacy_global_ratchet.py | 65 -- .../unit/test_model_override_split.py | 120 --- .../unit/test_module_state_ratchet.py | 61 -- .../unit/test_parallel_adoption_ratchet.py | 73 -- .../unit/test_platform_address_not_frozen.py | 81 -- .../unit/test_pre_publish_readers.py | 56 +- .../unit/test_publish_precedes_bag_reads.py | 479 --------- .../unit/test_ray_driver_reads_the_bags.py | 54 +- test/registered/unit/test_runtime_context.py | 111 --- .../unit/test_runtime_context_config_bags.py | 35 - .../unit/test_server_args_cli_metadata.py | 66 -- .../unit/test_server_args_mutation_ratchet.py | 81 -- ..._server_args_no_instance_mutation_entry.py | 119 --- .../test_split_attention_backend_decisions.py | 85 +- .../unit/utils/test_torch_npu_patch_utils.py | 41 - test/run_suite.py | 9 + 199 files changed, 1185 insertions(+), 11812 deletions(-) create mode 100644 python/sglang/multimodal_gen/test/runner/diffusion_suite_runner.py delete mode 100644 python/sglang/multimodal_gen/test/unit/realtime/test_realtime_webui.py delete mode 100644 python/sglang/multimodal_gen/test/unit/test_diffusion_benchmark_skill.py delete mode 100644 python/sglang/multimodal_gen/test/unit/test_server_warmup_progress.py delete mode 100644 python/sglang/multimodal_gen/test/unit/test_wan_attention_backend.py create mode 100644 python/sglang/test/ci/diffusion_suite_bridge.py delete mode 100644 sgl-model-gateway/bindings/python/tests/test_router_config.py delete mode 100644 sgl-model-gateway/bindings/python/tests/test_validation.py rename test/registered/{eval => accuracy/models}/test_text_models_gsm8k_eval.py (100%) rename test/registered/{eval => accuracy/models}/test_vlms_mmmu_eval.py (100%) delete mode 100644 test/registered/amd/accuracy/mi30x/test_grok_eval_amd.py delete mode 100644 test/registered/cpu/test_spec_eagle_cpu.py delete mode 100644 test/registered/cpu/test_spec_eagle_parity_cpu.py delete mode 100644 test/registered/cpu/test_spec_eagle_topk_cpu.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_1_gpu.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_1_gpu_5090.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_1_gpu_b200.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_2_gpu.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_bcg.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_component_accuracy.py create mode 100644 test/registered/e2e/diffusion/test_diffusion_unit.py rename test/registered/{models_e2e => e2e/models}/test_compressed_tensors_models.py (100%) rename test/registered/{models_e2e => e2e/models}/test_deepseek_v3_fp4.py (100%) rename test/registered/{models_e2e => e2e/models}/test_deepseek_v3_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_deepseek_v4_flash_fp4_b200.py (100%) rename test/registered/{models_e2e => e2e/models}/test_deepseek_v4_flash_fp4_h200.py (100%) rename test/registered/{models_e2e => e2e/models}/test_deepseek_v4_flash_fp4_megamoe_b200.py (100%) rename test/registered/{models_e2e => e2e/models}/test_deepseek_v4_flash_fp8_h200.py (100%) rename test/registered/{models_e2e => e2e/models}/test_dsa_glm52_dp_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_dsa_glm52_hisparse.py (100%) rename test/registered/{models_e2e => e2e/models}/test_dsa_glm52_nvfp4_dp_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_dsa_glm52_nvfp4_tp_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_dsa_glm52_pd_mtp_cp_layersplit.py (100%) rename test/registered/{models_e2e => e2e/models}/test_dsa_glm52_tp_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_gemma4_fp8_per_expert_loading.py (100%) rename test/registered/{models_e2e => e2e/models}/test_generation_models.py (100%) rename test/registered/{models_e2e => e2e/models}/test_glm53_flash_b200.py (100%) rename test/registered/{models_e2e => e2e/models}/test_glm53_flash_h200.py (100%) rename test/registered/{models_e2e => e2e/models}/test_gpt_oss_4gpu_mxfp4.py (100%) rename test/registered/{models_e2e => e2e/models}/test_gpt_oss_sm120.py (100%) rename test/registered/{models_e2e => e2e/models}/test_inkling.py (100%) rename test/registered/{models_e2e => e2e/models}/test_inkling_small_nvfp4.py (100%) rename test/registered/{models_e2e => e2e/models}/test_inkling_unified.py (100%) rename test/registered/{models_e2e => e2e/models}/test_kimi_k3_b300.py (100%) rename test/registered/{models_e2e => e2e/models}/test_kimi_k3_b300_low_latency.py (100%) rename test/registered/{models_e2e => e2e/models}/test_kimi_linear_models.py (100%) rename test/registered/{models_e2e => e2e/models}/test_kimi_linear_unified_memory.py (100%) rename test/registered/{models_e2e => e2e/models}/test_kimi_linear_unified_memory_dcp_blackwell.py (100%) rename test/registered/{models_e2e => e2e/models}/test_layernorm_sp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_mimo_v2.py (100%) rename test/registered/{models_e2e => e2e/models}/test_mimo_v2_flash.py (100%) rename test/registered/{models_e2e => e2e/models}/test_minimax_m25_basic.py (100%) rename test/registered/{models_e2e => e2e/models}/test_ministral4_models.py (100%) rename test/registered/{models_e2e => e2e/models}/test_nvidia_nemotron_3_nano.py (100%) rename test/registered/{models_e2e => e2e/models}/test_nvidia_nemotron_3_super_bf16.py (100%) rename test/registered/{models_e2e => e2e/models}/test_nvidia_nemotron_3_super_bf16_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_qwen35_fp4_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_qwen3_next_models.py (100%) rename test/registered/{models_e2e => e2e/models}/test_qwen3_next_models_extra.py (100%) rename test/registered/{models_e2e => e2e/models}/test_qwen3_next_models_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_step3p5_flash_chain_mtp.py (100%) rename test/registered/{models_e2e => e2e/models}/test_transformers_backend_eval.py (100%) rename test/registered/{models_e2e => e2e/models}/test_transformers_models.py (100%) rename test/registered/{models_e2e => e2e/models}/test_vlm_models.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_deepseek_v32_indexcache.py (100%) rename test/registered/{4-gpu-models => e2e/models_large}/test_deepseek_v3_cutedsl_4gpu.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_glm52_fp8.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_glm_46.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_gpt_oss_120b.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_inkling_nvfp4_nightly.py (98%) rename test/registered/{8-gpu-models => e2e/models_large}/test_kimi_k25.py (100%) rename test/registered/{4-gpu-models => e2e/models_large}/test_laguna_nvfp4_nightly.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_ling_2_6_flash.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_longcat_flash_lite_fp8.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_minimax_m25.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_mistral_large3.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_nvidia_nemotron_3_super_nightly.py (100%) rename test/registered/{4-gpu-models => e2e/models_large}/test_nvidia_nemotron_3_super_nvfp4.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_qwen35.py (100%) rename test/registered/{8-gpu-models => e2e/models_large}/test_ring_2_5_1t.py (100%) delete mode 100644 test/registered/kernels/ops/diffusion/test_import_surface.py delete mode 100644 test/registered/kernels/test_kernel_inventory.py delete mode 100644 test/registered/models_e2e/test_dummy_grok_models.py delete mode 100644 test/registered/models_e2e/test_ministral3_models.py delete mode 100644 test/registered/npu/interface/test_npu_openai_function_calling.py rename test/registered/perf/{ => models}/test_text_models_perf.py (100%) rename test/registered/perf/{ => models}/test_vlms_perf.py (100%) delete mode 100644 test/registered/profiling/test_profile_v2.py delete mode 100644 test/registered/rl/test_update_weights_from_disk.py rename test/registered/stress/{ => models}/test_stress_deepseek_v3.py (100%) rename test/registered/stress/{ => models}/test_stress_glm_4_6.py (100%) rename test/registered/stress/{ => models}/test_stress_kimi_k2.py (100%) rename test/registered/stress/{ => models}/test_stress_qwen3_235b.py (100%) delete mode 100644 test/registered/unit/distributed/test_pynccl_allocator_import.py delete mode 100644 test/registered/unit/entrypoints/test_effective_state_surfaces.py delete mode 100644 test/registered/unit/hardware_backend/mlx/test_runner_init_contract.py delete mode 100644 test/registered/unit/layers/moe/test_flashinfer_cutedsl_dispatch.py delete mode 100644 test/registered/unit/lora/test_deepseek_mla_correction.py delete mode 100644 test/registered/unit/model_executor/runner/test_flashinfer_autotune.py delete mode 100644 test/registered/unit/models/test_fusion_gate_coverage.py delete mode 100644 test/registered/unit/models/test_kimi_k25_mm_projection.py delete mode 100644 test/registered/unit/server_args/test_model_config_reads_resolved_input.py delete mode 100644 test/registered/unit/server_args/test_no_public_non_field_slot.py delete mode 100644 test/registered/unit/server_args/test_record_member_calls_resolve.py delete mode 100644 test/registered/unit/server_args/test_resolution_reads_no_bag.py delete mode 100644 test/registered/unit/test_bench_long_context.py delete mode 100644 test/registered/unit/test_context_accessor_shadowing.py delete mode 100644 test/registered/unit/test_dead_server_args_parameter_ratchet.py delete mode 100644 test/registered/unit/test_legacy_global_ratchet.py delete mode 100644 test/registered/unit/test_model_override_split.py delete mode 100644 test/registered/unit/test_module_state_ratchet.py delete mode 100644 test/registered/unit/test_parallel_adoption_ratchet.py delete mode 100644 test/registered/unit/test_platform_address_not_frozen.py delete mode 100644 test/registered/unit/test_publish_precedes_bag_reads.py delete mode 100644 test/registered/unit/test_server_args_mutation_ratchet.py delete mode 100644 test/registered/unit/test_server_args_no_instance_mutation_entry.py delete mode 100644 test/registered/unit/utils/test_torch_npu_patch_utils.py diff --git a/.claude/rules/unit-test-admission.md b/.claude/rules/unit-test-admission.md index fc573d5f1..4731a596e 100644 --- a/.claude/rules/unit-test-admission.md +++ b/.claude/rules/unit-test-admission.md @@ -1,6 +1,7 @@ --- paths: - "test/**/*.py" + - "python/sglang/multimodal_gen/test/**/*.py" --- # Unit Test Admission Criteria @@ -68,5 +69,13 @@ One strong case beats several weak ones: each additional case must guard a distinct failure mode. Ask "which bug escapes if I delete this case?" -- no answer means delete it. +New cases join an existing file in the same subsystem by default. Create a new +file only when it needs a different fixture, dependency, owner, or CI contract; +every file pays a separate interpreter-import cost in the CPU gate. + +Suite cadence is part of admission: if a failing run cannot be attributed to a +single PR's diff, the test belongs in a nightly or weekly suite rather than a +per-commit lane. + Test mechanics (placement, CI registration, fixtures) live in [`write-sglang-test`](../skills/write-sglang-test/SKILL.md). diff --git a/.claude/skills/write-sglang-test/SKILL.md b/.claude/skills/write-sglang-test/SKILL.md index 1faaf2d9b..2a4358b04 100644 --- a/.claude/skills/write-sglang-test/SKILL.md +++ b/.claude/skills/write-sglang-test/SKILL.md @@ -11,14 +11,14 @@ This skill covers **how to write and register tests**. For CI pipeline internals 1. **Always use `CustomTestCase`** — never raw `unittest.TestCase`. It ensures `tearDownClass` runs even when `setUpClass` fails, preventing resource leaks in CI. 2. **`tearDownClass` must be defensive** — use `hasattr`/null checks before accessing resources (e.g. `cls.process`) that `setUpClass` may not have finished allocating. -3. **Place tests in `test/registered//`** — including JIT kernel tests and benchmarks, which live in `test/registered/jit/` and `test/registered/jit/benchmark/` (nested subfolders are allowed) +3. **Place tests in `test/registered///`** — `` is `unit`, `kernel`, `e2e`, `accuracy`, `perf`, or `stress`; hardware belongs in registrations, not directory names 4. **Reuse server fixtures** — inherit from `DefaultServerBase` or write `setUpClass`/`tearDownClass` with `popen_launch_server` -5. **Prefer mock over real server** — when testing logic that doesn't need a server / engine launch (middleware, request routing, config validation, argument parsing), use `unittest.mock.patch` / `MagicMock` and place tests in `test/registered/unit/`. Only launch a real server when the test genuinely needs inference results or server lifecycle behavior. +5. **Mock boundaries, not SGLang behavior** — mock slow or external dependencies only when the assertion still checks an observable result, state transition, or error. A test whose evidence is only `assert_called*` mirrors its mock and is not admissible. Launch a real server only when inference results or lifecycle behavior are the contract under test. JIT kernel notes: - If the task is adding or updating code under `python/sglang/kernels/jit/`, prefer the `add-jit-kernel` skill first. -- JIT kernel correctness tests use `test/registered/jit/**/test_*.py`. -- JIT kernel benchmarks use `test/registered/jit/benchmark/**/bench_*.py`. +- New JIT kernel correctness tests use `test/registered/kernel/jit/**/test_*.py`. +- New JIT kernel benchmarks use `test/registered/kernel/jit/benchmark/**/bench_*.py`. - Those files are executed by `test/run_suite.py` through dedicated kernel suites (`base-b-kernel-*`); a `register_*_ci(...)` call placed under `python/sglang/` is rejected by the `check-no-registered-tests-in-package` pre-commit hook. --- @@ -27,7 +27,7 @@ JIT kernel notes: | Scenario | Model | CI Registration | Suite | |----------|-------|-----------------|-------| -| **Unit tests** (no server / engine launch) | None | `register_cpu_ci` (prefer) or `register_cuda_ci` | `base-a-test-cpu` or `base-b-test-1-gpu-small` | +| **Unit tests** (no server / engine launch) | None | `register_cpu_ci` | `base-a-test-cpu` | | **Common / backend-independent** (middleware, abort, routing, config, arg parsing) | `DEFAULT_SMALL_MODEL_NAME_FOR_TEST` (1B) | `register_cuda_ci` only | `base-b-test-1-gpu-small` | | **Model-agnostic functionality** (sampling, session, OpenAI API features) | `DEFAULT_SMALL_MODEL_NAME_FOR_TEST` (1B) | `register_cuda_ci` (+ AMD if relevant) | `base-b-test-1-gpu-small` | | **General performance** (single node, no spec/DP/parallelism) | `DEFAULT_MODEL_NAME_FOR_TEST` (8B) | `register_cuda_ci` | `base-b-test-1-gpu-large` | @@ -70,10 +70,10 @@ A per-commit suite name is **generated** from registration metadata as `{stage}- | `base-b-test-1-gpu-large` | `1-gpu-h100` | Tests that need H100-class memory or kernels (e.g. FA3) | | `base-b-test-2-gpu-large` | `2-gpu-h100` | Two-GPU correctness and parallelism (TP/PP) on H100 | | `base-b-test-4-gpu-b200` | `4-gpu-b200` | Early Blackwell coverage (SM100+ paths) on four GPUs | -| `base-b-kernel-unit-test-1-gpu-large` | `1-gpu-h100` | JIT kernel correctness tests under `test/registered/jit/` | +| `base-b-kernel-unit-test-1-gpu-large` | `1-gpu-h100` | JIT kernel correctness tests under `test/registered/kernel/jit/` | | `base-b-kernel-unit-test-4-gpu-b200` | `4-gpu-b200` | JIT kernel correctness tests for Blackwell / SM100-specific paths | -| `base-b-kernel-unit-test-8-gpu-h200` | `8-gpu-h200` | Multi-GPU JIT kernel correctness tests under `test/registered/jit/` | -| `base-b-kernel-benchmark-test-1-gpu-large` | `1-gpu-h100` | JIT kernel benchmark files under `test/registered/jit/benchmark/` | +| `base-b-kernel-unit-test-8-gpu-h200` | `8-gpu-h200` | Multi-GPU JIT kernel correctness tests under `test/registered/kernel/jit/` | +| `base-b-kernel-benchmark-test-1-gpu-large` | `1-gpu-h100` | JIT kernel benchmark files under `test/registered/kernel/jit/benchmark/` | | `base-c-test-4-gpu-h100` | `4-gpu-h100` | Large 4-GPU H100 integration and scaling tests | | `base-c-test-8-gpu-h200` | `8-gpu-h200` | Large 8-GPU H200 runs for big models and parallelism | | `base-c-test-8-gpu-h20` | `8-gpu-h20` | Large 8-GPU H20 runs for big models | @@ -162,7 +162,7 @@ from sglang.test.ci.ci_register import register_cpu_ci from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=5, suite="base-a-test-cpu") -# Prefer CPU. Only use register_cuda_ci when the test truly needs a GPU. +# Unit tests are CPU-only. GPU operator tests belong under `kernel/`. class TestTargetClass(CustomTestCase): def test_basic_behavior(self): @@ -180,7 +180,13 @@ if __name__ == "__main__": unittest.main() ``` -Use `unittest.mock.patch` / `MagicMock` to mock dependencies and isolate the logic under test. If the module transitively imports GPU-only packages (e.g. `sgl_kernel`), they can be stubbed so the test runs on CPU CI. Do not modify `sys.modules` at module level — use `patch.dict` (as a class decorator or with `start`/`stop`) to ensure cleanup and avoid cross-test pollution. See `test/registered/unit/README.md` for details and examples. +Use `unittest.mock.patch` / `MagicMock` only at dependency boundaries. Assert the +resulting value, state, protocol output, or error—not merely that the mock was +called. If the module transitively imports GPU-only packages (e.g. `sgl_kernel`), +they can be stubbed so the test runs on CPU CI. Do not modify `sys.modules` at +module level—use `patch.dict` (as a class decorator or with `start`/`stop`) to +ensure cleanup and avoid cross-test pollution. See +`test/registered/unit/README.md` for details and examples. **Quality bar** — test real logic (validation boundaries, state transitions, error paths, branching, etc.). Skip tests that just verify Python itself works (e.g., "does calling an abstract method raise `NotImplementedError`?", "does a dataclass store the field I assigned?"). Consolidate repetitive patterns into parameterized tests. No production code changes in test PRs. @@ -380,15 +386,12 @@ Every call generates a suite named `{stage}-test-{runner_config}`, e.g. `base-b- ``` test/ ├── registered/ # CI tests (auto-discovered by run_suite.py) -│ ├── unit/ # No server / engine launch (see test/registered/unit/README.md) -│ ├── kernels/ # CUDA kernel correctness (no server, GPU required) -│ ├── sampling/ # test_penalty.py, test_sampling_params.py ... -│ ├── sessions/ # test_session_control.py ... -│ ├── openai_server/ # basic/, features/, validation/ ... -│ ├── spec/ # eagle/, utils/ ... -│ ├── models/ # model-specific accuracy tests -│ ├── perf/ # performance benchmarks -│ └── / # create new category if needed +│ ├── unit// # CPU-only; no server or model weights +│ ├── kernel// # accelerator operator correctness/benchmarks +│ ├── e2e// # engine/server integration +│ ├── accuracy// # scheduled eval floors +│ ├── perf// # scheduled latency/throughput contracts +│ └── stress// # stress/weekly coverage ├── manual/ # Non-CI: debugging, one-off, manual verification └── run_suite.py # CI runner (scans registered/ plus jit_kernel test/benchmark files) @@ -398,10 +401,11 @@ python/sglang/kernels/jit/ ``` **Decision rule** (see also `test/registered/README.md`): -- Component logic, no server → `registered/unit/` -- JIT kernel correctness / benchmarks → `test/registered/jit/` or `test/registered/jit/benchmark/` -- Other kernel correctness → `registered/kernels/` -- Server needed → `registered//` +- CPU component logic, no server → `registered/unit//` +- JIT kernel correctness / benchmarks → `registered/kernel/jit/` +- Other accelerator operator correctness → `registered/kernel//` +- Server needed → `registered/e2e//` +- Eval floor / performance contract → `registered/{accuracy,perf}//` - Local debugging → `manual/` --- @@ -443,8 +447,8 @@ Before submitting a test: - [ ] Inherits from `CustomTestCase` (not `unittest.TestCase`) - [ ] Has `register_*_ci(...)` call at module level -- [ ] Placed in `test/registered//` (JIT kernel test/benchmark → `test/registered/jit/` or `test/registered/jit/benchmark/`) -- [ ] JIT kernel work: test files live in `test/registered/jit/`; only test-only helpers stay under `python/sglang/kernels/jit/` +- [ ] Placed in `test/registered///` +- [ ] JIT kernel work: test files live in `test/registered/kernel/jit/`; only test-only helpers stay under `python/sglang/kernels/jit/` - [ ] Backend-independent tests: `register_cuda_ci` only + smallest model - [ ] Logic that doesn't need a server / engine launch → unit test in `registered/unit/` (see Unit Tests section) - [ ] `setUpClass` launches server, `tearDownClass` kills it (if server-based) diff --git a/.github/workflows/pr-test-multimodal-gen.yml b/.github/workflows/pr-test-multimodal-gen.yml index 61cabb925..fb650d87e 100644 --- a/.github/workflows/pr-test-multimodal-gen.yml +++ b/.github/workflows/pr-test-multimodal-gen.yml @@ -119,17 +119,16 @@ jobs: timeout-minutes: 240 env: RUNAI_STREAMER_MEMORY_LIMIT: 0 - CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} - PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-1gpu }} + DIFFUSION_CONTINUE_ON_ERROR: ${{ inputs.continue_on_error }} + DIFFUSION_PARTITION_ID: ${{ matrix.part }} + DIFFUSION_TOTAL_PARTITIONS: ${{ needs.compute-diffusion-partitions.outputs['partition-count-1gpu'] }} + DIFFUSION_PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-1gpu }} SGLANG_DIFFUSION_ARTIFACT_DIR: ${{ github.workspace }}/diffusion-failures run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py \ - --suite 1-gpu \ - --partition-id ${{ matrix.part }} \ - --total-partitions ${{ needs.compute-diffusion-partitions.outputs['partition-count-1gpu'] }} \ - --partition-plan-json "$PARTITION_PLAN_JSON" \ - $CONTINUE_ON_ERROR_FLAG + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-1-gpu-h100 \ + --timeout-per-file 14400 - name: Upload execution report if: always() @@ -192,13 +191,13 @@ jobs: timeout-minutes: 120 env: RUNAI_STREAMER_MEMORY_LIMIT: 0 - CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} + DIFFUSION_CONTINUE_ON_ERROR: ${{ inputs.continue_on_error }} SGLANG_DIFFUSION_ARTIFACT_DIR: ${{ github.workspace }}/diffusion-failures run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py \ - --suite 1-gpu-5090 \ - $CONTINUE_ON_ERROR_FLAG + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-1-gpu-5090 \ + --timeout-per-file 7200 - name: Upload execution report if: always() @@ -261,13 +260,13 @@ jobs: timeout-minutes: 60 env: RUNAI_STREAMER_MEMORY_LIMIT: 0 - CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} + DIFFUSION_CONTINUE_ON_ERROR: ${{ inputs.continue_on_error }} SGLANG_DIFFUSION_ARTIFACT_DIR: ${{ github.workspace }}/diffusion-bcg-artifacts run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py \ - --suite bcg-diffusion \ - $CONTINUE_ON_ERROR_FLAG + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-bcg-1-gpu-h100 \ + --timeout-per-file 3600 - name: Upload BCG diffusion artifacts if: always() @@ -329,17 +328,16 @@ jobs: HF_TOKEN: ${{ secrets.SGLANG_DIFFUSION_CI_HF_TOKEN || secrets.HF_TOKEN }} HUGGING_FACE_HUB_TOKEN: ${{ secrets.SGLANG_DIFFUSION_CI_HF_TOKEN || secrets.HF_TOKEN }} RUNAI_STREAMER_MEMORY_LIMIT: 0 - CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} - PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-2gpu }} + DIFFUSION_CONTINUE_ON_ERROR: ${{ inputs.continue_on_error }} + DIFFUSION_PARTITION_ID: ${{ matrix.part }} + DIFFUSION_TOTAL_PARTITIONS: ${{ needs.compute-diffusion-partitions.outputs['partition-count-2gpu'] }} + DIFFUSION_PARTITION_PLAN_JSON: ${{ needs.compute-diffusion-partitions.outputs.plan-2gpu }} SGLANG_DIFFUSION_ARTIFACT_DIR: ${{ github.workspace }}/diffusion-failures run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py \ - --suite 2-gpu \ - --partition-id ${{ matrix.part }} \ - --total-partitions ${{ needs.compute-diffusion-partitions.outputs['partition-count-2gpu'] }} \ - --partition-plan-json "$PARTITION_PLAN_JSON" \ - $CONTINUE_ON_ERROR_FLAG + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-2-gpu-h100 \ + --timeout-per-file 14400 - name: Upload execution report if: always() @@ -402,12 +400,12 @@ jobs: timeout-minutes: 240 env: RUNAI_STREAMER_MEMORY_LIMIT: 0 - CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} + DIFFUSION_CONTINUE_ON_ERROR: ${{ inputs.continue_on_error }} run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py \ - --suite component-accuracy \ - $CONTINUE_ON_ERROR_FLAG + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-component-2-gpu-h100 \ + --timeout-per-file 14400 - uses: ./.github/actions/upload-cuda-coredumps if: always() @@ -453,13 +451,13 @@ jobs: timeout-minutes: 240 env: RUNAI_STREAMER_MEMORY_LIMIT: 0 - CONTINUE_ON_ERROR_FLAG: ${{ inputs.continue_on_error == 'true' && '--continue-on-error' || '' }} + DIFFUSION_CONTINUE_ON_ERROR: ${{ inputs.continue_on_error }} SGLANG_DIFFUSION_ARTIFACT_DIR: ${{ github.workspace }}/diffusion-failures run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py \ - --suite 1-gpu-b200 \ - $CONTINUE_ON_ERROR_FLAG + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-1-gpu-b200 \ + --timeout-per-file 14400 - name: Upload diffusion failure artifacts if: always() @@ -519,8 +517,10 @@ jobs: - name: Run diffusion unit tests timeout-minutes: 60 run: | - cd python - python3 sglang/multimodal_gen/test/run_suite.py --suite unit + python3 test/run_suite.py \ + --hw cuda \ + --suite base-b-test-diffusion-unit-1-gpu-h100 \ + --timeout-per-file 3600 diffusion-coverage-check: needs: [multimodal-gen-test-1-gpu, multimodal-gen-test-2-gpu] diff --git a/python/sglang/kernels/ops/diffusion/README.md b/python/sglang/kernels/ops/diffusion/README.md index c28c6fb6b..7543fb83f 100644 --- a/python/sglang/kernels/ops/diffusion/README.md +++ b/python/sglang/kernels/ops/diffusion/README.md @@ -17,7 +17,7 @@ from sglang.kernels.ops.diffusion import fused_rmsnorm_scale_shift_bitexact ``` **Import from the package, never from a submodule.** The internal layout is -free to move; the facade is not. `test_import_surface.py` enforces this, with +free to move; the facade is not. Callers should use the facade, with a small allowlist for tests that deliberately exercise one backend. Resolution is lazy (PEP 562): the backends have disjoint heavy dependencies @@ -189,7 +189,7 @@ inspecting model modules is its whole job. generated by the KDA workflow in `sglang.kernels.kda_kernels`, together with its source revision and any JIT CUDA source files. 2. Export it from `__init__.py` (`_EXPORTS`) and register a `KernelSpec` - (`_SPECS`) — `test_import_surface.py` checks both resolve. + (`_SPECS`). 3. Give it a `can_use_*` predicate; raise, don't return `None`. 4. State the numerical contract in the module docstring, including which shapes it was verified on. diff --git a/python/sglang/kernels/ops/diffusion/__init__.py b/python/sglang/kernels/ops/diffusion/__init__.py index 57de395de..468d11bbd 100644 --- a/python/sglang/kernels/ops/diffusion/__init__.py +++ b/python/sglang/kernels/ops/diffusion/__init__.py @@ -5,8 +5,8 @@ This module is the **only** supported import surface for these kernels:: from sglang.kernels.ops.diffusion import fused_rmsnorm_scale_shift_bitexact Importing a submodule directly (``...diffusion.norm.norm_triton``) couples the -caller to the file layout; ``test_import_surface.py`` guards against it. The -one exception is a test that deliberately exercises a single backend. +caller to the file layout. The one exception is a test that deliberately +exercises a single backend. Layout -- ordinary implementations use one subpackage per **operator domain** (``norm``, ``modulate``, ``rope``, ``activation``, ``attention``, ``routing``, diff --git a/python/sglang/multimodal_gen/test/run_suite.py b/python/sglang/multimodal_gen/test/run_suite.py index c5aec5f82..13246720d 100644 --- a/python/sglang/multimodal_gen/test/run_suite.py +++ b/python/sglang/multimodal_gen/test/run_suite.py @@ -1,691 +1,6 @@ -""" -Test runner for multimodal_gen that manages test suites and parallel execution. - -For diffusion 1-gpu/2-gpu suites, cases are partitioned by estimated runtime -using LPT so each CI shard has a similar total runtime. -""" - -import argparse -import copy -import json -import os -import random -import subprocess -import sys -import time -from dataclasses import dataclass -from pathlib import Path - -import tabulate - -from sglang.multimodal_gen.runtime.platforms import current_platform -from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger -from sglang.multimodal_gen.test.partitioning import PartitionItem, assign_partition -from sglang.multimodal_gen.test.runner.pytest_runner import ( - partition_items_by_index, - run_pytest, -) -from sglang.multimodal_gen.test.server.testcase_configs import ( - BASELINE_CONFIG, - DiffusionTestCase, -) - -# TODO: remove duplicated code -if current_platform.is_npu(): - from sglang.multimodal_gen.test.server.ascend.testcase_configs_npu import ( - _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE, - COMPONENT_ACCURACY_SUITES, - DEFAULT_EST_TIME_SECONDS, - DEFAULT_STANDALONE_EST_TIME_SECONDS, - FILE_SUITES, - PARAMETRIZED_CASE_GROUPS, - STANDALONE_FILE_EST_TIMES, - STANDALONE_FILES, - STARTUP_OVERHEAD_SECONDS, - SUITES, - ) -else: - from sglang.multimodal_gen.test.server.gpu_cases import ( # noqa: F401 It is used by ci scripts - _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE, - _UPDATE_WEIGHTS_MODEL_PAIR_ENV, - _UPDATE_WEIGHTS_MODEL_PAIR_IDS, - COMPONENT_ACCURACY_FILE_NUM_GPUS, - COMPONENT_ACCURACY_SUITES, - DEFAULT_EST_TIME_SECONDS, - DEFAULT_STANDALONE_EST_TIME_SECONDS, - FILE_SUITES, - ONE_GPU_CASES, - PARAMETRIZED_CASE_GROUPS, - STANDALONE_FILE_EST_TIMES, - STANDALONE_FILES, - STARTUP_OVERHEAD_SECONDS, - STRICT_SUITES, - SUITES, - TWO_GPU_CASES, - ) - - -logger = init_logger(__name__) - - -@dataclass(frozen=True) -class PartitionAssignment: - case_ids: list[str] - standalone_files: list[str] - estimated_time: float | None = None - missing_standalone_estimates: list[str] | None = None - - -def get_case_est_time(case_id: str) -> float: - scenario = BASELINE_CONFIG.scenarios.get(case_id) - if scenario is None: - return DEFAULT_EST_TIME_SECONDS - if scenario.estimated_full_test_time_s is not None: - return scenario.estimated_full_test_time_s - return scenario.expected_e2e_ms / 1000.0 + STARTUP_OVERHEAD_SECONDS - - -def get_standalone_file_est_time( - suite: str, standalone_file: str -) -> tuple[float, bool]: - suite_est_times = STANDALONE_FILE_EST_TIMES.get(suite, {}) - if standalone_file not in suite_est_times: - return DEFAULT_STANDALONE_EST_TIME_SECONDS, True - return suite_est_times[standalone_file], False - - -def get_all_standalone_file_est_times() -> dict[str, dict[str, float]]: - return copy.deepcopy(STANDALONE_FILE_EST_TIMES) - - -def validate_standalone_file_est_times() -> dict[str, list[str]]: - missing_by_suite: dict[str, list[str]] = {} - for suite, standalone_files in STANDALONE_FILES.items(): - suite_est_times = STANDALONE_FILE_EST_TIMES.get(suite, {}) - missing = [ - standalone_file - for standalone_file in standalone_files - if standalone_file not in suite_est_times - ] - if missing: - missing_by_suite[suite] = missing - return missing_by_suite - - -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: - return f"standalone:{standalone_file}" - - -def parse_partition_plan( - suite: str, - partition_id: int, - total_partitions: int, - plan_json: str, -) -> PartitionAssignment: - plan = json.loads(plan_json) - if plan.get("suite") != suite: - raise ValueError( - f"Partition plan suite mismatch: expected {suite!r}, " - f"got {plan.get('suite')!r}" - ) - - partition_count = plan.get("partition_count") - if partition_count != total_partitions: - raise ValueError( - f"Partition count mismatch for suite {suite!r}: " - f"plan={partition_count}, matrix={total_partitions}" - ) - - partitions = plan.get("partitions", []) - selected_partition = None - for partition in partitions: - if partition.get("part") == partition_id: - selected_partition = partition - break - - if selected_partition is None: - raise ValueError( - f"Partition {partition_id} not found in plan for suite {suite!r}" - ) - - return PartitionAssignment( - case_ids=list(selected_partition.get("case_ids", [])), - standalone_files=list(selected_partition.get("standalone_files", [])), - estimated_time=selected_partition.get("estimated_time"), - missing_standalone_estimates=list( - selected_partition.get("missing_standalone_estimates", []) - ), - ) - - -def build_local_partition_assignment( - suite: str, - partition_id: int, - total_partitions: int, -) -> PartitionAssignment: - """Assign this shard's work when CI did not precompute a partition plan. - - Lanes with a hardcoded ``--total-partitions`` (the AMD ones) cannot give - every standalone file a shard of its own, so standalone files are LPT - balanced together with the parametrized cases instead. - """ - items = [ - PartitionItem(kind="case", item_id=case.id, est_time=get_case_est_time(case.id)) - for case in _get_dynamic_suite_cases(suite) - ] - for standalone_file in STANDALONE_FILES.get(suite, []): - items.append( - PartitionItem( - kind="standalone", - item_id=standalone_file, - est_time=get_standalone_file_est_time(suite, standalone_file)[0], - ) - ) - - my_items = assign_partition(items, partition_id, total_partitions) - return PartitionAssignment( - case_ids=[item.item_id for item in my_items if item.kind == "case"], - standalone_files=[ - item.item_id for item in my_items if item.kind == "standalone" - ], - ) - - -def _merge_execution_results( - executed_cases: list[str], - case_results: dict[str, str], - new_executed_cases: list[str], - new_case_results: dict[str, str], -) -> None: - executed_cases.extend( - case_id for case_id in new_executed_cases if case_id not in executed_cases - ) - case_results.update(new_case_results) - - -def _format_standalone_estimate_snippet( - suite: str, standalone_file: str, measured_full_test_time_s: float -) -> str: - return ( - f'"{suite}": {{\n "{standalone_file}": {measured_full_test_time_s:.1f},\n}}' - ) - - -def _print_missing_standalone_estimate_message( - suite: str, - standalone_file: str, - measured_full_test_time_s: float, -) -> None: - snippet = _format_standalone_estimate_snippet( - suite, standalone_file, measured_full_test_time_s - ) - logger.error( - f"\n{'=' * 60}\n" - f'Add standalone estimate for suite "{suite}" and file "{standalone_file}":\n\n' - f"File: python/sglang/multimodal_gen/test/run_suite.py\n\n" - f"Current partition used fallback estimate: " - f"{DEFAULT_STANDALONE_EST_TIME_SECONDS:.1f}s\n\n" - f"{snippet}\n" - f"{'=' * 60}\n" - ) - - -def _run_standalone_file( - suite: str, - standalone_rel: str, - target_dir: Path, - extra_filter: str | None = None, -) -> tuple[int, list[str], dict[str, str], dict]: - if standalone_rel == _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE: - _maybe_pin_update_weights_model_pair([standalone_rel]) - - est_time, used_fallback_estimate = get_standalone_file_est_time( - suite, standalone_rel - ) - standalone_file = _resolve_suite_files(target_dir, [standalone_rel], strict=True)[0] - junit_xml_path = str( - target_dir / f"junit_results_{suite}_{Path(standalone_rel).stem}.xml" - ) - start_time = time.perf_counter() - exit_code, _, _ = run_pytest( - [standalone_file], - filter_expr=extra_filter, - junit_xml_path=junit_xml_path, - ) - measured_full_test_time_s = round(time.perf_counter() - start_time, 1) - standalone_key = _normalize_standalone_key(standalone_rel) - measurement = { - "suite": suite, - "standalone_file": standalone_rel, - "measured_full_test_time_s": measured_full_test_time_s, - "used_fallback_estimate": used_fallback_estimate, - "fallback_estimate_s": DEFAULT_STANDALONE_EST_TIME_SECONDS, - "had_configured_estimate": not used_fallback_estimate, - "configured_or_fallback_estimate_s": est_time, - } - if used_fallback_estimate: - _print_missing_standalone_estimate_message( - suite, standalone_rel, measured_full_test_time_s - ) - return ( - exit_code, - [standalone_key], - {standalone_key: "pass" if exit_code == 0 else "fail"}, - measurement, - ) - - -def parse_args(): - suite_choices = sorted(set(FILE_SUITES) | set(PARAMETRIZED_CASE_GROUPS)) - parser = argparse.ArgumentParser(description="Run multimodal_gen test suite") - parser.add_argument( - "--suite", - type=str, - required=True, - choices=suite_choices, - help="The test suite to run.", - ) - parser.add_argument( - "--partition-id", - type=int, - default=0, - help="Index of the current partition (for parallel execution)", - ) - parser.add_argument( - "--total-partitions", - type=int, - default=1, - help="Total number of partitions", - ) - parser.add_argument( - "--base-dir", - type=str, - default="server", - help="Base directory for tests relative to this script's parent", - ) - parser.add_argument( - "-k", - "--filter", - type=str, - default=None, - help="Pytest filter expression (passed to pytest -k)", - ) - parser.add_argument( - "--continue-on-error", - action="store_true", - default=False, - help="Continue running remaining tests even if one fails.", - ) - parser.add_argument( - "--partition-plan-json", - type=str, - default=None, - help="Full partition plan JSON for the current suite.", - ) - return parser.parse_args() - - -def write_execution_report( - suite: str, - partition_id: int, - total_partitions: int, - executed_cases: list[str], - is_standalone: bool = False, - standalone_file: str | None = None, - case_results: dict[str, str] | None = None, - missing_standalone_estimates: list[str] | None = None, - standalone_measurements: list[dict] | None = None, -) -> str: - report = { - "suite": suite, - "partition_id": partition_id, - "total_partitions": total_partitions, - "is_standalone": is_standalone, - "standalone_file": standalone_file, - "executed_cases": executed_cases, - "case_results": case_results or {}, - "missing_standalone_estimates": missing_standalone_estimates or [], - "standalone_measurements": standalone_measurements or [], - } - - report_filename = f"execution_report_{suite}_{partition_id}.json" - report_path = Path(__file__).parent / report_filename - with open(report_path, "w", encoding="utf-8") as f: - json.dump(report, f, indent=2) - - logger.info("Execution report written to: %s", report_path) - return str(report_path) - - -def run_component_accuracy_files(files, filter_expr=None, continue_on_error=False): - exit_code = 0 - for file_path in files: - file_name = Path(file_path).name - num_gpus = COMPONENT_ACCURACY_FILE_NUM_GPUS.get(file_name, 1) - if num_gpus > 1: - cmd = [ - sys.executable, - "-m", - "torch.distributed.run", - f"--nproc_per_node={num_gpus}", - "-m", - "pytest", - "-s", - "-v", - ] - else: - cmd = [sys.executable, "-m", "pytest", "-s", "-v"] - - if filter_expr: - cmd.extend(["-k", filter_expr]) - cmd.append(file_path) - - print(f"Running command: {' '.join(cmd)}") - file_exit_code = subprocess.call(cmd) - if file_exit_code == 5: - print( - "No tests collected (exit code 5). This is expected when filters " - "deselect all tests in a file. Treating as success." - ) - file_exit_code = 0 - if file_exit_code != 0 and exit_code == 0: - exit_code = file_exit_code - if file_exit_code != 0 and not continue_on_error: - return file_exit_code - return exit_code - - -def _is_in_ci() -> bool: - return os.environ.get("SGLANG_IS_IN_CI", "").lower() in ("1", "true", "yes", "on") - - -def _maybe_pin_update_weights_model_pair(suite_files_rel: list[str]) -> None: - if not _is_in_ci(): - return - if _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE not in suite_files_rel: - return - if os.environ.get(_UPDATE_WEIGHTS_MODEL_PAIR_ENV): - print( - f"Using preset {_UPDATE_WEIGHTS_MODEL_PAIR_ENV}=" - f"{os.environ[_UPDATE_WEIGHTS_MODEL_PAIR_ENV]}" - ) - return - - selected_pair = random.choice(_UPDATE_WEIGHTS_MODEL_PAIR_IDS) - os.environ[_UPDATE_WEIGHTS_MODEL_PAIR_ENV] = selected_pair - print(f"Selected {_UPDATE_WEIGHTS_MODEL_PAIR_ENV}={selected_pair} for this CI run") - - -def _resolve_suite_files( - target_dir: Path, suite_files_rel: list[str], strict: bool -) -> list[str]: - suite_files_abs = [] - for f_rel in suite_files_rel: - f_abs = target_dir / f_rel - if not f_abs.exists(): - msg = f"Test file {f_rel} not found in {target_dir}." - if strict: - print(f"Error: {msg}") - sys.exit(1) - print(f"Warning: {msg} Skipping.") - continue - suite_files_abs.append(str(f_abs)) - return suite_files_abs - - -def _run_file_suite(args, target_dir: Path) -> int: - suite_files_rel = FILE_SUITES[args.suite] - _maybe_pin_update_weights_model_pair(suite_files_rel) - suite_files_abs = _resolve_suite_files( - target_dir, suite_files_rel, args.suite in STRICT_SUITES - ) - - if not suite_files_abs: - print(f"No valid test files found for suite '{args.suite}'.") - return 1 if args.suite in STRICT_SUITES else 0 - - exit_code, _, _ = run_pytest( - suite_files_abs, - filter_expr=args.filter, - junit_xml_path=None, - ) - return exit_code - - -def _get_dynamic_suite_cases(suite: str) -> list[DiffusionTestCase]: - cases = [] - for _, case_group in PARAMETRIZED_CASE_GROUPS[suite]: - cases.extend(case_group) - return cases - - -def _get_parametrized_files_for_case_ids( - suite: str, case_ids: set[str], target_dir: Path -) -> list[str]: - files = [] - for filename, case_group in PARAMETRIZED_CASE_GROUPS[suite]: - if any(case.id in case_ids for case in case_group): - file_path = target_dir / filename - if file_path.exists(): - files.append(str(file_path)) - else: - logger.warning("Test file %s not found in %s", filename, target_dir) - return files - - -def _run_dynamic_suite(args, target_dir: Path) -> int: - if args.partition_plan_json: - assignment = parse_partition_plan( - suite=args.suite, - partition_id=args.partition_id, - total_partitions=args.total_partitions, - plan_json=args.partition_plan_json, - ) - else: - assignment = build_local_partition_assignment( - suite=args.suite, - partition_id=args.partition_id, - total_partitions=args.total_partitions, - ) - return _run_partition_assignment(args, target_dir, assignment) - - -def _run_partition_assignment( - args, target_dir: Path, assignment: PartitionAssignment -) -> int: - rows = [[args.suite, f"{args.partition_id + 1}/{args.total_partitions}"]] - print(tabulate.tabulate(rows, headers=["Suite", "Partition"], tablefmt="psql")) - - total_est_time = 0.0 - executed_cases: list[str] = [] - case_results: dict[str, str] = {} - missing_standalone_estimates: list[str] = [] - standalone_measurements: list[dict] = [] - overall_exit_code = 0 - - if assignment.case_ids: - case_id_set = set(assignment.case_ids) - total_est_time += sum( - get_case_est_time(case_id) for case_id in assignment.case_ids - ) - suite_files = _get_parametrized_files_for_case_ids( - args.suite, case_id_set, target_dir - ) - if not suite_files: - print(f"No valid parametrized test files found for suite '{args.suite}'.") - return 0 - - partition_filter = " or ".join( - f"[{case_id}]" for case_id in assignment.case_ids - ) - filter_expr = ( - f"({partition_filter}) and ({args.filter})" - if args.filter - else partition_filter - ) - - print( - f"Running {len(assignment.case_ids)} parametrized cases with estimated total " - f"{sum(get_case_est_time(case_id) for case_id in assignment.case_ids):.1f}s:" - ) - for case_id in assignment.case_ids: - print(f" - case: {case_id} ({get_case_est_time(case_id):.1f}s)") - print(f"Test files: {[Path(f).name for f in suite_files]}") - print(f"Filter expression: {filter_expr}") - - junit_xml_path = str( - target_dir / f"junit_results_{args.suite}_{args.partition_id}.xml" - ) - exit_code, new_executed_cases, new_case_results = run_pytest( - suite_files, - filter_expr=filter_expr, - junit_xml_path=junit_xml_path, - ) - _merge_execution_results( - executed_cases, case_results, new_executed_cases, new_case_results - ) - # A failing case must not swallow this shard's standalone files: they - # are separate pytest runs, and they only share a shard because the - # shard count is fixed. --continue-on-error still decides whether a - # failing standalone file stops the ones queued behind it. - if exit_code != 0 and overall_exit_code == 0: - overall_exit_code = exit_code - - if assignment.standalone_files: - standalone_estimate = sum( - get_standalone_file_est_time(args.suite, standalone_file)[0] - for standalone_file in assignment.standalone_files - ) - total_est_time += standalone_estimate - print( - f"Running {len(assignment.standalone_files)} standalone file(s) with estimated total " - f"{standalone_estimate:.1f}s:" - ) - for standalone_file in assignment.standalone_files: - est_time, used_fallback_estimate = get_standalone_file_est_time( - args.suite, standalone_file - ) - fallback_suffix = ( - f", fallback estimate {DEFAULT_STANDALONE_EST_TIME_SECONDS:.1f}s" - if used_fallback_estimate - else "" - ) - print( - f" - standalone: {standalone_file} ({est_time:.1f}s{fallback_suffix})" - ) - - for standalone_file in assignment.standalone_files: - exit_code, new_executed_cases, new_case_results, measurement = ( - _run_standalone_file( - args.suite, - standalone_file, - target_dir, - extra_filter=args.filter, - ) - ) - if measurement["used_fallback_estimate"]: - missing_standalone_estimates.append(standalone_file) - standalone_measurements.append(measurement) - _merge_execution_results( - executed_cases, - case_results, - new_executed_cases, - new_case_results, - ) - if exit_code != 0 and overall_exit_code == 0: - overall_exit_code = exit_code - if exit_code != 0 and not args.continue_on_error: - break - - if not assignment.case_ids and not assignment.standalone_files: - print(f"No work assigned to partition {args.partition_id}. Exiting success.") - - print(f"Partition estimated total time: {total_est_time:.1f}s") - write_execution_report( - suite=args.suite, - partition_id=args.partition_id, - total_partitions=args.total_partitions, - executed_cases=executed_cases, - is_standalone=False, - standalone_file=None, - case_results=case_results, - missing_standalone_estimates=missing_standalone_estimates, - standalone_measurements=standalone_measurements, - ) - return overall_exit_code - - -def main(): - args = parse_args() - validate_standalone_file_est_times() - test_root_dir = Path(__file__).resolve().parent - target_dir = test_root_dir / args.base_dir - - if not target_dir.exists(): - print(f"Error: Target directory {target_dir} does not exist.") - sys.exit(1) - - if args.suite in COMPONENT_ACCURACY_SUITES: - suite_files_rel = FILE_SUITES[args.suite] - suite_files_abs = _resolve_suite_files( - target_dir, suite_files_rel, args.suite in STRICT_SUITES - ) - - if not suite_files_abs: - print(f"No valid test files found for suite '{args.suite}'.") - sys.exit(1 if args.suite in STRICT_SUITES else 0) - - my_files = partition_items_by_index( - suite_files_abs, args.partition_id, args.total_partitions - ) - partition_info = ( - f"{args.partition_id + 1}/{args.total_partitions} " - f"(0-based id={args.partition_id})" - ) - headers = ["Suite", "Partition"] - rows = [[args.suite, partition_info]] - msg = tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n" - msg += f"Enabled {len(my_files)} file(s):\n" - for file_path in my_files: - msg += f" - {file_path}\n" - print(msg, flush=True) - print( - f"Suite: {args.suite} | Partition: {args.partition_id}/{args.total_partitions}" - ) - print(f"Selected {len(suite_files_abs)} files:") - for f in suite_files_abs: - print(f" - {os.path.basename(f)}") - - if not my_files: - print("No files assigned to this partition. Exiting success.") - sys.exit(0) - - print(f"Running {len(my_files)} files in this shard: {', '.join(my_files)}") - - exit_code = run_component_accuracy_files( - my_files, - filter_expr=args.filter, - continue_on_error=args.continue_on_error, - ) - - msg = "\n" + tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n" - msg += f"Executed {len(my_files)} file(s):\n" - for file_path in my_files: - msg += f" - {file_path}\n" - print(msg, flush=True) - elif args.suite in PARAMETRIZED_CASE_GROUPS: - exit_code = _run_dynamic_suite(args, target_dir) - else: - exit_code = _run_file_suite(args, target_dir) - - sys.exit(exit_code) +"""Compatibility entry point; CI dispatches diffusion via ``test/run_suite.py``.""" +from sglang.multimodal_gen.test.runner.diffusion_suite_runner import main if __name__ == "__main__": main() diff --git a/python/sglang/multimodal_gen/test/runner/diffusion_suite_runner.py b/python/sglang/multimodal_gen/test/runner/diffusion_suite_runner.py new file mode 100644 index 000000000..634803629 --- /dev/null +++ b/python/sglang/multimodal_gen/test/runner/diffusion_suite_runner.py @@ -0,0 +1,690 @@ +"""Internal diffusion-suite adapter used by ``test/run_suite.py`` bridges. + +For diffusion 1-gpu/2-gpu suites, cases are partitioned by estimated runtime +using LPT so each CI shard has a similar total runtime. +""" + +import argparse +import copy +import json +import os +import random +import subprocess +import sys +import time +from dataclasses import dataclass +from pathlib import Path + +import tabulate + +from sglang.multimodal_gen.runtime.platforms import current_platform +from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger +from sglang.multimodal_gen.test.partitioning import PartitionItem, assign_partition +from sglang.multimodal_gen.test.runner.pytest_runner import ( + partition_items_by_index, + run_pytest, +) +from sglang.multimodal_gen.test.server.testcase_configs import ( + BASELINE_CONFIG, + DiffusionTestCase, +) + +# TODO: remove duplicated code +if current_platform.is_npu(): + from sglang.multimodal_gen.test.server.ascend.testcase_configs_npu import ( + _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE, + COMPONENT_ACCURACY_SUITES, + DEFAULT_EST_TIME_SECONDS, + DEFAULT_STANDALONE_EST_TIME_SECONDS, + FILE_SUITES, + PARAMETRIZED_CASE_GROUPS, + STANDALONE_FILE_EST_TIMES, + STANDALONE_FILES, + STARTUP_OVERHEAD_SECONDS, + SUITES, + ) +else: + from sglang.multimodal_gen.test.server.gpu_cases import ( # noqa: F401 It is used by ci scripts + _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE, + _UPDATE_WEIGHTS_MODEL_PAIR_ENV, + _UPDATE_WEIGHTS_MODEL_PAIR_IDS, + COMPONENT_ACCURACY_FILE_NUM_GPUS, + COMPONENT_ACCURACY_SUITES, + DEFAULT_EST_TIME_SECONDS, + DEFAULT_STANDALONE_EST_TIME_SECONDS, + FILE_SUITES, + ONE_GPU_CASES, + PARAMETRIZED_CASE_GROUPS, + STANDALONE_FILE_EST_TIMES, + STANDALONE_FILES, + STARTUP_OVERHEAD_SECONDS, + STRICT_SUITES, + SUITES, + TWO_GPU_CASES, + ) + + +logger = init_logger(__name__) + + +@dataclass(frozen=True) +class PartitionAssignment: + case_ids: list[str] + standalone_files: list[str] + estimated_time: float | None = None + missing_standalone_estimates: list[str] | None = None + + +def get_case_est_time(case_id: str) -> float: + scenario = BASELINE_CONFIG.scenarios.get(case_id) + if scenario is None: + return DEFAULT_EST_TIME_SECONDS + if scenario.estimated_full_test_time_s is not None: + return scenario.estimated_full_test_time_s + return scenario.expected_e2e_ms / 1000.0 + STARTUP_OVERHEAD_SECONDS + + +def get_standalone_file_est_time( + suite: str, standalone_file: str +) -> tuple[float, bool]: + suite_est_times = STANDALONE_FILE_EST_TIMES.get(suite, {}) + if standalone_file not in suite_est_times: + return DEFAULT_STANDALONE_EST_TIME_SECONDS, True + return suite_est_times[standalone_file], False + + +def get_all_standalone_file_est_times() -> dict[str, dict[str, float]]: + return copy.deepcopy(STANDALONE_FILE_EST_TIMES) + + +def validate_standalone_file_est_times() -> dict[str, list[str]]: + missing_by_suite: dict[str, list[str]] = {} + for suite, standalone_files in STANDALONE_FILES.items(): + suite_est_times = STANDALONE_FILE_EST_TIMES.get(suite, {}) + missing = [ + standalone_file + for standalone_file in standalone_files + if standalone_file not in suite_est_times + ] + if missing: + missing_by_suite[suite] = missing + return missing_by_suite + + +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: + return f"standalone:{standalone_file}" + + +def parse_partition_plan( + suite: str, + partition_id: int, + total_partitions: int, + plan_json: str, +) -> PartitionAssignment: + plan = json.loads(plan_json) + if plan.get("suite") != suite: + raise ValueError( + f"Partition plan suite mismatch: expected {suite!r}, " + f"got {plan.get('suite')!r}" + ) + + partition_count = plan.get("partition_count") + if partition_count != total_partitions: + raise ValueError( + f"Partition count mismatch for suite {suite!r}: " + f"plan={partition_count}, matrix={total_partitions}" + ) + + partitions = plan.get("partitions", []) + selected_partition = None + for partition in partitions: + if partition.get("part") == partition_id: + selected_partition = partition + break + + if selected_partition is None: + raise ValueError( + f"Partition {partition_id} not found in plan for suite {suite!r}" + ) + + return PartitionAssignment( + case_ids=list(selected_partition.get("case_ids", [])), + standalone_files=list(selected_partition.get("standalone_files", [])), + estimated_time=selected_partition.get("estimated_time"), + missing_standalone_estimates=list( + selected_partition.get("missing_standalone_estimates", []) + ), + ) + + +def build_local_partition_assignment( + suite: str, + partition_id: int, + total_partitions: int, +) -> PartitionAssignment: + """Assign this shard's work when CI did not precompute a partition plan. + + Lanes with a hardcoded ``--total-partitions`` (the AMD ones) cannot give + every standalone file a shard of its own, so standalone files are LPT + balanced together with the parametrized cases instead. + """ + items = [ + PartitionItem(kind="case", item_id=case.id, est_time=get_case_est_time(case.id)) + for case in _get_dynamic_suite_cases(suite) + ] + for standalone_file in STANDALONE_FILES.get(suite, []): + items.append( + PartitionItem( + kind="standalone", + item_id=standalone_file, + est_time=get_standalone_file_est_time(suite, standalone_file)[0], + ) + ) + + my_items = assign_partition(items, partition_id, total_partitions) + return PartitionAssignment( + case_ids=[item.item_id for item in my_items if item.kind == "case"], + standalone_files=[ + item.item_id for item in my_items if item.kind == "standalone" + ], + ) + + +def _merge_execution_results( + executed_cases: list[str], + case_results: dict[str, str], + new_executed_cases: list[str], + new_case_results: dict[str, str], +) -> None: + executed_cases.extend( + case_id for case_id in new_executed_cases if case_id not in executed_cases + ) + case_results.update(new_case_results) + + +def _format_standalone_estimate_snippet( + suite: str, standalone_file: str, measured_full_test_time_s: float +) -> str: + return ( + f'"{suite}": {{\n "{standalone_file}": {measured_full_test_time_s:.1f},\n}}' + ) + + +def _print_missing_standalone_estimate_message( + suite: str, + standalone_file: str, + measured_full_test_time_s: float, +) -> None: + snippet = _format_standalone_estimate_snippet( + suite, standalone_file, measured_full_test_time_s + ) + logger.error( + f"\n{'=' * 60}\n" + f'Add standalone estimate for suite "{suite}" and file "{standalone_file}":\n\n' + "File: python/sglang/multimodal_gen/test/server/gpu_cases.py\n\n" + f"Current partition used fallback estimate: " + f"{DEFAULT_STANDALONE_EST_TIME_SECONDS:.1f}s\n\n" + f"{snippet}\n" + f"{'=' * 60}\n" + ) + + +def _run_standalone_file( + suite: str, + standalone_rel: str, + target_dir: Path, + extra_filter: str | None = None, +) -> tuple[int, list[str], dict[str, str], dict]: + if standalone_rel == _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE: + _maybe_pin_update_weights_model_pair([standalone_rel]) + + est_time, used_fallback_estimate = get_standalone_file_est_time( + suite, standalone_rel + ) + standalone_file = _resolve_suite_files(target_dir, [standalone_rel], strict=True)[0] + junit_xml_path = str( + target_dir / f"junit_results_{suite}_{Path(standalone_rel).stem}.xml" + ) + start_time = time.perf_counter() + exit_code, _, _ = run_pytest( + [standalone_file], + filter_expr=extra_filter, + junit_xml_path=junit_xml_path, + ) + measured_full_test_time_s = round(time.perf_counter() - start_time, 1) + standalone_key = _normalize_standalone_key(standalone_rel) + measurement = { + "suite": suite, + "standalone_file": standalone_rel, + "measured_full_test_time_s": measured_full_test_time_s, + "used_fallback_estimate": used_fallback_estimate, + "fallback_estimate_s": DEFAULT_STANDALONE_EST_TIME_SECONDS, + "had_configured_estimate": not used_fallback_estimate, + "configured_or_fallback_estimate_s": est_time, + } + if used_fallback_estimate: + _print_missing_standalone_estimate_message( + suite, standalone_rel, measured_full_test_time_s + ) + return ( + exit_code, + [standalone_key], + {standalone_key: "pass" if exit_code == 0 else "fail"}, + measurement, + ) + + +def parse_args(): + suite_choices = sorted(set(FILE_SUITES) | set(PARAMETRIZED_CASE_GROUPS)) + parser = argparse.ArgumentParser(description="Run multimodal_gen test suite") + parser.add_argument( + "--suite", + type=str, + required=True, + choices=suite_choices, + help="The test suite to run.", + ) + parser.add_argument( + "--partition-id", + type=int, + default=0, + help="Index of the current partition (for parallel execution)", + ) + parser.add_argument( + "--total-partitions", + type=int, + default=1, + help="Total number of partitions", + ) + parser.add_argument( + "--base-dir", + type=str, + default="server", + help="Base directory for tests relative to this script's parent", + ) + parser.add_argument( + "-k", + "--filter", + type=str, + default=None, + help="Pytest filter expression (passed to pytest -k)", + ) + parser.add_argument( + "--continue-on-error", + action="store_true", + default=False, + help="Continue running remaining tests even if one fails.", + ) + parser.add_argument( + "--partition-plan-json", + type=str, + default=None, + help="Full partition plan JSON for the current suite.", + ) + return parser.parse_args() + + +def write_execution_report( + suite: str, + partition_id: int, + total_partitions: int, + executed_cases: list[str], + is_standalone: bool = False, + standalone_file: str | None = None, + case_results: dict[str, str] | None = None, + missing_standalone_estimates: list[str] | None = None, + standalone_measurements: list[dict] | None = None, +) -> str: + report = { + "suite": suite, + "partition_id": partition_id, + "total_partitions": total_partitions, + "is_standalone": is_standalone, + "standalone_file": standalone_file, + "executed_cases": executed_cases, + "case_results": case_results or {}, + "missing_standalone_estimates": missing_standalone_estimates or [], + "standalone_measurements": standalone_measurements or [], + } + + report_filename = f"execution_report_{suite}_{partition_id}.json" + report_path = Path(__file__).resolve().parents[1] / report_filename + with open(report_path, "w", encoding="utf-8") as f: + json.dump(report, f, indent=2) + + logger.info("Execution report written to: %s", report_path) + return str(report_path) + + +def run_component_accuracy_files(files, filter_expr=None, continue_on_error=False): + exit_code = 0 + for file_path in files: + file_name = Path(file_path).name + num_gpus = COMPONENT_ACCURACY_FILE_NUM_GPUS.get(file_name, 1) + if num_gpus > 1: + cmd = [ + sys.executable, + "-m", + "torch.distributed.run", + f"--nproc_per_node={num_gpus}", + "-m", + "pytest", + "-s", + "-v", + ] + else: + cmd = [sys.executable, "-m", "pytest", "-s", "-v"] + + if filter_expr: + cmd.extend(["-k", filter_expr]) + cmd.append(file_path) + + print(f"Running command: {' '.join(cmd)}") + file_exit_code = subprocess.call(cmd) + if file_exit_code == 5: + print( + "No tests collected (exit code 5). This is expected when filters " + "deselect all tests in a file. Treating as success." + ) + file_exit_code = 0 + if file_exit_code != 0 and exit_code == 0: + exit_code = file_exit_code + if file_exit_code != 0 and not continue_on_error: + return file_exit_code + return exit_code + + +def _is_in_ci() -> bool: + return os.environ.get("SGLANG_IS_IN_CI", "").lower() in ("1", "true", "yes", "on") + + +def _maybe_pin_update_weights_model_pair(suite_files_rel: list[str]) -> None: + if not _is_in_ci(): + return + if _UPDATE_WEIGHTS_FROM_DISK_TEST_FILE not in suite_files_rel: + return + if os.environ.get(_UPDATE_WEIGHTS_MODEL_PAIR_ENV): + print( + f"Using preset {_UPDATE_WEIGHTS_MODEL_PAIR_ENV}=" + f"{os.environ[_UPDATE_WEIGHTS_MODEL_PAIR_ENV]}" + ) + return + + selected_pair = random.choice(_UPDATE_WEIGHTS_MODEL_PAIR_IDS) + os.environ[_UPDATE_WEIGHTS_MODEL_PAIR_ENV] = selected_pair + print(f"Selected {_UPDATE_WEIGHTS_MODEL_PAIR_ENV}={selected_pair} for this CI run") + + +def _resolve_suite_files( + target_dir: Path, suite_files_rel: list[str], strict: bool +) -> list[str]: + suite_files_abs = [] + for f_rel in suite_files_rel: + f_abs = target_dir / f_rel + if not f_abs.exists(): + msg = f"Test file {f_rel} not found in {target_dir}." + if strict: + print(f"Error: {msg}") + sys.exit(1) + print(f"Warning: {msg} Skipping.") + continue + suite_files_abs.append(str(f_abs)) + return suite_files_abs + + +def _run_file_suite(args, target_dir: Path) -> int: + suite_files_rel = FILE_SUITES[args.suite] + _maybe_pin_update_weights_model_pair(suite_files_rel) + suite_files_abs = _resolve_suite_files( + target_dir, suite_files_rel, args.suite in STRICT_SUITES + ) + + if not suite_files_abs: + print(f"No valid test files found for suite '{args.suite}'.") + return 1 if args.suite in STRICT_SUITES else 0 + + exit_code, _, _ = run_pytest( + suite_files_abs, + filter_expr=args.filter, + junit_xml_path=None, + ) + return exit_code + + +def _get_dynamic_suite_cases(suite: str) -> list[DiffusionTestCase]: + cases = [] + for _, case_group in PARAMETRIZED_CASE_GROUPS[suite]: + cases.extend(case_group) + return cases + + +def _get_parametrized_files_for_case_ids( + suite: str, case_ids: set[str], target_dir: Path +) -> list[str]: + files = [] + for filename, case_group in PARAMETRIZED_CASE_GROUPS[suite]: + if any(case.id in case_ids for case in case_group): + file_path = target_dir / filename + if file_path.exists(): + files.append(str(file_path)) + else: + logger.warning("Test file %s not found in %s", filename, target_dir) + return files + + +def _run_dynamic_suite(args, target_dir: Path) -> int: + if args.partition_plan_json: + assignment = parse_partition_plan( + suite=args.suite, + partition_id=args.partition_id, + total_partitions=args.total_partitions, + plan_json=args.partition_plan_json, + ) + else: + assignment = build_local_partition_assignment( + suite=args.suite, + partition_id=args.partition_id, + total_partitions=args.total_partitions, + ) + return _run_partition_assignment(args, target_dir, assignment) + + +def _run_partition_assignment( + args, target_dir: Path, assignment: PartitionAssignment +) -> int: + rows = [[args.suite, f"{args.partition_id + 1}/{args.total_partitions}"]] + print(tabulate.tabulate(rows, headers=["Suite", "Partition"], tablefmt="psql")) + + total_est_time = 0.0 + executed_cases: list[str] = [] + case_results: dict[str, str] = {} + missing_standalone_estimates: list[str] = [] + standalone_measurements: list[dict] = [] + overall_exit_code = 0 + + if assignment.case_ids: + case_id_set = set(assignment.case_ids) + total_est_time += sum( + get_case_est_time(case_id) for case_id in assignment.case_ids + ) + suite_files = _get_parametrized_files_for_case_ids( + args.suite, case_id_set, target_dir + ) + if not suite_files: + print(f"No valid parametrized test files found for suite '{args.suite}'.") + return 0 + + partition_filter = " or ".join( + f"[{case_id}]" for case_id in assignment.case_ids + ) + filter_expr = ( + f"({partition_filter}) and ({args.filter})" + if args.filter + else partition_filter + ) + + print( + f"Running {len(assignment.case_ids)} parametrized cases with estimated total " + f"{sum(get_case_est_time(case_id) for case_id in assignment.case_ids):.1f}s:" + ) + for case_id in assignment.case_ids: + print(f" - case: {case_id} ({get_case_est_time(case_id):.1f}s)") + print(f"Test files: {[Path(f).name for f in suite_files]}") + print(f"Filter expression: {filter_expr}") + + junit_xml_path = str( + target_dir / f"junit_results_{args.suite}_{args.partition_id}.xml" + ) + exit_code, new_executed_cases, new_case_results = run_pytest( + suite_files, + filter_expr=filter_expr, + junit_xml_path=junit_xml_path, + ) + _merge_execution_results( + executed_cases, case_results, new_executed_cases, new_case_results + ) + # A failing case must not swallow this shard's standalone files: they + # are separate pytest runs, and they only share a shard because the + # shard count is fixed. --continue-on-error still decides whether a + # failing standalone file stops the ones queued behind it. + if exit_code != 0 and overall_exit_code == 0: + overall_exit_code = exit_code + + if assignment.standalone_files: + standalone_estimate = sum( + get_standalone_file_est_time(args.suite, standalone_file)[0] + for standalone_file in assignment.standalone_files + ) + total_est_time += standalone_estimate + print( + f"Running {len(assignment.standalone_files)} standalone file(s) with estimated total " + f"{standalone_estimate:.1f}s:" + ) + for standalone_file in assignment.standalone_files: + est_time, used_fallback_estimate = get_standalone_file_est_time( + args.suite, standalone_file + ) + fallback_suffix = ( + f", fallback estimate {DEFAULT_STANDALONE_EST_TIME_SECONDS:.1f}s" + if used_fallback_estimate + else "" + ) + print( + f" - standalone: {standalone_file} ({est_time:.1f}s{fallback_suffix})" + ) + + for standalone_file in assignment.standalone_files: + exit_code, new_executed_cases, new_case_results, measurement = ( + _run_standalone_file( + args.suite, + standalone_file, + target_dir, + extra_filter=args.filter, + ) + ) + if measurement["used_fallback_estimate"]: + missing_standalone_estimates.append(standalone_file) + standalone_measurements.append(measurement) + _merge_execution_results( + executed_cases, + case_results, + new_executed_cases, + new_case_results, + ) + if exit_code != 0 and overall_exit_code == 0: + overall_exit_code = exit_code + if exit_code != 0 and not args.continue_on_error: + break + + if not assignment.case_ids and not assignment.standalone_files: + print(f"No work assigned to partition {args.partition_id}. Exiting success.") + + print(f"Partition estimated total time: {total_est_time:.1f}s") + write_execution_report( + suite=args.suite, + partition_id=args.partition_id, + total_partitions=args.total_partitions, + executed_cases=executed_cases, + is_standalone=False, + standalone_file=None, + case_results=case_results, + missing_standalone_estimates=missing_standalone_estimates, + standalone_measurements=standalone_measurements, + ) + return overall_exit_code + + +def main(): + args = parse_args() + validate_standalone_file_est_times() + test_root_dir = Path(__file__).resolve().parents[1] + target_dir = test_root_dir / args.base_dir + + if not target_dir.exists(): + print(f"Error: Target directory {target_dir} does not exist.") + sys.exit(1) + + if args.suite in COMPONENT_ACCURACY_SUITES: + suite_files_rel = FILE_SUITES[args.suite] + suite_files_abs = _resolve_suite_files( + target_dir, suite_files_rel, args.suite in STRICT_SUITES + ) + + if not suite_files_abs: + print(f"No valid test files found for suite '{args.suite}'.") + sys.exit(1 if args.suite in STRICT_SUITES else 0) + + my_files = partition_items_by_index( + suite_files_abs, args.partition_id, args.total_partitions + ) + partition_info = ( + f"{args.partition_id + 1}/{args.total_partitions} " + f"(0-based id={args.partition_id})" + ) + headers = ["Suite", "Partition"] + rows = [[args.suite, partition_info]] + msg = tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n" + msg += f"Enabled {len(my_files)} file(s):\n" + for file_path in my_files: + msg += f" - {file_path}\n" + print(msg, flush=True) + print( + f"Suite: {args.suite} | Partition: {args.partition_id}/{args.total_partitions}" + ) + print(f"Selected {len(suite_files_abs)} files:") + for f in suite_files_abs: + print(f" - {os.path.basename(f)}") + + if not my_files: + print("No files assigned to this partition. Exiting success.") + sys.exit(0) + + print(f"Running {len(my_files)} files in this shard: {', '.join(my_files)}") + + exit_code = run_component_accuracy_files( + my_files, + filter_expr=args.filter, + continue_on_error=args.continue_on_error, + ) + + msg = "\n" + tabulate.tabulate(rows, headers=headers, tablefmt="psql") + "\n" + msg += f"Executed {len(my_files)} file(s):\n" + for file_path in my_files: + msg += f" - {file_path}\n" + print(msg, flush=True) + elif args.suite in PARAMETRIZED_CASE_GROUPS: + exit_code = _run_dynamic_suite(args, target_dir) + else: + exit_code = _run_file_suite(args, target_dir) + + sys.exit(exit_code) + + +if __name__ == "__main__": + main() diff --git a/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py b/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py index ac7ffc8ae..a8f27cea2 100755 --- a/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py +++ b/python/sglang/multimodal_gen/test/scripts/gen_diffusion_ci_outputs.py @@ -2,7 +2,7 @@ """ Generate diffusion CI outputs for consistency testing. -This script reuses the CI test code by calling run_suite.py with SGLANG_GEN_GT=1, +This script reuses the CI test adapter with SGLANG_GEN_GT=1, ensuring that GT generation uses exactly the same code path as CI tests. Usage: @@ -21,7 +21,7 @@ from sglang.multimodal_gen.test.partitioning import ( PartitionItem, partition_items_by_lpt, ) -from sglang.multimodal_gen.test.run_suite import ( +from sglang.multimodal_gen.test.runner.diffusion_suite_runner import ( SUITES, _maybe_pin_update_weights_model_pair, get_case_est_time, diff --git a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_webui.py b/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_webui.py deleted file mode 100644 index 711683385..000000000 --- a/python/sglang/multimodal_gen/test/unit/realtime/test_realtime_webui.py +++ /dev/null @@ -1,112 +0,0 @@ -# SPDX-License-Identifier: Apache-2.0 - -from pathlib import Path - - -def test_realtime_webui_presets_do_not_emit_camera_scripts(): - repo_root = Path(__file__).resolve().parents[6] - app_js = ( - repo_root / "python/sglang/multimodal_gen/apps/realtime_webui/app.js" - ).read_text() - index_html = ( - repo_root / "python/sglang/multimodal_gen/apps/realtime_webui/index.html" - ).read_text() - styles_css = ( - repo_root / "python/sglang/multimodal_gen/apps/realtime_webui/styles.css" - ).read_text() - playback_js = ( - repo_root - / "python/sglang/multimodal_gen/apps/realtime_webui/playback_controller.js" - ).read_text() - - assert "preset.actions" not in app_js - assert "repeatActions" not in app_js - assert 'id="eventFrames"' not in index_html - assert "ControlStateController" in app_js - assert 'const DEFAULT_PREVIEW_OUTPUT_FORMAT = "webp";' in app_js - assert 'id="transportFormat"' in index_html - assert 'id="fps" type="number" value="25"' in index_html - assert 'id="superResolution" type="checkbox"' in index_html - assert 'id="upscalingScale"' in index_html - assert 'class="workspace"' in index_html - assert 'class="preview-frame"' in index_html - assert 'id="previewOverlay" class="preview-overlay"' in index_html - assert 'id="previewScale" type="range" min="80" max="170" value="120"' in index_html - assert 'id="previewScaleText"' in index_html - assert 'id="outputSizeText"' in index_html - assert 'id="frameInterpolation" type="checkbox" />' in index_html - assert ( - 'id="serverUrl" value="ws://127.0.0.1:30000/v1/realtime_video/generate"' - in index_html - ) - assert '' in index_html - assert 'id="serverSendText"' in index_html - assert 'id="theoreticalFpsText"' in index_html - assert 'id="renderFps"' in index_html - assert 'id="stageRenderFps"' not in index_html - assert "sglang-diffusion Realtime Studio" in index_html - assert "SGLD" not in index_html - assert 'class="tabs"' not in index_html - assert "Recordings" not in index_html - assert "API" not in index_html - assert "Info" not in index_html - assert 'id="steps" type="number" value="4"' in index_html - assert 'id="guidance" type="number" value="1"' in index_html - assert "styles.css?v=realtime-record-v49" in index_html - assert "app.js?v=realtime-record-v75" in index_html - assert ( - 'const DECODER_WORKER_URL = "./decoder_worker.js?v=rgb-worker-v10";' in app_js - ) - assert "const DEFAULT_TARGET_FPS = 25;" in app_js - assert "const DEFAULT_FRAME_INTERPOLATION_EXP = 1;" in app_js - assert "const DEFAULT_FRAME_INTERPOLATION_SCALE = 1.0;" in app_js - assert "const DEFAULT_UPSCALING_SCALE = 2;" in app_js - assert "const DEFAULT_PREVIEW_SCALE = 120;" in app_js - assert 'setPreviewState("waiting")' in app_js - assert "stage.dataset.previewState = state" in app_js - assert "previewProgressSpin" in styles_css - assert "previewDotPulse" not in styles_css - assert 'document.querySelector(".preview-frame")' in app_js - assert 'previewFrame.style.setProperty("--preview-scale"' in app_js - assert "cancelAnimationFrame(previewScaleFrame)" in app_js - assert "enable_frame_interpolation: true" in app_js - assert "frame_interpolation_exp: DEFAULT_FRAME_INTERPOLATION_EXP" in app_js - assert "frame_interpolation_scale: DEFAULT_FRAME_INTERPOLATION_SCALE" in app_js - assert "readSuperResolutionParams()" in app_js - assert "enable_upscaling: true" in app_js - assert "upscaling_scale: readUpscalingScale()" in app_js - assert "updateOutputSizeFromHeader(header)" in app_js - assert "setPreviewScale(DEFAULT_PREVIEW_SCALE)" in app_js - assert "preview_scale" in app_js - assert "sr_scale" in app_js - assert "elapsedMs < targetMs" in playback_js - assert "queuedDecodeFrames > maxQueuedFrames" in app_js - assert ( - 'const REACTOR_PRESET_BASE_URL = "https://www.reactor.inc/lingbot-world-fast-v1";' - in app_js - ) - assert "Dragon Dolly" in app_js - assert "no creature morphing" in app_js - assert "the Plastic Beach island stays centered" in app_js - assert "no camera descent, no push-in, no orbit" in app_js - assert "Ziggy Stardust" in app_js - assert "blue K. West sign" in app_js - assert "wet pavement reflecting a yellow streetlamp" in app_js - assert "ZiggyStardust.jpg" in app_js - assert "A slow aerial orbit around a pastel floating island hotel" not in app_js - assert app_js.index("Dragon Ride") < app_js.index("Dragon Dolly") - assert app_js.index("Ziggy Stardust") < app_js.index("Plastic Beach") - assert app_js.index("Dragon Dolly") < app_js.index("Kid A") - assert "dragon-ride.jpg" in app_js - assert "stageRenderFps" not in app_js - assert 'setStatus("Receiving", "live")' in app_js - assert "decodeQueue.push(" in app_js - assert "receiveChain" not in app_js - assert 'message.type === "chunk_stats"' in app_js - assert "chunkTotal > 0 ? numFrames / chunkTotal" in app_js - assert ".stage-stat" in styles_css - assert ".workspace" in styles_css - assert ".preview-frame" in styles_css - assert ".preview-overlay" in styles_css - assert ".preview-scale-control" in styles_css - assert "--preview-scale" in styles_css diff --git a/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py b/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py index 1383b38db..572f8aa78 100644 --- a/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py +++ b/python/sglang/multimodal_gen/test/unit/test_consistency_metrics.py @@ -47,18 +47,6 @@ def test_consistency_gt_urls_are_pinned_to_ci_data_revision(): assert pinned_revision_path in test_utils.SGL_TEST_FILES_SGLANG_CONSISTENCY_GT_BASE -def test_remote_file_exists_returns_false_for_definitive_404(monkeypatch): - class Response: - status_code = 404 - - def close(self): - pass - - monkeypatch.setattr(test_utils.requests, "head", lambda *args, **kwargs: Response()) - - assert test_utils._remote_file_exists("https://example.com/missing.png") is False - - def test_remote_video_gt_candidates_survive_inconclusive_probe(monkeypatch): monkeypatch.setenv(test_utils.CONSISTENCY_PLATFORM_ENV, "h100") monkeypatch.setattr(test_utils, "_remote_file_exists", lambda url: None) diff --git a/python/sglang/multimodal_gen/test/unit/test_cosmos3.py b/python/sglang/multimodal_gen/test/unit/test_cosmos3.py index 288365c12..de8ca2b17 100644 --- a/python/sglang/multimodal_gen/test/unit/test_cosmos3.py +++ b/python/sglang/multimodal_gen/test/unit/test_cosmos3.py @@ -1907,18 +1907,6 @@ class TestCosmos3ModalitySamplingParams(unittest.TestCase): self.assertEqual(sp.condition_frame_indexes, [0, 1]) self.assertEqual(sp.condition_video_keep, "first") - def test_action_fields_default_none(self): - sp = Cosmos3SamplingParams(prompt="t") - for field in ( - "action_mode", - "domain_id", - "domain_name", - "raw_action_dim", - "action_fps", - "action", - ): - self.assertIsNone(getattr(sp, field)) - class TestCosmos3CaptionMetadata(unittest.TestCase): """Structured captions get generation metadata; prose prompts opt out.""" diff --git a/python/sglang/multimodal_gen/test/unit/test_diffusion_benchmark_skill.py b/python/sglang/multimodal_gen/test/unit/test_diffusion_benchmark_skill.py deleted file mode 100644 index 1b4c1bd50..000000000 --- a/python/sglang/multimodal_gen/test/unit/test_diffusion_benchmark_skill.py +++ /dev/null @@ -1,658 +0,0 @@ -import importlib.util -import json -import sys -import tempfile -import types -import unittest -from pathlib import Path -from unittest.mock import patch - - -def _load_benchmark_module(temp_root: Path): - multimodal_gen_root = Path(__file__).resolve().parents[2] - script_path = ( - multimodal_gen_root - / ".claude" - / "skills" - / "sglang-diffusion-benchmark-profile" - / "scripts" - / "bench_diffusion_denoise.py" - ) - fake_env = types.ModuleType("diffusion_skill_env") - fake_env.ensure_dir = lambda path: ( - Path(path).mkdir(parents=True, exist_ok=True) or Path(path) - ) - fake_env.get_assets_dir = lambda _root: temp_root / "assets" - fake_env.get_output_dir = lambda _kind, _root: temp_root / "outputs" - fake_env.get_repo_root = lambda: temp_root / "repo" - fake_env.pick_idle_gpus = lambda count: list(range(count)) - - spec = importlib.util.spec_from_file_location( - "test_bench_diffusion_denoise", script_path - ) - assert spec is not None and spec.loader is not None - module = importlib.util.module_from_spec(spec) - with patch.dict(sys.modules, {"diffusion_skill_env": fake_env}): - spec.loader.exec_module(module) - return module - - -def _load_skill_env_module(): - script_path = ( - Path(__file__).resolve().parents[2] - / ".claude" - / "skills" - / "sglang-diffusion-benchmark-profile" - / "scripts" - / "diffusion_skill_env.py" - ) - spec = importlib.util.spec_from_file_location( - "test_diffusion_skill_env", script_path - ) - assert spec is not None and spec.loader is not None - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module - - -class TestDiffusionBenchmarkSkill(unittest.TestCase): - def test_skill_env_prefers_own_worktree_over_installed_package(self): - module = _load_skill_env_module() - installed = types.ModuleType("sglang") - installed.__file__ = "/sgl-workspace/sglang/python/sglang/__init__.py" - - with patch.dict(sys.modules, {"sglang": installed}): - repo_root = module.get_repo_root() - - self.assertEqual(repo_root, Path(__file__).resolve().parents[5]) - - def test_nightly_presets_remain_aligned(self): - with tempfile.TemporaryDirectory() as tmpdir: - module = _load_benchmark_module(Path(tmpdir)) - repo_root = Path(__file__).resolve().parents[5] - module.NIGHTLY_CONFIG_PATH = ( - repo_root - / "scripts" - / "ci" - / "utils" - / "diffusion" - / "comparison_configs.json" - ) - - self.assertEqual(module.validate_nightly_alignment(), 0) - - def test_recent_model_presets_are_eager_by_default(self): - with tempfile.TemporaryDirectory() as tmpdir: - module = _load_benchmark_module(Path(tmpdir)) - - expected = { - "longcat-image", - "longcat-image-edit", - "longcat-image-edit-turbo", - "qwen-edit-base", - "qwen-image-layered", - "stable-diffusion-3.5-medium", - "sana-video", - "sana-wm-bidirectional", - "sana-wm-streaming", - "lingbot-video-moe", - "lingbot-world", - "lingbot-world-v2", - "fastwan21-t2v-1.3b", - "fasth3-t2va-vsa", - "wan22-t2v-nvfp4", - "krea2-turbo", - "krea2-raw", - "ideogram4-fast", - "ideogram4-instant", - "longlive2-t2v", - "longlive2-i2v", - "fast-hunyuan", - "turbowan21-t2v-1.3b", - "helios-mid", - "helios-distilled", - "joy-echo", - "cosmos3-edge-t2i", - "cosmos3-super-t2v-cfg2tp2", - "cosmos3-super-i2v", - "cosmos3-super-t2i-distilled", - "ltx25", - "ltx25-diffusion-decoder", - } - self.assertTrue(expected.issubset(module.MODELS)) - - eager_cmd = module.build_sglang_cmd("longcat-image") - self.assertNotIn("--enable-torch-compile", eager_cmd) - self.assertIn("--enable-prompt-rewrite=false", eager_cmd) - self.assertIn("--quality=lossless", eager_cmd) - - compiled_cmd = module.build_sglang_cmd("longcat-image", torch_compile=True) - self.assertIn("--enable-torch-compile", compiled_cmd) - - longcat_edit_cmd = module.build_sglang_cmd("longcat-image-edit") - self.assertIn( - "--model-path=meituan-longcat/LongCat-Image-Edit", - longcat_edit_cmd, - ) - self.assertTrue( - any(arg.startswith("--image-path=") for arg in longcat_edit_cmd) - ) - self.assertIn("--enable-prompt-rewrite=false", longcat_edit_cmd) - - longcat_edit_bcg_cmd = module.build_sglang_cmd( - "longcat-image-edit", breakable_cuda_graph=True - ) - resolution_index = longcat_edit_bcg_cmd.index("--warmup-resolutions") - self.assertEqual(longcat_edit_bcg_cmd[resolution_index + 1], "1264x848") - - longcat_edit_turbo_cmd = module.build_sglang_cmd("longcat-image-edit-turbo") - self.assertIn( - "--model-path=meituan-longcat/LongCat-Image-Edit-Turbo", - longcat_edit_turbo_cmd, - ) - - layered_cmd = module.build_sglang_cmd("qwen-image-layered") - self.assertIn("--model-path=Qwen/Qwen-Image-Layered", layered_cmd) - self.assertIn("--num-frames=4", layered_cmd) - - sd35_cmd = module.build_sglang_cmd("stable-diffusion-3.5-medium") - self.assertIn( - "--model-path=stabilityai/stable-diffusion-3.5-medium-diffusers", - sd35_cmd, - ) - self.assertIn("stable-diffusion-3.5-medium", module.GATED_MODELS) - - h3_cmd = module.build_sglang_cmd("minimax-h3-t2va", torch_compile=True) - self.assertNotIn("--enable-torch-compile", h3_cmd) - - fastwan_cmd = module.build_sglang_cmd("fastwan21-t2v-1.3b") - self.assertIn("--num-frames=61", fastwan_cmd) - self.assertIn("--num-inference-steps=3", fastwan_cmd) - self.assertIn("--dit-layerwise-offload=false", fastwan_cmd) - - wan_nvfp4_cmd = module.build_sglang_cmd("wan22-t2v-nvfp4") - self.assertIn( - "--model-path=nvidia/Wan2.2-T2V-A14B-Diffusers-NVFP4", - wan_nvfp4_cmd, - ) - self.assertIn("--num-frames=81", wan_nvfp4_cmd) - self.assertIn("--dit-layerwise-offload=false", wan_nvfp4_cmd) - self.assertEqual(module.required_gpus_for_model("wan22-t2v-nvfp4"), 1) - - krea_raw_cmd = module.build_sglang_cmd("krea2-raw") - self.assertIn("--num-inference-steps=50", krea_raw_cmd) - self.assertIn("--guidance-scale=4.5", krea_raw_cmd) - - cosmos_i2v_cmd = module.build_sglang_cmd("cosmos3-super-i2v") - self.assertIn( - "--model-path=nvidia/Cosmos3-Super-Image2Video", cosmos_i2v_cmd - ) - self.assertIn("--num-gpus=2", cosmos_i2v_cmd) - self.assertIn("--tp-size=2", cosmos_i2v_cmd) - self.assertIn("--num-frames=81", cosmos_i2v_cmd) - - cosmos_cfg_cmd = module.build_sglang_cmd("cosmos3-super-t2v-cfg2tp2") - self.assertIn("--model-path=nvidia/Cosmos3-Super", cosmos_cfg_cmd) - self.assertIn("--num-gpus=4", cosmos_cfg_cmd) - self.assertIn("--tp-size=2", cosmos_cfg_cmd) - - sana_wm_dense_cmd = module.build_sglang_cmd("sana-wm-bidirectional") - self.assertIn( - "--model-path=Efficient-Large-Model/SANA-WM_bidirectional", - sana_wm_dense_cmd, - ) - self.assertIn("--num-inference-steps=20", sana_wm_dense_cmd) - self.assertNotIn("--streaming", sana_wm_dense_cmd) - - sana_wm_streaming_cmd = module.build_sglang_cmd("sana-wm-streaming") - self.assertIn("--streaming", sana_wm_streaming_cmd) - self.assertIn("--refiner-chunked", sana_wm_streaming_cmd) - self.assertIn("--action=w-16,wl-16,l-16", sana_wm_streaming_cmd) - - lingbot_world_cmd = module.build_sglang_cmd("lingbot-world") - self.assertIn( - "--model-path=robbyant/lingbot-world-fast-diffusers", - lingbot_world_cmd, - ) - self.assertIn("--num-frames=9", lingbot_world_cmd) - self.assertIn("--warmup-mode=off", lingbot_world_cmd) - self.assertIn("--config=", " ".join(lingbot_world_cmd)) - self.assertEqual( - module.MODELS["lingbot-world"]["config_overrides"]["actions"], - [["w"] for _ in range(9)], - ) - - lingbot_world_v2_cmd = module.build_sglang_cmd("lingbot-world-v2") - self.assertIn( - "--model-path=robbyant/lingbot-world-v2-14b-causal-fast-diffusers", - lingbot_world_v2_cmd, - ) - self.assertIn("--num-frames=9", lingbot_world_v2_cmd) - self.assertIn("--num-inference-steps=4", lingbot_world_v2_cmd) - - ideogram_cmd = module.build_sglang_cmd("ideogram4-instant") - self.assertFalse( - any(arg.startswith("--num-inference-steps") for arg in ideogram_cmd) - ) - - longlive_i2v_cmd = module.build_sglang_cmd("longlive2-i2v") - self.assertIn("--num-frames=61", longlive_i2v_cmd) - self.assertTrue( - any(arg.startswith("--image-path=") for arg in longlive_i2v_cmd) - ) - - joy_echo_cmd = module.build_sglang_cmd("joy-echo") - self.assertIn("--num-gpus=2", joy_echo_cmd) - self.assertIn("--ulysses-degree=2", joy_echo_cmd) - config_arg = next( - arg for arg in joy_echo_cmd if arg.startswith("--config=") - ) - config = json.loads(Path(config_arg.removeprefix("--config=")).read_text()) - self.assertFalse(config["enable_memory_bank"]) - - def test_quality_and_bcg_comparators_are_explicit_and_exclusive(self): - with tempfile.TemporaryDirectory() as tmpdir: - module = _load_benchmark_module(Path(tmpdir)) - - high_cmd = module.build_sglang_cmd("longcat-image", quality="high") - self.assertIn("--quality=high", high_cmd) - self.assertNotIn("--enable-breakable-cuda-graph", high_cmd) - - extra_high_cmd = module.build_sglang_cmd( - "longcat-image", quality="extra-high" - ) - self.assertIn("--quality=extra-high", extra_high_cmd) - self.assertNotIn("--enable-breakable-cuda-graph", extra_high_cmd) - - bcg_cmd = module.build_sglang_cmd( - "longcat-image", - breakable_cuda_graph=True, - bcg_text_buckets=[256, 512], - ) - self.assertIn("--enable-breakable-cuda-graph", bcg_cmd) - self.assertEqual( - bcg_cmd[bcg_cmd.index("--warmup-resolutions") + 1], "1024x1024" - ) - bucket_index = bcg_cmd.index("--bcg-text-buckets") - self.assertEqual( - bcg_cmd[bucket_index + 1 : bucket_index + 3], ["256", "512"] - ) - - sana_video_bcg_cmd = module.build_sglang_cmd( - "sana-video", breakable_cuda_graph=True - ) - self.assertEqual( - sana_video_bcg_cmd[sana_video_bcg_cmd.index("--warmup-num-frames") + 1], - "17", - ) - - for _, quality, breakable_cuda_graph in module.QUALITY_BCG_ABBA_MATRIX: - module.build_sglang_cmd( - "longcat-image", - quality=quality, - breakable_cuda_graph=breakable_cuda_graph, - bcg_text_buckets=[256, 512] if breakable_cuda_graph else None, - ) - - with self.assertRaisesRegex(ValueError, "comparators"): - module.build_sglang_cmd( - "longcat-image", - torch_compile=True, - breakable_cuda_graph=True, - ) - with self.assertRaisesRegex(ValueError, "requires"): - module.build_sglang_cmd("longcat-image", bcg_text_buckets=[256]) - - def test_isolated_cache_cleanup_writes_zero_residual_ledger(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - cache_root = temp_root / "model-caches" - cache_dir = module._prepare_model_cache( - cache_root, "longcat-image", "baseline" - ) - - weight_path = cache_dir / "huggingface" / "hub" / "model.safetensors" - weight_path.parent.mkdir(parents=True) - weight_path.write_bytes(b"weights") - env = module._model_cache_env(cache_dir) - self.assertTrue(env["HF_HOME"].startswith(str(cache_dir))) - self.assertTrue(env["HF_XET_CACHE"].startswith(str(cache_dir))) - self.assertTrue(env["TRANSFORMERS_CACHE"].startswith(str(cache_dir))) - self.assertTrue(env["MODELSCOPE_CACHE"].startswith(str(cache_dir))) - - ledger_path = temp_root / "artifacts" / "cleanup.jsonl" - record = module._cleanup_model_cache( - cache_root, - cache_dir, - ledger_path, - "longcat-image", - "baseline", - "success", - ) - - self.assertFalse(cache_dir.exists()) - self.assertEqual(record["before"]["weight_file_count"], 1) - self.assertEqual(record["after"]["file_count"], 0) - ledger = json.loads(ledger_path.read_text(encoding="utf-8")) - self.assertEqual(ledger["exit_reason"], "success") - self.assertEqual(ledger["after"]["weight_file_count"], 0) - - def test_isolated_cache_refuses_to_reuse_existing_run_directory(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - cache_root = temp_root / "model-caches" - module._prepare_model_cache(cache_root, "sana-video", "baseline") - - with self.assertRaises(FileExistsError): - module._prepare_model_cache(cache_root, "sana-video", "baseline") - - def test_isolated_cache_seeds_read_only_hf_cache_with_writable_overlay(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - seed_root = temp_root / "shared-hf" - source_model = seed_root / "hub" / "models--org--model" - source_weight = source_model / "snapshots" / "abc" / "model.safetensors" - source_weight.parent.mkdir(parents=True) - source_weight.write_bytes(b"shared weights") - source_ref = source_model / "refs" / "main" - source_ref.parent.mkdir() - source_ref.write_text("abc") - - cache_root = temp_root / "model-caches" - cache_dir = module._prepare_model_cache( - cache_root, - "sana-video", - "baseline", - seed_model_cache_roots=[seed_root], - ) - seeded_model = cache_dir / "huggingface" / "hub" / "models--org--model" - self.assertTrue(seeded_model.is_dir()) - self.assertFalse(seeded_model.is_symlink()) - seeded_weight = seeded_model / "snapshots" / "abc" / "model.safetensors" - self.assertTrue(seeded_weight.is_symlink()) - self.assertEqual( - seeded_weight.read_bytes(), - b"shared weights", - ) - - new_blob = seeded_model / "blobs" / "downloaded" - new_blob.parent.mkdir(exist_ok=True) - new_blob.write_bytes(b"new download") - (seeded_model / "refs" / "main").write_text("new-revision") - self.assertEqual(new_blob.read_bytes(), b"new download") - self.assertEqual(source_ref.read_text(), "abc") - - module._cleanup_model_cache( - cache_root, - cache_dir, - temp_root / "cleanup.jsonl", - "sana-video", - "baseline", - "success", - ) - self.assertFalse(cache_dir.exists()) - self.assertEqual(source_weight.read_bytes(), b"shared weights") - - def test_interrupted_run_cleans_isolated_cache_in_finally(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - cache_root = temp_root / "model-caches" - output_dir = temp_root / "outputs" - output_dir.mkdir() - - with ( - patch.object( - module, "_run_benchmark_once_impl", side_effect=KeyboardInterrupt - ), - self.assertRaises(KeyboardInterrupt), - ): - module.run_benchmark_once( - "sana-video", - "baseline", - output_dir, - model_cache_root=cache_root, - cleanup_model_cache=True, - ) - - self.assertFalse((cache_root / "sana-video-baseline").exists()) - ledger = json.loads( - (output_dir / "cleanup.jsonl").read_text(encoding="utf-8") - ) - self.assertEqual(ledger["exit_reason"], "interrupted") - self.assertEqual(ledger["after"]["weight_file_count"], 0) - - def test_failed_run_is_recorded_as_error_and_cleaned(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - cache_root = temp_root / "model-caches" - output_dir = temp_root / "outputs" - output_dir.mkdir() - - with ( - patch.object( - module, "_run_benchmark_once_impl", side_effect=RuntimeError("boom") - ), - self.assertRaisesRegex(RuntimeError, "boom"), - ): - module.run_benchmark_once( - "sana-video", - "baseline", - output_dir, - model_cache_root=cache_root, - cleanup_model_cache=True, - ) - - self.assertFalse((cache_root / "sana-video-baseline").exists()) - ledger = json.loads( - (output_dir / "cleanup.jsonl").read_text(encoding="utf-8") - ) - self.assertEqual(ledger["exit_reason"], "error") - - def test_zero_exit_without_artifacts_is_invalid(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - output_dir = temp_root / "outputs" - output_dir.mkdir() - - with patch.object(module.subprocess, "Popen") as popen: - popen.return_value.stdout = iter(()) - popen.return_value.wait.return_value = 0 - result = module._run_benchmark_once_impl( - "sana-video", - "missing-artifacts", - output_dir, - warmup=False, - cuda_visible_devices="0", - ) - - command = popen.call_args.args[0] - self.assertIn("--output-path", command) - self.assertIn("--output-file-name", command) - self.assertTrue(result["error"]) - self.assertEqual( - result["missing_artifacts"], ["perf dump", "generated output"] - ) - - def test_mesh_artifacts_are_accepted_and_hashed(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - output_dir = temp_root / "outputs" - output_dir.mkdir() - - def finish_run(): - (output_dir / "hunyuan3d-shape_mesh-output.json").write_text( - json.dumps({"total_duration_ms": 1000, "steps": []}), - encoding="utf-8", - ) - (output_dir / "hunyuan3d-shape-mesh-output.obj").write_bytes( - b"v 0 0 0\n" - ) - return 0 - - with patch.object(module.subprocess, "Popen") as popen: - popen.return_value.stdout = iter(()) - popen.return_value.wait.side_effect = finish_run - result = module._run_benchmark_once_impl( - "hunyuan3d-shape", - "mesh-output", - output_dir, - warmup=False, - cuda_visible_devices="0", - ) - - self.assertFalse(result["error"]) - self.assertEqual( - result["output_artifacts"], - [str(output_dir / "hunyuan3d-shape-mesh-output.obj")], - ) - self.assertEqual(len(result["output_sha256"]), 1) - - def test_high_bcg_rejects_quality_fusion_mounted_after_capture(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - output_dir = temp_root / "outputs" - output_dir.mkdir() - - with patch.object(module.subprocess, "Popen") as popen: - popen.return_value.stdout = iter( - ( - "[Diffusion BCG] captured 3 segment(s)\n", - "Mounted LTX-2 fused RMSNorm+modulate for quality=high\n", - ) - ) - popen.return_value.wait.return_value = 0 - result = module._run_benchmark_once_impl( - "longcat-image", - "bcg-high", - output_dir, - warmup=False, - quality="high", - breakable_cuda_graph=True, - cuda_visible_devices="0", - ) - - self.assertTrue(result["error"]) - self.assertEqual( - result["bcg_invalid_signals"], - [module.BCG_LATE_QUALITY_FUSION_SIGNAL], - ) - - def test_extra_high_bcg_rejects_quality_fusion_mounted_after_capture(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - output_dir = temp_root / "outputs" - output_dir.mkdir() - - with patch.object(module.subprocess, "Popen") as popen: - popen.return_value.stdout = iter( - ( - "[Diffusion BCG] captured 3 segment(s)\n", - "Mounted Qwen fused added-QKV for quality=extra-high\n", - ) - ) - popen.return_value.wait.return_value = 0 - result = module._run_benchmark_once_impl( - "longcat-image", - "bcg-extra-high", - output_dir, - warmup=False, - quality="extra-high", - breakable_cuda_graph=True, - cuda_visible_devices="0", - ) - - self.assertTrue(result["error"]) - self.assertEqual( - result["bcg_invalid_signals"], - [module.BCG_LATE_QUALITY_FUSION_SIGNAL], - ) - - def test_quality_bcg_matrix_reuses_one_gpu_set_and_cleans_once(self): - with tempfile.TemporaryDirectory() as tmpdir: - temp_root = Path(tmpdir) - module = _load_benchmark_module(temp_root) - cache_root = temp_root / "model-caches" - output_dir = temp_root / "outputs" - output_dir.mkdir() - calls = [] - - def fake_run(model_key, label, _output_dir, **kwargs): - calls.append((model_key, label, kwargs)) - cache_dir = kwargs["model_cache_dir"] - weight_path = cache_dir / "hub" / "model.safetensors" - weight_path.parent.mkdir(parents=True, exist_ok=True) - weight_path.write_bytes(b"weights") - return {"model": model_key, "label": label, "error": False} - - with patch.object(module, "_run_benchmark_once_impl", side_effect=fake_run): - results = module.run_quality_bcg_matrix( - "sana-video", - "h200", - output_dir, - model_cache_root=cache_root, - cleanup_model_cache=True, - ) - - self.assertEqual(len(results), 12) - self.assertEqual( - [ - (call[2]["quality"], call[2]["breakable_cuda_graph"]) - for call in calls - ], - [ - (quality, breakable_cuda_graph) - for _, quality, breakable_cuda_graph in module.QUALITY_BCG_ABBA_MATRIX - ], - ) - self.assertEqual({call[2]["cuda_visible_devices"] for call in calls}, {"0"}) - self.assertEqual( - {call[2]["model_cache_dir"] for call in calls}, - {calls[0][2]["model_cache_dir"]}, - ) - self.assertFalse(calls[0][2]["model_cache_dir"].exists()) - ledger = json.loads( - (output_dir / "cleanup.jsonl").read_text(encoding="utf-8") - ) - self.assertEqual(ledger["exit_reason"], "success") - self.assertEqual(ledger["before"]["weight_file_count"], 1) - self.assertEqual(ledger["after"]["weight_file_count"], 0) - - def test_quality_bcg_matrix_rejects_output_hash_mismatch(self): - with tempfile.TemporaryDirectory() as tmpdir: - module = _load_benchmark_module(Path(tmpdir)) - results = [ - { - "quality": "lossless", - "breakable_cuda_graph": False, - "output_sha256": ["eager"], - "error": False, - }, - { - "quality": "lossless", - "breakable_cuda_graph": True, - "output_sha256": ["bcg"], - "error": False, - }, - ] - - module._validate_quality_bcg_output_hashes(results) - - self.assertFalse(results[0]["error"]) - self.assertTrue(results[1]["error"]) - self.assertEqual( - results[1]["output_hash_error"], - "BCG lossless output hash differs from eager", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py index 1237b0720..d7d5cc94a 100644 --- a/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py +++ b/python/sglang/multimodal_gen/test/unit/test_disagg_roles.py @@ -617,13 +617,6 @@ class TestStageAffinityAndValidation(_GlobalStageArgsMixin, unittest.TestCase): pipeline.create_pipeline_stages(pipeline.server_args) self.assertEqual(list(pipeline._stage_name_mapping.keys()), stage_names) - def test_hunyuan3d_shape_stage_no_longer_stores_model_dtype(self): - pipeline = self._make_hunyuan_pipeline(RoleType.ENCODER, paint_enable=False) - pipeline.create_pipeline_stages(pipeline.server_args) - stage = pipeline._stage_name_mapping["shape_before_denoising"] - self.assertIsInstance(stage, Hunyuan3DShapeBeforeDenoisingStage) - self.assertFalse(hasattr(stage, "model_dtype")) - def test_ltx2_refinement_stage_keeps_class_name_stage_key(self): stage = object.__new__(LTX2RefinementStage) self.assertEqual( diff --git a/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py b/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py index 864b62509..6033d7203 100644 --- a/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py +++ b/python/sglang/multimodal_gen/test/unit/test_nvtx_pytorch_hooks.py @@ -29,17 +29,6 @@ from sglang.multimodal_gen.runtime.utils.nvtx_pytorch_hooks import ( class TestMaybeNvtxRange(unittest.TestCase): - def test_disabled_returns_noop_context_manager(self) -> None: - ran = False - with maybe_nvtx_range("never", enabled=False): - ran = True - self.assertTrue(ran) - - def test_disabled_propagates_exception(self) -> None: - with self.assertRaises(RuntimeError): - with maybe_nvtx_range("never", enabled=False): - raise RuntimeError("boom") - def test_disabled_does_not_call_nvtx(self) -> None: with ( patch.object(nvtx_pytorch_hooks.nvtx, "range_push") as push, diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py b/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py index 13b6f444a..cc58a9876 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen2_5vl_generation.py @@ -15,8 +15,6 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl import ( Qwen2_5_VLAttention, Qwen2_5_VLForConditionalGeneration, _apply_repetition_penalty, - _make_column_linear, - _make_row_linear, _select_next_token, ) from sglang.multimodal_gen.runtime.models.encoders.qwen2_5vl_vision import ( @@ -29,14 +27,7 @@ from sglang.multimodal_gen.runtime.pipelines.longcat_image import LongCatImagePi from sglang.srt.layers.linear import ( ColumnParallelLinear, ReplicatedLinear, - RowParallelLinear, ) -from sglang.srt.models.qwen2_5_vl import ( - Qwen2_5_VisionPatchEmbed, - Qwen2_5_VisionPatchMerger, - Qwen2_5_VLMLP, -) -from sglang.srt.runtime_context import get_parallel class _StubQwen2_5VL(Qwen2_5_VLForConditionalGeneration): @@ -79,49 +70,6 @@ class _AttentionRecorder(nn.Module): return query -def test_native_vision_reuses_srt_modules(): - config = SimpleNamespace( - hidden_size=16, - intermediate_size=24, - hidden_act="silu", - num_heads=2, - depth=0, - patch_size=2, - temporal_patch_size=1, - in_channels=3, - spatial_merge_size=2, - out_hidden_size=12, - fullatt_block_indexes=[], - window_size=8, - ) - with get_parallel().override(tp_size=1, tp_rank=0): - model = Qwen2_5VLVisionTransformer(config) - mlp = Qwen2_5_VLMLP( - 16, - 24, - fuse_gate_up=False, - ) - fused_mlp = Qwen2_5_VLMLP(16, 24) - - assert isinstance(model.patch_embed, Qwen2_5_VisionPatchEmbed) - assert isinstance(model.merger, Qwen2_5_VisionPatchMerger) - assert not mlp.fuse_gate_up - assert isinstance(mlp.gate_proj, ColumnParallelLinear) - assert isinstance(mlp.up_proj, ColumnParallelLinear) - assert mlp.gate_proj.tp_size == mlp.up_proj.tp_size == 1 - assert isinstance(mlp.down_proj, ReplicatedLinear) - assert isinstance(fused_mlp.down_proj, RowParallelLinear) - assert mlp.act is not None - assert isinstance( - _make_column_linear(16, 24, bias=False, use_tensor_parallel=False), - ReplicatedLinear, - ) - assert isinstance( - _make_row_linear(24, 16, bias=False, use_tensor_parallel=False), - ReplicatedLinear, - ) - - def test_text_mlp_uses_single_rank_when_intermediate_size_is_not_tp_divisible( monkeypatch, ): diff --git a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py index e71794a48..aea564332 100644 --- a/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py +++ b/python/sglang/multimodal_gen/test/unit/test_qwen3vl_vision.py @@ -12,7 +12,6 @@ from sglang.multimodal_gen.runtime.models.encoders.minimax_h3_qwen3vl import ( ) from sglang.multimodal_gen.runtime.models.encoders.qwen3vl import ( Qwen3VLForConditionalGeneration, - _make_text_rms_norm, ) from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import ( Qwen3VLVisionRotaryEmbedding, @@ -20,7 +19,6 @@ from sglang.multimodal_gen.runtime.models.encoders.qwen3vl_vision import ( _vision_cu_seqlens, _vision_position_ids, ) -from sglang.srt.layers.layernorm import RMSNorm from sglang.srt.models.qwen3_vl import ( Qwen3VLMoeVisionPatchMerger, Qwen3VLVisionPatchEmbed, @@ -48,13 +46,6 @@ def test_native_vision_layout_matches_qwen3_merge_order(): assert cu_seqlens.tolist() == [0, 24, 32, 40] -def test_qwen3vl_text_reuses_srt_rms_norm(): - norm = _make_text_rms_norm(16, 1e-6) - - assert isinstance(norm, RMSNorm) - assert norm.cast_x_before_out_mul - - def test_native_vision_keeps_checkpoint_parameter_names(): config = SimpleNamespace( hidden_size=16, diff --git a/python/sglang/multimodal_gen/test/unit/test_server_args.py b/python/sglang/multimodal_gen/test/unit/test_server_args.py index 5bd23f436..00106d08d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_server_args.py +++ b/python/sglang/multimodal_gen/test/unit/test_server_args.py @@ -770,11 +770,6 @@ class TestWarmupModeNormalization(unittest.TestCase): sa = self._resolve(warmup_mode="server") self.assertEqual(sa.warmup_mode, "server") - def test_defaulted_mode_applies_without_legacy_flags(self): - # Bare `sglang serve` defaults to server-based warmup. - sa = self._resolve(warmup_mode="server") - self.assertEqual(sa.warmup_mode, "server") - def test_resolutions_force_warmup_on(self): sa = self._resolve( warmup_mode="off", diff --git a/python/sglang/multimodal_gen/test/unit/test_server_warmup_progress.py b/python/sglang/multimodal_gen/test/unit/test_server_warmup_progress.py deleted file mode 100644 index 19c1e2a81..000000000 --- a/python/sglang/multimodal_gen/test/unit/test_server_warmup_progress.py +++ /dev/null @@ -1,36 +0,0 @@ -"""Unit tests for server warmup progress reporting.""" - -import unittest -from unittest.mock import MagicMock, patch - -from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch -from sglang.multimodal_gen.runtime.server_warmup import SchedulerWarmupMixin - - -class TestServerWarmupProgress(unittest.TestCase): - def test_ci_progress_uses_scheduler_counter_when_tqdm_is_disabled(self): - scheduler = SchedulerWarmupMixin() - scheduler._show_warmup_progress = True - scheduler._warmup_total = 1 - scheduler._warmup_processed = 1 - progress_bar = MagicMock(total=1, n=0) - scheduler._warmup_progress_bar = progress_bar - - with ( - patch( - "sglang.multimodal_gen.runtime.server_warmup._is_ci_log_env", - return_value=True, - ), - patch("sglang.multimodal_gen.runtime.server_warmup.logger") as logger, - ): - scheduler._advance_warmup_progress_bar(object(), OutputBatch()) - - logger.info.assert_called_once_with( - "Warmup requests: %s/%s %s", 1, 1, "warmup req" - ) - progress_bar.close.assert_called_once_with() - self.assertIsNone(scheduler._warmup_progress_bar) - - -if __name__ == "__main__": - unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_single_rank_device_group.py b/python/sglang/multimodal_gen/test/unit/test_single_rank_device_group.py index f10da5283..4b03d827d 100644 --- a/python/sglang/multimodal_gen/test/unit/test_single_rank_device_group.py +++ b/python/sglang/multimodal_gen/test/unit/test_single_rank_device_group.py @@ -25,11 +25,6 @@ class TestSingleRankDeviceGroup(unittest.TestCase): new_device_group(ranks, requested) new_group.assert_called_once_with(ranks, backend=requested) - def test_backend_defaults_to_none_for_multi_rank(self): - with patch(NEW_GROUP_PATH) as new_group: - new_device_group([0, 1]) - new_group.assert_called_once_with([0, 1], backend=None) - if __name__ == "__main__": unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py b/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py index b3910dfa4..549527c68 100644 --- a/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py +++ b/python/sglang/multimodal_gen/test/unit/test_sol_attn_backend.py @@ -1,46 +1,14 @@ -import importlib.util import unittest from unittest.mock import MagicMock, patch -import torch - from sglang.multimodal_gen.runtime.layers.attention.backends.sol_attn import ( - SolAttnBackend, SolAttnImpl, _get_sol_attn_runtime_config, _parse_layer_ranges, ) -from sglang.multimodal_gen.runtime.platforms.cuda import CudaPlatformBase -from sglang.multimodal_gen.runtime.platforms.interface import AttentionBackendEnum - - -class FakeCudaPlatform(CudaPlatformBase): - is_sm120_device = False - is_blackwell_device = False - supports_flash_attention = True - - @classmethod - def is_sm120(cls): - return cls.is_sm120_device - - @classmethod - def is_blackwell(cls): - return cls.is_blackwell_device - - @classmethod - def has_device_capability( - cls, - capability: tuple[int, int] | int, - device_id: int = 0, - ) -> bool: - return cls.supports_flash_attention class TestSolAttnBackend(unittest.TestCase): - def test_enum_name(self): - self.assertEqual(str(AttentionBackendEnum.SOL_ATTN), "sol_attn") - self.assertTrue(AttentionBackendEnum.SOL_ATTN.is_sparse) - def test_parse_layer_ranges(self): self.assertEqual(_parse_layer_ranges("0,1,3-5"), frozenset({0, 1, 3, 4, 5})) @@ -70,9 +38,6 @@ class TestSolAttnBackend(unittest.TestCase): ): _get_sol_attn_runtime_config() - def test_backend_head_size(self): - self.assertEqual(SolAttnBackend.get_supported_head_sizes(), [128]) - def test_dense_guard_uses_early_steps(self): impl = SolAttnImpl( num_heads=8, @@ -127,19 +92,6 @@ class TestSolAttnBackend(unittest.TestCase): ): self.assertFalse(impl._should_use_dense()) - def test_cuda_resolver(self): - if importlib.util.find_spec("sol_attn") is None: - self.skipTest("sol_attn package is not available") - cls_str = FakeCudaPlatform.get_attn_backend_cls_str( - selected_backend=AttentionBackendEnum.SOL_ATTN, - head_size=128, - dtype=torch.bfloat16, - ) - self.assertTrue(cls_str.endswith("SolAttnBackend")) - - def test_supports_packed_varlen(self): - self.assertTrue(SolAttnBackend.supports_packed_varlen()) - if __name__ == "__main__": unittest.main() diff --git a/python/sglang/multimodal_gen/test/unit/test_sp_shard.py b/python/sglang/multimodal_gen/test/unit/test_sp_shard.py index 6edcae21d..50e16b367 100644 --- a/python/sglang/multimodal_gen/test/unit/test_sp_shard.py +++ b/python/sglang/multimodal_gen/test/unit/test_sp_shard.py @@ -156,12 +156,6 @@ def test_strategy_sp1_replicates(monkeypatch): assert sps.plan_text_strategy(100) == "replicate" -def test_strategy_shard_when_legal(monkeypatch): - _fake_sp(monkeypatch, 2) - assert sps.plan_text_strategy(15) == "shard" - assert sps.plan_text_strategy(16) == "shard" - - def test_strategy_replicates_when_padding_spans_multiple_shards(monkeypatch): _fake_sp(monkeypatch, 8) assert sps.plan_text_strategy(1) == "replicate" diff --git a/python/sglang/multimodal_gen/test/unit/test_srt_clip_reuse.py b/python/sglang/multimodal_gen/test/unit/test_srt_clip_reuse.py index 638e192c7..a850f00b2 100644 --- a/python/sglang/multimodal_gen/test/unit/test_srt_clip_reuse.py +++ b/python/sglang/multimodal_gen/test/unit/test_srt_clip_reuse.py @@ -34,12 +34,6 @@ class _FakeProjection(nn.Module): return hidden_states, None -def test_mmgen_clip_reuses_srt_components(): - assert mmgen_clip.CLIPEncoder is srt_clip.CLIPEncoder - assert mmgen_clip.CLIPTextEmbeddings is srt_clip.CLIPTextEmbeddings - assert mmgen_clip.CLIPVisionEmbeddings is srt_clip.CLIPVisionEmbeddings - - def test_clip_encoder_propagates_causal_semantics(): with ( patch.object(srt_clip, "CLIPAttention", return_value=nn.Identity()) as attn, diff --git a/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py b/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py index b6b11bade..9dddc90d4 100644 --- a/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py +++ b/python/sglang/multimodal_gen/test/unit/test_subblock_sparse_attention.py @@ -24,7 +24,6 @@ from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse.rou _snap_up_to_8, ) from sglang.multimodal_gen.runtime.layers.attention.backends.subblock_sparse_attn import ( - SubBlockSparseAttentionBackend, SubBlockSparseAttentionImpl, SubBlockSparseSchedule, _dit_layer_index, @@ -187,17 +186,6 @@ class TestBudgetGranularity(unittest.TestCase): class TestSubBlockSparseBackend(unittest.TestCase): - def test_the_advertised_builder_can_be_built(self): - """`AttentionMetadataBuilder.__init__` is abstract; a builder that does - not override it makes `get_builder_cls()()` a TypeError.""" - builder = SubBlockSparseAttentionBackend.get_builder_cls()() - builder.prepare() - metadata = builder.build(current_timestep=7) - self.assertIsInstance( - metadata, SubBlockSparseAttentionBackend.get_metadata_cls() - ) - self.assertEqual(metadata.current_timestep, 7) - def test_sm90_adapter_uses_presorted_indices_and_64x64_blocks(self): captured = {} diff --git a/python/sglang/multimodal_gen/test/unit/test_suite_partitioning.py b/python/sglang/multimodal_gen/test/unit/test_suite_partitioning.py index 1d2647a1a..6c4214177 100644 --- a/python/sglang/multimodal_gen/test/unit/test_suite_partitioning.py +++ b/python/sglang/multimodal_gen/test/unit/test_suite_partitioning.py @@ -11,9 +11,9 @@ from types import SimpleNamespace import pytest -from sglang.multimodal_gen.test import run_suite from sglang.multimodal_gen.test.partitioning import PartitionItem, assign_partition -from sglang.multimodal_gen.test.run_suite import ( +from sglang.multimodal_gen.test.runner import diffusion_suite_runner as run_suite +from sglang.multimodal_gen.test.runner.diffusion_suite_runner import ( PartitionAssignment, build_local_partition_assignment, ) diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py index 163ac4752..f385875b7 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_loader.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_loader.py @@ -71,9 +71,6 @@ class _FakeServerArgs: def should_start_component_on_cpu(self, _component_name): return False - def should_use_fsdp_for_component(self, _component_name): - return False - def should_configure_layerwise_offload_for_lazy_component(self, component_name): return component_name in self.layerwise_components @@ -615,12 +612,6 @@ class TestVAELoader(unittest.TestCase): native_load.assert_not_called() - def test_pipeline_config_declares_an_empty_native_only_default(self): - loader = vae_loader.VAELoader() - server_args = _FakeServerArgs(QwenImagePipelineConfig()) - - self.assertFalse(loader.should_raise_customized_load_error(server_args, "vae")) - def test_backfill_ltx2_audio_vae_latent_stats_maps_official_keys(self): loaded = { "per_channel_statistics.mean-of-means": torch.tensor([1.0, 2.0]), diff --git a/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py b/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py index cffd746e8..0df66e936 100644 --- a/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py +++ b/python/sglang/multimodal_gen/test/unit/test_vae_spatial_parallel_decode.py @@ -77,12 +77,6 @@ class _DispatchProbeVAE(ParallelTiledVAE): class TestVAESpatialParallelDecode(unittest.TestCase): - def test_base_vae_config_defaults_to_auto_parallel_decode(self): - config = VAEConfig() - - self.assertTrue(config.use_parallel_decode) - self.assertEqual(config.parallel_decode_mode, "auto") - def test_image_video_vae_configs_default_to_auto_parallel_decode(self): configs = ( ErnieImageVAEConfig(), diff --git a/python/sglang/multimodal_gen/test/unit/test_wan_attention_backend.py b/python/sglang/multimodal_gen/test/unit/test_wan_attention_backend.py deleted file mode 100644 index d1125813c..000000000 --- a/python/sglang/multimodal_gen/test/unit/test_wan_attention_backend.py +++ /dev/null @@ -1,36 +0,0 @@ -import unittest -from unittest.mock import patch - -from torch import nn - -from sglang.multimodal_gen.runtime.models.dits.wanvideo import WanSelfAttention -from sglang.multimodal_gen.runtime.platforms import AttentionBackendEnum - -_WAN = "sglang.multimodal_gen.runtime.models.dits.wanvideo" - - -class TestWanAttentionBackendRole(unittest.TestCase): - def test_cross_attention_role_is_forwarded_to_usp(self): - with ( - patch(f"{_WAN}.ColumnParallelLinear", return_value=nn.Identity()), - patch(f"{_WAN}.RowParallelLinear", return_value=nn.Identity()), - patch(f"{_WAN}.get_tp_world_size", return_value=1), - patch(f"{_WAN}.USPAttention") as usp_attention, - ): - WanSelfAttention( - dim=128, - num_heads=1, - qk_norm=False, - is_cross_attention=True, - supported_attention_backends={ - AttentionBackendEnum.FA, - AttentionBackendEnum.TORCH_SDPA, - }, - ) - - self.assertTrue(usp_attention.call_args.kwargs["is_cross_attention"]) - self.assertTrue(usp_attention.call_args.kwargs["skip_sequence_parallel"]) - - -if __name__ == "__main__": - unittest.main() diff --git a/python/sglang/srt/arg_groups/model_overrides/__init__.py b/python/sglang/srt/arg_groups/model_overrides/__init__.py index b9431a2aa..88dcb05e6 100644 --- a/python/sglang/srt/arg_groups/model_overrides/__init__.py +++ b/python/sglang/srt/arg_groups/model_overrides/__init__.py @@ -5,8 +5,7 @@ Importing this package is what registers them. An architecture may be claimed by more than one module here -- one supplies its attention shape, another its MoE runner -- but two of them must never declare the *same* field for it: nobody would own that value, and which module supplied it would come down to -the order of the imports below. ``test_model_override_split.py`` forbids the -overlap, which is why this list needs no particular order. +the order of the imports below. Keep each field owned by one family module. """ from sglang.srt.arg_groups.model_overrides import cohere2_moe # noqa: F401 diff --git a/python/sglang/srt/arg_groups/model_overrides/minicpm.py b/python/sglang/srt/arg_groups/model_overrides/minicpm.py index 76fb4fef5..6cd91ef3f 100644 --- a/python/sglang/srt/arg_groups/model_overrides/minicpm.py +++ b/python/sglang/srt/arg_groups/model_overrides/minicpm.py @@ -33,8 +33,8 @@ def _minicpm_sala_overrides(server_args: Any, hf_config: Any) -> dict: "minicpm_flashattn": ("fa4" if get_platform().is_blackwell else "fa3"), "minicpm_flashinfer": "flashinfer", } - # Literal keys keep the written-field set statically derivable; a loop - # variable hides it from the census in test_chain_read_ratchet.py. + # Keep the three backend decisions explicit so each resolved field is + # easy to review independently. dense_attention = dense_backends.get(cfg.attention_backend) if dense_attention is not None: overrides["attention_backend"] = dense_attention diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 652bfead8..124805eb2 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -797,8 +797,7 @@ def m3_fp8_attn_gemm_enabled(args) -> bool: # NOTE: The process-wide ServerArgs is owned by the runtime context # (sglang.srt.runtime_context). The two functions below are LEGACY shims kept # for the existing call-sites; they publish/read the same live object by -# reference. Do not add new call-sites — the counts are ratcheted -# (decrease-only) by test/registered/unit/test_legacy_global_ratchet.py. +# reference. Do not add new call-sites. # Imports are in-function so the two modules stay cycle-free at import time. @functools.lru_cache(maxsize=1) def _underscore_field_names() -> frozenset: diff --git a/python/sglang/test/ci/ci_register.py b/python/sglang/test/ci/ci_register.py index 57bee1290..227fd5701 100644 --- a/python/sglang/test/ci/ci_register.py +++ b/python/sglang/test/ci/ci_register.py @@ -98,19 +98,6 @@ def register_amd_ci( return None -def register_musa_ci( - est_time: float, - suite: Optional[str] = None, - nightly: bool = False, - disabled: Optional[str] = None, - *, - stage: Optional[str] = None, - runner_config: Optional[str] = None, -): - """Marker for MUSA CI registration (parsed via AST; runtime no-op).""" - return None - - def register_npu_ci( est_time: float, suite: Optional[str] = None, diff --git a/python/sglang/test/ci/diffusion_suite_bridge.py b/python/sglang/test/ci/diffusion_suite_bridge.py new file mode 100644 index 000000000..1c6cb8062 --- /dev/null +++ b/python/sglang/test/ci/diffusion_suite_bridge.py @@ -0,0 +1,35 @@ +"""Bridge registered diffusion suites to their case-aware pytest adapter.""" + +import os +import sys +from pathlib import Path + + +def _enabled(name: str) -> bool: + return os.environ.get(name, "").lower() in {"1", "true", "yes", "on"} + + +def run_diffusion_suite(suite: str) -> None: + """Run one legacy-named diffusion suite without exposing a second CI CLI.""" + + from sglang.multimodal_gen.test.runner.diffusion_suite_runner import main + + args = [sys.argv[0], "--suite", suite] + optional_values = ( + ("DIFFUSION_PARTITION_ID", "--partition-id"), + ("DIFFUSION_TOTAL_PARTITIONS", "--total-partitions"), + ("DIFFUSION_PARTITION_PLAN_JSON", "--partition-plan-json"), + ("DIFFUSION_PYTEST_FILTER", "--filter"), + ) + for environment_name, option in optional_values: + value = os.environ.get(environment_name) + if value: + args.extend([option, value]) + if _enabled("DIFFUSION_CONTINUE_ON_ERROR"): + args.append("--continue-on-error") + + # Preserve the historical cwd: several diffusion fixtures emit artifacts + # relative to ``python/`` rather than to their source file. + os.chdir(Path(__file__).resolve().parents[4] / "python") + sys.argv = args + main() diff --git a/scripts/ci/utils/diffusion/diffusion_case_parser.py b/scripts/ci/utils/diffusion/diffusion_case_parser.py index 81af8756c..fbbc69d8a 100755 --- a/scripts/ci/utils/diffusion/diffusion_case_parser.py +++ b/scripts/ci/utils/diffusion/diffusion_case_parser.py @@ -44,7 +44,9 @@ STARTUP_OVERHEAD_SECONDS = 120.0 # Paths relative to repository root BASELINE_REL_PATH = "python/sglang/multimodal_gen/test/server/perf_baselines" BASELINE_PLATFORM_ORDER = ("h100", "b200", "5090") -RUN_SUITE_REL_PATH = "python/sglang/multimodal_gen/test/run_suite.py" +RUN_SUITE_REL_PATH = ( + "python/sglang/multimodal_gen/test/runner/diffusion_suite_runner.py" +) USE_NPU_CONFIGS = os.getenv("USE_NPU_CONFIGS", "0").lower() in ("1", "true") diff --git a/scripts/ci/utils/diffusion/verify_diffusion_coverage.py b/scripts/ci/utils/diffusion/verify_diffusion_coverage.py index 925b99c26..2643391fc 100755 --- a/scripts/ci/utils/diffusion/verify_diffusion_coverage.py +++ b/scripts/ci/utils/diffusion/verify_diffusion_coverage.py @@ -155,7 +155,8 @@ def print_missing_standalone_estimates_summary( print("\n" + "=" * 60) print( - "Add standalone estimate(s) to python/sglang/multimodal_gen/test/run_suite.py" + "Add standalone estimate(s) to " + "python/sglang/multimodal_gen/test/runner/diffusion_suite_runner.py" ) print("=" * 60) print("The following standalone file(s) used fallback estimate 300.0s.") diff --git a/scripts/lint/check_registered_tests.py b/scripts/lint/check_registered_tests.py index 32828e802..8699b67c3 100755 --- a/scripts/lint/check_registered_tests.py +++ b/scripts/lint/check_registered_tests.py @@ -26,6 +26,7 @@ import glob import importlib.util import os import re +import subprocess import sys # Suite names of the form `{stage}-test-{runner_config}` are exactly what the @@ -38,6 +39,8 @@ _MODERN_SHAPE = re.compile(r"^(.+)-test-(.+)$") # no suite any workflow invokes and the test silently never runs. _LEGACY_CUDA_PREFIXES = ("stress",) +_TEST_KINDS = {"unit", "kernel", "e2e", "accuracy", "perf", "stress"} + def _defines_testcase(tree: ast.AST) -> bool: """True if the file defines unittest classes, statically or via type().""" @@ -70,6 +73,99 @@ def _main_runs_tests(tree: ast.Module) -> bool: return False +def _git_lines(*args: str) -> list[str] | None: + result = subprocess.run(["git", *args], capture_output=True, text=True, check=False) + if result.returncode != 0: + return None + return [line for line in result.stdout.splitlines() if line] + + +def _changed_registered_files() -> set[str]: + """Return added, copied, or renamed registered-test destinations.""" + + lines = _git_lines("diff", "--cached", "--name-status", "--diff-filter=ACR") + if not lines: + base_ref = os.environ.get("GITHUB_BASE_REF", "main") + for candidate in (f"origin/{base_ref}", base_ref): + if _git_lines("rev-parse", "--verify", candidate) is None: + continue + merge_base = _git_lines("merge-base", candidate, "HEAD") + if not merge_base: + continue + lines = _git_lines( + "diff", + "--name-status", + "--diff-filter=ACR", + merge_base[0], + "HEAD", + ) + break + + selected = set() + for line in lines or []: + fields = line.split("\t") + destination = fields[-1] + if destination.startswith("test/registered/") and destination.endswith(".py"): + selected.add(destination) + return selected + + +def _contains_call(tree: ast.AST, name: str) -> bool: + return any( + isinstance(node, ast.Call) + and ( + (isinstance(node.func, ast.Name) and node.func.id == name) + or (isinstance(node.func, ast.Attribute) and node.func.attr == name) + ) + for node in ast.walk(tree) + ) + + +def taxonomy_errors(path: str, registries: list, tree: ast.AST) -> list[str]: + """Validate the kind/subsystem contract for a newly admitted path.""" + + parts = path.split("/") + relative_parts = parts[2:] if parts[:2] == ["test", "registered"] else [] + if len(relative_parts) < 3 or relative_parts[0] not in _TEST_KINDS: + return [ + f"{path}: registered tests must live under " + "test/registered///; kind must be one of " + + ", ".join(sorted(_TEST_KINDS)) + ] + + kind = relative_parts[0] + errors = [] + if kind == "unit": + non_cpu = [r for r in registries if r.backend.name != "CPU"] + if non_cpu: + errors.append(f"{path}: unit tests may register only CPU suites") + if any(r.est_time > 60 for r in registries): + errors.append(f"{path}: unit test est_time must be <= 60 seconds") + if _contains_call(tree, "popen_launch_server"): + errors.append(f"{path}: unit tests may not launch a server") + elif kind == "kernel": + if any("-kernel-" not in (r.effective_suite or "") for r in registries): + errors.append(f"{path}: kernel tests must use a *-kernel-* suite") + elif kind in {"accuracy", "perf"}: + invalid = [ + r + for r in registries + if not (r.effective_suite or "").startswith(("nightly-", "weekly-")) + ] + if invalid: + errors.append(f"{path}: {kind} tests must use nightly/weekly suites") + elif kind == "stress": + invalid = [ + r + for r in registries + if (r.effective_suite or "") != "stress" + and not (r.effective_suite or "").startswith("weekly-") + ] + if invalid: + errors.append(f"{path}: stress tests must use stress/weekly suites") + return errors + + def main() -> int: # Import ci_register directly to avoid pulling in all of sglang spec = importlib.util.spec_from_file_location( @@ -93,6 +189,8 @@ def main() -> int: legacy_shape = [] # (file, suite, stage, runner_config) -- has a -test- split non_dispatchable = [] # (file, suite) -- legacy CUDA suite no workflow invokes dead_tests = [] # (file) -- TestCase classes that `python3 file.py` never runs + taxonomy_violations = [] + changed_files = _changed_registered_files() for f in files: try: registries, _has_main_entry = ci_register.ut_parse_one_file(f) @@ -106,6 +204,8 @@ def main() -> int: # `python3 file.py`); the ERROR text below explains the fix. with open(f, "r", encoding="utf-8") as fh: tree = ast.parse(fh.read(), filename=f) + if f in changed_files: + taxonomy_violations.extend(taxonomy_errors(f, registries, tree)) if _defines_testcase(tree) and not _main_runs_tests(tree): dead_tests.append(f) for r in registries: @@ -172,6 +272,12 @@ def main() -> int: print(f" {f}") print() exit_code = 1 + if taxonomy_violations: + print("ERROR: Registered-test taxonomy violations:") + for error in taxonomy_violations: + print(f" {error}") + print() + exit_code = 1 return exit_code diff --git a/sgl-model-gateway/bindings/python/README.md b/sgl-model-gateway/bindings/python/README.md index 5f913e241..024fff0ca 100644 --- a/sgl-model-gateway/bindings/python/README.md +++ b/sgl-model-gateway/bindings/python/README.md @@ -18,9 +18,8 @@ bindings/python/ │ └── mini_lb.py ├── tests/ # Python unit tests │ ├── conftest.py -│ ├── test_validation.py │ ├── test_arg_parser.py -│ ├── test_router_config.py +│ ├── test_pyo3_binding.py │ └── test_startup_sequence.py ├── Cargo.toml # Rust package configuration for bindings ├── pyproject.toml # Python package configuration diff --git a/sgl-model-gateway/bindings/python/tests/test_router_config.py b/sgl-model-gateway/bindings/python/tests/test_router_config.py deleted file mode 100644 index 5ba91cb1d..000000000 --- a/sgl-model-gateway/bindings/python/tests/test_router_config.py +++ /dev/null @@ -1,423 +0,0 @@ -""" -Unit tests for router configuration validation and setup. - -These tests focus on testing the router configuration logic in isolation, -including validation of configuration parameters and their interactions. -""" - -from unittest.mock import MagicMock, patch - -import pytest -from sglang_router.launch_router import RouterArgs, launch_router -from sglang_router.router import policy_from_str -from sglang_router.sglang_router_rs import PolicyType - - -class TestRouterConfigValidation: - """Test router configuration validation logic.""" - - def test_valid_basic_config(self): - """Test that a valid basic configuration passes validation.""" - args = RouterArgs( - host="127.0.0.1", - port=30000, - worker_urls=["http://worker1:8000", "http://worker2:8000"], - policy="cache_aware", - ) - - # Should not raise any exceptions - assert args.host == "127.0.0.1" - assert args.port == 30000 - assert args.worker_urls == ["http://worker1:8000", "http://worker2:8000"] - assert args.policy == "cache_aware" - - def test_valid_pd_config(self): - """Test that a valid PD configuration passes validation.""" - args = RouterArgs( - host="127.0.0.1", - port=30000, - pd_disaggregation=True, - prefill_urls=[ - ("http://prefill1:8000", 9000), - ("http://prefill2:8000", None), - ], - decode_urls=["http://decode1:8001", "http://decode2:8001"], - policy="cache_aware", - ) - - assert args.pd_disaggregation is True - assert args.prefill_urls == [ - ("http://prefill1:8000", 9000), - ("http://prefill2:8000", None), - ] - assert args.decode_urls == ["http://decode1:8001", "http://decode2:8001"] - assert args.policy == "cache_aware" - - def test_pd_config_without_urls_allowed(self): - """Test that PD mode without URLs is now allowed (URLs are optional).""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=False, - ) - - # Should not raise validation error - URLs are now optional - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # This should succeed without raising an error - launch_router(args) - router_mod.from_args.assert_called_once() - - def test_pd_config_with_service_discovery_allows_empty_urls(self): - """Test that PD mode with service discovery allows empty URLs.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=True, - ) - - # Should not raise validation error when service discovery is enabled - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_regular_mode_without_workers_allows_empty_urls(self): - """Test that regular mode allows empty worker URLs.""" - args = RouterArgs(worker_urls=[], service_discovery=False) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_cache_threshold_validation(self): - """Test cache threshold validation.""" - # Valid cache threshold - args = RouterArgs(cache_threshold=0.5) - assert args.cache_threshold == 0.5 - - # Edge cases - args = RouterArgs(cache_threshold=0.0) - assert args.cache_threshold == 0.0 - - args = RouterArgs(cache_threshold=1.0) - assert args.cache_threshold == 1.0 - - def test_balance_threshold_validation(self): - """Test load balancing threshold validation.""" - # Valid thresholds - args = RouterArgs(balance_abs_threshold=64, balance_rel_threshold=1.5) - assert args.balance_abs_threshold == 64 - assert args.balance_rel_threshold == 1.5 - - # Edge cases - args = RouterArgs(balance_abs_threshold=0, balance_rel_threshold=1.0) - assert args.balance_abs_threshold == 0 - assert args.balance_rel_threshold == 1.0 - - def test_timeout_validation(self): - """Test timeout parameter validation.""" - # Valid timeouts - args = RouterArgs( - worker_startup_timeout_secs=600, - worker_startup_check_interval=30, - request_timeout_secs=1800, - queue_timeout_secs=60, - ) - assert args.worker_startup_timeout_secs == 600 - assert args.worker_startup_check_interval == 30 - assert args.request_timeout_secs == 1800 - assert args.queue_timeout_secs == 60 - - def test_retry_config_validation(self): - """Test retry configuration validation.""" - # Valid retry config - args = RouterArgs( - retry_max_retries=5, - retry_initial_backoff_ms=50, - retry_max_backoff_ms=30000, - retry_backoff_multiplier=1.5, - retry_jitter_factor=0.2, - disable_retries=False, - ) - assert args.retry_max_retries == 5 - assert args.retry_initial_backoff_ms == 50 - assert args.retry_max_backoff_ms == 30000 - assert args.retry_backoff_multiplier == 1.5 - assert args.retry_jitter_factor == 0.2 - assert args.disable_retries is False - - def test_circuit_breaker_config_validation(self): - """Test circuit breaker configuration validation.""" - # Valid circuit breaker config - args = RouterArgs( - cb_failure_threshold=10, - cb_success_threshold=3, - cb_timeout_duration_secs=60, - cb_window_duration_secs=120, - disable_circuit_breaker=False, - ) - assert args.cb_failure_threshold == 10 - assert args.cb_success_threshold == 3 - assert args.cb_timeout_duration_secs == 60 - assert args.cb_window_duration_secs == 120 - assert args.disable_circuit_breaker is False - - def test_health_check_config_validation(self): - """Test health check configuration validation.""" - # Valid health check config - args = RouterArgs( - health_failure_threshold=3, - health_success_threshold=2, - health_check_timeout_secs=5, - health_check_interval_secs=60, - health_check_endpoint="/health", - ) - assert args.health_failure_threshold == 3 - assert args.health_success_threshold == 2 - assert args.health_check_timeout_secs == 5 - assert args.health_check_interval_secs == 60 - assert args.health_check_endpoint == "/health" - - def test_rate_limiting_config_validation(self): - """Test rate limiting configuration validation.""" - # Valid rate limiting config - args = RouterArgs( - max_concurrent_requests=256, - queue_size=100, - queue_timeout_secs=60, - rate_limit_tokens_per_second=100, - ) - assert args.max_concurrent_requests == 256 - assert args.queue_size == 100 - assert args.queue_timeout_secs == 60 - assert args.rate_limit_tokens_per_second == 100 - - def test_service_discovery_config_validation(self): - """Test service discovery configuration validation.""" - # Valid service discovery config - args = RouterArgs( - service_discovery=True, - selector={"app": "worker", "env": "prod"}, - service_discovery_port=8080, - service_discovery_namespace="default", - ) - assert args.service_discovery is True - assert args.selector == {"app": "worker", "env": "prod"} - assert args.service_discovery_port == 8080 - assert args.service_discovery_namespace == "default" - - def test_pd_service_discovery_config_validation(self): - """Test PD service discovery configuration validation.""" - # Valid PD service discovery config - args = RouterArgs( - pd_disaggregation=True, - service_discovery=True, - prefill_selector={"app": "prefill"}, - decode_selector={"app": "decode"}, - bootstrap_port_annotation="sglang.ai/bootstrap-port", - ) - assert args.pd_disaggregation is True - assert args.service_discovery is True - assert args.prefill_selector == {"app": "prefill"} - assert args.decode_selector == {"app": "decode"} - assert args.bootstrap_port_annotation == "sglang.ai/bootstrap-port" - - def test_prometheus_config_validation(self): - """Test Prometheus configuration validation.""" - # Valid Prometheus config - args = RouterArgs(prometheus_port=29000, prometheus_host="127.0.0.1") - assert args.prometheus_port == 29000 - assert args.prometheus_host == "127.0.0.1" - - def test_cors_config_validation(self): - """Test CORS configuration validation.""" - # Valid CORS config - args = RouterArgs( - cors_allowed_origins=["http://localhost:3000", "https://example.com"] - ) - assert args.cors_allowed_origins == [ - "http://localhost:3000", - "https://example.com", - ] - - def test_tokenizer_config_validation(self): - """Test tokenizer configuration validation.""" - # Note: model_path and tokenizer_path are not available in current RouterArgs - pytest.skip("Tokenizer configuration not available in current implementation") - - def test_dp_aware_config_validation(self): - """Test data parallelism aware configuration validation.""" - # Valid DP aware config - args = RouterArgs(dp_aware=True, api_key="test-api-key") - assert args.dp_aware is True - assert args.api_key == "test-api-key" - - def test_request_id_headers_validation(self): - """Test request ID headers configuration validation.""" - # Valid request ID headers config - args = RouterArgs( - request_id_headers=["x-request-id", "x-trace-id", "x-correlation-id"] - ) - assert args.request_id_headers == [ - "x-request-id", - "x-trace-id", - "x-correlation-id", - ] - - def test_policy_consistency_validation(self): - """Test policy consistency validation in PD mode.""" - # Test with both prefill and decode policies specified - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", None)], - decode_urls=["http://decode1:8001"], - policy="cache_aware", - prefill_policy="power_of_two", - decode_policy="round_robin", - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_policy_fallback_validation(self): - """Test policy fallback validation in PD mode.""" - # Test with only prefill policy specified - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", None)], - decode_urls=["http://decode1:8001"], - policy="cache_aware", - prefill_policy="power_of_two", - decode_policy=None, - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_policy_enum_conversion(self): - """Test policy string to enum conversion.""" - # Test all valid policy conversions - assert policy_from_str("random") == PolicyType.Random - assert policy_from_str("round_robin") == PolicyType.RoundRobin - assert policy_from_str("cache_aware") == PolicyType.CacheAware - assert policy_from_str("power_of_two") == PolicyType.PowerOfTwo - - def test_invalid_policy_enum_conversion(self): - """Test invalid policy string to enum conversion.""" - with pytest.raises(KeyError): - policy_from_str("invalid_policy") - - def test_config_immutability(self): - """Test that configuration objects are properly immutable.""" - args = RouterArgs( - host="127.0.0.1", port=30000, worker_urls=["http://worker1:8000"] - ) - - # Test that we can't modify the configuration after creation - # (This is more of a design test - dataclasses are mutable by default) - original_host = args.host - args.host = "0.0.0.0" - assert args.host == "0.0.0.0" # Dataclasses are mutable - assert args.host != original_host - - def test_config_defaults_consistency(self): - """Test that configuration defaults are consistent.""" - args1 = RouterArgs() - args2 = RouterArgs() - - # Both instances should have the same defaults - assert args1.host == args2.host - assert args1.port == args2.port - assert args1.policy == args2.policy - assert args1.worker_urls == args2.worker_urls - assert args1.pd_disaggregation == args2.pd_disaggregation - - def test_config_serialization(self): - """Test that configuration can be serialized/deserialized.""" - args = RouterArgs( - host="127.0.0.1", - port=30000, - worker_urls=["http://worker1:8000"], - policy="cache_aware", - cache_threshold=0.5, - ) - - # Test that we can access all attributes - assert hasattr(args, "host") - assert hasattr(args, "port") - assert hasattr(args, "worker_urls") - assert hasattr(args, "policy") - assert hasattr(args, "cache_threshold") - - def test_config_with_none_values(self): - """Test configuration with None values.""" - args = RouterArgs( - api_key=None, - log_dir=None, - log_level=None, - prometheus_port=None, - prometheus_host=None, - request_id_headers=None, - rate_limit_tokens_per_second=None, - service_discovery_namespace=None, - ) - - # All None values should be preserved - assert args.api_key is None - assert args.log_dir is None - assert args.log_level is None - assert args.prometheus_port is None - assert args.prometheus_host is None - assert args.request_id_headers is None - assert args.rate_limit_tokens_per_second is None - assert args.service_discovery_namespace is None - - def test_config_with_empty_lists(self): - """Test configuration with empty lists.""" - args = RouterArgs( - worker_urls=[], prefill_urls=[], decode_urls=[], cors_allowed_origins=[] - ) - - # All empty lists should be preserved - assert args.worker_urls == [] - assert args.prefill_urls == [] - assert args.decode_urls == [] - assert args.cors_allowed_origins == [] - - def test_config_with_empty_dicts(self): - """Test configuration with empty dictionaries.""" - args = RouterArgs(selector={}, prefill_selector={}, decode_selector={}) - - # All empty dictionaries should be preserved - assert args.selector == {} - assert args.prefill_selector == {} - assert args.decode_selector == {} diff --git a/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py b/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py index e2b9567d2..8a937a284 100644 --- a/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py +++ b/sgl-model-gateway/bindings/python/tests/test_startup_sequence.py @@ -1,642 +1,4 @@ -""" -Unit tests for startup sequence logic in sglang_router. - -These tests focus on testing the startup sequence logic in isolation, -including router initialization, configuration validation, and startup flow. -""" - -import logging -from unittest.mock import MagicMock, patch - -import pytest -from sglang_router.launch_router import RouterArgs, launch_router -from sglang_router.router import policy_from_str - - -# Local helper mirroring the router logger setup used in production -def setup_logger(): - logger = logging.getLogger("router") - logger.setLevel(logging.INFO) - if not logger.handlers: - formatter = logging.Formatter( - "[Router (Python)] %(asctime)s - %(levelname)s - %(message)s", - datefmt="%Y-%m-%d %H:%M:%S", - ) - handler = logging.StreamHandler() - handler.setFormatter(formatter) - logger.addHandler(handler) - return logger - - -from sglang_router.sglang_router_rs import PolicyType - - -class TestSetupLogger: - """Test logger setup functionality.""" - - def test_setup_logger_returns_logger(self): - """Test that setup_logger returns a logger instance.""" - logger = setup_logger() - - assert isinstance(logger, logging.Logger) - assert logger.name == "router" - assert logger.level == logging.INFO - - def test_setup_logger_has_handler(self): - """Test that setup_logger configures a handler.""" - logger = setup_logger() - - assert len(logger.handlers) > 0 - handler = logger.handlers[0] - assert isinstance(handler, logging.StreamHandler) - - def test_setup_logger_has_formatter(self): - """Test that setup_logger configures a formatter.""" - logger = setup_logger() - - handler = logger.handlers[0] - formatter = handler.formatter - - assert formatter is not None - assert "[Router (Python)]" in formatter._fmt - - def test_setup_logger_multiple_calls(self): - """Test that multiple calls to setup_logger work correctly.""" - logger1 = setup_logger() - logger2 = setup_logger() - - # Should return the same logger instance - assert logger1 is logger2 - - -class TestPolicyFromStr: - """Test policy string to enum conversion in startup context.""" - - def test_policy_conversion_in_startup(self): - """Test policy conversion during startup sequence.""" - # Test all valid policies - policies = ["random", "round_robin", "cache_aware", "power_of_two"] - expected_enums = [ - PolicyType.Random, - PolicyType.RoundRobin, - PolicyType.CacheAware, - PolicyType.PowerOfTwo, - ] - - for policy_str, expected_enum in zip(policies, expected_enums): - result = policy_from_str(policy_str) - assert result == expected_enum - - def test_invalid_policy_in_startup(self): - """Test handling of invalid policy during startup.""" - with pytest.raises(KeyError): - policy_from_str("invalid_policy") - - -class TestRouterInitialization: - """Test router initialization logic.""" - - def test_router_initialization_basic(self): - """Test basic router initialization.""" - args = RouterArgs( - host="127.0.0.1", - port=30000, - worker_urls=["http://worker1:8000"], - policy="cache_aware", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - # capture needed fields from RouterArgs - captured_args.update( - dict( - host=router_args.host, - port=router_args.port, - worker_urls=router_args.worker_urls, - policy=policy_from_str(router_args.policy), - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify Router.from_args was called and captured fields match - router_mod.from_args.assert_called_once() - assert captured_args["host"] == "127.0.0.1" - assert captured_args["port"] == 30000 - assert captured_args["worker_urls"] == ["http://worker1:8000"] - assert captured_args["policy"] == PolicyType.CacheAware - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_pd_mode(self): - """Test router initialization in PD mode.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", 9000)], - decode_urls=["http://decode1:8001"], - policy="power_of_two", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - pd_disaggregation=router_args.pd_disaggregation, - prefill_urls=router_args.prefill_urls, - decode_urls=router_args.decode_urls, - policy=policy_from_str(router_args.policy), - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify Router.from_args was called with PD parameters - router_mod.from_args.assert_called_once() - assert captured_args["pd_disaggregation"] is True - assert captured_args["prefill_urls"] == [("http://prefill1:8000", 9000)] - assert captured_args["decode_urls"] == ["http://decode1:8001"] - assert captured_args["policy"] == PolicyType.PowerOfTwo - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_service_discovery(self): - """Test router initialization with service discovery.""" - args = RouterArgs( - service_discovery=True, - selector={"app": "worker", "env": "prod"}, - service_discovery_port=8080, - service_discovery_namespace="default", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - service_discovery=router_args.service_discovery, - selector=router_args.selector, - service_discovery_port=router_args.service_discovery_port, - service_discovery_namespace=router_args.service_discovery_namespace, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify Router.from_args was called with service discovery parameters - router_mod.from_args.assert_called_once() - assert captured_args["service_discovery"] is True - assert captured_args["selector"] == {"app": "worker", "env": "prod"} - assert captured_args["service_discovery_port"] == 8080 - assert captured_args["service_discovery_namespace"] == "default" - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_retry_config(self): - """Test router initialization with retry configuration.""" - args = RouterArgs( - retry_max_retries=3, - retry_initial_backoff_ms=100, - retry_max_backoff_ms=10000, - retry_backoff_multiplier=2.0, - retry_jitter_factor=0.1, - disable_retries=False, - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - retry_max_retries=router_args.retry_max_retries, - retry_initial_backoff_ms=router_args.retry_initial_backoff_ms, - retry_max_backoff_ms=router_args.retry_max_backoff_ms, - retry_backoff_multiplier=router_args.retry_backoff_multiplier, - retry_jitter_factor=router_args.retry_jitter_factor, - disable_retries=router_args.disable_retries, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify router was created with retry parameters - router_mod.from_args.assert_called_once() - assert captured_args["retry_max_retries"] == 3 - assert captured_args["retry_initial_backoff_ms"] == 100 - assert captured_args["retry_max_backoff_ms"] == 10000 - assert captured_args["retry_backoff_multiplier"] == 2.0 - assert captured_args["retry_jitter_factor"] == 0.1 - assert captured_args["disable_retries"] is False - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_circuit_breaker_config(self): - """Test router initialization with circuit breaker configuration.""" - args = RouterArgs( - cb_failure_threshold=5, - cb_success_threshold=2, - cb_timeout_duration_secs=30, - cb_window_duration_secs=60, - disable_circuit_breaker=False, - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - cb_failure_threshold=router_args.cb_failure_threshold, - cb_success_threshold=router_args.cb_success_threshold, - cb_timeout_duration_secs=router_args.cb_timeout_duration_secs, - cb_window_duration_secs=router_args.cb_window_duration_secs, - disable_circuit_breaker=router_args.disable_circuit_breaker, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify router was created with circuit breaker parameters - router_mod.from_args.assert_called_once() - assert captured_args["cb_failure_threshold"] == 5 - assert captured_args["cb_success_threshold"] == 2 - assert captured_args["cb_timeout_duration_secs"] == 30 - assert captured_args["cb_window_duration_secs"] == 60 - assert captured_args["disable_circuit_breaker"] is False - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_rate_limiting_config(self): - """Test router initialization with rate limiting configuration.""" - args = RouterArgs( - max_concurrent_requests=512, - queue_size=200, - queue_timeout_secs=120, - rate_limit_tokens_per_second=100, - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - max_concurrent_requests=router_args.max_concurrent_requests, - queue_size=router_args.queue_size, - queue_timeout_secs=router_args.queue_timeout_secs, - rate_limit_tokens_per_second=router_args.rate_limit_tokens_per_second, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify router was created with rate limiting parameters - router_mod.from_args.assert_called_once() - assert captured_args["max_concurrent_requests"] == 512 - assert captured_args["queue_size"] == 200 - assert captured_args["queue_timeout_secs"] == 120 - assert captured_args["rate_limit_tokens_per_second"] == 100 - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_health_check_config(self): - """Test router initialization with health check configuration.""" - args = RouterArgs( - health_failure_threshold=2, - health_success_threshold=1, - health_check_timeout_secs=3, - health_check_interval_secs=30, - health_check_endpoint="/healthz", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - health_failure_threshold=router_args.health_failure_threshold, - health_success_threshold=router_args.health_success_threshold, - health_check_timeout_secs=router_args.health_check_timeout_secs, - health_check_interval_secs=router_args.health_check_interval_secs, - health_check_endpoint=router_args.health_check_endpoint, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify router was created with health check parameters - router_mod.from_args.assert_called_once() - assert captured_args["health_failure_threshold"] == 2 - assert captured_args["health_success_threshold"] == 1 - assert captured_args["health_check_timeout_secs"] == 3 - assert captured_args["health_check_interval_secs"] == 30 - assert captured_args["health_check_endpoint"] == "/healthz" - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_prometheus_config(self): - """Test router initialization with Prometheus configuration.""" - args = RouterArgs(prometheus_port=29000, prometheus_host="127.0.0.1") - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - prometheus_port=router_args.prometheus_port, - prometheus_host=router_args.prometheus_host, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify router was created with Prometheus parameters - router_mod.from_args.assert_called_once() - assert captured_args["prometheus_port"] == 29000 - assert captured_args["prometheus_host"] == "127.0.0.1" - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_cors_config(self): - """Test router initialization with CORS configuration.""" - args = RouterArgs( - cors_allowed_origins=["http://localhost:3000", "https://example.com"] - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict(cors_allowed_origins=router_args.cors_allowed_origins) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify router was created with CORS parameters - router_mod.from_args.assert_called_once() - assert captured_args["cors_allowed_origins"] == [ - "http://localhost:3000", - "https://example.com", - ] - - # Verify router.start() was called - mock_router_instance.start.assert_called_once() - - # Function returns None; ensure start was invoked - - def test_router_initialization_with_tokenizer_config(self): - """Test router initialization with tokenizer configuration.""" - # Note: model_path and tokenizer_path are not available in current RouterArgs - pytest.skip("Tokenizer configuration not available in current implementation") - - -class TestStartupValidation: - """Test startup validation logic.""" - - def test_pd_mode_validation_during_startup(self): - """Test PD mode validation during startup.""" - # PD mode without URLs is now allowed (URLs are optional) - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=False, - ) - - # Should not raise validation error - URLs are now optional - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # This should succeed without raising an error - launch_router(args) - router_mod.from_args.assert_called_once() - - def test_pd_mode_with_service_discovery_validation(self): - """Test PD mode with service discovery validation during startup.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=True, - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - result = launch_router(args) - - # Should create router instance - router_mod.from_args.assert_called_once() - - def test_policy_warning_during_startup(self): - """Test policy warning during startup in PD mode.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", None)], - decode_urls=["http://decode1:8001"], - policy="cache_aware", - prefill_policy="power_of_two", - decode_policy="round_robin", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # The policy messages are emitted by router_args logger - with patch("sglang_router.router_args.logger") as mock_logger: - result = launch_router(args) - - # Should log warning about policy usage - mock_logger.warning.assert_called_once() - warning_call = mock_logger.warning.call_args[0][0] - assert ( - "Both --prefill-policy and --decode-policy are specified" - in warning_call - ) - - # Should create router instance - router_mod.from_args.assert_called_once() - - def test_policy_info_during_startup(self): - """Test policy info logging during startup in PD mode.""" - # Test with only prefill policy specified - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", None)], - decode_urls=["http://decode1:8001"], - policy="cache_aware", - prefill_policy="power_of_two", - decode_policy=None, - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # The policy messages are emitted by router_args logger - with patch("sglang_router.router_args.logger") as mock_logger: - result = launch_router(args) - - # Should log info about policy usage - mock_logger.info.assert_called_once() - info_call = mock_logger.info.call_args[0][0] - assert "Using --prefill-policy 'power_of_two'" in info_call - assert "and --policy 'cache_aware'" in info_call - - # Should create router instance - router_mod.from_args.assert_called_once() - - def test_policy_info_decode_only_during_startup(self): - """Test policy info logging during startup with only decode policy specified.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", None)], - decode_urls=["http://decode1:8001"], - policy="cache_aware", - prefill_policy=None, - decode_policy="round_robin", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # The policy messages are emitted by router_args logger - with patch("sglang_router.router_args.logger") as mock_logger: - result = launch_router(args) - - # Should log info about policy usage - mock_logger.info.assert_called_once() - info_call = mock_logger.info.call_args[0][0] - assert "Using --policy 'cache_aware'" in info_call - assert "and --decode-policy 'round_robin'" in info_call - - # Should create router instance - router_mod.from_args.assert_called_once() - - -class TestStartupErrorHandling: - """Test startup error handling logic.""" - - def test_router_creation_error_handling(self): - """Test error handling when router creation fails.""" - args = RouterArgs( - host="127.0.0.1", port=30000, worker_urls=["http://worker1:8000"] - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - # Simulate router creation failure in from_args - router_mod.from_args = MagicMock( - side_effect=Exception("Router creation failed") - ) - - with patch("sglang_router.launch_router.logger") as mock_logger: - with pytest.raises(Exception, match="Router creation failed"): - launch_router(args) - - # Should log error - mock_logger.error.assert_called_once() - error_call = mock_logger.error.call_args[0][0] - assert "Error starting router: Router creation failed" in error_call - - def test_router_start_error_handling(self): - """Test error handling when router start fails.""" - args = RouterArgs( - host="127.0.0.1", port=30000, worker_urls=["http://worker1:8000"] - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # Simulate router start failure - mock_router_instance.start.side_effect = Exception("Router start failed") - - with patch("sglang_router.launch_router.logger") as mock_logger: - with pytest.raises(Exception, match="Router start failed"): - launch_router(args) - - # Should log error - mock_logger.error.assert_called_once() - error_call = mock_logger.error.call_args[0][0] - assert "Error starting router: Router start failed" in error_call - - -# --- Added unit tests for Router wrapper and launch_server helpers --- +"""Focused runtime tests for the Python router and server-launch helpers.""" def _install_sglang_stubs(monkeypatch): @@ -827,7 +189,7 @@ def test_launch_server_process_and_cleanup(monkeypatch): sa = SA() sa.tp_size = 2 - proc = ls.launch_server_process(sa, worker_port=31001, dp_id=3) + ls.launch_server_process(sa, worker_port=31001, dp_id=3) assert created.get("started") is True targ, targ_args = created["target"], created["args"] assert targ is ls.run_server @@ -911,206 +273,3 @@ def test_launch_server_process_declares_on_a_resolved_record(monkeypatch): worker = created["args"][0] assert (worker.port, worker.base_gpu_id, worker.dp_size) == (31002, 6, 1) assert (parent.port, parent.base_gpu_id, parent.dp_size) == (30000, 0, 4) - - def test_validation_error_handling(self): - """Test error handling when validation fails.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=False, - ) - - with patch("sglang_router.launch_router.logger") as mock_logger: - with pytest.raises( - ValueError, match="PD disaggregation mode requires --prefill" - ): - launch_router(args) - - # Should log error for validation failures - mock_logger.error.assert_called_once() - - -class TestStartupFlow: - """Test complete startup flow.""" - - def test_complete_startup_flow_basic(self): - """Test complete startup flow for basic configuration.""" - args = RouterArgs( - host="127.0.0.1", - port=30000, - worker_urls=["http://worker1:8000", "http://worker2:8000"], - policy="cache_aware", - cache_threshold=0.5, - balance_abs_threshold=32, - balance_rel_threshold=1.5, - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - result = launch_router(args) - - # Verify complete flow - router_mod.from_args.assert_called_once() - mock_router_instance.start.assert_called_once() - - def test_complete_startup_flow_pd_mode(self): - """Test complete startup flow for PD mode configuration.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[ - ("http://prefill1:8000", 9000), - ("http://prefill2:8000", None), - ], - decode_urls=["http://decode1:8001", "http://decode2:8001"], - policy="power_of_two", - prefill_policy="cache_aware", - decode_policy="round_robin", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - with patch("sglang_router.router_args.logger") as mock_logger: - result = launch_router(args) - - # Verify complete flow - router_mod.from_args.assert_called_once() - mock_router_instance.start.assert_called_once() - - # Verify policy warning was logged - mock_logger.warning.assert_called_once() - - def test_complete_startup_flow_with_all_features(self): - """Test complete startup flow with all features enabled.""" - args = RouterArgs( - host="0.0.0.0", - port=30001, - worker_urls=["http://worker1:8000"], - policy="round_robin", - service_discovery=True, - selector={"app": "worker"}, - service_discovery_port=8080, - service_discovery_namespace="default", - dp_aware=True, - api_key="test-key", - log_dir="/tmp/logs", - log_level="debug", - prometheus_port=29000, - prometheus_host="0.0.0.0", - request_id_headers=["x-request-id", "x-trace-id"], - request_timeout_secs=1200, - max_concurrent_requests=512, - queue_size=200, - queue_timeout_secs=120, - rate_limit_tokens_per_second=100, - cors_allowed_origins=["http://localhost:3000"], - retry_max_retries=3, - retry_initial_backoff_ms=100, - retry_max_backoff_ms=10000, - retry_backoff_multiplier=2.0, - retry_jitter_factor=0.1, - cb_failure_threshold=5, - cb_success_threshold=2, - cb_timeout_duration_secs=30, - cb_window_duration_secs=60, - health_failure_threshold=2, - health_success_threshold=1, - health_check_timeout_secs=3, - health_check_interval_secs=30, - health_check_endpoint="/healthz", - ) - - with patch("sglang_router.launch_router.Router") as router_mod: - captured_args = {} - mock_router_instance = MagicMock() - - def fake_from_args(router_args): - captured_args.update( - dict( - host=router_args.host, - port=router_args.port, - worker_urls=router_args.worker_urls, - policy=policy_from_str(router_args.policy), - service_discovery=router_args.service_discovery, - selector=router_args.selector, - service_discovery_port=router_args.service_discovery_port, - service_discovery_namespace=router_args.service_discovery_namespace, - dp_aware=router_args.dp_aware, - api_key=router_args.api_key, - log_dir=router_args.log_dir, - log_level=router_args.log_level, - prometheus_port=router_args.prometheus_port, - prometheus_host=router_args.prometheus_host, - request_id_headers=router_args.request_id_headers, - request_timeout_secs=router_args.request_timeout_secs, - max_concurrent_requests=router_args.max_concurrent_requests, - queue_size=router_args.queue_size, - queue_timeout_secs=router_args.queue_timeout_secs, - rate_limit_tokens_per_second=router_args.rate_limit_tokens_per_second, - cors_allowed_origins=router_args.cors_allowed_origins, - retry_max_retries=router_args.retry_max_retries, - retry_initial_backoff_ms=router_args.retry_initial_backoff_ms, - retry_max_backoff_ms=router_args.retry_max_backoff_ms, - retry_backoff_multiplier=router_args.retry_backoff_multiplier, - retry_jitter_factor=router_args.retry_jitter_factor, - cb_failure_threshold=router_args.cb_failure_threshold, - cb_success_threshold=router_args.cb_success_threshold, - cb_timeout_duration_secs=router_args.cb_timeout_duration_secs, - cb_window_duration_secs=router_args.cb_window_duration_secs, - health_failure_threshold=router_args.health_failure_threshold, - health_success_threshold=router_args.health_success_threshold, - health_check_timeout_secs=router_args.health_check_timeout_secs, - health_check_interval_secs=router_args.health_check_interval_secs, - health_check_endpoint=router_args.health_check_endpoint, - ) - ) - return mock_router_instance - - router_mod.from_args = MagicMock(side_effect=fake_from_args) - - result = launch_router(args) - - # Verify complete flow - router_mod.from_args.assert_called_once() - mock_router_instance.start.assert_called_once() - - # Verify key parameters were propagated into RouterArgs - assert captured_args["host"] == "0.0.0.0" - assert captured_args["port"] == 30001 - assert captured_args["worker_urls"] == ["http://worker1:8000"] - assert captured_args["policy"] == PolicyType.RoundRobin - assert captured_args["service_discovery"] is True - assert captured_args["selector"] == {"app": "worker"} - assert captured_args["service_discovery_port"] == 8080 - assert captured_args["service_discovery_namespace"] == "default" - assert captured_args["dp_aware"] is True - assert captured_args["api_key"] == "test-key" - assert captured_args["log_dir"] == "/tmp/logs" - assert captured_args["log_level"] == "debug" - assert captured_args["prometheus_port"] == 29000 - assert captured_args["prometheus_host"] == "0.0.0.0" - assert captured_args["request_id_headers"] == ["x-request-id", "x-trace-id"] - assert captured_args["request_timeout_secs"] == 1200 - assert captured_args["max_concurrent_requests"] == 512 - assert captured_args["queue_size"] == 200 - assert captured_args["queue_timeout_secs"] == 120 - assert captured_args["rate_limit_tokens_per_second"] == 100 - assert captured_args["cors_allowed_origins"] == ["http://localhost:3000"] - assert captured_args["retry_max_retries"] == 3 - assert captured_args["retry_initial_backoff_ms"] == 100 - assert captured_args["retry_max_backoff_ms"] == 10000 - assert captured_args["retry_backoff_multiplier"] == 2.0 - assert captured_args["retry_jitter_factor"] == 0.1 - assert captured_args["cb_failure_threshold"] == 5 - assert captured_args["cb_success_threshold"] == 2 - assert captured_args["cb_timeout_duration_secs"] == 30 - assert captured_args["cb_window_duration_secs"] == 60 - assert captured_args["health_failure_threshold"] == 2 - assert captured_args["health_success_threshold"] == 1 - assert captured_args["health_check_timeout_secs"] == 3 - assert captured_args["health_check_interval_secs"] == 30 - assert captured_args["health_check_endpoint"] == "/healthz" diff --git a/sgl-model-gateway/bindings/python/tests/test_validation.py b/sgl-model-gateway/bindings/python/tests/test_validation.py deleted file mode 100644 index 587cd9504..000000000 --- a/sgl-model-gateway/bindings/python/tests/test_validation.py +++ /dev/null @@ -1,509 +0,0 @@ -""" -Unit tests for validation logic in sglang_router. - -These tests focus on testing the validation logic in isolation, -including parameter validation, URL validation, and configuration validation. -""" - -from unittest.mock import MagicMock, patch - -import pytest -from sglang_router.launch_router import RouterArgs, launch_router - - -class TestURLValidation: - """Test URL validation logic.""" - - def test_valid_worker_urls(self): - """Test validation of valid worker URLs.""" - valid_urls = [ - "http://worker1:8000", - "https://worker2:8000", - "http://localhost:8000", - "http://127.0.0.1:8000", - "http://192.168.1.100:8000", - "http://worker.example.com:8000", - ] - - for url in valid_urls: - args = RouterArgs(worker_urls=[url]) - # Should not raise any validation errors - assert url in args.worker_urls - - def test_valid_prefill_urls(self): - """Test validation of valid prefill URLs.""" - valid_prefill_urls = [ - ("http://prefill1:8000", 9000), - ("https://prefill2:8000", None), - ("http://localhost:8000", 9000), - ("http://127.0.0.1:8000", None), - ] - - for url, bootstrap_port in valid_prefill_urls: - args = RouterArgs(prefill_urls=[(url, bootstrap_port)]) - # Should not raise any validation errors - assert (url, bootstrap_port) in args.prefill_urls - - def test_valid_decode_urls(self): - """Test validation of valid decode URLs.""" - valid_decode_urls = [ - "http://decode1:8001", - "https://decode2:8001", - "http://localhost:8001", - "http://127.0.0.1:8001", - ] - - for url in valid_decode_urls: - args = RouterArgs(decode_urls=[url]) - # Should not raise any validation errors - assert url in args.decode_urls - - def test_malformed_urls(self): - """Test handling of malformed URLs.""" - # Note: The current implementation doesn't validate URL format - # This test documents the current behavior - malformed_urls = [ - "not-a-url", - "ftp://worker1:8000", # Wrong protocol - "http://", # Missing host - ":8000", # Missing protocol and host - "http://worker1", # Missing port - ] - - for url in malformed_urls: - args = RouterArgs(worker_urls=[url]) - # Currently, malformed URLs are accepted - # This might be something to improve in the future - assert url in args.worker_urls - - -class TestPortValidation: - """Test port validation logic.""" - - def test_valid_ports(self): - """Test validation of valid port numbers.""" - valid_ports = [1, 80, 8000, 30000, 65535] - - for port in valid_ports: - args = RouterArgs(port=port) - assert args.port == port - - def test_invalid_ports(self): - """Test handling of invalid port numbers.""" - # Note: The current implementation doesn't validate port ranges - # This test documents the current behavior - invalid_ports = [0, -1, 65536, 70000] - - for port in invalid_ports: - args = RouterArgs(port=port) - # Currently, invalid ports are accepted - # This might be something to improve in the future - assert args.port == port - - def test_bootstrap_port_validation(self): - """Test validation of bootstrap ports in PD mode.""" - valid_bootstrap_ports = [1, 80, 9000, 30000, 65535, None] - - for bootstrap_port in valid_bootstrap_ports: - args = RouterArgs(prefill_urls=[("http://prefill1:8000", bootstrap_port)]) - assert args.prefill_urls[0][1] == bootstrap_port - - -class TestParameterValidation: - """Test parameter validation logic.""" - - def test_cache_threshold_validation(self): - """Test cache threshold parameter validation.""" - # Valid cache thresholds - valid_thresholds = [0.0, 0.1, 0.5, 0.9, 1.0] - - for threshold in valid_thresholds: - args = RouterArgs(cache_threshold=threshold) - assert args.cache_threshold == threshold - - def test_balance_threshold_validation(self): - """Test load balancing threshold parameter validation.""" - # Valid absolute thresholds - valid_abs_thresholds = [0, 1, 32, 64, 128, 1000] - for threshold in valid_abs_thresholds: - args = RouterArgs(balance_abs_threshold=threshold) - assert args.balance_abs_threshold == threshold - - # Valid relative thresholds - valid_rel_thresholds = [1.0, 1.1, 1.5, 2.0, 10.0] - for threshold in valid_rel_thresholds: - args = RouterArgs(balance_rel_threshold=threshold) - assert args.balance_rel_threshold == threshold - - def test_timeout_validation(self): - """Test timeout parameter validation.""" - # Valid timeouts - valid_timeouts = [1, 30, 60, 300, 600, 1800, 3600] - - for timeout in valid_timeouts: - args = RouterArgs( - worker_startup_timeout_secs=timeout, - worker_startup_check_interval=timeout, - request_timeout_secs=timeout, - queue_timeout_secs=timeout, - ) - assert args.worker_startup_timeout_secs == timeout - assert args.worker_startup_check_interval == timeout - assert args.request_timeout_secs == timeout - assert args.queue_timeout_secs == timeout - - def test_retry_parameter_validation(self): - """Test retry parameter validation.""" - # Valid retry parameters - valid_retry_counts = [0, 1, 3, 5, 10] - for count in valid_retry_counts: - args = RouterArgs(retry_max_retries=count) - assert args.retry_max_retries == count - - # Valid backoff parameters - valid_backoff_ms = [1, 50, 100, 1000, 30000] - for backoff in valid_backoff_ms: - args = RouterArgs( - retry_initial_backoff_ms=backoff, retry_max_backoff_ms=backoff - ) - assert args.retry_initial_backoff_ms == backoff - assert args.retry_max_backoff_ms == backoff - - # Valid multiplier parameters - valid_multipliers = [1.0, 1.5, 2.0, 3.0] - for multiplier in valid_multipliers: - args = RouterArgs(retry_backoff_multiplier=multiplier) - assert args.retry_backoff_multiplier == multiplier - - # Valid jitter parameters - valid_jitter = [0.0, 0.1, 0.2, 0.5] - for jitter in valid_jitter: - args = RouterArgs(retry_jitter_factor=jitter) - assert args.retry_jitter_factor == jitter - - def test_circuit_breaker_parameter_validation(self): - """Test circuit breaker parameter validation.""" - # Valid failure thresholds - valid_failure_thresholds = [1, 3, 5, 10, 20] - for threshold in valid_failure_thresholds: - args = RouterArgs(cb_failure_threshold=threshold) - assert args.cb_failure_threshold == threshold - - # Valid success thresholds - valid_success_thresholds = [1, 2, 3, 5] - for threshold in valid_success_thresholds: - args = RouterArgs(cb_success_threshold=threshold) - assert args.cb_success_threshold == threshold - - # Valid timeout durations - valid_timeouts = [10, 30, 60, 120, 300] - for timeout in valid_timeouts: - args = RouterArgs( - cb_timeout_duration_secs=timeout, cb_window_duration_secs=timeout - ) - assert args.cb_timeout_duration_secs == timeout - assert args.cb_window_duration_secs == timeout - - def test_health_check_parameter_validation(self): - """Test health check parameter validation.""" - # Valid failure thresholds - valid_failure_thresholds = [1, 2, 3, 5, 10] - for threshold in valid_failure_thresholds: - args = RouterArgs(health_failure_threshold=threshold) - assert args.health_failure_threshold == threshold - - # Valid success thresholds - valid_success_thresholds = [1, 2, 3, 5] - for threshold in valid_success_thresholds: - args = RouterArgs(health_success_threshold=threshold) - assert args.health_success_threshold == threshold - - # Valid timeouts and intervals - valid_times = [1, 5, 10, 30, 60, 120] - for time_val in valid_times: - args = RouterArgs( - health_check_timeout_secs=time_val, health_check_interval_secs=time_val - ) - assert args.health_check_timeout_secs == time_val - assert args.health_check_interval_secs == time_val - - def test_rate_limiting_parameter_validation(self): - """Test rate limiting parameter validation.""" - # Valid concurrent request limits - valid_limits = [1, 10, 64, 256, 512, 1000] - for limit in valid_limits: - args = RouterArgs(max_concurrent_requests=limit) - assert args.max_concurrent_requests == limit - - # Valid queue sizes - valid_queue_sizes = [0, 10, 50, 100, 500, 1000] - for size in valid_queue_sizes: - args = RouterArgs(queue_size=size) - assert args.queue_size == size - - # Valid token rates - valid_rates = [1, 10, 50, 100, 500, 1000] - for rate in valid_rates: - args = RouterArgs(rate_limit_tokens_per_second=rate) - assert args.rate_limit_tokens_per_second == rate - - def test_tree_size_validation(self): - """Test tree size parameter validation.""" - # Valid tree sizes (powers of 2) - valid_sizes = [2**10, 2**20, 2**24, 2**26, 2**28, 2**30] - - for size in valid_sizes: - args = RouterArgs(max_tree_size=size) - assert args.max_tree_size == size - - def test_payload_size_validation(self): - """Test payload size parameter validation.""" - # Valid payload sizes - valid_sizes = [ - 1024, # 1KB - 1024 * 1024, # 1MB - 10 * 1024 * 1024, # 10MB - 100 * 1024 * 1024, # 100MB - 512 * 1024 * 1024, # 512MB - 1024 * 1024 * 1024, # 1GB - ] - - for size in valid_sizes: - args = RouterArgs(max_payload_size=size) - assert args.max_payload_size == size - - -class TestConfigurationValidation: - """Test configuration validation logic.""" - - def test_pd_mode_validation(self): - """Test PD mode configuration validation.""" - # Valid PD configuration - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", 9000)], - decode_urls=["http://decode1:8001"], - ) - - assert args.pd_disaggregation is True - assert len(args.prefill_urls) > 0 - assert len(args.decode_urls) > 0 - - def test_service_discovery_validation(self): - """Test service discovery configuration validation.""" - # Valid service discovery configuration - args = RouterArgs( - service_discovery=True, - selector={"app": "worker", "env": "prod"}, - service_discovery_port=8080, - service_discovery_namespace="default", - ) - - assert args.service_discovery is True - assert args.selector == {"app": "worker", "env": "prod"} - assert args.service_discovery_port == 8080 - assert args.service_discovery_namespace == "default" - - def test_pd_service_discovery_validation(self): - """Test PD service discovery configuration validation.""" - # Valid PD service discovery configuration - args = RouterArgs( - pd_disaggregation=True, - service_discovery=True, - prefill_selector={"app": "prefill"}, - decode_selector={"app": "decode"}, - ) - - assert args.pd_disaggregation is True - assert args.service_discovery is True - assert args.prefill_selector == {"app": "prefill"} - assert args.decode_selector == {"app": "decode"} - - def test_policy_validation(self): - """Test policy configuration validation.""" - # Valid policies - valid_policies = ["random", "round_robin", "cache_aware", "power_of_two"] - - for policy in valid_policies: - args = RouterArgs(policy=policy) - assert args.policy == policy - - def test_pd_policy_validation(self): - """Test PD policy configuration validation.""" - # Valid PD policies - valid_policies = ["random", "round_robin", "cache_aware", "power_of_two"] - - for prefill_policy in valid_policies: - for decode_policy in valid_policies: - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", None)], - decode_urls=["http://decode1:8001"], - prefill_policy=prefill_policy, - decode_policy=decode_policy, - ) - assert args.prefill_policy == prefill_policy - assert args.decode_policy == decode_policy - - def test_cors_validation(self): - """Test CORS configuration validation.""" - # Valid CORS origins - valid_origins = [ - [], - ["http://localhost:3000"], - ["https://example.com"], - ["http://localhost:3000", "https://example.com"], - ["*"], # Wildcard (if supported) - ] - - for origins in valid_origins: - args = RouterArgs(cors_allowed_origins=origins) - assert args.cors_allowed_origins == origins - - def test_logging_validation(self): - """Test logging configuration validation.""" - # Valid log levels - valid_log_levels = ["debug", "info", "warning", "error", "critical"] - - for level in valid_log_levels: - args = RouterArgs(log_level=level) - assert args.log_level == level - - def test_prometheus_validation(self): - """Test Prometheus configuration validation.""" - # Valid Prometheus configuration - args = RouterArgs(prometheus_port=29000, prometheus_host="127.0.0.1") - - assert args.prometheus_port == 29000 - assert args.prometheus_host == "127.0.0.1" - - def test_tokenizer_validation(self): - """Test tokenizer configuration validation.""" - # Note: model_path and tokenizer_path are not available in current RouterArgs - pytest.skip("Tokenizer configuration not available in current implementation") - - def test_request_id_headers_validation(self): - """Test request ID headers configuration validation.""" - # Valid request ID headers - valid_headers = [ - ["x-request-id"], - ["x-request-id", "x-trace-id"], - ["x-request-id", "x-trace-id", "x-correlation-id"], - ["custom-header"], - ] - - for headers in valid_headers: - args = RouterArgs(request_id_headers=headers) - assert args.request_id_headers == headers - - -class TestLaunchValidation: - """Test launch-time validation logic.""" - - def test_pd_mode_allows_empty_urls(self): - """Test that PD mode now allows empty URLs (URLs are optional).""" - # PD mode without URLs is now allowed - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=False, - ) - - # Should not raise validation error - URLs are now optional - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - # This should succeed without raising an error - launch_router(args) - router_mod.from_args.assert_called_once() - - def test_pd_mode_with_service_discovery_allows_empty_urls(self): - """Test that PD mode with service discovery allows empty URLs.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[], - decode_urls=[], - service_discovery=True, - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_regular_mode_allows_empty_worker_urls(self): - """Test that regular mode allows empty worker URLs.""" - args = RouterArgs(worker_urls=[], service_discovery=False) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_launch_with_valid_config(self): - """Test launching with valid configuration.""" - args = RouterArgs( - host="127.0.0.1", - port=30000, - worker_urls=["http://worker1:8000"], - policy="cache_aware", - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_launch_with_pd_config(self): - """Test launching with valid PD configuration.""" - args = RouterArgs( - pd_disaggregation=True, - prefill_urls=[("http://prefill1:8000", 9000)], - decode_urls=["http://decode1:8001"], - policy="cache_aware", - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() - - def test_launch_with_service_discovery_config(self): - """Test launching with valid service discovery configuration.""" - args = RouterArgs( - service_discovery=True, - selector={"app": "worker"}, - service_discovery_port=8080, - ) - - # Should not raise validation error - with patch("sglang_router.launch_router.Router") as router_mod: - mock_router_instance = MagicMock() - router_mod.from_args = MagicMock(return_value=mock_router_instance) - - launch_router(args) - - # Should create router instance via from_args - router_mod.from_args.assert_called_once() diff --git a/test/README.md b/test/README.md index ddc2c358e..1656837cd 100644 --- a/test/README.md +++ b/test/README.md @@ -72,9 +72,28 @@ Parameters: `est_time` (seconds), `stage` + `runner_config` (target stage and ru Keep `est_time`, `stage`, `runner_config` as **literal values** — `run_suite.py` collects them by AST parsing. -JIT kernel correctness tests and benchmarks live under `test/registered/jit/`, same as other registered tests (their helpers stay alongside the kernel source under `python/sglang/kernels/jit/` and are imported by absolute path): -- Correctness tests: `test/registered/jit/test_*.py` → `base-b-kernel-unit-test-1-gpu-large` -- Benchmarks: `test/registered/jit/benchmark/bench_*.py` → `base-b-kernel-benchmark-test-1-gpu-large` +New and renamed tests use this layout: + +```text +test/registered///test_*.py +``` + +`` is one of `unit`, `kernel`, `e2e`, `accuracy`, `perf`, or `stress`. +Hardware is expressed by one or more `register_*_ci` calls, never by creating a +new top-level hardware directory. The admission checker applies the layout and +kind/suite contract incrementally while legacy paths are migrated. + +Diffusion workflows also enter through `test/run_suite.py`; registered bridge +files preserve their case-level pytest partitioning until the remaining +diffusion cases are moved out of the package test-support tree. + +New JIT kernel correctness tests and benchmarks live under +`test/registered/kernel/jit/`; legacy `test/registered/jit/` files are migrated +incrementally. Helpers stay alongside the kernel source under +`python/sglang/kernels/jit/` and are imported by absolute path: + +- Correctness tests: `test/registered/kernel/jit/test_*.py` → `base-b-kernel-unit-test-1-gpu-large` +- Benchmarks: `test/registered/kernel/jit/benchmark/bench_*.py` → `base-b-kernel-benchmark-test-1-gpu-large` ## Choosing a Suite @@ -94,6 +113,20 @@ Use the lightest suite that meets your test's needs. Full suite tables are in th See the [write-sglang-test skill](../.claude/skills/write-sglang-test/SKILL.md) for templates, fixtures, model selection, and a complete checklist. +Before adding a registered test, identify the production change that would make +it fail. Prefer extending an existing fixture/server launch over adding another +file. The incremental admission check applies these ratchets to new or modified +registered tests: + +- Temporary `disabled=` registrations and unconditional skips must reference an + issue and include `until YYYY-MM-DD`; expired entries fail lint. +- A file registered on CUDA plus another accelerator must place a nearby + `backend-specific:` comment above the extra registration and name the path or + failure mode that only that backend can catch. +- Default PR registrations are limited to 1,200 estimated weighted accelerator-seconds + per backend (`est_time * GPU count`). Move larger matrices to extra/nightly, + or document a nearby `ci-cost-override:` rationale. + ## Multi-Hardware Backends This README mostly describes the NVIDIA GPU CI pipeline. Other hardware backends (AMD, NPU) follow the same practices and use the multi-backend registry system. A scheduled job summarizes test coverage across all backends; [here is an example run](https://github.com/sgl-project/sglang/actions/runs/23424304300). @@ -111,4 +144,4 @@ This README mostly describes the NVIDIA GPU CI pipeline. Other hardware backends ### Adding New Models to Nightly CI - **Text models**: Extend the [global model list variables](https://github.com/sgl-project/sglang/blob/85c1f7937781199203b38bb46325a2840f353a04/python/sglang/test/test_utils.py#L104) in `test_utils.py`. -- **VLMs**: Extend the `MODEL_THRESHOLDS` dictionary in `test/registered/eval/test_vlms_mmmu_eval.py`. +- **VLMs**: Extend the `MODEL_THRESHOLDS` dictionary in `test/registered/accuracy/models/test_vlms_mmmu_eval.py`. diff --git a/test/registered/eval/test_text_models_gsm8k_eval.py b/test/registered/accuracy/models/test_text_models_gsm8k_eval.py similarity index 100% rename from test/registered/eval/test_text_models_gsm8k_eval.py rename to test/registered/accuracy/models/test_text_models_gsm8k_eval.py diff --git a/test/registered/eval/test_vlms_mmmu_eval.py b/test/registered/accuracy/models/test_vlms_mmmu_eval.py similarity index 100% rename from test/registered/eval/test_vlms_mmmu_eval.py rename to test/registered/accuracy/models/test_vlms_mmmu_eval.py diff --git a/test/registered/amd/accuracy/mi30x/test_grok_eval_amd.py b/test/registered/amd/accuracy/mi30x/test_grok_eval_amd.py deleted file mode 100644 index 89ed2ae66..000000000 --- a/test/registered/amd/accuracy/mi30x/test_grok_eval_amd.py +++ /dev/null @@ -1,290 +0,0 @@ -"""AMD GROK GSM8K Completion Evaluation Test (8-GPU) - -Tests GROK models (Grok-1 FP8, Grok-1 INT4, Grok-2) using -few-shot completion benchmark on MI300X. - -Registry: nightly-amd-8-gpu-grok suite -""" - -import ast -import os -import re -import time -import unittest -from dataclasses import dataclass -from typing import List, Optional, Tuple - -import numpy as np - -from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci -from sglang.test.test_utils import ( - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - DEFAULT_URL_FOR_TEST, - is_in_ci, - popen_launch_server, - write_github_step_summary, -) -from sglang.utils import download_and_cache_file, read_jsonl - -# DISABLED: Split into individual files for each model variant -# See: test_grok1_fp8_eval_amd.py, test_grok1_int4_eval_amd.py, test_grok2_eval_amd.py -register_amd_ci( - est_time=2700, - suite="nightly-amd-8-gpu-grok", - nightly=True, - disabled="Split into test_grok1_fp8_eval_amd.py, test_grok1_int4_eval_amd.py, test_grok2_eval_amd.py", -) - -INVALID = -9999999 - - -@dataclass -class ModelConfig: - """Configuration for a model to test.""" - - model_path: str - tp_size: int = 8 - accuracy_threshold: float = 0.50 - other_args: Optional[List[str]] = None - env_vars: Optional[dict] = None - tokenizer_path: Optional[str] = None - timeout: Optional[int] = None - - def __post_init__(self): - if self.other_args is None: - self.other_args = [] - if self.env_vars is None: - self.env_vars = {} - - -# GROK models for MI300X -GROK_MODELS = [ - # GROK1-FP8 - ModelConfig( - model_path="lmzheng/grok-1", - tp_size=8, - accuracy_threshold=0.80, - timeout=3600, - tokenizer_path="Xenova/grok-1-tokenizer", - other_args=[ - "--quantization", - "fp8", - "--attention-backend", - "aiter", - "--mem-fraction-static", - "0.85", - "--trust-remote-code", - ], - env_vars={ - "RCCL_MSCCL_ENABLE": "0", - "SGLANG_USE_AITER": "1", - "SGLANG_INT4_WEIGHT": "0", - }, - ), - # GROK1-INT4 - ModelConfig( - model_path="amd/grok-1-W4A8KV8", - tp_size=8, - accuracy_threshold=0.80, - timeout=3600, - tokenizer_path="Xenova/grok-1-tokenizer", - other_args=[ - "--quantization", - "fp8", - "--attention-backend", - "aiter", - "--mem-fraction-static", - "0.85", - "--trust-remote-code", - ], - env_vars={ - "RCCL_MSCCL_ENABLE": "0", - "SGLANG_USE_AITER": "1", - "SGLANG_INT4_WEIGHT": "1", - }, - ), - # GROK2 - ModelConfig( - model_path="xai-org/grok-2", - tp_size=8, - accuracy_threshold=0.915, - timeout=3600, - tokenizer_path="alvarobartt/grok-2-tokenizer", - other_args=[ - "--quantization", - "fp8", - "--attention-backend", - "aiter", - "--mem-fraction-static", - "0.85", - "--trust-remote-code", - ], - env_vars={ - "RCCL_MSCCL_ENABLE": "0", - "SGLANG_USE_AITER": "1", - "SGLANG_INT4_WEIGHT": "0", - }, - ), -] - - -def get_one_example(lines, i, include_answer): - """Format a single GSM8K example.""" - ret = "Question: " + lines[i]["question"] + "\nAnswer:" - if include_answer: - ret += " " + lines[i]["answer"] - return ret - - -def get_few_shot_examples(lines, k): - """Get k few-shot examples for prompting.""" - ret = "" - for i in range(k): - ret += get_one_example(lines, i, True) + "\n\n" - return ret - - -def get_answer_value(answer_str): - """Extract numerical answer from response.""" - answer_str = answer_str.replace(",", "") - numbers = re.findall(r"\d+", answer_str) - if len(numbers) < 1: - return INVALID - try: - return ast.literal_eval(numbers[-1]) - except SyntaxError: - return INVALID - - -def run_gsm8k_benchmark( - base_url: str, - num_questions: int = 200, - num_shots: int = 5, - parallel: int = 64, -) -> Tuple[float, float, float]: - """Run GSM8K few-shot completion benchmark.""" - import sglang as sgl - from sglang.lang.backend.runtime_endpoint import RuntimeEndpoint - - url = "https://raw.githubusercontent.com/openai/grade-school-math/master/grade_school_math/data/test.jsonl" - data_path = download_and_cache_file(url) - lines = list(read_jsonl(data_path)) - - few_shot_examples = get_few_shot_examples(lines, num_shots) - - questions = [] - labels = [] - for i in range(len(lines[:num_questions])): - questions.append(get_one_example(lines, i, False)) - labels.append(get_answer_value(lines[i]["answer"])) - assert all(l != INVALID for l in labels) - arguments = [{"question": q} for q in questions] - - @sgl.function - def few_shot_gsm8k(s, question): - s += few_shot_examples + question - s += sgl.gen( - "answer", max_tokens=512, stop=["Question", "Assistant:", "<|separator|>"] - ) - - backend = RuntimeEndpoint(base_url) - sgl.set_default_backend(backend) - - tic = time.perf_counter() - states = few_shot_gsm8k.run_batch( - arguments, temperature=0, num_threads=parallel, progress_bar=True - ) - latency = time.perf_counter() - tic - - preds = [get_answer_value(states[i]["answer"]) for i in range(len(states))] - acc = np.mean(np.array(preds) == np.array(labels)) - invalid = np.mean(np.array(preds) == INVALID) - - return float(acc), float(invalid), float(latency) - - -class TestGrokEvalAMD(unittest.TestCase): - """GROK GSM8K Completion Evaluation Test for AMD MI300X.""" - - @classmethod - def setUpClass(cls): - cls.models = GROK_MODELS - cls.base_url = DEFAULT_URL_FOR_TEST - cls.num_questions = int(os.environ.get("GSM8K_NUM_QUESTIONS", "200")) - - def test_grok_accuracy(self): - """Test GROK models with GSM8K completion benchmark.""" - all_results = [] - summary = "### GROK Models (MI300X)\n\n" - summary += "| Model | TP | Accuracy | Threshold | Status |\n" - summary += "| ----- | -- | -------- | --------- | ------ |\n" - - for config in self.models: - with self.subTest(model=config.model_path): - print(f"\n{'=' * 60}") - print(f"Testing: {config.model_path}") - print(f"{'=' * 60}") - - env = os.environ.copy() - for key, value in config.env_vars.items(): - env[key] = value - - other_args = list(config.other_args) - other_args.extend(["--tp", str(config.tp_size)]) - if config.tokenizer_path: - other_args.extend(["--tokenizer-path", config.tokenizer_path]) - timeout = config.timeout or DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH - - try: - process = popen_launch_server( - model=config.model_path, - base_url=self.base_url, - timeout=timeout, - other_args=other_args, - env=env, - ) - - try: - acc, invalid, latency = run_gsm8k_benchmark( - self.base_url, num_questions=self.num_questions - ) - passed = acc >= config.accuracy_threshold - status = "✅ PASS" if passed else "❌ FAIL" - print( - f" accuracy={acc:.3f} threshold={config.accuracy_threshold} {status}" - ) - - all_results.append( - { - "model": config.model_path, - "accuracy": acc, - "passed": passed, - } - ) - summary += f"| {config.model_path} | {config.tp_size} | {acc:.3f} | {config.accuracy_threshold} | {status} |\n" - - finally: - kill_process_tree(process.pid) - - except Exception as e: - summary += f"| {config.model_path} | {config.tp_size} | N/A | {config.accuracy_threshold} | ❌ ERROR |\n" - all_results.append( - { - "model": config.model_path, - "accuracy": None, - "passed": False, - "error": str(e), - } - ) - - if is_in_ci(): - write_github_step_summary(summary) - - failed = [r for r in all_results if not r["passed"]] - if failed: - raise AssertionError(f"Failed models: {[r['model'] for r in failed]}") - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/amd/accuracy/mi35x/test_glm52_fp8_eval_mi35x.py b/test/registered/amd/accuracy/mi35x/test_glm52_fp8_eval_mi35x.py index 4a0ccd0b5..ebe5181bd 100644 --- a/test/registered/amd/accuracy/mi35x/test_glm52_fp8_eval_mi35x.py +++ b/test/registered/amd/accuracy/mi35x/test_glm52_fp8_eval_mi35x.py @@ -23,7 +23,7 @@ which is below the >=3.5.0 the aiter gluon DSA kernels need, so it logs what loses the accuracy. Re-add a 7.0 job once its image ships Triton >=3.5.0, or once the legacy DSA fallback is fixed on gfx950. -The eval matches the CUDA GLM-5.2-FP8 nightly (`test/registered/8-gpu-models/ +The eval matches the CUDA GLM-5.2-FP8 nightly (`test/registered/e2e/models_large/ test_glm52_fp8.py`): same dataset and same 0.92 baseline, so a red run here means AMD diverged from CUDA rather than the harness diverging. diff --git a/test/registered/core/test_engine_child_pids.py b/test/registered/core/test_engine_child_pids.py index 288a2fea9..5b38585c5 100644 --- a/test/registered/core/test_engine_child_pids.py +++ b/test/registered/core/test_engine_child_pids.py @@ -14,14 +14,13 @@ import unittest import psutil import sglang as sgl -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, CustomTestCase, ) -register_cuda_ci(est_time=38, stage="base-b", runner_config="1-gpu-small") -register_amd_ci(est_time=77, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=77, stage="base-b", runner_config="1-gpu-small") class TestEngineChildPids(CustomTestCase): diff --git a/test/registered/core/test_request_queue_validation.py b/test/registered/core/test_request_queue_validation.py index 65fa4bc91..6bd35a636 100644 --- a/test/registered/core/test_request_queue_validation.py +++ b/test/registered/core/test_request_queue_validation.py @@ -4,7 +4,7 @@ import re import unittest from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -17,8 +17,7 @@ from sglang.test.test_utils import ( send_generate_requests, ) -register_cuda_ci(est_time=65, stage="base-b", runner_config="1-gpu-small") -register_amd_ci(est_time=70, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=53, stage="base-b", runner_config="1-gpu-small") class TestMaxQueuedRequests(CustomTestCase): diff --git a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py index d30d00678..33e676a5c 100644 --- a/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py +++ b/test/registered/cp/test_deepseek_v4_flash_fp4_b200_cp.py @@ -2,7 +2,7 @@ Balanced recipe (TP=4, DeepEP, EAGLE) plus --attn-cp-size=4 with the DSA prefill-CP interleave strategy. Split out of -models_e2e/test_deepseek_v4_flash_fp4_b200.py so the `cp` group covers +e2e/models/test_deepseek_v4_flash_fp4_b200.py so the `cp` group covers all context-parallel tests. Registry: extra-b-test-4-gpu-b200 (label-gated extra CI, 4x B200) diff --git a/test/registered/cpu/test_spec_eagle_cpu.py b/test/registered/cpu/test_spec_eagle_cpu.py deleted file mode 100644 index b0e70b3e9..000000000 --- a/test/registered/cpu/test_spec_eagle_cpu.py +++ /dev/null @@ -1,71 +0,0 @@ -"""EAGLE spec-decoding core on CPU: the standard config (topk=1, page_size=1) -on the synchronous (non-overlap) path. topk > 1 tree drafting is covered in -test_spec_eagle_topk_cpu.py (split to stay under the per-file CI timeout). -""" - -import unittest - -from sglang.srt.environ import envs -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.kits.matched_stop_kit import MatchedStopMixin -from sglang.test.kits.spec_server_kits import ( - SpecAccuracyKit, - SpecCorrectnessKit, - SpecFeatureKit, - SpecLogprobKit, - SpecPenaltyKit, -) -from sglang.test.server_fixtures.spec_eagle_fixture import EagleLlama2Base - -# Measured 780s all-green on a 40-core GNR socket (1 launch + 18 methods). -register_cpu_ci( - est_time=800, - suite="stage-a-test-cpu-intel", - disabled="EagleLlama2Base needs gated meta-llama/Llama-2-7b-chat-hf", -) - -_KITS = ( - SpecCorrectnessKit, - SpecAccuracyKit, - SpecLogprobKit, - SpecPenaltyKit, - SpecFeatureKit, - MatchedStopMixin, -) - - -class _Core(EagleLlama2Base): - """EAGLE (Llama-2) preset on CPU.""" - - attention_backend = "intel_amx" - disable_overlap = True - mem_fraction_static = 0.3 - # CPU decode is compute-bound; a wider batch buys nothing here. - max_running_requests = 8 - gsm8k_num_examples = 64 - env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) - - -class TestEagleLlama2NoOverlap(_Core, *_KITS): - """Spec v1 (overlap scheduler off) -- the only mode reachable on CPU.""" - - # Standard chain config (topk=1, page_size=1), same shape as the CUDA core. - spec_steps = 5 - spec_topk = 1 - spec_tokens = 6 - # EAGLE/Llama-2 topk=1 accepts modestly; tune against CI if needed. - acc_length_thres = 1.6 - batch_accept_len_thres = 1.3 - gsm8k_accept_len_thres = 1.3 - - @unittest.skip( - "constrained decoding on CPU needs a vocab-mask CPU branch in the " - "xgrammar backend (upstream gap, not spec-specific); the other grammar " - "backends lack the rollback spec verification requires" - ) - def test_constrained_decoding(self): - pass - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/cpu/test_spec_eagle_parity_cpu.py b/test/registered/cpu/test_spec_eagle_parity_cpu.py deleted file mode 100644 index 1f446f2fa..000000000 --- a/test/registered/cpu/test_spec_eagle_parity_cpu.py +++ /dev/null @@ -1,29 +0,0 @@ -import unittest - -from sglang.srt.environ import envs -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.kits.spec_server_kits import SpecParityKit -from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base - -# Estimated: 2 sequential 8B server launches + one 4-prompt greedy method -# (CUDA sibling: 360); tune from CI TIMINGS once it has run there. -register_cpu_ci( - est_time=480, - suite="stage-a-test-cpu-intel", - disabled="EAGLE3 numerical parity mismatches on CPU intel_amx", -) - - -class TestEagle3ParityCPU(SpecParityKit, Eagle3Base): - """EAGLE3 spec (intel_amx) greedy output == non-spec reference.""" - - attention_backend = "intel_amx" - disable_overlap = True - mem_fraction_static = 0.3 - # CPU decode is compute-bound; a wider batch buys nothing here. - max_running_requests = 8 - env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/cpu/test_spec_eagle_topk_cpu.py b/test/registered/cpu/test_spec_eagle_topk_cpu.py deleted file mode 100644 index ecd02af4b..000000000 --- a/test/registered/cpu/test_spec_eagle_topk_cpu.py +++ /dev/null @@ -1,68 +0,0 @@ -"""EAGLE topk > 1 tree drafting on CPU (Llama-2 topk=4, synchronous path). - -Split from test_spec_eagle_cpu.py, mirroring the CUDA test_spec_eagle.py / -test_spec_eagle_topk.py layout, so each file stays under the per-file CI -timeout. -""" - -import unittest - -from sglang.srt.environ import envs -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.kits.spec_server_kits import ( - SpecAccuracyKit, - SpecCorrectnessKit, - SpecFeatureKit, - SpecLogprobKit, - SpecPenaltyKit, -) -from sglang.test.server_fixtures.spec_eagle_fixture import EagleLlama2Base - -# Measured 830s all-green on a 40-core GNR socket (1 launch + 14 methods). -register_cpu_ci( - est_time=850, - suite="stage-a-test-cpu-intel", - disabled="EagleLlama2Base needs gated meta-llama/Llama-2-7b-chat-hf", -) - - -class _Core(EagleLlama2Base): - """EAGLE (Llama-2) preset on CPU.""" - - attention_backend = "intel_amx" - disable_overlap = True - mem_fraction_static = 0.3 - # CPU decode is compute-bound; a wider batch buys nothing here. - max_running_requests = 8 - gsm8k_num_examples = 64 - env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),) - - -class TestEagleLlama2Topk4( - _Core, - SpecCorrectnessKit, - SpecAccuracyKit, - SpecLogprobKit, - SpecPenaltyKit, - SpecFeatureKit, -): - """EAGLE/Llama-2 topk=4 tree coverage (kits listed in bases).""" - - spec_steps = 3 - spec_topk = 4 - spec_tokens = 8 - acc_length_thres = 2.4 - batch_accept_len_thres = 1.6 - gsm8k_accept_len_thres = 2.0 - - @unittest.skip( - "constrained decoding on CPU needs a vocab-mask CPU branch in the " - "xgrammar backend (upstream gap, not spec-specific); the other grammar " - "backends lack the rollback spec verification requires" - ) - def test_constrained_decoding(self): - pass - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py b/test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py index 0bc7bec89..01c2cdfbf 100644 --- a/test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py +++ b/test/registered/debug_utils/comparator/aligner/unsharder/test_executor.py @@ -802,269 +802,6 @@ class TestReduceSum: without_dim_names(unsharder_result.tensors[0]), without_dim_names(expected) ) - def test_recompute_pseudo_mismatch(self) -> None: - """_verify_replicated_group returns failed check for RECOMPUTE_PSEUDO axis mismatch.""" - tensor_a = torch.ones(4) - tensor_b = torch.ones(4) + 0.1 - - checks: list[ReplicatedCheckResult] = _verify_replicated_group( - [tensor_a, tensor_b], - axis=ParallelAxis.RECOMPUTE_PSEUDO, - group_index=0, - ) - assert len(checks) == 1 - assert checks[0].axis == "recompute_pseudo" - assert checks[0].group_index == 0 - assert checks[0].compared_index == 1 - assert checks[0].baseline_index == 0 - assert not checks[0].passed - assert checks[0].diff.max_abs_diff == pytest.approx(0.1, abs=1e-5) - - -class TestThdCpConcat: - def test_single_seq(self) -> None: - """Single seq THD unshard: 2 ranks → per-seq concat.""" - rank0 = apply_dim_names(torch.tensor([1, 2, 3]), ["t"]) - rank1 = apply_dim_names(torch.tensor([4, 5, 6]), ["t"]) - - plan = UnsharderPlan( - axis=ParallelAxis.CP, - params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3]), - groups=[[0, 1]], - ) - unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1]) - - assert len(unsharder_result.tensors) == 1 - expected = torch.tensor([1, 2, 3, 4, 5, 6]) - assert torch.equal(without_dim_names(unsharder_result.tensors[0]), expected) - - def test_multi_seq(self) -> None: - """Multi-seq THD unshard: 2 ranks, seq_lens=[50, 32, 46].""" - # rank0: [seqA_r0(50) | seqB_r0(32) | pad_r0(46)] - # rank1: [seqA_r1(50) | seqB_r1(32) | pad_r1(46)] - seq_a_r0 = torch.arange(0, 50) - seq_b_r0 = torch.arange(100, 132) - pad_r0 = torch.full((46,), -1) - rank0 = apply_dim_names(torch.cat([seq_a_r0, seq_b_r0, pad_r0]), ["t"]) - - seq_a_r1 = torch.arange(50, 100) - seq_b_r1 = torch.arange(132, 164) - pad_r1 = torch.full((46,), -2) - rank1 = apply_dim_names(torch.cat([seq_a_r1, seq_b_r1, pad_r1]), ["t"]) - - plan = UnsharderPlan( - axis=ParallelAxis.CP, - params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[50, 32, 46]), - groups=[[0, 1]], - ) - unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1]) - - assert len(unsharder_result.tensors) == 1 - unsharded: torch.Tensor = without_dim_names(unsharder_result.tensors[0]) - - # seqA: r0(50) + r1(50) = 100 tokens, values 0..99 - assert torch.equal(unsharded[:100], torch.cat([seq_a_r0, seq_a_r1])) - # seqB: r0(32) + r1(32) = 64 tokens - assert torch.equal(unsharded[100:164], torch.cat([seq_b_r0, seq_b_r1])) - # pad: r0(46) + r1(46) = 92 tokens - assert torch.equal(unsharded[164:256], torch.cat([pad_r0, pad_r1])) - - def test_with_hidden_dim(self) -> None: - """THD unshard with trailing hidden dim: shape [T, H].""" - torch.manual_seed(42) - hidden: int = 4 - # rank0: [seqA_r0(3, 4) | seqB_r0(2, 4)] - # rank1: [seqA_r1(3, 4) | seqB_r1(2, 4)] - seq_a_r0 = torch.randn(3, hidden) - seq_b_r0 = torch.randn(2, hidden) - rank0 = apply_dim_names(torch.cat([seq_a_r0, seq_b_r0]), ["t", "h"]) - - seq_a_r1 = torch.randn(3, hidden) - seq_b_r1 = torch.randn(2, hidden) - rank1 = apply_dim_names(torch.cat([seq_a_r1, seq_b_r1]), ["t", "h"]) - - plan = UnsharderPlan( - axis=ParallelAxis.CP, - params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]), - groups=[[0, 1]], - ) - unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1]) - - assert len(unsharder_result.tensors) == 1 - unsharded: torch.Tensor = without_dim_names(unsharder_result.tensors[0]) - - assert unsharded.shape == (10, hidden) - assert torch.equal(unsharded[:6], torch.cat([seq_a_r0, seq_a_r1])) - assert torch.equal(unsharded[6:10], torch.cat([seq_b_r0, seq_b_r1])) - - def test_with_leading_batch_dim(self) -> None: - """THD unshard with leading batch dim: shape [B, T, H], t is dim=1.""" - torch.manual_seed(42) - batch: int = 2 - hidden: int = 4 - # rank0: [seqA_r0(3) | seqB_r0(2)] per batch item - # rank1: [seqA_r1(3) | seqB_r1(2)] per batch item - seq_a_r0 = torch.randn(batch, 3, hidden) - seq_b_r0 = torch.randn(batch, 2, hidden) - rank0 = apply_dim_names(torch.cat([seq_a_r0, seq_b_r0], dim=1), ["b", "t", "h"]) - - seq_a_r1 = torch.randn(batch, 3, hidden) - seq_b_r1 = torch.randn(batch, 2, hidden) - rank1 = apply_dim_names(torch.cat([seq_a_r1, seq_b_r1], dim=1), ["b", "t", "h"]) - - plan = UnsharderPlan( - axis=ParallelAxis.CP, - params=CpThdConcatParams(dim_name="t", seq_lens_per_rank=[3, 2]), - groups=[[0, 1]], - ) - unsharder_result: UnsharderResult = execute_unsharder_plan(plan, [rank0, rank1]) - - assert len(unsharder_result.tensors) == 1 - unsharded: torch.Tensor = without_dim_names(unsharder_result.tensors[0]) - - assert unsharded.shape == (batch, 10, hidden) - # seqA: r0(3) + r1(3) = 6 tokens per batch - assert torch.equal(unsharded[:, :6, :], torch.cat([seq_a_r0, seq_a_r1], dim=1)) - # seqB: r0(2) + r1(2) = 4 tokens per batch - assert torch.equal( - unsharded[:, 6:10, :], torch.cat([seq_b_r0, seq_b_r1], dim=1) - ) - - -class TestReduceSum: - def test_basic_tp2_reduce(self) -> None: - """2 partial tensors sum to full tensor.""" - torch.manual_seed(42) - full_tensor = torch.randn(4, 8) - part_a = full_tensor * 0.6 - part_b = full_tensor * 0.4 - - dim_specs = parse_dims("h[tp:partial] d").dims - parallel_infos = [ - {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=2)} for i in range(2) - ] - plans = compute_unsharder_plan(dim_specs, parallel_infos) - assert len(plans) == 1 - assert isinstance(plans[0].params, ReduceSumParams) - - named_parts: list[torch.Tensor] = _name_tensors([part_a, part_b], dim_specs) - unsharder_result: UnsharderResult = execute_unsharder_plan( - plans[0], named_parts - ) - - assert len(unsharder_result.tensors) == 1 - assert torch.allclose( - without_dim_names(unsharder_result.tensors[0]), full_tensor - ) - - def test_tp4_reduce(self) -> None: - """4 partial tensors sum to full tensor.""" - torch.manual_seed(42) - full_tensor = torch.randn(4, 8) - parts: list[torch.Tensor] = [full_tensor * 0.25 for _ in range(4)] - - dim_specs = parse_dims("h[tp:partial] d").dims - parallel_infos = [ - {ParallelAxis.TP: AxisInfo(axis_rank=i, axis_size=4)} for i in range(4) - ] - plans = compute_unsharder_plan(dim_specs, parallel_infos) - assert len(plans) == 1 - - named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs) - unsharder_result: UnsharderResult = execute_unsharder_plan( - plans[0], named_parts - ) - - assert len(unsharder_result.tensors) == 1 - assert torch.allclose( - without_dim_names(unsharder_result.tensors[0]), full_tensor - ) - - def test_multi_axis_concat_then_reduce(self) -> None: - """CP concat + TP reduce end-to-end.""" - torch.manual_seed(42) - full_tensor = torch.randn(4, 8, 16) - - cp_chunks = list(full_tensor.chunk(2, dim=1)) - # Each CP chunk is held as partial sums across TP ranks - tensors: list[torch.Tensor] = [] - parallel_infos: list[dict[ParallelAxis, AxisInfo]] = [] - for cp_rank in range(2): - for tp_rank in range(2): - tensors.append(cp_chunks[cp_rank] * 0.5) - parallel_infos.append( - { - ParallelAxis.CP: AxisInfo(axis_rank=cp_rank, axis_size=2), - ParallelAxis.TP: AxisInfo(axis_rank=tp_rank, axis_size=2), - } - ) - - dim_specs = parse_dims("b s[cp] h[tp:partial]").dims - plans = compute_unsharder_plan(dim_specs, parallel_infos) - assert len(plans) == 2 - - current: list[torch.Tensor] = _name_tensors(tensors, dim_specs) - for plan in plans: - unsharder_result: UnsharderResult = execute_unsharder_plan(plan, current) - current = unsharder_result.tensors - - assert len(current) == 1 - assert torch.allclose(without_dim_names(current[0]), full_tensor) - - def test_reduce_scrambled_ranks(self) -> None: - """Scrambled rank order — sum is commutative so result is the same.""" - torch.manual_seed(42) - full_tensor = torch.randn(4, 8) - parts: list[torch.Tensor] = [ - full_tensor * 0.1, - full_tensor * 0.2, - full_tensor * 0.3, - full_tensor * 0.4, - ] - - parallel_infos = [ - {ParallelAxis.TP: AxisInfo(axis_rank=2, axis_size=4)}, - {ParallelAxis.TP: AxisInfo(axis_rank=0, axis_size=4)}, - {ParallelAxis.TP: AxisInfo(axis_rank=3, axis_size=4)}, - {ParallelAxis.TP: AxisInfo(axis_rank=1, axis_size=4)}, - ] - dim_specs = parse_dims("h[tp:partial] d").dims - plans = compute_unsharder_plan(dim_specs, parallel_infos) - - named_parts: list[torch.Tensor] = _name_tensors(parts, dim_specs) - unsharder_result: UnsharderResult = execute_unsharder_plan( - plans[0], named_parts - ) - - assert len(unsharder_result.tensors) == 1 - assert torch.allclose( - without_dim_names(unsharder_result.tensors[0]), full_tensor - ) - - def test_reduce_preserves_named_dims(self) -> None: - """Named tensor dimensions are preserved through reduce_sum.""" - dim_specs = parse_dims("h[tp:partial] d").dims - part_a = apply_dim_names(torch.randn(4, 8), ["h", "d"]) - part_b = apply_dim_names(torch.randn(4, 8), ["h", "d"]) - - plan = UnsharderPlan( - axis=ParallelAxis.TP, - params=ReduceSumParams(), - groups=[[0, 1]], - ) - unsharder_result: UnsharderResult = execute_unsharder_plan( - plan, [part_a, part_b] - ) - - assert len(unsharder_result.tensors) == 1 - assert get_dim_names(unsharder_result.tensors[0]) == ("h", "d") - expected = apply_dim_names( - without_dim_names(part_a) + without_dim_names(part_b), ["h", "d"] - ) - assert torch.allclose( - without_dim_names(unsharder_result.tensors[0]), without_dim_names(expected) - ) - class TestFusedDimExecutor: def test_fused_tp2_concat(self) -> None: diff --git a/test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py b/test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py index 9b3b6b0af..d9f48ebb5 100644 --- a/test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py +++ b/test/registered/debug_utils/comparator/tensor_comparator/test_comparator.py @@ -19,79 +19,6 @@ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=20, stage="weekly", runner_config="cpu") -class TestComputeTensorInfo: - def test_basic_tensor_returns_correct_shape_and_dtype(self) -> None: - tensor = torch.randn(2, 3) - info = compute_tensor_info(tensor) - assert info.shape == [2, 3] - assert info.dtype == "torch.float32" - assert info.stats.mean == pytest.approx(tensor.float().mean().item(), abs=1e-4) - - def test_include_sample_false_returns_none_sample(self) -> None: - tensor = torch.randn(2, 3) - info = compute_tensor_info(tensor, include_sample=False) - assert info.sample is None - - def test_include_sample_true_returns_string_sample(self) -> None: - tensor = torch.randn(2, 3) - info = compute_tensor_info(tensor, include_sample=True) - assert info.sample is not None - assert isinstance(info.sample, str) - - def test_empty_tensor_stats_are_zero(self) -> None: - tensor = torch.tensor([]) - info = compute_tensor_info(tensor) - assert info.stats.mean == 0.0 - assert info.stats.std == 0.0 - assert info.shape == [0] - - def test_integer_tensor_converted_to_float_for_stats(self) -> None: - """Integer tensors should be cast to float internally for stats computation.""" - tensor = torch.tensor([1, 2, 3, 4], dtype=torch.int32) - info = compute_tensor_info(tensor) - assert info.dtype == "torch.int32" - assert info.stats.mean == pytest.approx(2.5, abs=1e-4) - assert info.stats.min == pytest.approx(1.0, abs=1e-4) - assert info.stats.max == pytest.approx(4.0, abs=1e-4) - - def test_bfloat16_tensor_shape_and_stats(self) -> None: - """bfloat16 tensors produce correct shape and dtype string.""" - tensor = torch.ones(3, 4, dtype=torch.bfloat16) - info = compute_tensor_info(tensor) - assert info.shape == [3, 4] - assert info.dtype == "torch.bfloat16" - assert info.stats.mean == pytest.approx(1.0, abs=1e-2) - - def test_multidimensional_shape(self) -> None: - """Shape is preserved for high-rank tensors.""" - tensor = torch.randn(2, 3, 4, 5) - info = compute_tensor_info(tensor) - assert info.shape == [2, 3, 4, 5] - - def test_scalar_tensor(self) -> None: - """Scalar (0-dim) tensor produces empty shape list.""" - tensor = torch.tensor(3.14) - info = compute_tensor_info(tensor) - assert info.shape == [] - assert info.stats.mean == pytest.approx(3.14, abs=1e-4) - assert info.stats.min == pytest.approx(3.14, abs=1e-4) - assert info.stats.max == pytest.approx(3.14, abs=1e-4) - - def test_include_sample_true_contains_tensor_representation(self) -> None: - """Sample string should contain some recognizable tensor content.""" - tensor = torch.tensor([1.0, 2.0]) - info = compute_tensor_info(tensor, include_sample=True) - assert info.sample is not None - assert "1." in info.sample or "2." in info.sample - - def test_percentiles_present_for_small_tensor(self) -> None: - """Small tensors (< threshold) should have percentile data.""" - tensor = torch.randn(100) - info = compute_tensor_info(tensor) - assert len(info.stats.percentiles) > 0 - assert 50 in info.stats.percentiles - - class TestComputeTensorInfo: def test_basic_tensor_returns_correct_shape_and_dtype(self) -> None: tensor = torch.randn(2, 3) diff --git a/test/registered/debug_utils/test_dumper.py b/test/registered/debug_utils/test_dumper.py index 639c12053..145891a78 100644 --- a/test/registered/debug_utils/test_dumper.py +++ b/test/registered/debug_utils/test_dumper.py @@ -389,47 +389,6 @@ class TestTorchSave: assert "skip the tensor" in captured.out -class TestLog: - def test_log_format(self): - with _capture_stdout() as captured: - _log("hello") - out = captured.getvalue() - assert "hello" in out, out - assert "[Dumper, rank=" in out, out - assert ", t=" in out, out - - -class TestCompareTensorsQuick: - def test_identical(self): - a = torch.tensor([1.0, 2.0, 3.0]) - s = _compare_tensors_quick(a, a.clone()) - assert "rel_diff=0" in s, s - assert "max_abs=0" in s, s - - def test_diverged(self): - a = torch.tensor([1.0, 2.0, 3.0]) - b = torch.tensor([1.0, 2.0, 4.0]) # last element differs by 1 - s = _compare_tensors_quick(a, b) - assert "max_abs=1" in s, s - assert "rel_diff=" in s, s - - def test_shape_mismatch(self): - s = _compare_tensors_quick(torch.zeros(3), torch.zeros(4)) - assert "shape mismatch" in s, s - - def test_dtype_unified(self): - s = _compare_tensors_quick( - torch.zeros(3, dtype=torch.float32), - torch.zeros(3, dtype=torch.float64), - ) - assert "rel_diff=" in s, s - assert "max_abs=" in s, s - - def test_empty(self): - s = _compare_tensors_quick(torch.zeros(0), torch.zeros(0)) - assert s == "empty" - - class TestCollectiveTimeout: def test_watchdog_fires_on_timeout(self): block_event = threading.Event() diff --git a/test/registered/debug_utils/test_tensor_dump_forward_hook.py b/test/registered/debug_utils/test_tensor_dump_forward_hook.py index 1661df8c3..7a9351620 100644 --- a/test/registered/debug_utils/test_tensor_dump_forward_hook.py +++ b/test/registered/debug_utils/test_tensor_dump_forward_hook.py @@ -1,4 +1,6 @@ +import tempfile import unittest +from pathlib import Path import torch from torch import nn @@ -12,20 +14,11 @@ from sglang.srt.layers.linear import LinearBase from sglang.srt.models.qwen2 import Qwen2MLP from sglang.srt.server_args import ServerArgs, set_global_server_args_for_scheduler from sglang.srt.utils import add_prefix, get_device -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.layer_ut_utils import init_single_process_dist +from sglang.test.test_utils import CustomTestCase -register_cuda_ci( - est_time=9, - stage="base-b", - runner_config="1-gpu-small", - disabled="Test uses pytest-style function without TestCase class - see #17145", -) -register_amd_ci( - est_time=15, - suite="stage-b-test-1-gpu-small-amd", - disabled="Test uses pytest-style function without TestCase class - see #17145", -) +register_cuda_ci(est_time=9, stage="base-b", runner_config="1-gpu-small") TEST_HIDDEN_SIZE = 32 @@ -73,26 +66,29 @@ def init_weights(module): torch.nn.init.ones_(module.weight) -def test_model_forward_dump(tmp_path): - set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) - device = get_device() - init_single_process_dist(backend=get_default_distributed_backend(device)) - model = MockCausalLM() - model.apply(init_weights) - model = model.to(device=device, dtype=torch.bfloat16) - dumper = register_forward_hook_for_model( - model, tmp_path / "sglang_dump", [0], 0, 0, 0 - ) +class TestTensorDumpForwardHook(CustomTestCase): + def test_model_forward_dump(self): + set_global_server_args_for_scheduler(ServerArgs(model_path="dummy")) + device = get_device() + init_single_process_dist(backend=get_default_distributed_backend(device)) + model = MockCausalLM() + model.apply(init_weights) + model = model.to(device=device, dtype=torch.bfloat16) - dir_path = dumper.get_dump_dir() - inp = torch.randn(4, TEST_HIDDEN_SIZE, dtype=torch.bfloat16) * 0.01 - result = model(inp.to(device)) - data = torch.load(f"{dir_path}/Pass00000.pt") - assert "model.layernorm" in data - assert "model.mlp.down_proj" in data - assert torch.allclose( - data["model.mlp.down_proj"], result.cpu(), rtol=1e-5, atol=1e-5 - ) + with tempfile.TemporaryDirectory() as temp_dir: + dumper = register_forward_hook_for_model( + model, Path(temp_dir) / "sglang_dump", [0], 0, 0, 0 + ) + dir_path = dumper.get_dump_dir() + inp = torch.randn(4, TEST_HIDDEN_SIZE, dtype=torch.bfloat16) * 0.01 + result = model(inp.to(device)) + data = torch.load(f"{dir_path}/Pass00000.pt") + + self.assertIn("model.layernorm", data) + self.assertIn("model.mlp.down_proj", data) + torch.testing.assert_close( + data["model.mlp.down_proj"], result.cpu(), rtol=1e-5, atol=1e-5 + ) if __name__ == "__main__": diff --git a/test/registered/e2e/diffusion/test_diffusion_1_gpu.py b/test/registered/e2e/diffusion/test_diffusion_1_gpu.py new file mode 100644 index 000000000..e66a8563f --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_1_gpu.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=14400, + stage="base-b", + runner_config="diffusion-1-gpu-h100", +) + + +if __name__ == "__main__": + run_diffusion_suite("1-gpu") diff --git a/test/registered/e2e/diffusion/test_diffusion_1_gpu_5090.py b/test/registered/e2e/diffusion/test_diffusion_1_gpu_5090.py new file mode 100644 index 000000000..f8374cc25 --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_1_gpu_5090.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=7200, + stage="base-b", + runner_config="diffusion-1-gpu-5090", +) + + +if __name__ == "__main__": + run_diffusion_suite("1-gpu-5090") diff --git a/test/registered/e2e/diffusion/test_diffusion_1_gpu_b200.py b/test/registered/e2e/diffusion/test_diffusion_1_gpu_b200.py new file mode 100644 index 000000000..4c98a6839 --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_1_gpu_b200.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=14400, + stage="base-b", + runner_config="diffusion-1-gpu-b200", +) + + +if __name__ == "__main__": + run_diffusion_suite("1-gpu-b200") diff --git a/test/registered/e2e/diffusion/test_diffusion_2_gpu.py b/test/registered/e2e/diffusion/test_diffusion_2_gpu.py new file mode 100644 index 000000000..a94986052 --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_2_gpu.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=14400, + stage="base-b", + runner_config="diffusion-2-gpu-h100", +) + + +if __name__ == "__main__": + run_diffusion_suite("2-gpu") diff --git a/test/registered/e2e/diffusion/test_diffusion_bcg.py b/test/registered/e2e/diffusion/test_diffusion_bcg.py new file mode 100644 index 000000000..dff5d62f6 --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_bcg.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=3600, + stage="base-b", + runner_config="diffusion-bcg-1-gpu-h100", +) + + +if __name__ == "__main__": + run_diffusion_suite("bcg-diffusion") diff --git a/test/registered/e2e/diffusion/test_diffusion_component_accuracy.py b/test/registered/e2e/diffusion/test_diffusion_component_accuracy.py new file mode 100644 index 000000000..4d28fb4d8 --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_component_accuracy.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=14400, + stage="base-b", + runner_config="diffusion-component-2-gpu-h100", +) + + +if __name__ == "__main__": + run_diffusion_suite("component-accuracy") diff --git a/test/registered/e2e/diffusion/test_diffusion_unit.py b/test/registered/e2e/diffusion/test_diffusion_unit.py new file mode 100644 index 000000000..9c46e22ef --- /dev/null +++ b/test/registered/e2e/diffusion/test_diffusion_unit.py @@ -0,0 +1,13 @@ +from sglang.test.ci.ci_register import register_cuda_ci +from sglang.test.ci.diffusion_suite_bridge import run_diffusion_suite + +# ci-cost-override: compatibility bridge preserves the existing diffusion lane. +register_cuda_ci( + est_time=3600, + stage="base-b", + runner_config="diffusion-unit-1-gpu-h100", +) + + +if __name__ == "__main__": + run_diffusion_suite("unit") diff --git a/test/registered/models_e2e/test_compressed_tensors_models.py b/test/registered/e2e/models/test_compressed_tensors_models.py similarity index 100% rename from test/registered/models_e2e/test_compressed_tensors_models.py rename to test/registered/e2e/models/test_compressed_tensors_models.py diff --git a/test/registered/models_e2e/test_deepseek_v3_fp4.py b/test/registered/e2e/models/test_deepseek_v3_fp4.py similarity index 100% rename from test/registered/models_e2e/test_deepseek_v3_fp4.py rename to test/registered/e2e/models/test_deepseek_v3_fp4.py diff --git a/test/registered/models_e2e/test_deepseek_v3_mtp.py b/test/registered/e2e/models/test_deepseek_v3_mtp.py similarity index 100% rename from test/registered/models_e2e/test_deepseek_v3_mtp.py rename to test/registered/e2e/models/test_deepseek_v3_mtp.py diff --git a/test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200.py b/test/registered/e2e/models/test_deepseek_v4_flash_fp4_b200.py similarity index 100% rename from test/registered/models_e2e/test_deepseek_v4_flash_fp4_b200.py rename to test/registered/e2e/models/test_deepseek_v4_flash_fp4_b200.py diff --git a/test/registered/models_e2e/test_deepseek_v4_flash_fp4_h200.py b/test/registered/e2e/models/test_deepseek_v4_flash_fp4_h200.py similarity index 100% rename from test/registered/models_e2e/test_deepseek_v4_flash_fp4_h200.py rename to test/registered/e2e/models/test_deepseek_v4_flash_fp4_h200.py diff --git a/test/registered/models_e2e/test_deepseek_v4_flash_fp4_megamoe_b200.py b/test/registered/e2e/models/test_deepseek_v4_flash_fp4_megamoe_b200.py similarity index 100% rename from test/registered/models_e2e/test_deepseek_v4_flash_fp4_megamoe_b200.py rename to test/registered/e2e/models/test_deepseek_v4_flash_fp4_megamoe_b200.py diff --git a/test/registered/models_e2e/test_deepseek_v4_flash_fp8_h200.py b/test/registered/e2e/models/test_deepseek_v4_flash_fp8_h200.py similarity index 100% rename from test/registered/models_e2e/test_deepseek_v4_flash_fp8_h200.py rename to test/registered/e2e/models/test_deepseek_v4_flash_fp8_h200.py diff --git a/test/registered/models_e2e/test_dsa_glm52_dp_mtp.py b/test/registered/e2e/models/test_dsa_glm52_dp_mtp.py similarity index 100% rename from test/registered/models_e2e/test_dsa_glm52_dp_mtp.py rename to test/registered/e2e/models/test_dsa_glm52_dp_mtp.py diff --git a/test/registered/models_e2e/test_dsa_glm52_hisparse.py b/test/registered/e2e/models/test_dsa_glm52_hisparse.py similarity index 100% rename from test/registered/models_e2e/test_dsa_glm52_hisparse.py rename to test/registered/e2e/models/test_dsa_glm52_hisparse.py diff --git a/test/registered/models_e2e/test_dsa_glm52_nvfp4_dp_mtp.py b/test/registered/e2e/models/test_dsa_glm52_nvfp4_dp_mtp.py similarity index 100% rename from test/registered/models_e2e/test_dsa_glm52_nvfp4_dp_mtp.py rename to test/registered/e2e/models/test_dsa_glm52_nvfp4_dp_mtp.py diff --git a/test/registered/models_e2e/test_dsa_glm52_nvfp4_tp_mtp.py b/test/registered/e2e/models/test_dsa_glm52_nvfp4_tp_mtp.py similarity index 100% rename from test/registered/models_e2e/test_dsa_glm52_nvfp4_tp_mtp.py rename to test/registered/e2e/models/test_dsa_glm52_nvfp4_tp_mtp.py diff --git a/test/registered/models_e2e/test_dsa_glm52_pd_mtp_cp_layersplit.py b/test/registered/e2e/models/test_dsa_glm52_pd_mtp_cp_layersplit.py similarity index 100% rename from test/registered/models_e2e/test_dsa_glm52_pd_mtp_cp_layersplit.py rename to test/registered/e2e/models/test_dsa_glm52_pd_mtp_cp_layersplit.py diff --git a/test/registered/models_e2e/test_dsa_glm52_tp_mtp.py b/test/registered/e2e/models/test_dsa_glm52_tp_mtp.py similarity index 100% rename from test/registered/models_e2e/test_dsa_glm52_tp_mtp.py rename to test/registered/e2e/models/test_dsa_glm52_tp_mtp.py diff --git a/test/registered/models_e2e/test_gemma4_fp8_per_expert_loading.py b/test/registered/e2e/models/test_gemma4_fp8_per_expert_loading.py similarity index 100% rename from test/registered/models_e2e/test_gemma4_fp8_per_expert_loading.py rename to test/registered/e2e/models/test_gemma4_fp8_per_expert_loading.py diff --git a/test/registered/models_e2e/test_generation_models.py b/test/registered/e2e/models/test_generation_models.py similarity index 100% rename from test/registered/models_e2e/test_generation_models.py rename to test/registered/e2e/models/test_generation_models.py diff --git a/test/registered/models_e2e/test_glm53_flash_b200.py b/test/registered/e2e/models/test_glm53_flash_b200.py similarity index 100% rename from test/registered/models_e2e/test_glm53_flash_b200.py rename to test/registered/e2e/models/test_glm53_flash_b200.py diff --git a/test/registered/models_e2e/test_glm53_flash_h200.py b/test/registered/e2e/models/test_glm53_flash_h200.py similarity index 100% rename from test/registered/models_e2e/test_glm53_flash_h200.py rename to test/registered/e2e/models/test_glm53_flash_h200.py diff --git a/test/registered/models_e2e/test_gpt_oss_4gpu_mxfp4.py b/test/registered/e2e/models/test_gpt_oss_4gpu_mxfp4.py similarity index 100% rename from test/registered/models_e2e/test_gpt_oss_4gpu_mxfp4.py rename to test/registered/e2e/models/test_gpt_oss_4gpu_mxfp4.py diff --git a/test/registered/models_e2e/test_gpt_oss_sm120.py b/test/registered/e2e/models/test_gpt_oss_sm120.py similarity index 100% rename from test/registered/models_e2e/test_gpt_oss_sm120.py rename to test/registered/e2e/models/test_gpt_oss_sm120.py diff --git a/test/registered/models_e2e/test_inkling.py b/test/registered/e2e/models/test_inkling.py similarity index 100% rename from test/registered/models_e2e/test_inkling.py rename to test/registered/e2e/models/test_inkling.py diff --git a/test/registered/models_e2e/test_inkling_small_nvfp4.py b/test/registered/e2e/models/test_inkling_small_nvfp4.py similarity index 100% rename from test/registered/models_e2e/test_inkling_small_nvfp4.py rename to test/registered/e2e/models/test_inkling_small_nvfp4.py diff --git a/test/registered/models_e2e/test_inkling_unified.py b/test/registered/e2e/models/test_inkling_unified.py similarity index 100% rename from test/registered/models_e2e/test_inkling_unified.py rename to test/registered/e2e/models/test_inkling_unified.py diff --git a/test/registered/models_e2e/test_kimi_k3_b300.py b/test/registered/e2e/models/test_kimi_k3_b300.py similarity index 100% rename from test/registered/models_e2e/test_kimi_k3_b300.py rename to test/registered/e2e/models/test_kimi_k3_b300.py diff --git a/test/registered/models_e2e/test_kimi_k3_b300_low_latency.py b/test/registered/e2e/models/test_kimi_k3_b300_low_latency.py similarity index 100% rename from test/registered/models_e2e/test_kimi_k3_b300_low_latency.py rename to test/registered/e2e/models/test_kimi_k3_b300_low_latency.py diff --git a/test/registered/models_e2e/test_kimi_linear_models.py b/test/registered/e2e/models/test_kimi_linear_models.py similarity index 100% rename from test/registered/models_e2e/test_kimi_linear_models.py rename to test/registered/e2e/models/test_kimi_linear_models.py diff --git a/test/registered/models_e2e/test_kimi_linear_unified_memory.py b/test/registered/e2e/models/test_kimi_linear_unified_memory.py similarity index 100% rename from test/registered/models_e2e/test_kimi_linear_unified_memory.py rename to test/registered/e2e/models/test_kimi_linear_unified_memory.py diff --git a/test/registered/models_e2e/test_kimi_linear_unified_memory_dcp_blackwell.py b/test/registered/e2e/models/test_kimi_linear_unified_memory_dcp_blackwell.py similarity index 100% rename from test/registered/models_e2e/test_kimi_linear_unified_memory_dcp_blackwell.py rename to test/registered/e2e/models/test_kimi_linear_unified_memory_dcp_blackwell.py diff --git a/test/registered/models_e2e/test_layernorm_sp.py b/test/registered/e2e/models/test_layernorm_sp.py similarity index 100% rename from test/registered/models_e2e/test_layernorm_sp.py rename to test/registered/e2e/models/test_layernorm_sp.py diff --git a/test/registered/models_e2e/test_mimo_v2.py b/test/registered/e2e/models/test_mimo_v2.py similarity index 100% rename from test/registered/models_e2e/test_mimo_v2.py rename to test/registered/e2e/models/test_mimo_v2.py diff --git a/test/registered/models_e2e/test_mimo_v2_flash.py b/test/registered/e2e/models/test_mimo_v2_flash.py similarity index 100% rename from test/registered/models_e2e/test_mimo_v2_flash.py rename to test/registered/e2e/models/test_mimo_v2_flash.py diff --git a/test/registered/models_e2e/test_minimax_m25_basic.py b/test/registered/e2e/models/test_minimax_m25_basic.py similarity index 100% rename from test/registered/models_e2e/test_minimax_m25_basic.py rename to test/registered/e2e/models/test_minimax_m25_basic.py diff --git a/test/registered/models_e2e/test_ministral4_models.py b/test/registered/e2e/models/test_ministral4_models.py similarity index 100% rename from test/registered/models_e2e/test_ministral4_models.py rename to test/registered/e2e/models/test_ministral4_models.py diff --git a/test/registered/models_e2e/test_nvidia_nemotron_3_nano.py b/test/registered/e2e/models/test_nvidia_nemotron_3_nano.py similarity index 100% rename from test/registered/models_e2e/test_nvidia_nemotron_3_nano.py rename to test/registered/e2e/models/test_nvidia_nemotron_3_nano.py diff --git a/test/registered/models_e2e/test_nvidia_nemotron_3_super_bf16.py b/test/registered/e2e/models/test_nvidia_nemotron_3_super_bf16.py similarity index 100% rename from test/registered/models_e2e/test_nvidia_nemotron_3_super_bf16.py rename to test/registered/e2e/models/test_nvidia_nemotron_3_super_bf16.py diff --git a/test/registered/models_e2e/test_nvidia_nemotron_3_super_bf16_mtp.py b/test/registered/e2e/models/test_nvidia_nemotron_3_super_bf16_mtp.py similarity index 100% rename from test/registered/models_e2e/test_nvidia_nemotron_3_super_bf16_mtp.py rename to test/registered/e2e/models/test_nvidia_nemotron_3_super_bf16_mtp.py diff --git a/test/registered/models_e2e/test_qwen35_fp4_mtp.py b/test/registered/e2e/models/test_qwen35_fp4_mtp.py similarity index 100% rename from test/registered/models_e2e/test_qwen35_fp4_mtp.py rename to test/registered/e2e/models/test_qwen35_fp4_mtp.py diff --git a/test/registered/models_e2e/test_qwen3_next_models.py b/test/registered/e2e/models/test_qwen3_next_models.py similarity index 100% rename from test/registered/models_e2e/test_qwen3_next_models.py rename to test/registered/e2e/models/test_qwen3_next_models.py diff --git a/test/registered/models_e2e/test_qwen3_next_models_extra.py b/test/registered/e2e/models/test_qwen3_next_models_extra.py similarity index 100% rename from test/registered/models_e2e/test_qwen3_next_models_extra.py rename to test/registered/e2e/models/test_qwen3_next_models_extra.py diff --git a/test/registered/models_e2e/test_qwen3_next_models_mtp.py b/test/registered/e2e/models/test_qwen3_next_models_mtp.py similarity index 100% rename from test/registered/models_e2e/test_qwen3_next_models_mtp.py rename to test/registered/e2e/models/test_qwen3_next_models_mtp.py diff --git a/test/registered/models_e2e/test_step3p5_flash_chain_mtp.py b/test/registered/e2e/models/test_step3p5_flash_chain_mtp.py similarity index 100% rename from test/registered/models_e2e/test_step3p5_flash_chain_mtp.py rename to test/registered/e2e/models/test_step3p5_flash_chain_mtp.py diff --git a/test/registered/models_e2e/test_transformers_backend_eval.py b/test/registered/e2e/models/test_transformers_backend_eval.py similarity index 100% rename from test/registered/models_e2e/test_transformers_backend_eval.py rename to test/registered/e2e/models/test_transformers_backend_eval.py diff --git a/test/registered/models_e2e/test_transformers_models.py b/test/registered/e2e/models/test_transformers_models.py similarity index 100% rename from test/registered/models_e2e/test_transformers_models.py rename to test/registered/e2e/models/test_transformers_models.py diff --git a/test/registered/models_e2e/test_vlm_models.py b/test/registered/e2e/models/test_vlm_models.py similarity index 100% rename from test/registered/models_e2e/test_vlm_models.py rename to test/registered/e2e/models/test_vlm_models.py diff --git a/test/registered/8-gpu-models/test_deepseek_v32_indexcache.py b/test/registered/e2e/models_large/test_deepseek_v32_indexcache.py similarity index 100% rename from test/registered/8-gpu-models/test_deepseek_v32_indexcache.py rename to test/registered/e2e/models_large/test_deepseek_v32_indexcache.py diff --git a/test/registered/4-gpu-models/test_deepseek_v3_cutedsl_4gpu.py b/test/registered/e2e/models_large/test_deepseek_v3_cutedsl_4gpu.py similarity index 100% rename from test/registered/4-gpu-models/test_deepseek_v3_cutedsl_4gpu.py rename to test/registered/e2e/models_large/test_deepseek_v3_cutedsl_4gpu.py diff --git a/test/registered/8-gpu-models/test_glm52_fp8.py b/test/registered/e2e/models_large/test_glm52_fp8.py similarity index 100% rename from test/registered/8-gpu-models/test_glm52_fp8.py rename to test/registered/e2e/models_large/test_glm52_fp8.py diff --git a/test/registered/8-gpu-models/test_glm_46.py b/test/registered/e2e/models_large/test_glm_46.py similarity index 100% rename from test/registered/8-gpu-models/test_glm_46.py rename to test/registered/e2e/models_large/test_glm_46.py diff --git a/test/registered/8-gpu-models/test_gpt_oss_120b.py b/test/registered/e2e/models_large/test_gpt_oss_120b.py similarity index 100% rename from test/registered/8-gpu-models/test_gpt_oss_120b.py rename to test/registered/e2e/models_large/test_gpt_oss_120b.py diff --git a/test/registered/8-gpu-models/test_inkling_nvfp4_nightly.py b/test/registered/e2e/models_large/test_inkling_nvfp4_nightly.py similarity index 98% rename from test/registered/8-gpu-models/test_inkling_nvfp4_nightly.py rename to test/registered/e2e/models_large/test_inkling_nvfp4_nightly.py index d2e9e8460..e564d1957 100644 --- a/test/registered/8-gpu-models/test_inkling_nvfp4_nightly.py +++ b/test/registered/e2e/models_large/test_inkling_nvfp4_nightly.py @@ -80,7 +80,7 @@ class TestInklingNVFP4Nightly(unittest.TestCase): class TestInklingSmallCacheConsistencyNightly(unittest.TestCase): """Bitwise version of the per-commit check in - ``test/registered/models_e2e/test_inkling.py``, on the real checkpoint and + ``test/registered/e2e/models/test_inkling.py``, on the real checkpoint and at a batch shape the tiny checkpoint never reaches.""" @unittest.skipIf(not is_blackwell_system(), "NVFP4 requires Blackwell") diff --git a/test/registered/8-gpu-models/test_kimi_k25.py b/test/registered/e2e/models_large/test_kimi_k25.py similarity index 100% rename from test/registered/8-gpu-models/test_kimi_k25.py rename to test/registered/e2e/models_large/test_kimi_k25.py diff --git a/test/registered/4-gpu-models/test_laguna_nvfp4_nightly.py b/test/registered/e2e/models_large/test_laguna_nvfp4_nightly.py similarity index 100% rename from test/registered/4-gpu-models/test_laguna_nvfp4_nightly.py rename to test/registered/e2e/models_large/test_laguna_nvfp4_nightly.py diff --git a/test/registered/8-gpu-models/test_ling_2_6_flash.py b/test/registered/e2e/models_large/test_ling_2_6_flash.py similarity index 100% rename from test/registered/8-gpu-models/test_ling_2_6_flash.py rename to test/registered/e2e/models_large/test_ling_2_6_flash.py diff --git a/test/registered/8-gpu-models/test_longcat_flash_lite_fp8.py b/test/registered/e2e/models_large/test_longcat_flash_lite_fp8.py similarity index 100% rename from test/registered/8-gpu-models/test_longcat_flash_lite_fp8.py rename to test/registered/e2e/models_large/test_longcat_flash_lite_fp8.py diff --git a/test/registered/8-gpu-models/test_minimax_m25.py b/test/registered/e2e/models_large/test_minimax_m25.py similarity index 100% rename from test/registered/8-gpu-models/test_minimax_m25.py rename to test/registered/e2e/models_large/test_minimax_m25.py diff --git a/test/registered/8-gpu-models/test_mistral_large3.py b/test/registered/e2e/models_large/test_mistral_large3.py similarity index 100% rename from test/registered/8-gpu-models/test_mistral_large3.py rename to test/registered/e2e/models_large/test_mistral_large3.py diff --git a/test/registered/8-gpu-models/test_nvidia_nemotron_3_super_nightly.py b/test/registered/e2e/models_large/test_nvidia_nemotron_3_super_nightly.py similarity index 100% rename from test/registered/8-gpu-models/test_nvidia_nemotron_3_super_nightly.py rename to test/registered/e2e/models_large/test_nvidia_nemotron_3_super_nightly.py diff --git a/test/registered/4-gpu-models/test_nvidia_nemotron_3_super_nvfp4.py b/test/registered/e2e/models_large/test_nvidia_nemotron_3_super_nvfp4.py similarity index 100% rename from test/registered/4-gpu-models/test_nvidia_nemotron_3_super_nvfp4.py rename to test/registered/e2e/models_large/test_nvidia_nemotron_3_super_nvfp4.py diff --git a/test/registered/8-gpu-models/test_qwen35.py b/test/registered/e2e/models_large/test_qwen35.py similarity index 100% rename from test/registered/8-gpu-models/test_qwen35.py rename to test/registered/e2e/models_large/test_qwen35.py diff --git a/test/registered/8-gpu-models/test_ring_2_5_1t.py b/test/registered/e2e/models_large/test_ring_2_5_1t.py similarity index 100% rename from test/registered/8-gpu-models/test_ring_2_5_1t.py rename to test/registered/e2e/models_large/test_ring_2_5_1t.py diff --git a/test/registered/kernels/ops/diffusion/test_import_surface.py b/test/registered/kernels/ops/diffusion/test_import_surface.py deleted file mode 100644 index 7f739406b..000000000 --- a/test/registered/kernels/ops/diffusion/test_import_surface.py +++ /dev/null @@ -1,228 +0,0 @@ -"""Guards that keep the ``diffusion`` package's import surface from eroding. - -The reorganization only stays useful if two invariants hold: - -1. runtime code imports from ``sglang.kernels.ops.diffusion`` and not from a - submodule, so the internal layout can move without touching call sites; -2. the facade's ``_EXPORTS`` table and the registry's ``_SPECS`` table both - point at symbols that actually exist. - -Neither is checkable by the type system, and both fail silently -- a stale -``_EXPORTS`` entry only raises when some model happens to call that kernel, on -a GPU, at serving time. These are pure-CPU tests: they read the tables and -resolve them with ``importlib``/``ast`` without importing torch backends. -""" - -import ast -import functools -import importlib -import pathlib -import subprocess -import sys - -import pytest - -from sglang.kernels.ops.diffusion import _EXPORTS, _SPECS -from sglang.kernels.registry import registry -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=16, suite="base-a-test-cpu") - -PACKAGE = "sglang.kernels.ops.diffusion" -_PACKAGE_DIR = pathlib.Path(importlib.import_module(PACKAGE).__file__ or "").parent -_REPO_ROOT = _PACKAGE_DIR.parents[4] # /python/sglang/kernels/ops/diffusion - -# Backend-specific test files may name a leaf module on purpose; everything -# else -- all runtime code -- must go through the facade. -_DEEP_IMPORT_ALLOWLIST = { - "python/sglang/multimodal_gen/test/unit/test_latent_upsampler_group_norm_silu.py", - "test/registered/kernels/ops/diffusion/test_model_fast_paths.py", - "test/registered/kernels/ops/diffusion/test_sites.py", - # This test exercises the pure-Torch fallback implementation directly. - "test/registered/unit/utils/test_diffusion_torch_fallback.py", -} - - -def _module_defines(module_path: str) -> set[str]: - """Top-level names bound by a submodule, without importing it. - - Importing would pull in Triton / CuTe-DSL / FlyDSL, none of which are - installed on the CPU CI lane -- so this reads the source instead. - """ - if module_path.startswith("sglang."): - spec = importlib.util.find_spec(module_path) - assert spec is not None and spec.origin is not None, module_path - path = pathlib.Path(spec.origin) - else: - path = _PACKAGE_DIR / (module_path.replace(".", "/") + ".py") - if not path.exists(): - path = _PACKAGE_DIR / module_path.replace(".", "/") / "__init__.py" - assert path.exists(), f"{PACKAGE}.{module_path} does not exist" - - names: set[str] = set() - for node in ast.parse(path.read_text(encoding="utf-8")).body: - if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): - names.add(node.name) - elif isinstance(node, ast.Assign): - names.update(t.id for t in node.targets if isinstance(t, ast.Name)) - elif isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name): - names.add(node.target.id) - elif isinstance(node, (ast.Import, ast.ImportFrom)): - names.update((a.asname or a.name).split(".")[0] for a in node.names) - elif isinstance(node, (ast.If, ast.Try)): - # Platform-conditional rebinds (``x = select_impl(...)``) and - # guarded defs still bind a public name. - for inner in ast.walk(node): - if isinstance(inner, (ast.FunctionDef, ast.ClassDef)): - names.add(inner.name) - elif isinstance(inner, ast.Assign): - names.update(t.id for t in inner.targets if isinstance(t, ast.Name)) - return names - - -@functools.lru_cache(maxsize=None) -def _scan_root(root: str) -> tuple[frozenset[str], tuple[str, ...]]: - unexported: set[str] = set() - offenders: list[str] = [] - root_dir = _REPO_ROOT / root - if not root_dir.exists(): - return frozenset(), () - - for path in root_dir.rglob("*.py"): - rel = path.relative_to(_REPO_ROOT).as_posix() - if rel.startswith( - ( - "python/sglang/kernels/ops/diffusion/", - "python/sglang/kernels/kda_kernels/", - ) - ): - continue - try: - source = path.read_text(encoding="utf-8") - except UnicodeDecodeError: - continue - if PACKAGE not in source: - continue - try: - tree = ast.parse(source) - except SyntaxError: - continue - allowlisted = rel in _DEEP_IMPORT_ALLOWLIST - for node in ast.walk(tree): - if isinstance(node, ast.ImportFrom): - if node.module == PACKAGE: - unexported.update( - a.name - for a in node.names - if a.name not in _EXPORTS and not a.name.startswith("_") - ) - elif ( - not allowlisted - and node.module - and node.module.startswith(f"{PACKAGE}.") - ): - offenders.append(f"{rel}:{node.lineno} imports {node.module}") - elif isinstance(node, ast.Import) and not allowlisted: - offenders.extend( - f"{rel}:{node.lineno} imports {a.name}" - for a in node.names - if a.name.startswith(f"{PACKAGE}.") - ) - return frozenset(unexported), tuple(offenders) - - -def test_every_export_resolves_to_a_real_symbol(): - missing = [ - f"{symbol} -> {module}" - for symbol, module in sorted(_EXPORTS.items()) - if symbol not in _module_defines(module) - ] - assert not missing, f"stale _EXPORTS entries: {missing}" - - -def test_every_symbol_imported_from_the_facade_is_exported(): - """The reverse of the check above, and the one that actually bites. - - A missing ``_EXPORTS`` entry raises ``ImportError`` at module import, so a - module-level ``from ...diffusion import x`` fails loudly. A *function-local* - one -- the pattern used for optional backends -- fails only when that test - or code path runs, on the platform that has the backend. Enumerating the - call sites catches it here instead. - """ - unexported: set[str] = set() - for root in ("python/sglang", "test", "benchmark"): - unexported.update(_scan_root(root)[0]) - assert not unexported, f"imported but not in _EXPORTS: {sorted(unexported)}" - - -def test_every_registered_spec_target_resolves(): - missing = [] - for _op, _backend, target, _caps, _description in _SPECS: - module, _, attr = target.partition(":") - if attr not in _module_defines(module): - missing.append(target) - assert not missing, f"stale _SPECS targets: {missing}" - - -def test_registry_holds_the_diffusion_ops(): - # Registration happens at package import, is metadata-only, and is what - # ``select_kernel`` / the tracing tools read. - registered = {op for op in registry.ops() if op.startswith("diffusion.")} - assert {op for op, *_ in _SPECS} <= registered - - -def test_facade_rejects_unknown_attributes(): - module = sys.modules[PACKAGE] - with pytest.raises(AttributeError): - module.definitely_not_a_kernel - assert set(module.__all__) == set(_EXPORTS) - assert set(_EXPORTS) <= set(dir(module)) - - -def test_importing_the_package_does_not_import_any_leaf_module(): - """The reason ``__getattr__`` is lazy rather than a block of re-exports. - - The backends have disjoint, heavy, mutually-exclusive dependencies -- - Triton (CUDA/ROCm), CUTLASS/CuTe-DSL, and FlyDSL (gfx950). If - ``_EXPORTS`` ever degrades into eager ``from .norm.x import y`` lines, all - of them become import-time requirements on every platform, which is how a - CPU-only or Apple install starts failing at ``import sglang``. - - Asserted on this package's own leaf modules rather than on ``triton`` in - ``sys.modules``: sibling operator groups import Triton for their own - reasons, so a global check would not isolate this package's behavior. - Run in a fresh interpreter because this process has already resolved - exports through the facade. - """ - code = ( - "import importlib, sys\n" - f"importlib.import_module('{PACKAGE}')\n" - f"prefix = '{PACKAGE}.'\n" - "leaves = [m for m in sys.modules if m.startswith(prefix)" - " and not m.endswith('__init__')]\n" - "print(','.join(sorted(m for m in leaves if '.' in m[len(prefix):]" - " or sys.modules[m].__file__ and not sys.modules[m].__file__" - ".endswith('__init__.py'))))\n" - ) - result = subprocess.run( - [sys.executable, "-c", code], capture_output=True, text=True, timeout=600 - ) - assert result.returncode == 0, result.stderr - leaked = [m for m in result.stdout.strip().split(",") if m] - assert not leaked, f"importing {PACKAGE} eagerly imported: {leaked}" - - -@pytest.mark.parametrize("root", ["python/sglang", "test", "benchmark"]) -def test_runtime_code_imports_only_through_the_facade(root): - if not (_REPO_ROOT / root).exists(): # source checkouts only - pytest.skip(f"{root} not present in this install") - - offenders = _scan_root(root)[1] - assert not offenders, ( - "import from sglang.kernels.ops.diffusion instead of a submodule:\n " - + "\n ".join(offenders) - ) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py index 9c1ceb12b..e715a0066 100644 --- a/test/registered/kernels/ops/layernorm/test_kernels_namespace.py +++ b/test/registered/kernels/ops/layernorm/test_kernels_namespace.py @@ -1,7 +1,6 @@ """GPU-free import / registry / selector tests for ``sglang.kernels`` (RFC #29630).""" import importlib -import importlib.util import subprocess import sys @@ -9,127 +8,19 @@ import pytest import sglang.kernels as K import sglang.kernels.fused_op as fo -import sglang.kernels.ops # noqa: F401 -- populate the registry import sglang.kernels.selector as sel -from sglang.kernels import DeviceType, KernelBackend, PlatformInfo +from sglang.kernels import KernelBackend, PlatformInfo from sglang.kernels.spec import CapabilityRequirement as Cap from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=24, suite="base-a-test-cpu") -GROUPS = K.ops.__all__ - -# Representative ops checked as a subset (the registry holds many more). -EXPECTED = { - "activation.silu_and_mul": {"aot", "jit", "aiter", "torch", "torch_compile"}, - "activation.relu2": {"jit", "torch", "torch_compile"}, - "layernorm.rmsnorm": {"aot", "jit", "aiter", "torch_npu", "torch", "torch_compile"}, - "layernorm.gemma_rmsnorm": {"aot", "jit", "torch_npu", "torch", "torch_compile"}, - "gemm.fp8_scaled_mm": {"aot", "torch", "torch_compile"}, - "moe.moe_align_block_size": {"aot", "jit"}, - "quantization.nvfp4_gemm_swiglu_nvfp4_quant": {"cute_dsl"}, - "kvcache.reshape_and_cache_flash": {"triton"}, - "diffusion.apply_group_norm_silu": {"triton"}, - "diffusion.norm_scale_shift": {"KDA", "cute_dsl", "flydsl"}, - "diffusion.scale_residual_norm_scale_shift": { - "KDA", - "triton", - "cute_dsl", - "flydsl", - }, - "diffusion.residual_gate_add": {"KDA"}, - "diffusion.ltx2_qknorm_split_rope": {"KDA"}, - "diffusion.causal_conv3d_cat_pad": {"KDA", "triton"}, - "diffusion.flux2_layernorm_modulate_fp8_quant": {"KDA"}, - "diffusion.flux2_qkv_epilogue": {"KDA"}, - "diffusion.flux2_token_cat_fp8": {"KDA"}, - "gemm.qwen3x_nvfp4": {"KDA"}, - "gemm.sm120_fp8_linear": {"KDA"}, -} - _CPU = PlatformInfo(device_type="cpu") _SM90 = PlatformInfo(device_type="cuda", cuda_arch_major=9, cuda_arch_minor=0) _SM100 = PlatformInfo(device_type="cuda", cuda_arch_major=10, cuda_arch_minor=0) _HIP = PlatformInfo(device_type="hip") -def test_top_level_exports(): - for name in ( - "KernelSpec", - "KernelBackend", - "FormatSignature", - "CapabilityRequirement", - "PlatformInfo", - "registry", - "get_kernel", - "select_kernel", - ): - assert hasattr(K, name), name - - -@pytest.mark.parametrize("group", GROUPS) -def test_group_importable(group): - assert importlib.import_module(f"sglang.kernels.ops.{group}") is not None - - -@pytest.mark.parametrize("op, backends", list(EXPECTED.items())) -def test_registry_backends(op, backends): - assert {s.backend.value for s in K.registry.get(op)} == backends - - -def test_specs_well_formed(): - for spec in K.registry.all_specs(): - assert spec.op == f"{spec.group}.{spec.name}" - mod, sep, attr = spec.target.partition(":") - assert sep == ":" and mod and attr, spec.target - - -def test_internal_registry_target_modules_exist(): - for spec in K.registry.all_specs(): - module, _, _ = spec.target.partition(":") - if module.startswith("sglang.kernels."): - assert importlib.util.find_spec(module) is not None, spec.target - - -def test_sparse_linear_attention_registry_targets_forward_kernel(): - spec = K.registry.get_backend( - "diffusion.sparse_linear_attn_fwd", KernelBackend.TRITON - ) - assert spec.target.endswith(":_attn_fwd") - - -@pytest.mark.parametrize( - "op, target_suffix", - ( - ("diffusion.norm_scale_shift", ":kda_norm_scale_shift"), - ( - "diffusion.scale_residual_norm_scale_shift", - ":kda_scale_residual_norm_scale_shift", - ), - ("diffusion.residual_gate_add", ":residual_gate_add"), - ( - "diffusion.ltx2_qknorm_split_rope", - ":ltx2_qknorm_split_rope_cuda", - ), - ( - "diffusion.causal_conv3d_cat_pad", - ":fused_causal_conv3d_cat_pad_cuda", - ), - ), -) -def test_merged_diffusion_kda_provenance_backend(op, target_suffix): - spec = K.registry.get_backend(op, KernelBackend.KDA) - assert spec.target.endswith(target_suffix) - - -def test_kda_backend_implementations_live_in_kda_home(): - specs = [ - spec for spec in K.registry.all_specs() if spec.backend is KernelBackend.KDA - ] - assert specs - assert all(spec.target.startswith("sglang.kernels.kda_kernels.") for spec in specs) - - def test_single_backend_resolves_without_backend(): assert ( K.select_kernel("kvcache.reshape_and_cache_flash").backend @@ -193,15 +84,6 @@ def test_layernorm_default_backend(monkeypatch, op_attr, device, expect): assert getattr(ln, op_attr).auto_selected_backend().value == expect -def test_per_op_backend_subset(): - # silu_and_mul ships an aiter (HIP) kernel; the gelu siblings deliberately - # do not -- ROCm coverage is a per-(op, backend) subset. - from sglang.kernels.ops.activation import _GELU_AND_MUL, _SILU_AND_MUL - - assert KernelBackend.AITER in _SILU_AND_MUL.available_backends() - assert KernelBackend.AITER not in _GELU_AND_MUL.available_backends() - - @pytest.mark.parametrize( "req, plat, ok", [ @@ -227,20 +109,6 @@ def test_capabilities_or_semantics(): assert K.capabilities_satisfied(Cap.CUDA, _SM90) # single tolerated -def test_capability_shortcuts(): - assert Cap.CUDA == Cap(device=DeviceType.CUDA) - assert Cap.HIP == Cap(device=DeviceType.HIP) - assert Cap.NPU == Cap(device=DeviceType.NPU) - assert {Cap.CUDA, Cap.HIP} == {Cap.HIP, Cap.CUDA} - assert Cap.cuda(min_sm=(10, 0)) == Cap( - device=DeviceType.CUDA, min_cuda_arch=(10, 0) - ) - - -def test_platform_detect_does_not_raise(): - assert PlatformInfo.detect().device_type in ("cpu", "cuda", "hip", "npu") - - @pytest.mark.parametrize( "relative_path", ( diff --git a/test/registered/kernels/test_kernel_inventory.py b/test/registered/kernels/test_kernel_inventory.py deleted file mode 100644 index f3854471a..000000000 --- a/test/registered/kernels/test_kernel_inventory.py +++ /dev/null @@ -1,239 +0,0 @@ -"""CPU-only structural checks for the unified kernel tree.""" - -from __future__ import annotations - -import ast -import importlib.util -import sys -from pathlib import Path - -import pytest - -import sglang.kernels as kernels -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=13, suite="base-a-test-cpu") - -REPO_ROOT = Path(__file__).resolve().parents[3] -KERNELS_ROOT = REPO_ROOT / "python" / "sglang" / "kernels" -OPS_ROOT = KERNELS_ROOT / "ops" -JIT_CSRC_ROOT = KERNELS_ROOT / "jit" / "csrc" -AOT_ROOT = KERNELS_ROOT / "aot" - - -def _directory_names(root: Path) -> set[str]: - return { - path.name - for path in root.iterdir() - if path.is_dir() - and not path.name.startswith((".", "__")) - and any(path.rglob("*.py")) - } - - -def _target_names(target: ast.expr) -> set[str]: - if isinstance(target, ast.Name): - return {target.id} - if isinstance(target, (ast.List, ast.Tuple)): - return {name for element in target.elts for name in _target_names(element)} - return set() - - -def _bound_names(statements: list[ast.stmt]) -> set[str]: - """Collect names a module can bind without importing it.""" - names: set[str] = set() - for statement in statements: - if isinstance(statement, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): - names.add(statement.name) - elif isinstance(statement, ast.Assign): - for target in statement.targets: - names.update(_target_names(target)) - elif isinstance(statement, (ast.AnnAssign, ast.AugAssign)): - names.update(_target_names(statement.target)) - elif isinstance(statement, (ast.Import, ast.ImportFrom)): - for alias in statement.names: - names.add(alias.asname or alias.name.split(".", 1)[0]) - elif isinstance(statement, (ast.For, ast.AsyncFor)): - names.update(_target_names(statement.target)) - names.update(_bound_names(statement.body)) - names.update(_bound_names(statement.orelse)) - elif isinstance(statement, ast.If): - names.update(_bound_names(statement.body)) - names.update(_bound_names(statement.orelse)) - elif isinstance(statement, (ast.With, ast.AsyncWith)): - names.update(_bound_names(statement.body)) - elif isinstance(statement, ast.Try): - names.update(_bound_names(statement.body)) - names.update(_bound_names(statement.orelse)) - names.update(_bound_names(statement.finalbody)) - for handler in statement.handlers: - names.update(_bound_names(handler.body)) - elif isinstance(statement, ast.Match): - for case in statement.cases: - names.update(_bound_names(case.body)) - return names - - -def _module_string_constants(tree: ast.Module) -> dict[str, str]: - constants: dict[str, str] = {} - for statement in tree.body: - if not isinstance(statement, (ast.Assign, ast.AnnAssign)): - continue - value = statement.value - if not isinstance(value, ast.Constant) or not isinstance(value.value, str): - continue - targets = ( - statement.targets - if isinstance(statement, ast.Assign) - else [statement.target] - ) - for target in targets: - for name in _target_names(target): - constants[name] = value.value - return constants - - -def _source_patterns(expression: ast.expr, constants: dict[str, str]) -> list[str]: - if isinstance(expression, (ast.List, ast.Tuple)): - return [ - pattern - for element in expression.elts - for pattern in _source_patterns(element, constants) - ] - if isinstance(expression, ast.Constant) and isinstance(expression.value, str): - return [expression.value] - if isinstance(expression, ast.Name) and expression.id in constants: - return [constants[expression.id]] - if isinstance(expression, ast.JoinedStr): - parts = [] - for value in expression.values: - if isinstance(value, ast.Constant): - parts.append(str(value.value)) - elif isinstance(value, ast.FormattedValue): - parts.append("*") - else: - raise AssertionError(f"Unsupported f-string segment: {ast.dump(value)}") - return ["".join(parts)] - raise AssertionError( - f"Unsupported JIT source declaration: {ast.unparse(expression)}" - ) - - -def test_declared_operator_groups_match_packages(): - assert set(kernels.ops.__all__) == _directory_names(OPS_ROOT) - - -def test_registered_kernel_test_groups_are_known(): - declared_groups = set(kernels.ops.__all__) - registered_root = REPO_ROOT / "test" / "registered" / "kernels" - for kind in ("ops", "benchmark"): - unknown = _directory_names(registered_root / kind) - declared_groups - assert not unknown, ( - f"Unknown {kind} kernel group directories: {sorted(unknown)}" - ) - - -def test_internal_registry_target_attributes_are_declared(): - missing = [] - for spec in kernels.registry.all_specs(): - module_name, _, attribute_path = spec.target.partition(":") - if not module_name.startswith("sglang.kernels."): - continue - module_spec = importlib.util.find_spec(module_name) - if ( - module_spec is None - or module_spec.origin is None - or not module_spec.origin.endswith(".py") - ): - continue - tree = ast.parse(Path(module_spec.origin).read_text()) - root_attribute = attribute_path.split(".", 1)[0] - if root_attribute not in _bound_names(tree.body): - missing.append(spec.target) - assert not missing, f"KernelSpec targets missing attributes: {missing}" - - -# `load_jit` takes in-tree names and absolute paths on the same keyword, so this -# check can only reach the declarations spelled out in the source. A module that -# assembles its file list at runtime from a package outside `jit/csrc` has no -# in-tree name to verify and belongs here; there is none at the moment. -_RUNTIME_JIT_SOURCE_MODULES: set[str] = set() - - -def test_jit_source_declarations_exist(): - missing = [] - unsupported = [] - for python_file in OPS_ROOT.rglob("*.py"): - if python_file.relative_to(OPS_ROOT).as_posix() in _RUNTIME_JIT_SOURCE_MODULES: - continue - tree = ast.parse(python_file.read_text()) - constants = _module_string_constants(tree) - for call in (node for node in ast.walk(tree) if isinstance(node, ast.Call)): - function_name = ( - call.func.id - if isinstance(call.func, ast.Name) - else call.func.attr - if isinstance(call.func, ast.Attribute) - else None - ) - if function_name != "load_jit": - continue - for keyword in call.keywords: - if keyword.arg not in {"cpp_files", "cuda_files"}: - continue - try: - patterns = _source_patterns(keyword.value, constants) - except AssertionError as exc: - unsupported.append(f"{python_file.relative_to(REPO_ROOT)}: {exc}") - continue - for pattern in patterns: - matches = list(JIT_CSRC_ROOT.glob(pattern)) - if not matches: - missing.append( - f"{python_file.relative_to(REPO_ROOT)} -> {pattern}" - ) - assert not unsupported, "Unsupported JIT source declarations:\n" + "\n".join( - unsupported - ) - assert not missing, "Missing JIT sources:\n" + "\n".join(missing) - - -def test_aot_compilation_units_are_accounted_for(): - manifests = [ - AOT_ROOT / "CMakeLists.txt", - AOT_ROOT / "setup_metal.py", - AOT_ROOT / "setup_musa.py", - AOT_ROOT / "setup_rocm.py", - AOT_ROOT / "csrc" / "cpu" / "CMakeLists.txt", - *sorted((AOT_ROOT / "cmake").rglob("*.cmake")), - ] - manifest_text = "\n".join(path.read_text() for path in manifests) - source_text = { - path: path.read_text(errors="ignore") - for path in (AOT_ROOT / "csrc").rglob("*") - if path.is_file() - } - compilation_suffixes = {".cc", ".cpp", ".cu", ".hip", ".metal", ".mu"} - missing = [] - for source in source_text: - if source.suffix not in compilation_suffixes: - continue - if AOT_ROOT / "csrc" / "cpu" in source.parents: - # The CPU build intentionally uses file(GLOB_RECURSE ... *.cpp). - continue - relative_path = source.relative_to(AOT_ROOT).as_posix() - if relative_path in manifest_text: - continue - if any( - source.name in text - for other_source, text in source_text.items() - if other_source != source - ): - # Some CUDA translation units are included by another source. - continue - missing.append(relative_path) - assert not missing, f"AOT compilation units missing from build manifests: {missing}" - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py b/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py index 43f5529a6..587a052f3 100644 --- a/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py +++ b/test/registered/mlx/models_e2e/test_gpt_oss_mlx_correctness.py @@ -42,7 +42,7 @@ import unittest import requests from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import ( DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_URL_FOR_TEST, @@ -54,7 +54,6 @@ from sglang.test.test_utils import ( # Registered on the CPU suite but skipped wherever mlx is absent; runs for real # only on Apple Silicon. Also registered under stage-b-e2e-mlx, which the # macOS CI lane (pr-test-mlx.yml) only dispatches via a gated workflow_dispatch. -register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-b-e2e-mlx") _HAS_MLX = ( diff --git a/test/registered/mlx/models_e2e/test_qwen2_moe_mlx_correctness.py b/test/registered/mlx/models_e2e/test_qwen2_moe_mlx_correctness.py index a4722ee6d..641643631 100644 --- a/test/registered/mlx/models_e2e/test_qwen2_moe_mlx_correctness.py +++ b/test/registered/mlx/models_e2e/test_qwen2_moe_mlx_correctness.py @@ -5,7 +5,7 @@ import unittest import requests from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import ( DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_URL_FOR_TEST, @@ -17,7 +17,6 @@ from sglang.test.test_utils import ( # Registered on the CPU suite but skipped wherever mlx is absent; runs for real # only on Apple Silicon. Also registered under stage-b-e2e-mlx, which the # macOS CI lane (pr-test-mlx.yml) only dispatches via a gated workflow_dispatch. -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-b-e2e-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/mlx/models_e2e/test_qwen3_moe_mlx_correctness.py b/test/registered/mlx/models_e2e/test_qwen3_moe_mlx_correctness.py index 430969a49..818370ed4 100644 --- a/test/registered/mlx/models_e2e/test_qwen3_moe_mlx_correctness.py +++ b/test/registered/mlx/models_e2e/test_qwen3_moe_mlx_correctness.py @@ -5,7 +5,7 @@ import unittest import requests from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import ( DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, DEFAULT_URL_FOR_TEST, @@ -17,7 +17,6 @@ from sglang.test.test_utils import ( # Registered on the CPU suite but skipped wherever mlx is absent; runs for real # only on Apple Silicon. Also registered under stage-b-e2e-mlx, which the # macOS CI lane (pr-test-mlx.yml) only dispatches via a gated workflow_dispatch. -register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-b-e2e-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/models_e2e/test_dummy_grok_models.py b/test/registered/models_e2e/test_dummy_grok_models.py deleted file mode 100644 index f252d6ad3..000000000 --- a/test/registered/models_e2e/test_dummy_grok_models.py +++ /dev/null @@ -1,41 +0,0 @@ -import unittest - -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import CustomTestCase, is_in_ci, run_bench_one_batch - -register_cuda_ci( - est_time=120, - stage="base-b", - runner_config="2-gpu-large", - disabled="Temporarily disabled", -) - - -class TestDummyGrok1(CustomTestCase): - def test_dummy_grok_1(self): - _, output_throughput, _ = run_bench_one_batch( - None, - [ - "--model", - "/dummy-grok", - "--tokenizer-path", - "Xenova/grok-1-tokenizer", - "--batch-size", - "2", - "--tp", - "2", - "--quantization", - "fp8", - "--load-format", - "dummy", - "--json-model-override-args", - '{"num_hidden_layers": 2}', - ], - ) - - if is_in_ci(): - self.assertGreater(output_throughput, 0) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/models_e2e/test_ministral3_models.py b/test/registered/models_e2e/test_ministral3_models.py deleted file mode 100644 index f99a2bc26..000000000 --- a/test/registered/models_e2e/test_ministral3_models.py +++ /dev/null @@ -1,34 +0,0 @@ -import unittest - -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.kits.eval_accuracy_kit import GSM8KMixin -from sglang.test.kits.mmmu_vlm_kit import MMMUMixin -from sglang.test.server_fixtures.default_fixture import DefaultServerBase -from sglang.test.server_fixtures.mmmu_fixture import MMMUServerBase - -register_cuda_ci( - est_time=200, - stage="base-b", - runner_config="1-gpu-small", - disabled="Temporarily disabled", -) - -MODEL = "mistralai/Ministral-3-3B-Instruct-2512" - - -class TestMinistral3TextOnly(GSM8KMixin, DefaultServerBase): - gsm8k_accuracy_thres = 0.6 - model = MODEL - other_args = ["--trust-remote-code"] - - -class TestMinistral3MMMU(MMMUMixin, MMMUServerBase): - accuracy = 0.3 - model = MODEL - other_args = ["--trust-remote-code"] - mmmu_args = ["--limit=0.1"] - """`--limit=0.1`: 10 percent of each task - this is fine for testing since the nominal result isn't interesting - this run is just to prevent relative regressions.""" - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/npu/interface/test_npu_openai_function_calling.py b/test/registered/npu/interface/test_npu_openai_function_calling.py deleted file mode 100644 index e41d18dbb..000000000 --- a/test/registered/npu/interface/test_npu_openai_function_calling.py +++ /dev/null @@ -1,940 +0,0 @@ -import json -import unittest - -import openai - -from sglang.srt.utils import kill_process_tree -from sglang.srt.utils.hf_transformers_utils import get_tokenizer -from sglang.test.ascend.test_ascend_utils import LLAMA_3_2_1B_INSTRUCT_WEIGHTS_PATH -from sglang.test.ci.ci_register import register_npu_ci -from sglang.test.test_utils import ( - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - DEFAULT_URL_FOR_TEST, - CustomTestCase, - popen_launch_server, -) - -register_npu_ci( - est_time=400, - suite="full-1-npu-a3", - nightly=True, -) - - -class TestOpenAIServerFunctionCalling(CustomTestCase): - """Testcase:Verify the correctness of full-scenario OpenAI-style function calling with llama3 parser for Llama-3.2-1B-Instruct model. - Cover: Single/multi-turn calls, streaming/non-streaming returns, multi-parameter verification of tool_choice, and JSON parsing validity of function parameters. - - [Test Category] Interface - [Test Target] /v1/chat/completions - """ - - # NOTE: this system_message is for Llama3.2 system prompt. Without this, - # sometimes Llama3.2 gives a different tool call format such as: - # '<|python_tag|>{"type": "function", "function": "add", "parameters": {"a": "3", "b": "5"}}' - SYSTEM_MESSAGE = ( - "You are a helpful assistant with tool calling capabilities. " - "Only reply with a tool call if the function exists in the library provided by the user. " - "If it doesn't exist, just reply directly in natural language. " - "When you receive a tool call response, use the output to format an answer to the original user question. " - "You have access to the following functions. " - "To call a function, please respond with JSON for a function call. " - 'Respond in the format {"name": function name, "parameters": dictionary of argument name and its value}. ' - "Do not use variables.\n\n" - ) - - @classmethod - def setUpClass(cls): - # Replace with the model name needed for testing - cls.model = LLAMA_3_2_1B_INSTRUCT_WEIGHTS_PATH - cls.base_url = DEFAULT_URL_FOR_TEST - cls.api_key = "sk-123456" - - # Start the local OpenAI Server. If necessary, you can add other parameters such as --enable-tools. - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - api_key=cls.api_key, - other_args=[ - # If your server needs extra parameters to test function calling, please add them here. - "--attention-backend", - "ascend", - "--disable-cuda-graph", - "--tool-call-parser", - "llama3", - ], - ) - cls.base_url += "/v1" - cls.tokenizer = get_tokenizer(cls.model) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_function_calling_format(self): - """ - Test: Whether the function call format returned by the AI is correct. - When returning a tool call, message.content should be None, and tool_calls should be a list. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "add", - "description": "Compute the sum of two numbers", - "parameters": { - "type": "object", - "properties": { - "a": { - "type": "integer", - "description": "A number", - }, - "b": { - "type": "integer", - "description": "A number", - }, - }, - "required": ["a", "b"], - }, - }, - } - ] - - messages = [ - {"role": "system", "content": self.SYSTEM_MESSAGE}, - {"role": "user", "content": "Compute (3+5)"}, - ] - response = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=False, - tools=tools, - ) - - tool_calls = response.choices[0].message.tool_calls - - assert isinstance(tool_calls, list) and len(tool_calls) > 0, ( - "tool_calls should be a non-empty list" - ) - - function_name = tool_calls[0].function.name - assert function_name == "add", "Function name should be 'add'" - - # This unit test is too difficult for default model. Mark it as optional unit tests so it won't trigger unless specified. - def _test_function_calling_multiturn(self): - """ - Test: Whether the function call format returned by the AI is correct. - When returning a tool call, message.content should be None, and tool_calls should be a list. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "add", - "description": "Compute the sum of two numbers", - "parameters": { - "type": "object", - "properties": { - "a": { - "type": "integer", - "description": "A number", - }, - "b": { - "type": "integer", - "description": "A number", - }, - }, - "required": ["a", "b"], - }, - }, - } - ] - - messages = [{"role": "user", "content": "Compute (3+5)"}] - - response = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=False, - tools=tools, - ) - - tool_call = response.choices[0].message.tool_calls[0] - function_name = tool_call.function.name - assert function_name == "add", "Function name should be 'add'" - function_arguments = json.loads(tool_call.function.arguments) - assert function_arguments in [ - {"a": 3, "b": 5}, - {"a": "3", "b": "5"}, - ], f"Unexpected function arguments: {function_arguments}" - - messages.append(response.choices[0].message) - messages.append( - { - "role": "tool", - "tool_call_id": tool_call.id, - "content": "8", - "name": function_name, - } - ) - - final_response = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=False, - tools=tools, - ) - - assert "8" in final_response.choices[0].message.content, ( - "tool_call response should have the sum 8 in the content" - ) - - def test_function_calling_streaming_simple(self): - """ - Test: Whether the function name can be correctly recognized in streaming mode. - - Expect a function call to be found, and the function name to be correct. - - Verify that streaming mode returns at least multiple chunks. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "city": { - "type": "string", - "description": "The city to find the weather for", - }, - "unit": { - "type": "string", - "description": "Weather unit (celsius or fahrenheit)", - "enum": ["celsius", "fahrenheit"], - }, - }, - "required": ["city", "unit"], - }, - }, - } - ] - - messages = [ - {"role": "system", "content": self.SYSTEM_MESSAGE}, - { - "role": "user", - "content": "What is the temperature in Paris in celsius??", - }, - ] - - response_stream = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=True, - tools=tools, - ) - - chunks = list(response_stream) - self.assertTrue(len(chunks) > 0, "Streaming should return at least one chunk") - - found_function_name = False - for chunk in chunks: - choice = chunk.choices[0] - # Check whether the current chunk contains tool_calls - if choice.delta.tool_calls: - tool_call = choice.delta.tool_calls[0] - if tool_call.function.name: - self.assertEqual( - tool_call.function.name, - "get_current_weather", - "Function name should be 'get_current_weather'", - ) - found_function_name = True - break - - self.assertTrue( - found_function_name, - "Target function name 'get_current_weather' was not found in the streaming chunks", - ) - - finish_reason = chunks[-1].choices[0].finish_reason - self.assertEqual( - finish_reason, - "tool_calls", - "Final response of function calling should have finish_reason 'tool_calls'", - ) - - def test_function_calling_streaming_args_parsing(self): - """ - Test: Whether the function call arguments returned in streaming mode can be correctly concatenated into valid JSON. - - The user request requires multiple parameters. - - AI may return the arguments in chunks that need to be concatenated. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "add", - "description": "Compute the sum of two integers", - "parameters": { - "type": "object", - "properties": { - "a": { - "type": "integer", - "description": "First integer", - }, - "b": { - "type": "integer", - "description": "Second integer", - }, - }, - "required": ["a", "b"], - }, - "strict": True, # Llama-3.2-1B is flaky in tool call. It won't always respond with parameters unless we set strict. - }, - } - ] - - messages = [ - {"role": "system", "content": self.SYSTEM_MESSAGE}, - {"role": "user", "content": "Please sum 5 and 7, just call the function."}, - ] - - response_stream = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.9, - top_p=0.9, - stream=True, - tools=tools, - ) - - argument_fragments = [] - chunks = list(response_stream) - function_name = None - for chunk in chunks: - choice = chunk.choices[0] - if choice.delta.tool_calls: - tool_call = choice.delta.tool_calls[0] - # Record the function name on first occurrence - function_name = tool_call.function.name or function_name - # In case of multiple chunks, JSON fragments may need to be concatenated - if tool_call.function.arguments is not None: - argument_fragments.append(tool_call.function.arguments) - - self.assertEqual(function_name, "add", "Function name should be 'add'") - joined_args = "".join(argument_fragments) - self.assertTrue( - len(joined_args) > 0, - "No parameter fragments were returned in the function call", - ) - - finish_reason = chunks[-1].choices[0].finish_reason - self.assertEqual( - finish_reason, - "tool_calls", - "Final response of function calling should have finish_reason 'tool_calls'", - ) - - # Check whether the concatenated JSON is valid - try: - args_obj = json.loads(joined_args) - except json.JSONDecodeError: - self.fail( - "The concatenated tool call arguments are not valid JSON, parsing failed" - ) - - self.assertIn("a", args_obj, "Missing parameter 'a'") - self.assertIn("b", args_obj, "Missing parameter 'b'") - self.assertEqual(str(args_obj["a"]), "5", "Parameter a should be 5") - self.assertEqual(str(args_obj["b"]), "7", "Parameter b should be 7") - - def test_function_call_strict(self): - """ - Test: Whether the strict mode of function calling works as expected. - - When strict mode is enabled, the AI should not return a function call if the function name is not recognized. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "sub", - "description": "Compute the difference of two integers", - "parameters": { - "type": "object", - "properties": { - "int_a": { - "type": "integer", - "description": "First integer", - }, - "int_b": { - "type": "integer", - "description": "Second integer", - }, - }, - "required": ["int_a", "int_b"], - }, - "strict": True, - }, - } - ] - - messages = [ - {"role": "user", "content": "Please compute 5 - 7, using your tool."} - ] - response = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=False, - tools=tools, - ) - - tool_calls = response.choices[0].message.tool_calls - function_name = tool_calls[0].function.name - arguments = tool_calls[0].function.arguments - args_obj = json.loads(arguments) - - self.assertEqual(function_name, "sub", "Function name should be 'sub'") - self.assertEqual(str(args_obj["int_a"]), "5", "Parameter int_a should be 5") - self.assertEqual(str(args_obj["int_b"]), "7", "Parameter int_b should be 7") - - def test_function_call_required(self): - """ - Test: Whether tool_choice: "required" works as expected. - - When tool_choice == "required", the model MUST return one or more tool_calls. - - The model may choose ANY of the provided tools; we only verify that - a tool call exists and the selected name is among the candidates. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "sub", - "description": "Compute the difference of two integers", - "parameters": { - "type": "object", - "properties": { - "int_a": { - "type": "integer", - "description": "First integer", - }, - "int_b": { - "type": "integer", - "description": "Second integer", - }, - }, - "required": ["int_a", "int_b"], - }, - "strict": True, - }, - }, - { - "type": "function", - "function": { - "name": "get_weather", - "description": "use this to get latest weather information for a city given its name", - "parameters": { - "type": "object", - "properties": { - "city": { - "type": "string", - "description": "name of the city to get weather for", - } - }, - "required": ["city"], - }, - "strict": True, - }, - }, - ] - - valid_tool_names = {t["function"]["name"] for t in tools} - - messages = [{"role": "user", "content": "Tell me about Paris"}] - response = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0, - stream=False, - tools=tools, - tool_choice="required", - ) - - tool_calls = response.choices[0].message.tool_calls - self.assertIsNotNone( - tool_calls, "tool_choice='required' must produce tool_calls" - ) - self.assertGreater(len(tool_calls), 0, "tool_calls list should be non-empty") - - function_name = tool_calls[0].function.name - self.assertIn( - function_name, - valid_tool_names, - f"Function name '{function_name}' is not among the provided tools: {valid_tool_names}", - ) - - # Verify the arguments are parseable JSON - arguments = tool_calls[0].function.arguments - args_obj = json.loads(arguments) - self.assertIsInstance( - args_obj, dict, "Function arguments should be a JSON object" - ) - - def test_function_call_specific(self): - """ - Test: Whether tool_choice: ToolChoice works as expected - - When tool_choice is a specific ToolChoice, the model should return one or more tool_calls. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "sub", - "description": "Compute the difference of two integers", - "parameters": { - "type": "object", - "properties": { - "int_a": { - "type": "integer", - "description": "First integer", - }, - "int_b": { - "type": "integer", - "description": "Second integer", - }, - }, - "required": ["int_a", "int_b"], - }, - "strict": True, - }, - }, - { - "type": "function", - "function": { - "name": "get_weather", - "description": "use this to get latest weather information for a city given its name", - "parameters": { - "type": "object", - "properties": { - "city": { - "type": "string", - "description": "name of the city to get weather for", - } - }, - "required": ["city"], - }, - "strict": True, - }, - }, - ] - - messages = [{"role": "user", "content": "What is the capital of France?"}] - response = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=False, - tools=tools, - tool_choice={"type": "function", "function": {"name": "get_weather"}}, - ) - - tool_calls = response.choices[0].message.tool_calls - self.assertIsNotNone(tool_calls, "No tool_calls in the response") - function_name = tool_calls[0].function.name - arguments = tool_calls[0].function.arguments - args_obj = json.loads(arguments) - - self.assertEqual( - function_name, "get_weather", "Function name should be 'get_weather'" - ) - self.assertIn("city", args_obj, "Function arguments should have 'city'") - - def test_streaming_multiple_choices_finish_reason(self): - """ - Test: Verify that each choice gets its own finish_reason chunk in streaming mode with n > 1. - This tests the fix for the bug where only the last index got a finish_reason chunk. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The city and state, e.g. San Francisco, CA", - }, - "unit": { - "type": "string", - "enum": ["celsius", "fahrenheit"], - }, - }, - "required": ["location"], - }, - }, - } - ] - - messages = [ - {"role": "user", "content": "What is the weather like in Los Angeles?"} - ] - - # Request with n=2 to get multiple choices - response_stream = client.chat.completions.create( - model=self.model, - messages=messages, - max_tokens=2048, - temperature=0.8, - stream=True, - tools=tools, - tool_choice="required", # Force tool calls - n=2, # Multiple choices - ) - - chunks = list(response_stream) - - # Track finish_reason chunks for each index - finish_reason_chunks = {} - for chunk in chunks: - if chunk.choices: - for choice in chunk.choices: - if choice.finish_reason is not None: - index = choice.index - if index not in finish_reason_chunks: - finish_reason_chunks[index] = [] - finish_reason_chunks[index].append(choice.finish_reason) - - # Verify we got finish_reason chunks for both indices - self.assertEqual( - len(finish_reason_chunks), - 2, - f"Expected finish_reason chunks for 2 indices, got {len(finish_reason_chunks)}", - ) - - # Verify both index 0 and 1 have finish_reason - self.assertIn( - 0, finish_reason_chunks, "Missing finish_reason chunk for index 0" - ) - self.assertIn( - 1, finish_reason_chunks, "Missing finish_reason chunk for index 1" - ) - - # Verify the finish_reason is "tool_calls" since we forced tool calls - for index, reasons in finish_reason_chunks.items(): - self.assertEqual( - reasons[-1], # Last finish_reason for this index - "tool_calls", - f"Expected finish_reason 'tool_calls' for index {index}, got {reasons[-1]}", - ) - - def test_function_calling_streaming_no_tool_call(self): - """ - Test: Whether the finish_reason is stop in streaming mode when no tool call is given. - - Expect no function call to be found. - - Verify that finish_reason is stop - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - tools = [ - { - "type": "function", - "function": { - "name": "get_current_weather", - "description": "Get the current weather in a given location", - "parameters": { - "type": "object", - "properties": { - "city": { - "type": "string", - "description": "The city to find the weather for", - }, - "unit": { - "type": "string", - "description": "Weather unit (celsius or fahrenheit)", - "enum": ["celsius", "fahrenheit"], - }, - }, - "required": ["city", "unit"], - }, - }, - } - ] - - messages = [{"role": "user", "content": "Who are you?"}] - - response_stream = client.chat.completions.create( - model=self.model, - max_tokens=2048, - messages=messages, - temperature=0.8, - top_p=0.8, - stream=True, - tools=tools, - tool_choice="none", - ) - - chunks = list(response_stream) - self.assertTrue(len(chunks) > 0, "Streaming should return at least one chunk") - - found_tool_call = False - for chunk in chunks: - choice = chunk.choices[0] - # Check whether the current chunk contains tool_calls - found_tool_call = choice.delta.tool_calls is not None - - self.assertFalse( - found_tool_call, - "Shouldn't have any tool_call in the streaming chunks", - ) - - finish_reason = chunks[-1].choices[0].finish_reason - self.assertEqual( - finish_reason, - "stop", - "Final response of no function calling should have finish_reason 'stop'", - ) - - def test_streaming_multiple_choices_without_tools(self): - """ - Test: Verify that each choice gets its own finish_reason chunk without tool calls. - This tests the fix for regular content streaming with multiple choices. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - - messages = [{"role": "user", "content": "Say hello in one word."}] - - # Request with n=2 to get multiple choices, no tools - response_stream = client.chat.completions.create( - model=self.model, - messages=messages, - temperature=0.8, - stream=True, - max_tokens=10, # Keep it short - n=2, # Multiple choices - ) - - chunks = list(response_stream) - - # Track finish_reason chunks for each index - finish_reason_chunks = {} - for chunk in chunks: - if chunk.choices: - for choice in chunk.choices: - if choice.finish_reason is not None: - index = choice.index - if index not in finish_reason_chunks: - finish_reason_chunks[index] = [] - finish_reason_chunks[index].append(choice.finish_reason) - - # Verify we got finish_reason chunks for both indices - self.assertEqual( - len(finish_reason_chunks), - 2, - f"Expected finish_reason chunks for 2 indices, got {len(finish_reason_chunks)}", - ) - - # Verify both index 0 and 1 have finish_reason - self.assertIn( - 0, finish_reason_chunks, "Missing finish_reason chunk for index 0" - ) - self.assertIn( - 1, finish_reason_chunks, "Missing finish_reason chunk for index 1" - ) - - # Verify the finish_reason is "stop" (regular completion) - for index, reasons in finish_reason_chunks.items(): - self.assertIn( - reasons[-1], - ["stop", "length"], # Could be either depending on how model responds - f"Expected finish_reason 'stop' or 'length' for index {index}, got {reasons[-1]}", - ) - - -class TestOpenAIPythonicFunctionCalling(CustomTestCase): - """Testcase:Verify the functionality of Python-style list-format function calling with pythonic parser for Llama-3.2-1B-Instruct model on Ascend NPU backend. - Cover: Explicit format prompt verification, streaming call index integrity, and return validity of parallel tool calls. - - [Test Category] Interface - [Test Target] /v1/chat/completions - """ - - PYTHONIC_TOOLS = [ - { - "type": "function", - "function": { - "name": "get_weather", - "description": "Get the current weather for a given location.", - "parameters": { - "type": "object", - "properties": { - "location": { - "type": "string", - "description": "The name of the city or location.", - } - }, - "required": ["location"], - }, - }, - }, - { - "type": "function", - "function": { - "name": "get_tourist_attractions", - "description": "Get a list of top tourist attractions for a given city.", - "parameters": { - "type": "object", - "properties": { - "city": { - "type": "string", - "description": "The name of the city to find attractions for.", - } - }, - "required": ["city"], - }, - }, - }, - ] - - PYTHONIC_MESSAGES = [ - { - "role": "system", - "content": ( - "You are a travel assistant. " - "When asked to call functions, ALWAYS respond ONLY with a python list of function calls, " - "using this format: [func_name1(param1=value1, param2=value2), func_name2(param=value)]. " - "Do NOT use JSON, do NOT use variables, do NOT use any other format. " - "Here is an example:\n" - '[get_weather(location="Paris"), get_tourist_attractions(city="Paris")]' - ), - }, - { - "role": "user", - "content": ( - "I'm planning a trip to Tokyo next week. What's the weather like and what are some top tourist attractions? " - "Propose parallel tool calls at once, using the python list of function calls format as shown above." - ), - }, - ] - - @classmethod - def setUpClass(cls): - cls.model = LLAMA_3_2_1B_INSTRUCT_WEIGHTS_PATH - cls.base_url = DEFAULT_URL_FOR_TEST - cls.api_key = "sk-123456" - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - api_key=cls.api_key, - other_args=[ - "--attention-backend", - "ascend", - "--disable-cuda-graph", - "--tool-call-parser", - "pythonic", - ], - ) - cls.base_url += "/v1" - cls.tokenizer = get_tokenizer(cls.model) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def test_pythonic_tool_call_prompt(self): - """ - Test: Explicit prompt for pythonic tool call format without chat template. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - response = client.chat.completions.create( - model=self.model, - messages=self.PYTHONIC_MESSAGES, - tools=self.PYTHONIC_TOOLS, - temperature=0.1, - stream=False, - ) - tool_calls = response.choices[0].message.tool_calls - self.assertIsInstance(tool_calls, list, "No tool_calls found") - self.assertGreaterEqual(len(tool_calls), 1) - names = [tc.function.name for tc in tool_calls] - self.assertTrue( - "get_weather" in names or "get_tourist_attractions" in names, - f"Function name '{names}' should container either 'get_weather' or 'get_tourist_attractions'", - ) - - def test_pythonic_tool_call_streaming(self): - """ - Test: Streaming pythonic tool call format; assert tool_call index is present. - """ - client = openai.Client(api_key=self.api_key, base_url=self.base_url) - response_stream = client.chat.completions.create( - model=self.model, - messages=self.PYTHONIC_MESSAGES, - tools=self.PYTHONIC_TOOLS, - temperature=0.1, - stream=True, - ) - found_tool_calls = False - found_index = False - found_names = set() - for chunk in response_stream: - choice = chunk.choices[0] - if getattr(choice.delta, "tool_calls", None): - found_tool_calls = True - tool_call = choice.delta.tool_calls[0] - if hasattr(tool_call, "index") or ( - isinstance(tool_call, dict) and "index" in tool_call - ): - found_index = True - found_names.add(str(tool_call.function.name)) - - self.assertTrue(found_tool_calls, "No tool_calls found in streaming response") - self.assertTrue(found_index, "No index field found in any streamed tool_call") - self.assertTrue( - "get_weather" in found_names or "get_tourist_attractions" in found_names, - f"Function name '{found_names}' should container either 'get_weather' or 'get_tourist_attractions'", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/observability/test_priority_metrics.py b/test/registered/observability/test_priority_metrics.py index 593d8b4eb..dd08c3eca 100644 --- a/test/registered/observability/test_priority_metrics.py +++ b/test/registered/observability/test_priority_metrics.py @@ -9,7 +9,6 @@ from prometheus_client.samples import Sample from sglang.srt.observability.metrics_collector import QueueCount from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import ( - register_amd_ci, register_cpu_ci, register_cuda_ci, ) @@ -25,7 +24,6 @@ register_cuda_ci( stage="base-b", runner_config="1-gpu-small", ) -register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd") register_cpu_ci(est_time=49, suite="stage-b-test-cpu-intel") _MODEL_NAME = "Qwen/Qwen3-0.6B" diff --git a/test/registered/observability/test_tracing.py b/test/registered/observability/test_tracing.py index 48dbe4b92..9eef37242 100644 --- a/test/registered/observability/test_tracing.py +++ b/test/registered/observability/test_tracing.py @@ -34,7 +34,7 @@ from sglang.srt.observability.trace import ( ) from sglang.srt.utils import kill_process_tree from sglang.srt.utils.network import get_zmq_socket -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -46,8 +46,7 @@ from sglang.test.test_utils import ( logger = logging.getLogger(__name__) # CI registration -register_cuda_ci(est_time=172, stage="extra-a", runner_config="1-gpu-small") -register_amd_ci(est_time=113, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=113, stage="extra-a", runner_config="1-gpu-small") # ============================================================================ diff --git a/test/registered/openai_server/basic/test_http2_server.py b/test/registered/openai_server/basic/test_http2_server.py index 7bea70a59..db2d5d71b 100644 --- a/test/registered/openai_server/basic/test_http2_server.py +++ b/test/registered/openai_server/basic/test_http2_server.py @@ -11,7 +11,7 @@ import unittest import requests from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -27,8 +27,7 @@ try: except ImportError: _HAS_GRANIAN = False -register_cuda_ci(est_time=108, stage="base-b", runner_config="1-gpu-small") -register_amd_ci(est_time=150, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=150, stage="base-b", runner_config="1-gpu-small") @unittest.skipUnless(_HAS_GRANIAN, "granian not installed (pip install sglang[http2])") diff --git a/test/registered/openai_server/function_call/test_anthropic_tool_use.py b/test/registered/openai_server/function_call/test_anthropic_tool_use.py index 4d3455041..f239b517b 100644 --- a/test/registered/openai_server/function_call/test_anthropic_tool_use.py +++ b/test/registered/openai_server/function_call/test_anthropic_tool_use.py @@ -17,7 +17,6 @@ import requests from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import ( - register_amd_ci, register_cpu_ci, register_cuda_ci, ) @@ -29,8 +28,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=51, stage="base-b", runner_config="1-gpu-large") -register_amd_ci(est_time=140, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=50, stage="base-b", runner_config="1-gpu-large") register_cpu_ci(est_time=54, suite="stage-b-test-cpu-intel") # System message to guide Llama3.2 to produce proper tool call format diff --git a/test/registered/openai_server/function_call/test_openai_function_calling.py b/test/registered/openai_server/function_call/test_openai_function_calling.py index 56288335c..851f84824 100644 --- a/test/registered/openai_server/function_call/test_openai_function_calling.py +++ b/test/registered/openai_server/function_call/test_openai_function_calling.py @@ -3,9 +3,9 @@ import unittest import openai -from sglang.srt.utils import kill_process_tree +from sglang.srt.utils import is_npu, kill_process_tree from sglang.srt.utils.hf_transformers_utils import get_tokenizer -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci, register_npu_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -15,8 +15,27 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=210, stage="base-b", runner_config="1-gpu-large") -register_amd_ci(est_time=73, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=100, stage="base-b", runner_config="1-gpu-large") +# Backend-specific: Ascend uses a local model mirror and its native +# attention backend, while sharing the protocol assertions below. +register_npu_ci(est_time=400, suite="full-1-npu-a3", nightly=True) + + +def _model_path(): + if is_npu(): + from sglang.test.ascend.test_ascend_utils import ( + LLAMA_3_2_1B_INSTRUCT_WEIGHTS_PATH, + ) + + return LLAMA_3_2_1B_INSTRUCT_WEIGHTS_PATH + return DEFAULT_SMALL_MODEL_NAME_FOR_TEST + + +def _server_args(parser): + args = ["--tool-call-parser", parser] + if is_npu(): + args[:0] = ["--attention-backend", "ascend", "--disable-cuda-graph"] + return args class TestOpenAIServerFunctionCalling(CustomTestCase): @@ -36,8 +55,7 @@ class TestOpenAIServerFunctionCalling(CustomTestCase): @classmethod def setUpClass(cls): - # Replace with the model name needed for testing; if not required, reuse DEFAULT_SMALL_MODEL_NAME_FOR_TEST - cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST + cls.model = _model_path() cls.base_url = DEFAULT_URL_FOR_TEST cls.api_key = "sk-123456" @@ -47,11 +65,7 @@ class TestOpenAIServerFunctionCalling(CustomTestCase): cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, api_key=cls.api_key, - other_args=[ - # If your server needs extra parameters to test function calling, please add them here. - "--tool-call-parser", - "llama3", - ], + other_args=_server_args("llama3"), ) cls.base_url += "/v1" cls.tokenizer = get_tokenizer(cls.model) @@ -97,7 +111,7 @@ class TestOpenAIServerFunctionCalling(CustomTestCase): {"role": "system", "content": self.SYSTEM_MESSAGE}, {"role": "user", "content": "Compute (3+5)"}, ] - response = client.chat.completions.create( + request = dict( model=self.model, max_tokens=2048, messages=messages, @@ -105,8 +119,12 @@ class TestOpenAIServerFunctionCalling(CustomTestCase): top_p=0.8, stream=False, tools=tools, - tool_choice="required", ) + # Ascend keeps the historical auto-choice coverage; CUDA forces the + # call so this assertion never depends on a stochastic model decision. + if not is_npu(): + request["tool_choice"] = "required" + response = client.chat.completions.create(**request) tool_calls = response.choices[0].message.tool_calls @@ -843,7 +861,7 @@ class TestOpenAIPythonicFunctionCalling(CustomTestCase): @classmethod def setUpClass(cls): - cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST + cls.model = _model_path() cls.base_url = DEFAULT_URL_FOR_TEST cls.api_key = "sk-123456" cls.process = popen_launch_server( @@ -851,10 +869,7 @@ class TestOpenAIPythonicFunctionCalling(CustomTestCase): cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, api_key=cls.api_key, - other_args=[ - "--tool-call-parser", - "pythonic", - ], + other_args=_server_args("pythonic"), ) cls.base_url += "/v1" cls.tokenizer = get_tokenizer(cls.model) @@ -922,6 +937,7 @@ class TestOpenAIPythonicFunctionCalling(CustomTestCase): is_rust_server_built(), "embedded rust server extension not built", ) +@unittest.skipIf(is_npu(), "the embedded Rust server is not an Ascend path") class TestOpenAIFunctionCallingWithRust(TestOpenAIServerFunctionCalling): """Run the registered unary/streaming function-call suite through Rust.""" @@ -946,6 +962,7 @@ class TestOpenAIFunctionCallingWithRust(TestOpenAIServerFunctionCalling): is_rust_server_built(), "embedded rust server extension not built", ) +@unittest.skipIf(is_npu(), "the embedded Rust server is not an Ascend path") class TestOpenAIPythonicFunctionCallingWithRust(TestOpenAIPythonicFunctionCalling): """Run Pythonic unary/streaming tool calls through Rust.""" diff --git a/test/registered/openai_server/validation/test_large_max_new_tokens.py b/test/registered/openai_server/validation/test_large_max_new_tokens.py index 8daa7a83b..2ce6f44c4 100644 --- a/test/registered/openai_server/validation/test_large_max_new_tokens.py +++ b/test/registered/openai_server/validation/test_large_max_new_tokens.py @@ -12,7 +12,6 @@ import openai from sglang.srt.utils import kill_process_tree from sglang.srt.utils.hf_transformers_utils import get_tokenizer from sglang.test.ci.ci_register import ( - register_amd_ci, register_cpu_ci, register_cuda_ci, ) @@ -26,8 +25,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=59, stage="base-b", runner_config="1-gpu-large") -register_amd_ci(est_time=41, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=58, stage="base-b", runner_config="1-gpu-large") register_cpu_ci(est_time=101, suite="stage-b-test-cpu-intel") diff --git a/test/registered/openai_server/validation/test_matched_stop.py b/test/registered/openai_server/validation/test_matched_stop.py index d2242d54b..4846103f8 100644 --- a/test/registered/openai_server/validation/test_matched_stop.py +++ b/test/registered/openai_server/validation/test_matched_stop.py @@ -2,7 +2,6 @@ import unittest from sglang.srt.utils import kill_process_tree from sglang.test.ci.ci_register import ( - register_amd_ci, register_cpu_ci, register_cuda_ci, ) @@ -14,8 +13,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=63, stage="base-b", runner_config="1-gpu-small") -register_amd_ci(est_time=60, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=52, stage="base-b", runner_config="1-gpu-small") register_cpu_ci(est_time=83, suite="stage-b-test-cpu-intel") diff --git a/test/registered/openai_server/validation/test_request_length_validation.py b/test/registered/openai_server/validation/test_request_length_validation.py index c53a1148e..f8810771c 100644 --- a/test/registered/openai_server/validation/test_request_length_validation.py +++ b/test/registered/openai_server/validation/test_request_length_validation.py @@ -4,7 +4,7 @@ import openai import requests from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.test_utils import ( DEFAULT_SMALL_MODEL_NAME_FOR_TEST, DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, @@ -13,8 +13,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=50, stage="base-b", runner_config="1-gpu-large") -register_amd_ci(est_time=31, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=49, stage="base-b", runner_config="1-gpu-large") class TestRequestLengthValidation(CustomTestCase): diff --git a/test/registered/perf/test_text_models_perf.py b/test/registered/perf/models/test_text_models_perf.py similarity index 100% rename from test/registered/perf/test_text_models_perf.py rename to test/registered/perf/models/test_text_models_perf.py diff --git a/test/registered/perf/test_vlms_perf.py b/test/registered/perf/models/test_vlms_perf.py similarity index 100% rename from test/registered/perf/test_vlms_perf.py rename to test/registered/perf/models/test_vlms_perf.py diff --git a/test/registered/profiling/test_profile_v2.py b/test/registered/profiling/test_profile_v2.py deleted file mode 100644 index ab21fefad..000000000 --- a/test/registered/profiling/test_profile_v2.py +++ /dev/null @@ -1,109 +0,0 @@ -import os -import shutil -import tempfile -import unittest -from pathlib import Path - -import requests - -from sglang.srt.environ import envs -from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_cuda_ci -from sglang.test.test_utils import ( - DEFAULT_SMALL_MODEL_NAME_FOR_TEST, - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - DEFAULT_URL_FOR_TEST, - CustomTestCase, - popen_launch_server, -) - -register_cuda_ci( - est_time=120, - stage="base-b", - runner_config="1-gpu-small", - disabled="Temporarily disabled", -) - - -class TestStartProfile(CustomTestCase): - @classmethod - def setUpClass(cls): - cls.output_dir = tempfile.mkdtemp() - envs.SGLANG_TORCH_PROFILER_DIR.set(cls.output_dir) - envs.SGLANG_PROFILE_V2.set(True) - cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def setUp(self): - self._clear_profile_dir() - - def test_profile_by_stage(self): - self._start_profile( - profile_by_stage=True, - num_steps=10, - ) - - self._post_request() - - self._check_profile_output(pattern="*-prefill*", expect_existence=True) - self._check_profile_output(pattern="*-decode*", expect_existence=True) - - def test_decode_only(self): - self._start_profile( - profile_by_stage=True, - profile_stages=["decode"], - num_steps=10, - ) - - self._post_request() - - self._check_profile_output(pattern="*-prefill*", expect_existence=False) # NOTE - self._check_profile_output(pattern="*-decode*", expect_existence=True) - - def _start_profile(self, **kwargs): - """Start profiling with optional parameters.""" - response = requests.post( - f"{DEFAULT_URL_FOR_TEST}/start_profile", - json=kwargs if kwargs else None, - ) - self.assertEqual(response.status_code, 200) - - def _post_request(self): - response = requests.post( - f"{DEFAULT_URL_FOR_TEST}/generate", - json={ - "text": "The capital of France is", - "sampling_params": { - "temperature": 0, - "max_new_tokens": 32, - }, - }, - ) - self.assertEqual(response.status_code, 200) - - def _clear_profile_dir(self): - if os.path.isdir(self.output_dir): - shutil.rmtree(self.output_dir) - - def _check_profile_output(self, pattern: str, expect_existence: bool): - self.assertTrue( - os.path.isdir(self.output_dir), "Output directory does not exist." - ) - self.assertEqual( - len(list(Path(self.output_dir).glob(pattern))) > 0, - expect_existence, - f"Does not find {pattern=} ({list(Path(self.output_dir).glob('**/*'))=})", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_hybrid_bitexact.py b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_hybrid_bitexact.py index f61553ae8..c7c5358ba 100644 --- a/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_hybrid_bitexact.py +++ b/test/registered/radix_cache/unified_radix_tree/test_unified_radix_cache_kl_hybrid_bitexact.py @@ -73,7 +73,7 @@ from sglang.test.test_utils import ( unified_radix_tree_server_env, ) -register_cuda_ci(est_time=790, stage="base-b", runner_config="1-gpu-large") +register_cuda_ci(est_time=2300, stage="extra-a", runner_config="1-gpu-large") _MODEL_PATH = os.environ.get("INKLING_TEST_MODEL_PATH", "thinkingmachines/Inkling") _MODEL_REVISION = os.environ.get("INKLING_TEST_MODEL_REVISION", "test") diff --git a/test/registered/reasoning/test_reasoning.py b/test/registered/reasoning/test_reasoning.py index 6af8da2a0..1a53b476f 100644 --- a/test/registered/reasoning/test_reasoning.py +++ b/test/registered/reasoning/test_reasoning.py @@ -4,7 +4,7 @@ import unittest import requests from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci +from sglang.test.ci.ci_register import register_cuda_ci from sglang.test.kits.reasoning_kit import ( ReasoningTokenUsageMixin, SeparateReasoningMixin, @@ -17,8 +17,7 @@ from sglang.test.test_utils import ( popen_launch_server, ) -register_cuda_ci(est_time=100, stage="base-b", runner_config="1-gpu-large") -register_amd_ci(est_time=200, suite="stage-b-test-1-gpu-small-amd") +register_cuda_ci(est_time=129, stage="base-b", runner_config="1-gpu-large") class TestEnableThinking( diff --git a/test/registered/rl/test_update_weights_from_disk.py b/test/registered/rl/test_update_weights_from_disk.py deleted file mode 100644 index ed7255916..000000000 --- a/test/registered/rl/test_update_weights_from_disk.py +++ /dev/null @@ -1,447 +0,0 @@ -import json -import random -import time -import unittest -from concurrent.futures import ThreadPoolExecutor, as_completed - -import requests - -import sglang as sgl -from sglang.srt.utils import kill_process_tree -from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci -from sglang.test.test_utils import ( - DEFAULT_SMALL_MODEL_NAME_FOR_TEST, - DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - DEFAULT_URL_FOR_TEST, - CustomTestCase, - is_in_ci, - popen_launch_server, -) - -register_amd_ci( - est_time=210, suite="stage-b-test-1-gpu-small-amd", disabled="see #14021" -) -register_cuda_ci( - est_time=210, stage="base-b", runner_config="1-gpu-large", disabled="see #14021" -) - - -############################################################################### -# Engine Mode Tests (Single-configuration) -############################################################################### -class TestEngineUpdateWeightsFromDisk(CustomTestCase): - def setUp(self): - self.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST - # Initialize the engine in offline (direct) mode. - self.engine = sgl.Engine(model_path=self.model) - - def tearDown(self): - self.engine.shutdown() - - def run_decode(self): - prompts = ["The capital of France is"] - sampling_params = {"temperature": 0, "max_new_tokens": 32} - outputs = self.engine.generate(prompts, sampling_params) - print("=" * 100) - print( - f"[Engine Mode] Prompt: {prompts[0]}\nGenerated text: {outputs[0]['text']}" - ) - return outputs[0]["text"] - - def run_update_weights(self, model_path): - ret = self.engine.update_weights_from_disk(model_path) - print(json.dumps(ret)) - return ret - - def test_update_weights(self): - origin_response = self.run_decode() - # Update weights: use new model (remove "-Instruct") - new_model_path = self.model.replace("-Instruct", "") - ret = self.run_update_weights(new_model_path) - self.assertTrue(ret[0]) # ret is a tuple; index 0 holds the success flag - - updated_response = self.run_decode() - self.assertNotEqual(origin_response[:32], updated_response[:32]) - - # Revert back to original weights - ret = self.run_update_weights(self.model) - self.assertTrue(ret[0]) - reverted_response = self.run_decode() - self.assertEqual(origin_response[:32], reverted_response[:32]) - - def test_update_weights_unexist_model(self): - origin_response = self.run_decode() - new_model_path = self.model.replace("-Instruct", "wrong") - ret = self.run_update_weights(new_model_path) - self.assertFalse(ret[0]) - updated_response = self.run_decode() - self.assertEqual(origin_response[:32], updated_response[:32]) - - -############################################################################### -# HTTP Server Mode Tests (Single-configuration) -############################################################################### -class TestServerUpdateWeightsFromDisk(CustomTestCase): - @classmethod - def setUpClass(cls): - cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, cls.base_url, timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def run_decode(self): - response = requests.post( - self.base_url + "/generate", - json={ - "text": "The capital of France is", - "sampling_params": {"temperature": 0, "max_new_tokens": 32}, - }, - ) - print("=" * 100) - print(f"[Server Mode] Generated text: {response.json()['text']}") - return response.json()["text"] - - def run_decode_random(self, max_new_tokens=32): - response = requests.post( - self.base_url + "/generate", - json={ - "text": f"Question: {random.randint(0, 100)},The capital of France is", - "sampling_params": { - "temperature": 0, - "max_new_tokens": max_new_tokens, - "ignore_eos": True, - }, - }, - ) - return response.json() - - def get_model_info(self): - response = requests.get(self.base_url + "/get_model_info") - model_path = response.json()["model_path"] - print(json.dumps(response.json())) - return model_path - - def run_update_weights(self, model_path, flush_cache=True): - response = requests.post( - self.base_url + "/update_weights_from_disk", - json={ - "model_path": model_path, - "flush_cache": flush_cache, - }, - ) - ret = response.json() - return ret - - def pause_generation(self, mode): - response = requests.post( - self.base_url + "/pause_generation", - json={"mode": mode}, - ) - ret = response.json() - return ret - - def continue_generation(self): - response = requests.post( - self.base_url + "/continue_generation", - json={}, - ) - ret = response.json() - return ret - - def test_update_weights(self): - origin_model_path = self.get_model_info() - print(f"[Server Mode] origin_model_path: {origin_model_path}") - origin_response = self.run_decode() - - new_model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST.replace("-Instruct", "") - ret = self.run_update_weights(new_model_path) - self.assertTrue(ret["success"]) - - updated_model_path = self.get_model_info() - print(f"[Server Mode] updated_model_path: {updated_model_path}") - self.assertEqual(updated_model_path, new_model_path) - self.assertNotEqual(updated_model_path, origin_model_path) - - updated_response = self.run_decode() - self.assertNotEqual(origin_response[:32], updated_response[:32]) - - ret = self.run_update_weights(origin_model_path) - self.assertTrue(ret["success"]) - updated_model_path = self.get_model_info() - self.assertEqual(updated_model_path, origin_model_path) - - updated_response = self.run_decode() - self.assertEqual(origin_response[:32], updated_response[:32]) - - def test_update_weights_non_blocking(self): - origin_model_path = self.get_model_info() - print(f"[Server Mode] origin_model_path: {origin_model_path}") - - pause_generation_modes = ["in_place", "retract"] - for pause_generation_mode in pause_generation_modes: - num_requests = 32 - with ThreadPoolExecutor(num_requests) as executor: - futures = [ - executor.submit(self.run_decode_random, 1600) - for _ in range(num_requests) - ] - - # ensure the decode has been started - time.sleep(2) - - new_model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST.replace( - "-Instruct", "" - ) - ret = self.pause_generation(pause_generation_mode) - ret = self.run_update_weights( - new_model_path, flush_cache=pause_generation_mode == "retract" - ) - self.assertTrue(ret["success"]) - ret = self.continue_generation() - - for future in as_completed(futures): - self.assertNotEqual( - future.result()["meta_info"]["finish_reason"]["type"], "abort" - ) - - updated_model_path = self.get_model_info() - print(f"[Server Mode] updated_model_path: {updated_model_path}") - self.assertEqual(updated_model_path, new_model_path) - self.assertNotEqual(updated_model_path, origin_model_path) - - def test_update_weights_unexist_model(self): - origin_model_path = self.get_model_info() - print(f"[Server Mode] origin_model_path: {origin_model_path}") - origin_response = self.run_decode() - - new_model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST.replace("-Instruct", "wrong") - ret = self.run_update_weights(new_model_path) - self.assertFalse(ret["success"]) - - updated_model_path = self.get_model_info() - print(f"[Server Mode] updated_model_path: {updated_model_path}") - self.assertEqual(updated_model_path, origin_model_path) - - updated_response = self.run_decode() - self.assertEqual(origin_response[:32], updated_response[:32]) - - -class TestServerUpdateWeightsFromDiskAbortAllRequests(CustomTestCase): - @classmethod - def setUpClass(cls): - cls.model = DEFAULT_SMALL_MODEL_NAME_FOR_TEST - cls.base_url = DEFAULT_URL_FOR_TEST - cls.process = popen_launch_server( - cls.model, - cls.base_url, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=["--max-running-requests", 8], - ) - - @classmethod - def tearDownClass(cls): - kill_process_tree(cls.process.pid) - - def run_decode(self, max_new_tokens=32): - response = requests.post( - self.base_url + "/generate", - json={ - "text": "The capital of France is", - "sampling_params": { - "temperature": 0, - "max_new_tokens": max_new_tokens, - "ignore_eos": True, - }, - }, - ) - return response.json() - - def get_model_info(self): - response = requests.get(self.base_url + "/get_model_info") - model_path = response.json()["model_path"] - print(json.dumps(response.json())) - return model_path - - def run_update_weights(self, model_path, abort_all_requests=False): - response = requests.post( - self.base_url + "/update_weights_from_disk", - json={ - "model_path": model_path, - "abort_all_requests": abort_all_requests, - }, - ) - ret = response.json() - print(json.dumps(ret)) - return ret - - def test_update_weights_abort_all_requests(self): - origin_model_path = self.get_model_info() - print(f"[Server Mode] origin_model_path: {origin_model_path}") - - num_requests = 32 - with ThreadPoolExecutor(num_requests) as executor: - futures = [ - executor.submit(self.run_decode, 16000) for _ in range(num_requests) - ] - - # ensure the decode has been started - time.sleep(2) - - new_model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST.replace("-Instruct", "") - ret = self.run_update_weights(new_model_path, abort_all_requests=True) - self.assertTrue(ret["success"]) - - for future in as_completed(futures): - self.assertEqual( - future.result()["meta_info"]["finish_reason"]["type"], "abort" - ) - - updated_model_path = self.get_model_info() - print(f"[Server Mode] updated_model_path: {updated_model_path}") - self.assertEqual(updated_model_path, new_model_path) - self.assertNotEqual(updated_model_path, origin_model_path) - - -############################################################################### -# Parameterized Tests for update_weights_from_disk -# Test coverage is determined based on the value of is_in_ci: -# - In a CI environment: randomly select one mode (Engine or Server) and test only with tp=1, dp=1. -# - In a non-CI environment: test both Engine and Server modes, and enumerate all combinations -# with tp and dp ranging from 1 to 2. -############################################################################### -class TestUpdateWeightsFromDiskParameterized(CustomTestCase): - def run_common_test(self, mode, tp, dp): - """ - Common test procedure for update_weights_from_disk. - For Engine mode, we instantiate the engine with tp_size=tp. - For Server mode, we launch the server with additional arguments for tp (dp is not used in server launch here). - """ - if mode == "Engine": - # Instantiate engine with additional parameter tp_size. - print(f"[Parameterized Engine] Testing with tp={tp}, dp={dp}") - engine = sgl.Engine( - model_path=DEFAULT_SMALL_MODEL_NAME_FOR_TEST, - random_seed=42, - tp_size=tp, - # dp parameter is not explicitly used in this API. - ) - try: - origin_response = self._engine_update_weights_test(engine) - finally: - engine.shutdown() - elif mode == "Server": - print(f"[Parameterized Server] Testing with tp={tp}, dp={dp}") - # Pass additional arguments to launch the server. - base_args = ["--tp-size", str(tp)] - process = popen_launch_server( - DEFAULT_SMALL_MODEL_NAME_FOR_TEST, - DEFAULT_URL_FOR_TEST, - timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH, - other_args=base_args, - ) - try: - origin_response = self._server_update_weights_test(DEFAULT_URL_FOR_TEST) - finally: - kill_process_tree(process.pid) - else: - raise ValueError(f"Unknown mode: {mode}") - - def _engine_update_weights_test(self, engine): - # Run the update weights test on the given engine instance. - def run_decode(): - prompts = ["The capital of France is"] - sampling_params = {"temperature": 0, "max_new_tokens": 32} - outputs = engine.generate(prompts, sampling_params) - print("=" * 100) - print( - f"[Parameterized Engine] Prompt: {prompts[0]}\nGenerated text: {outputs[0]['text']}" - ) - return outputs[0]["text"] - - def run_update_weights(model_path): - ret = engine.update_weights_from_disk(model_path) - print(json.dumps(ret)) - return ret - - origin_response = run_decode() - new_model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST.replace("-Instruct", "") - ret = run_update_weights(new_model_path) - self.assertTrue(ret[0]) - updated_response = run_decode() - self.assertNotEqual(origin_response[:32], updated_response[:32]) - ret = run_update_weights(DEFAULT_SMALL_MODEL_NAME_FOR_TEST) - self.assertTrue(ret[0]) - reverted_response = run_decode() - self.assertEqual(origin_response[:32], reverted_response[:32]) - return origin_response - - def _server_update_weights_test(self, base_url): - def run_decode(): - response = requests.post( - base_url + "/generate", - json={ - "text": "The capital of France is", - "sampling_params": {"temperature": 0, "max_new_tokens": 32}, - }, - ) - print("=" * 100) - print(f"[Parameterized Server] Generated text: {response.json()['text']}") - return response.json()["text"] - - def get_model_info(): - response = requests.get(base_url + "/get_model_info") - model_path = response.json()["model_path"] - print(json.dumps(response.json())) - return model_path - - def run_update_weights(model_path): - response = requests.post( - base_url + "/update_weights_from_disk", - json={"model_path": model_path}, - ) - ret = response.json() - print(json.dumps(ret)) - return ret - - origin_model_path = get_model_info() - origin_response = run_decode() - new_model_path = DEFAULT_SMALL_MODEL_NAME_FOR_TEST.replace("-Instruct", "") - ret = run_update_weights(new_model_path) - self.assertTrue(ret["success"]) - updated_model_path = get_model_info() - self.assertEqual(updated_model_path, new_model_path) - self.assertNotEqual(updated_model_path, origin_model_path) - updated_response = run_decode() - self.assertNotEqual(origin_response[:32], updated_response[:32]) - ret = run_update_weights(origin_model_path) - self.assertTrue(ret["success"]) - updated_model_path = get_model_info() - self.assertEqual(updated_model_path, origin_model_path) - reverted_response = run_decode() - self.assertEqual(origin_response[:32], reverted_response[:32]) - return origin_response - - def test_parameterized_update_weights(self): - if is_in_ci(): - # In CI, choose one random mode (Engine or Server) with tp=1, dp=1. - mode = random.choice(["Engine", "Server"]) - test_suits = [(1, 1, mode)] - else: - # Otherwise, test both modes and enumerate tp,dp combinations from 1 to 2. - test_suits = [] - for mode in ["Engine", "Server"]: - for tp in [1, 2]: - for dp in [1, 2]: - test_suits.append((tp, dp, mode)) - for tp, dp, mode in test_suits: - with self.subTest(mode=mode, tp=tp, dp=dp): - self.run_common_test(mode, tp, dp) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/spec/eagle/test_spec_eagle_stress.py b/test/registered/spec/eagle/test_spec_eagle_stress.py index cea766e22..a6b229cae 100644 --- a/test/registered/spec/eagle/test_spec_eagle_stress.py +++ b/test/registered/spec/eagle/test_spec_eagle_stress.py @@ -17,7 +17,7 @@ from sglang.test.kits.spec_server_kits import ( ) from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base -register_cuda_ci(est_time=675, stage="base-b", runner_config="1-gpu-large") +register_cuda_ci(est_time=440, stage="extra-a", runner_config="1-gpu-large") class TestEagle3Perf(Eagle3Base, SpecPerfKit): diff --git a/test/registered/stress/test_stress_deepseek_v3.py b/test/registered/stress/models/test_stress_deepseek_v3.py similarity index 100% rename from test/registered/stress/test_stress_deepseek_v3.py rename to test/registered/stress/models/test_stress_deepseek_v3.py diff --git a/test/registered/stress/test_stress_glm_4_6.py b/test/registered/stress/models/test_stress_glm_4_6.py similarity index 100% rename from test/registered/stress/test_stress_glm_4_6.py rename to test/registered/stress/models/test_stress_glm_4_6.py diff --git a/test/registered/stress/test_stress_kimi_k2.py b/test/registered/stress/models/test_stress_kimi_k2.py similarity index 100% rename from test/registered/stress/test_stress_kimi_k2.py rename to test/registered/stress/models/test_stress_kimi_k2.py diff --git a/test/registered/stress/test_stress_qwen3_235b.py b/test/registered/stress/models/test_stress_qwen3_235b.py similarity index 100% rename from test/registered/stress/test_stress_qwen3_235b.py rename to test/registered/stress/models/test_stress_qwen3_235b.py diff --git a/test/registered/unit/README.md b/test/registered/unit/README.md index 711e729f3..5a5fb66f9 100644 --- a/test/registered/unit/README.md +++ b/test/registered/unit/README.md @@ -1,7 +1,8 @@ # Unit Tests -Component-level tests that do **not** launch a server or load model weights. -Tests can use CPU or GPU — the key criterion is **no server process**. +CPU-only component tests that do **not** launch a server, load model weights, +or require an accelerator. GPU operator correctness belongs under +`test/registered/kernel//`. ## Quick Start @@ -15,11 +16,7 @@ Tests can use CPU or GPU — the key criterion is **no server process**. ```python from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=5, suite="base-a-test-cpu") - # or: register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small") ``` - CUDA suites whose names follow `{stage}-test-{runner_config}` use the - `stage=` + `runner_config=` form. Legacy `suite="..."` is kept for - nightly/stress/weekly + AMD/CPU/NPU suites that don't fit that shape. 4. Run locally: ```bash pytest test/registered/unit/ -v # all unit tests @@ -95,4 +92,6 @@ process. If you must stub, use `patch.dict("sys.modules", ...)` with proper clea - **No** `popen_launch_server()` or `Engine(...)`. - **No** model weight loading. - Use `CustomTestCase` (from `sglang.test.test_utils`, adds CI retry). -- Use `unittest.mock` for dependencies that are expensive to construct. +- Mock external or slow dependency boundaries only when the assertion still + checks a result, state transition, protocol output, or error. A test that + proves only that its mock was called is not sufficient. diff --git a/test/registered/unit/distributed/test_pynccl_allocator_import.py b/test/registered/unit/distributed/test_pynccl_allocator_import.py deleted file mode 100644 index 40ee390ca..000000000 --- a/test/registered/unit/distributed/test_pynccl_allocator_import.py +++ /dev/null @@ -1,62 +0,0 @@ -"""Regression test for https://github.com/sgl-project/sglang/issues/28999. - -``pynccl_allocator`` must not import private ``torch.cuda.memory`` symbols at -module scope: they are absent before torch 2.8 and abort startup on Ascend NPU. -""" - -import ast -import unittest -from pathlib import Path - -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=10, suite="base-a-test-cpu") - -# test/registered/unit/distributed/ -> repo root -REPO_ROOT = Path(__file__).resolve().parents[4] -SOURCE_PATH = ( - REPO_ROOT / "python/sglang/srt/distributed/device_communicators/pynccl_allocator.py" -) - - -def _import_time_nodes(tree: ast.Module): - """Yield nodes that run at import time, including ``try`` / ``if`` bodies.""" - stack = list(tree.body) - while stack: - node = stack.pop() - if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): - continue - yield node - stack.extend(ast.iter_child_nodes(node)) - - -class TestPyncclAllocatorImportGuard(CustomTestCase): - def test_no_import_time_private_cuda_memory_symbols(self): - self.assertTrue( - SOURCE_PATH.is_file(), - f"cannot locate pynccl_allocator.py at {SOURCE_PATH}; " - "update REPO_ROOT if the tree layout changed", - ) - tree = ast.parse(SOURCE_PATH.read_text(), filename=str(SOURCE_PATH)) - - offenders = [ - alias.name - for node in _import_time_nodes(tree) - if isinstance(node, ast.ImportFrom) and node.module == "torch.cuda.memory" - for alias in node.names - if alias.name.startswith("_cuda_") - ] - - self.assertEqual( - offenders, - [], - "pynccl_allocator must not import private torch.cuda.memory symbols " - f"at module scope (found {offenders}); these are absent on torch<2.8 " - "and break startup on Ascend NPU. Reach them via torch._C. at " - "the call site instead.", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/entrypoints/test_effective_state_surfaces.py b/test/registered/unit/entrypoints/test_effective_state_surfaces.py deleted file mode 100644 index 32805722b..000000000 --- a/test/registered/unit/entrypoints/test_effective_state_surfaces.py +++ /dev/null @@ -1,401 +0,0 @@ -"""Every serving surface can answer what is running, not only what was asked. - -`/server_info` and its gRPC and in-process twins report the startup record, -parsers included -- the launcher resolves `auto` into the record before -publishing. What changes after publication -- the model a weight update -swapped in, its load format, an operator-set weight version -- is reported by -the model-info surface, and there is one per entry point: HTTP, gRPC and -`Engine`. Adding a field to one and forgetting the others leaves that entry -point's users with no way to see it, which no test notices because each -surface passes its own tests. - -The required set has two halves. The derived half comes from the control-plane -writers: whatever a process writes with `override` after publication is exactly -what can differ from the record, and both the keyword and the `**`-expansion -shapes resolve statically here. The second half is a policy, not a derivation -- -what any one surface reports effectively, all of them owe their users -- so a -field every surface drops at once leaves the set with it. Each surface must both -carry the key and take its value from the effective config: reading -`server_args.` under the right key reports the startup value with a -straight face. -""" - -import ast -import inspect -import pathlib -import re -import unittest - -import sglang -from sglang.srt.entrypoints.engine import Engine -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=13, suite="base-a-test-cpu") - -_PACKAGE_ROOT = pathlib.Path(sglang.__file__).resolve().parent -_RUST_MODEL_INFO = ( - pathlib.Path(__file__).resolve().parents[4] - / "rust/sglang-server/src/api_server/common.rs" -) - -# The tokenizer process is the one whose control-plane writes a served request -# can observe; a field it overrides there is a post-launch fact. -_TOKENIZER_WRITERS = ( - "srt/managers/tokenizer_manager.py", - "srt/managers/tokenizer_control_mixin.py", - "srt/entrypoints/http_server.py", - "srt/entrypoints/engine.py", -) - -# The model path and the served name stay manager attributes: a weight update -# moves them on the manager rather than through `override`, so the writer -# derivation cannot see them. They are subtracted from the derived half and -# asserted directly instead -- an exemption whose premise ("every surface -# reports them") is checked, not assumed. It was not true when it was written: -# only `Engine` carried `served_model_name`, and the router had to read the -# launch record off `/server_info` to learn a name a weight update had moved. -_MANAGER_ATTRIBUTES = {"model_path", "served_model_name"} - - -def _hicache_status_fields() -> set: - """The fields `GET /hicache/storage-backend` answers with. - - The HiCache mirror is written post-publish and reported by its own - endpoint, so the model-info surfaces do not owe it. Taking the set from - that handler is what keeps the exemption honest: a field it stops - reporting falls back to them. - """ - tree = ast.parse((_PACKAGE_ROOT / "srt/entrypoints/http_server.py").read_text()) - for fn in ast.walk(tree): - if ( - not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)) - or fn.name != "hicache_storage_backend_status" - ): - continue - fields = { - elt.value - for comp in ast.walk(fn) - if isinstance(comp, ast.DictComp) - for gen in comp.generators - for elt in getattr(gen.iter, "elts", []) - if isinstance(elt, ast.Constant) and isinstance(elt.value, str) - } - assert fields, "the HiCache status handler names no field" - return fields - raise AssertionError( - "no `hicache_storage_backend_status` handler: the HiCache fields have " - "no endpoint of their own and fall to the model-info surfaces" - ) - - -def _expanded_write_keys(rel: str, tree: ast.AST, call: ast.Call, kw: ast.keyword): - """The field names behind a `**` at a control-plane writer call. - - Resolves a dict literal -- constant keys, or a key bound by an enclosing - literal `for` -- and a name bound to a dict literal in the enclosing - function, including constant-subscript stores onto it. A `**` that - forwards its own function's `**kwargs` names no field: its callers do. - Anything else raises, because a skipped expansion shrinks the required - set instead of failing. - """ - enclosing = None - for fn in ast.walk(tree): - if isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)) and ( - fn.lineno <= call.lineno <= (fn.end_lineno or fn.lineno) - ): - if enclosing is None or fn.lineno > enclosing.lineno: - enclosing = fn - - def loop_bound(name: str) -> set: - """The values a `for name, ... in ()` around the call binds.""" - values = set() - for node in ast.walk(tree): - if not isinstance(node, ast.For): - continue - target = node.target - names = ( - [target] - if isinstance(target, ast.Name) - else list(getattr(target, "elts", [])) - ) - if not names or not isinstance(names[0], ast.Name) or names[0].id != name: - continue - if not (node.lineno <= call.lineno <= (node.end_lineno or node.lineno)): - continue - for item in getattr(node.iter, "elts", []): - first = ( - item.elts[0] if isinstance(item, ast.Tuple) and item.elts else item - ) - if isinstance(first, ast.Constant) and isinstance(first.value, str): - values.add(first.value) - return values - - def dict_keys(node: ast.Dict) -> set: - keys = set() - for key in node.keys: - if isinstance(key, ast.Constant): - keys.add(key.value) - continue - assert isinstance(key, ast.Name), ( - f"non-literal dict key in a writer expansion at {rel}:{call.lineno}" - ) - bound = loop_bound(key.id) - assert bound, ( - f"dict key {key.id!r} at {rel}:{call.lineno} is not bound by a " - "literal loop; extend the resolver" - ) - keys |= bound - return keys - - if isinstance(kw.value, ast.Dict): - return dict_keys(kw.value) - assert isinstance(kw.value, ast.Name), ( - f"unresolvable writer expansion at {rel}:{call.lineno}" - ) - assert enclosing is not None, ( - f"writer expansion outside any function at {rel}:{call.lineno}" - ) - name = kw.value.id - if enclosing.args.kwarg is not None and enclosing.args.kwarg.arg == name: - return set() - keys = set() - found = False - for node in ast.walk(enclosing): - if not (isinstance(node, ast.Assign) and len(node.targets) == 1): - continue - target = node.targets[0] - if isinstance(target, ast.Name) and target.id == name: - assert isinstance(node.value, ast.Dict), ( - f"writer expansion {name!r} at {rel}:{call.lineno} is assigned " - "something other than a dict literal; extend the resolver" - ) - found = True - keys |= dict_keys(node.value) - elif ( - isinstance(target, ast.Subscript) - and isinstance(target.value, ast.Name) - and target.value.id == name - and isinstance(target.slice, ast.Constant) - ): - keys.add(target.slice.value) - assert found, ( - f"writer expansion {name!r} at {rel}:{call.lineno} has no dict-literal " - "assignment in its function; extend the resolver" - ) - return keys - - -def _overridden_fields() -> set: - """Fields the tokenizer process writes after publication.""" - fields = set() - for rel in _TOKENIZER_WRITERS: - tree = ast.parse((_PACKAGE_ROOT / rel).read_text()) - for node in ast.walk(tree): - if not isinstance(node, ast.Call): - continue - name = ( - node.func.attr - if isinstance(node.func, ast.Attribute) - else getattr(node.func, "id", "") - ) - if name not in ("override", "record_config_updates"): - continue - for kw in node.keywords: - if kw.arg == "source": - # The provenance label, not a config field. - continue - if kw.arg: - fields.add(kw.arg) - else: - fields |= _expanded_write_keys(rel, tree, node, kw) - # Only the mirror's own fields earn the endpoint exemption: adding an - # unrelated field to that handler must not buy it a pass here. - hicache = {f for f in _hicache_status_fields() if f.startswith("hicache_")} - return fields - _MANAGER_ATTRIBUTES - hicache - - -def _effective_reads_in(source: str, func_name: str) -> set: - """Fields the function reports *and* reads through the effective config. - - A key whose value comes off the `ServerArgs` record does not count: that is - the startup value under a name that promises the running one. - """ - tree = ast.parse(source) - for fn in ast.walk(tree): - if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - if fn.name != func_name: - continue - reported = set() - for node in ast.walk(fn): - if not isinstance(node, ast.Dict): - continue - for key, value in zip(node.keys, node.values): - if not (isinstance(key, ast.Constant) and isinstance(key.value, str)): - continue - for inner in ast.walk(value): - if ( - isinstance(inner, ast.Call) - and isinstance(inner.func, ast.Attribute) - and inner.func.attr in ("config_value", "config_leaf") - and inner.args - and isinstance(inner.args[0], ast.Constant) - and inner.args[0].value == key.value - ): - reported.add(key.value) - return reported - return set() - - -def _rust_model_info_keys() -> set: - """The keys the rust server's `/model_info` handler answers with. - - Scanned as text, from the handler's signature to the next item in the - file, with each line cut at its first `//`. A key is always left of its - value, so cutting inside a string can only drop keys, never invent one. - The signature must appear exactly once: a rust handler that moved or was - renamed would otherwise contribute an empty set and let the parity check - pass on nothing. - """ - source = _RUST_MODEL_INFO.read_text() - marker = "async fn model_info(" - found = source.count(marker) - assert found == 1, ( - f"{_RUST_MODEL_INFO.name} declares `{marker}` {found} times; the rust " - "/model_info surface is the one users reach under SGLANG_RUST_SERVER=1 " - "and is no longer being read" - ) - body = source[source.index(marker) + len(marker) :] - following = re.search(r"^(?:pub(?:\([^)]*\))?\s+)?(?:async\s+)?fn\s", body, re.M) - if following is not None: - body = body[: following.start()] - code = "\n".join(line.split("//")[0] for line in body.splitlines()) - keys = set(re.findall(r'"([A-Za-z_][A-Za-z0-9_]*)"\s*:', code)) - assert keys, "the rust /model_info handler names no field" - return keys - - -def _reported_keys_in(source: str, func_name: str) -> set: - """Every string key the function's response dicts carry, whatever the value. - - The manager-owned attributes are read off the manager, not through - `config_value`, so `_effective_reads_in` does not see them; this is how they - are checked. - """ - tree = ast.parse(source) - for fn in ast.walk(tree): - if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - if fn.name != func_name: - continue - return { - key.value - for node in ast.walk(fn) - if isinstance(node, ast.Dict) - for key in node.keys - if isinstance(key, ast.Constant) and isinstance(key.value, str) - } - return set() - - -def _model_info_sources() -> dict: - """`(source, function name)` per Python model-info surface.""" - return { - "http /model_info": ( - (_PACKAGE_ROOT / "srt/entrypoints/http_server.py").read_text(), - "model_info", - ), - "grpc get_model_info": ( - (_PACKAGE_ROOT / "srt/entrypoints/grpc_bridge.py").read_text(), - "get_model_info", - ), - "Engine.get_model_info": ( - inspect.getsource(Engine.get_model_info).lstrip(), - "get_model_info", - ), - } - - -def _model_info_surfaces() -> dict: - """Each Python entry point's model-info surface, by what it reports - effectively.""" - return { - name: _effective_reads_in(source, func_name) - for name, (source, func_name) in _model_info_sources().items() - } - - -class TestEffectiveStateSurfaces(CustomTestCase): - def test_each_entry_point_reports_the_post_launch_facts(self): - surfaces = _model_info_surfaces() - written = _overridden_fields() - # Both writer shapes are in reach of the scan: `load_format` is a - # literal keyword, the parsers arrive as `**{attr: ...}` under a key - # the enclosing loop binds. - self.assertLessEqual( - {"load_format", "reasoning_parser", "tool_call_parser"}, - written, - "the derivation stopped finding the control-plane writers", - ) - # What any surface reports effectively, all of them owe their users; - # what the control plane overrides, every surface owes regardless. - required = set().union(*surfaces.values()) | written - missing = { - name: sorted(required - reported) - for name, reported in surfaces.items() - if required - reported - } - self.assertEqual( - missing, - {}, - f"a serving surface cannot report what it is running: {missing}", - ) - - def test_every_surface_reports_the_manager_owned_identity(self): - """The identity a weight update moves is answered where it is read. - - `model_path` and `served_model_name` live on the tokenizer manager, so - the writer derivation cannot reach them and they are subtracted from the - required set. This is the assertion that pays for that subtraction. A - surface that drops one sends its clients back to the launch record -- - which is what the rust router had to read, under a name that had since - moved. - """ - surfaces = { - name: _reported_keys_in(source, func_name) - for name, (source, func_name) in _model_info_sources().items() - } - surfaces["rust /model_info"] = _rust_model_info_keys() - missing = { - name: sorted(_MANAGER_ATTRIBUTES - reported) - for name, reported in surfaces.items() - if _MANAGER_ATTRIBUTES - reported - } - self.assertEqual( - missing, - {}, - f"a model-info surface does not say which model it serves: {missing}", - ) - - def test_the_rust_model_info_answers_the_same_keys(self): - """`SGLANG_RUST_SERVER=1` swaps the whole HTTP server, not one handler. - - The keys are owed there too, or the endpoint's contract depends on - which server the operator launched. The values are the launch record: - that process parses `server_args` once and mounts no route that can - change weights or parsers, so this is a key-set check and the handler - states which it reports. - """ - required = set().union(*_model_info_surfaces().values()) | _overridden_fields() - missing = sorted(required - _rust_model_info_keys()) - self.assertEqual( - missing, - [], - "the rust /model_info answers a different contract than the Python " - f"one it replaces: {missing}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 4aa5bf528..6ecda2269 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -8,9 +8,8 @@ from collections import deque from types import SimpleNamespace from sglang.srt.managers.schedule_batch import ReqKvInfo -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci -register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-a-unit-test-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/unit/hardware_backend/mlx/test_attn_dp_request_capacity.py b/test/registered/unit/hardware_backend/mlx/test_attn_dp_request_capacity.py index ddf446e41..0f30c6991 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attn_dp_request_capacity.py +++ b/test/registered/unit/hardware_backend/mlx/test_attn_dp_request_capacity.py @@ -14,10 +14,9 @@ from types import SimpleNamespace from unittest import mock from sglang.srt.runtime_context import get_context -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-a-unit-test-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py b/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py index 040fe5d0a..760573085 100644 --- a/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py +++ b/test/registered/unit/hardware_backend/mlx/test_max_running_requests.py @@ -23,12 +23,9 @@ from unittest import mock from sglang.srt.managers.schedule_batch import ReqKvInfo from sglang.srt.runtime_context import get_context -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=11, suite="base-a-test-cpu") - - register_mlx_ci(est_time=1, suite="stage-a-unit-test-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py b/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py index bc11d7e89..2e38e67ba 100644 --- a/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py +++ b/test/registered/unit/hardware_backend/mlx/test_metal_profiler.py @@ -20,9 +20,8 @@ import unittest from pathlib import Path from unittest.mock import MagicMock, patch -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci -register_cpu_ci(est_time=6, suite="base-a-test-cpu") register_mlx_ci(est_time=5, suite="stage-a-unit-test-mlx") _IS_APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64" diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py b/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py index 1b98d4619..bba29e4f2 100644 --- a/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_pool_dtype.py @@ -19,10 +19,9 @@ from __future__ import annotations import importlib.util import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=4, suite="stage-a-unit-test-mlx") _HAS_MLX = ( diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py b/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py index 313d35096..88f46e9af 100644 --- a/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_reference_correctness.py @@ -40,10 +40,9 @@ import importlib.util import os import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-b-e2e-mlx") _HAS_MLX = ( diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_runner_pool_contract.py b/test/registered/unit/hardware_backend/mlx/test_mlx_runner_pool_contract.py index c2d579aa1..f35f4effb 100644 --- a/test/registered/unit/hardware_backend/mlx/test_mlx_runner_pool_contract.py +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_runner_pool_contract.py @@ -24,10 +24,9 @@ import importlib.util import inspect import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=1, suite="stage-a-unit-test-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py b/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py index 5043984fe..34c84e907 100644 --- a/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py +++ b/test/registered/unit/hardware_backend/mlx/test_mlx_sampling.py @@ -6,10 +6,9 @@ import importlib.util import unittest from collections import Counter -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_mlx_ci(est_time=20, suite="stage-a-unit-test-mlx") _HAS_MLX = importlib.util.find_spec("mlx") is not None diff --git a/test/registered/unit/hardware_backend/mlx/test_muse_glimmer_mlx_model.py b/test/registered/unit/hardware_backend/mlx/test_muse_glimmer_mlx_model.py index efb51a3b5..ae652f58b 100644 --- a/test/registered/unit/hardware_backend/mlx/test_muse_glimmer_mlx_model.py +++ b/test/registered/unit/hardware_backend/mlx/test_muse_glimmer_mlx_model.py @@ -20,10 +20,9 @@ from __future__ import annotations import importlib.util import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=6, suite="stage-a-unit-test-mlx") _HAS_MLX = ( diff --git a/test/registered/unit/hardware_backend/mlx/test_quantization.py b/test/registered/unit/hardware_backend/mlx/test_quantization.py index 3f6605cd7..d0f6d9609 100644 --- a/test/registered/unit/hardware_backend/mlx/test_quantization.py +++ b/test/registered/unit/hardware_backend/mlx/test_quantization.py @@ -17,7 +17,7 @@ import importlib.util import platform import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci # Registered on the CPU suite but skipped wherever mlx is absent; runs for real # only on Apple Silicon. Also registered under stage-b-e2e-mlx, not stage-a: @@ -27,7 +27,6 @@ from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci # fails with LocalEntryNotFoundError on a runner with no pre-warmed cache. The # macOS CI lane (pr-test-mlx.yml) only dispatches stage-b-e2e-mlx via a gated # workflow_dispatch, matching the models_e2e correctness tests' convention. -register_cpu_ci(est_time=6, suite="base-a-test-cpu") register_mlx_ci(est_time=10, suite="stage-b-e2e-mlx") _IS_APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64" diff --git a/test/registered/unit/hardware_backend/mlx/test_runner_init_contract.py b/test/registered/unit/hardware_backend/mlx/test_runner_init_contract.py deleted file mode 100644 index e150b5b0b..000000000 --- a/test/registered/unit/hardware_backend/mlx/test_runner_init_contract.py +++ /dev/null @@ -1,85 +0,0 @@ -"""Guard the MLX initialize override against ModelRunner contract drift. - -MlxModelRunnerStub.initialize must stay callable exactly as ModelRunner invokes -it. The check is signature-only and MLX-gated because importing the stub pulls in -mlx.core. -""" - -from __future__ import annotations - -import importlib.util -import inspect -import unittest - -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci - -register_cpu_ci(est_time=6, suite="base-a-test-cpu") -register_mlx_ci(est_time=1, suite="stage-a-unit-test-mlx") - -_HAS_MLX = importlib.util.find_spec("mlx") is not None -_SKIP_REASON = "requires mlx" - -if _HAS_MLX: - from sglang.srt.hardware_backend.mlx.model_runner_stub import MlxModelRunnerStub - from sglang.srt.model_executor.model_runner import ModelRunner - - -def _required_params_beyond_self(func) -> list[str]: - """Names of parameters (after ``self``) a caller MUST supply. - - Excludes ``self``, anything carrying a default, and ``*args`` / ``**kwargs``. - """ - params = list(inspect.signature(func).parameters.values())[1:] - return [ - p.name - for p in params - if p.default is inspect.Parameter.empty - and p.kind - in ( - inspect.Parameter.POSITIONAL_ONLY, - inspect.Parameter.POSITIONAL_OR_KEYWORD, - inspect.Parameter.KEYWORD_ONLY, - ) - ] - - -@unittest.skipUnless(_HAS_MLX, _SKIP_REASON) -class TestMlxRunnerInitContract(unittest.TestCase): - """``MlxModelRunnerStub.initialize`` must match how the base calls it.""" - - def test_base_initialize_takes_no_extra_args(self): - # The assumption this guard rests on: base ModelRunner.initialize is - # parameterless and is invoked as ``self.initialize()`` (model_runner.py). - # If the base re-introduces a required parameter, the override contract - # below must be revisited -- fail loudly here so that change is noticed. - required = _required_params_beyond_self(ModelRunner.initialize) - self.assertEqual( - required, - [], - msg=( - "Base ModelRunner.initialize gained required parameter(s) " - f"{required}. Since #23862 it is parameterless and called as " - "self.initialize(); if that changed, re-check the MLX override " - "(MlxModelRunnerStub.initialize, #28660)." - ), - ) - - def test_stub_initialize_binds_like_base_call(self): - # Core regression guard for #28660: the base calls self.initialize() with - # zero extra args, so the override must bind with the instance alone. The - # pre-#28660 signature (self, pre_model_load_memory) raises here. - sig = inspect.signature(MlxModelRunnerStub.initialize) - try: - sig.bind(object()) # stands in for ``self``; mirrors self.initialize() - except TypeError as exc: - self.fail( - "MlxModelRunnerStub.initialize is not call-compatible with the " - "base ModelRunner call site self.initialize() (no extra args): " - f"{exc}. Base initialize(self) has been parameterless since " - "#23862; the override must not require an argument the base no " - "longer passes (regression of #28660)." - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py index 457344116..6e4b9cbbd 100644 --- a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py +++ b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py @@ -16,9 +16,8 @@ import platform import unittest from unittest.mock import MagicMock, patch -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci -register_cpu_ci(est_time=6, suite="base-a-test-cpu") register_mlx_ci(est_time=5, suite="stage-a-unit-test-mlx") _IS_APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64" diff --git a/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py b/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py index a87a315ed..81bde66e6 100644 --- a/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py +++ b/test/registered/unit/hardware_backend/mlx/test_sliding_window_attention.py @@ -30,10 +30,9 @@ import importlib.util import unittest from types import SimpleNamespace -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=6, suite="stage-a-unit-test-mlx") _HAS_MLX = ( diff --git a/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py b/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py index 2acd4a2a7..bcf638bc0 100644 --- a/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py +++ b/test/registered/unit/hardware_backend/mlx/test_swa_radix_pool.py @@ -13,10 +13,9 @@ from __future__ import annotations import importlib.util import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=10, suite="stage-a-unit-test-mlx") _HAS_MLX = ( diff --git a/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py b/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py index 790982bb9..695f24608 100644 --- a/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py +++ b/test/registered/unit/hardware_backend/mlx/test_tp_worker_routing.py @@ -41,13 +41,12 @@ import torch from sglang.srt.managers.schedule_batch import ReqKvInfo from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.runtime_context import get_context -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase # CPU marker is AST-parsed "this test exists"; actual CPU-side execution is # gated by the @skipUnless guard below. MLX marker runs for real on the MLX # lane's stage-a (model-free: mocks the runner, loads no model). -register_cpu_ci(est_time=10, suite="base-a-test-cpu") register_mlx_ci(est_time=10, suite="stage-a-unit-test-mlx") _IS_APPLE_SILICON = platform.system() == "Darwin" and platform.machine() == "arm64" diff --git a/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py b/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py index b0ff3e4f1..950535e42 100644 --- a/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py +++ b/test/registered/unit/hardware_backend/mlx/test_windowed_kv_cache.py @@ -17,10 +17,9 @@ from __future__ import annotations import importlib.util import unittest -from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci +from sglang.test.ci.ci_register import register_mlx_ci from sglang.test.test_utils import CustomTestCase -register_cpu_ci(est_time=11, suite="base-a-test-cpu") register_mlx_ci(est_time=10, suite="stage-a-unit-test-mlx") _HAS_MLX = ( diff --git a/test/registered/unit/layers/moe/test_flashinfer_cutedsl_dispatch.py b/test/registered/unit/layers/moe/test_flashinfer_cutedsl_dispatch.py deleted file mode 100644 index 5a807ce35..000000000 --- a/test/registered/unit/layers/moe/test_flashinfer_cutedsl_dispatch.py +++ /dev/null @@ -1,64 +0,0 @@ -import sys -from types import SimpleNamespace -from unittest.mock import Mock, patch - -import pytest -import torch - -import sglang.srt.layers.moe.moe_runner.flashinfer_cutedsl as cutedsl_runner -from sglang.srt.layers.moe.token_dispatcher.standard import ( - StandardCombineInput, - StandardDispatchOutput, -) -from sglang.srt.layers.moe.topk import StandardTopKOutput -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=11, suite="base-a-test-cpu") - - -def test_flashinfer_prefill_returns_standard_combine_input(): - dispatch_output = StandardDispatchOutput( - hidden_states=torch.empty(2, 16, dtype=torch.bfloat16), - hidden_states_scale=None, - topk_output=StandardTopKOutput( - topk_weights=torch.empty(2, 1), - topk_ids=torch.zeros(2, 1, dtype=torch.int32), - router_logits=None, - ), - ) - expected_output = torch.empty(2, 16, dtype=torch.bfloat16) - wrapper = Mock() - wrapper.run.return_value = expected_output - quant_info = SimpleNamespace( - wrapper=wrapper, - quant_mode="w4a4", - use_per_token_activation=False, - a1_scale=torch.tensor(1.0), - a2_scale=torch.tensor(1.0), - w13_weight=object(), - w13_weight_sf=object(), - w1_alpha=object(), - w2_weight=object(), - w2_weight_sf=object(), - w2_alpha=object(), - ) - runner_config = SimpleNamespace(activation="silu") - - with patch( - "sglang.srt.layers.quantization.fp4_utils.fp4_quantize", - return_value=( - torch.empty(2, 8, dtype=torch.uint8), - torch.empty(2, 1, dtype=torch.float8_e4m3fn), - ), - ): - result = cutedsl_runner.fused_experts_flashinfer_to_flashinfer_cutedsl_fp4( - dispatch_output, quant_info, runner_config - ) - - assert isinstance(result, StandardCombineInput) - assert result.hidden_states is expected_output - wrapper.run.assert_called_once() - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/lora/test_deepseek_mla_correction.py b/test/registered/unit/lora/test_deepseek_mla_correction.py deleted file mode 100644 index 37af119d6..000000000 --- a/test/registered/unit/lora/test_deepseek_mla_correction.py +++ /dev/null @@ -1,24 +0,0 @@ -import unittest -from types import SimpleNamespace - -from sglang.srt.lora.deepseek_mla_correction import is_kv_b_lora_active -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=7, suite="base-a-test-cpu") - - -class TestDeepseekMLACorrection(unittest.TestCase): - def test_kv_b_lora_probe(self): - self.assertFalse(is_kv_b_lora_active(SimpleNamespace())) - self.assertFalse( - is_kv_b_lora_active(SimpleNamespace(kv_b_proj=SimpleNamespace())) - ) - self.assertTrue( - is_kv_b_lora_active( - SimpleNamespace(kv_b_proj=SimpleNamespace(set_lora=True)) - ) - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/mem_cache/test_session_unified_radix_cache.py b/test/registered/unit/mem_cache/test_session_unified_radix_cache.py index b85de9997..e50475985 100644 --- a/test/registered/unit/mem_cache/test_session_unified_radix_cache.py +++ b/test/registered/unit/mem_cache/test_session_unified_radix_cache.py @@ -4,10 +4,8 @@ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=10, suite="base-a-test-cpu") -import ast import unittest from array import array -from pathlib import Path from types import SimpleNamespace import torch @@ -25,77 +23,6 @@ from sglang.srt.mem_cache.unified_cache.components import ComponentType from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache from sglang.test.test_utils import CustomTestCase -REPO_ROOT = Path(__file__).resolve().parents[4] -MEM_CACHE_ROOT = REPO_ROOT / "python/sglang/srt/mem_cache" - - -def class_bases(path: Path, class_name: str) -> set[str]: - tree = ast.parse(path.read_text()) - class_node = next( - node - for node in tree.body - if isinstance(node, ast.ClassDef) and node.name == class_name - ) - return { - base.id if isinstance(base, ast.Name) else ast.unparse(base) - for base in class_node.bases - } - - -class TestSessionCacheOwnership(CustomTestCase): - def test_only_unified_radix_cache_owns_session_ref_tracker(self): - ordinary_mixin = MEM_CACHE_ROOT / "session_radix_cache.py" - radix_cache = MEM_CACHE_ROOT / "radix_cache.py" - hiradix_cache = MEM_CACHE_ROOT / "hiradix_cache.py" - evict_policy = MEM_CACHE_ROOT / "evict_policy.py" - unified_cache = MEM_CACHE_ROOT / "unified_radix_cache.py" - session_ref_tracker = ( - MEM_CACHE_ROOT / "unified_cache" / "session_ref_tracker.py" - ) - - self.assertFalse(ordinary_mixin.exists()) - ordinary_source = "\n".join( - path.read_text() for path in (radix_cache, hiradix_cache, evict_policy) - ) - for removed_symbol in ( - "SessionRadixCacheMixin", - "SessionAwareEvictionStrategy", - "session_ref", - "_session_on_", - "_session_forget_node", - "_account_new_evictable_node", - "_supports_session_radix_cache", - "enable_session_radix_cache", - ): - self.assertNotIn(removed_symbol, ordinary_source) - self.assertNotIn( - "SessionRadixCacheMixin", class_bases(radix_cache, "RadixCache") - ) - # Session behavior is composed, not mixed in (general-code-style rule). - self.assertEqual( - class_bases(unified_cache, "UnifiedRadixCache"), {"BasePrefixCache"} - ) - self.assertIn("UnifiedSessionRefTracker", session_ref_tracker.read_text()) - self.assertNotIn("SessionUnifiedRadixCacheMixin", unified_cache.read_text()) - - for component in ( - "full_component.py", - "swa_component.py", - "mamba_component.py", - ): - self.assertIn( - "session_ref", - ( - MEM_CACHE_ROOT / "unified_cache" / "components" / component - ).read_text(), - ) - - registry = MEM_CACHE_ROOT / "registry.py" - self.assertIn( - "--enable-session-radix-cache requires UnifiedRadixCache", - registry.read_text(), - ) - def make_params(enable_session: bool) -> CacheInitParams: dtype = torch.float16 diff --git a/test/registered/unit/model_executor/runner/test_flashinfer_autotune.py b/test/registered/unit/model_executor/runner/test_flashinfer_autotune.py deleted file mode 100644 index 47531e1cd..000000000 --- a/test/registered/unit/model_executor/runner/test_flashinfer_autotune.py +++ /dev/null @@ -1,97 +0,0 @@ -import sys -from types import SimpleNamespace -from unittest.mock import Mock, patch - -import pytest - -from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardMode -from sglang.srt.model_executor.runner import base_runner, flashinfer_autotune -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=12, suite="base-a-test-cpu") - - -@pytest.mark.parametrize( - "mode,error", - [ - ("prefill", "_dummy_run needs a static buffer"), - ("decode", "ordinary or PD-prefill target EXTEND dummy"), - ], -) -def test_packed_speculative_extend_is_limited_to_pd_prefill_target(mode, error): - runner = SimpleNamespace( - model_runner=SimpleNamespace( - is_draft_worker=False, - spec_algorithm=SimpleNamespace(is_speculative=lambda: True), - decode_num_tokens_per_req=lambda: 6, - ) - ) - with ( - patch.object( - base_runner, - "get_disagg", - return_value=SimpleNamespace(disaggregation_mode=mode), - ), - patch.object( - base_runner, - "get_server_return_hidden_states_mode", - return_value=CaptureHiddenMode.NULL, - ), - pytest.raises(AssertionError, match=error), - ): - base_runner.BaseRunner._dummy_run( - runner, - batch_size=1, - buffers=None, - forward_mode_override=ForwardMode.EXTEND, - extend_num_tokens_per_req=1, - ) - - -def test_chunked_prefill_disabled_uses_legacy_token_ceiling(): - model_runner = SimpleNamespace( - server_args=SimpleNamespace(), - is_generation=True, - is_draft_worker=False, - spec_algorithm=SimpleNamespace(is_speculative=lambda: False), - attn_backend=SimpleNamespace(extend_dummy_seqs_capped_by_req_pool=False), - canary_manager=None, - ) - runner = SimpleNamespace( - model_runner=model_runner, - _alloc_dummy_decode_buffers=Mock(return_value=object()), - _dummy_run=Mock(), - ) - with ( - patch.object( - flashinfer_autotune.envs.SGLANG_FLASHINFER_AUTOTUNE_EXTEND, - "get", - return_value=True, - ), - patch.object( - flashinfer_autotune, - "get_disagg", - return_value=SimpleNamespace(disaggregation_mode="prefill"), - ), - patch.object(flashinfer_autotune, "max_prefill_buffer_tokens", return_value=0), - patch.object( - flashinfer_autotune, - "get_schedule", - return_value=SimpleNamespace(max_prefill_tokens=32768), - ), - patch.object(flashinfer_autotune, "run_flashinfer_autotune_forward"), - patch.object(flashinfer_autotune.torch.cuda, "empty_cache"), - ): - flashinfer_autotune.maybe_flashinfer_autotune_extend( - runner, decode_num_tokens=128 - ) - - runner._alloc_dummy_decode_buffers.assert_called_once_with( - 32768, - num_tokens_per_req=1, - allocate_logits_buffer=False, - ) - - -if __name__ == "__main__": - sys.exit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/models/test_draft_entry_hook_parity.py b/test/registered/unit/models/test_draft_entry_hook_parity.py index cd8b03678..d1ebb4ce2 100644 --- a/test/registered/unit/models/test_draft_entry_hook_parity.py +++ b/test/registered/unit/models/test_draft_entry_hook_parity.py @@ -9,13 +9,6 @@ and its weights are laid out for the other layout. That shipped once: the DSV4 DSpark draft skipped its bundled shared-expert tensors until #33312 gave it the gate. -`test_fusion_gate_coverage.py` walks the same registry but asks a different -question: does this entry class *touch* the decision (read the flag, name a gated -class)? That catches a class after it starts consuming the decision. It cannot -catch one that should consume it and does not, which is what the DSpark class -looked like: it built the family's *layer* classes, so the flag reader lived in -another module and its own source named no gated class. - This case asks the invariant directly instead: **presence parity between a draft entry class and its target**. Identity is deliberately not required -- a draft that delegates with adapted arguments (the Qwen3.5 MTP unwraps `text_config` and diff --git a/test/registered/unit/models/test_fusion_gate_coverage.py b/test/registered/unit/models/test_fusion_gate_coverage.py deleted file mode 100644 index d37ba7ba1..000000000 --- a/test/registered/unit/models/test_fusion_gate_coverage.py +++ /dev/null @@ -1,181 +0,0 @@ -"""Every loader entry class that can reach a fusion-gated family answers for it. - -The loader installs the shared-experts-fusion decision for the class it -instantiates (`install_shared_experts_fusion_decision`). A model whose layers -read `is_shared_experts_fusion_disabled()` therefore gets whatever answer that -*entry* class produced — and an entry class with no -`shared_experts_fusion_disable_reason` falls back to the user's intent, silently -skipping the family's auto-disable conditions. - -That is easy to miss for a wrapper: `KimiVLForConditionalGeneration` is the -registered arch, but a DeepSeek body is built inside it, and the DeepSeek -conditions used to be evaluated during that nested construction. This case walks -the registry so a new wrapper (or a new MTP/nextn entry) cannot reintroduce the -gap. -""" - -import ast -import importlib -import inspect -import os -import sys -import unittest - -from sglang.srt.models.registry import ModelRegistry -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=22, suite="base-a-test-cpu") - -GATE = "shared_experts_fusion_disable_reason" -FLAG_READERS = ( - "is_shared_experts_fusion_disabled", - "determine_num_fused_shared_experts", -) - -# Archs that read the fusion flag but deliberately have no gate: nothing in -# their lineage carries auto-disable conditions, so they follow the user's -# intent — the behavior they had before the decision moved to the loader. -GATELESS_BY_DESIGN = { - # The in-tree class has no ``determine_num_fused_shared_experts`` at all - # (the call is guarded by ``hasattr`` for a downstream variant). - "BailingMoeForCausalLMNextN", - # Its target family (Glm4v) is dense; there is no gate to inherit. - "GlmOcrForConditionalGenerationNextN", - # The vision tower registered on its own: it shares a module with - # PixtralForConditionalGeneration (which does answer) but builds no language - # model, so there is nothing for a gate to decide. - "PixtralVisionModel", -} - - -def _gated_classes(source: str) -> set: - """Classes in this module that define or receive a fusion gate.""" - names = set() - tree = ast.parse(source) - for node in ast.walk(tree): - if isinstance(node, ast.ClassDef): - for body in node.body: - if ( - isinstance(body, (ast.FunctionDef, ast.AsyncFunctionDef)) - and body.name == GATE - ): - names.add(node.name) - if isinstance(body, ast.Assign) and any( - isinstance(t, ast.Name) and t.id == GATE for t in body.targets - ): - names.add(node.name) - # ``for cls in (A, B): cls. = ...`` - if ( - isinstance(node, ast.For) - and isinstance(node.iter, (ast.Tuple, ast.List)) - and GATE in ast.dump(node) - ): - names |= {e.id for e in node.iter.elts if isinstance(e, ast.Name)} - if isinstance(node, ast.Assign): - for target in node.targets: - if ( - isinstance(target, ast.Attribute) - and target.attr == GATE - and isinstance(target.value, ast.Name) - ): - names.add(target.value.id) - return names - - -def gated_class_names() -> set: - """Every model class that *resolves* a gate, inherited ones included. - - A subclass like `DeepseekV3ForCausalLM` inherits the gate without naming it, - so collecting names from class bodies alone would let a wrapper that builds - the subclass slip through. - """ - names = set() - for module_name, module in list(sys.modules.items()): - if not module_name.startswith("sglang.srt.models.") or module is None: - continue - for member in vars(module).values(): - # transformers re-exports lazy placeholders that raise on any - # attribute access when their optional backend is missing. - try: - if ( - inspect.isclass(member) - and (member.__module__ or "").startswith("sglang.srt.models.") - and hasattr(member, GATE) - ): - names.add(member.__name__) - except Exception: - continue - return names - - -class TestFusionGateCoverage(CustomTestCase): - def test_every_entry_class_reaching_a_gated_family_has_a_gate(self): - models_dir = list(importlib.import_module("sglang.srt.models").__path__)[0] - gates_by_module = {} - for name in sorted(os.listdir(models_dir)): - if not name.endswith(".py"): - continue - with open(os.path.join(models_dir, name), encoding="utf-8") as f: - try: - gates_by_module[f"sglang.srt.models.{name[:-3]}"] = _gated_classes( - f.read() - ) - except SyntaxError: - continue - - missing = [] - all_gated = None - for arch in sorted(ModelRegistry.get_supported_archs()): - try: - model_class, _ = ModelRegistry.resolve_model_cls(arch) - except Exception: - continue - if hasattr(model_class, GATE) or arch in GATELESS_BY_DESIGN: - continue - module = importlib.import_module(model_class.__module__) - try: - source = inspect.getsource(module) - except OSError: - continue - reasons = [] - if any(reader in source for reader in FLAG_READERS): - reasons.append("reads the fusion flag") - try: - tree = ast.parse(source) - except SyntaxError: - tree = None - if tree is not None: - # Any *use* of a gated class counts, whatever the shape: a - # direct call (`DeepseekV2ForCausalLM(...)`), a module attribute - # (`qwen3_5.Qwen3_5MoeForCausalLM`), or a class attribute the - # constructor later calls (`body_cls = qwen3_5.Qwen3_5...`). - # Only matching calls would miss the last two. - if all_gated is None: - all_gated = gated_class_names() - used = set() - for node in ast.walk(tree): - if isinstance(node, ast.Name) and node.id in all_gated: - used.add(node.id) - elif isinstance(node, ast.Attribute) and node.attr in all_gated: - used.add(node.attr) - for name in sorted(used): - if name != model_class.__name__: - reasons.append(f"references {name}") - if reasons: - missing.append( - f"{arch} ({model_class.__module__}): {', '.join(reasons)}" - ) - - self.assertEqual( - [], - missing, - "these entry classes reach a fusion-gated family but resolve no " - f"{GATE}, so the loader falls back to the user's intent for them and " - "the family's auto-disable conditions never run:\n " - + "\n ".join(missing), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/models/test_kimi_k25_mm_projection.py b/test/registered/unit/models/test_kimi_k25_mm_projection.py deleted file mode 100644 index 09082ed6f..000000000 --- a/test/registered/unit/models/test_kimi_k25_mm_projection.py +++ /dev/null @@ -1,41 +0,0 @@ -"""CPU coverage for the Kimi vision-projector packing fast path.""" - -import pytest -import torch -import torch.nn as nn - -from sglang.srt.models.kimi_k25 import mm_projection_auto -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=12, suite="base-a-test-cpu") - - -class _FlattenProjector(nn.Module): - def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return hidden_states.flatten(start_dim=1) - - -def test_mm_projection_auto_packs_variable_image_outputs_once(): - outputs = [ - torch.arange(2 * 4 * 3, dtype=torch.float32).reshape(2, 4, 3), - torch.arange(3 * 4 * 3, dtype=torch.float32).reshape(3, 4, 3), - ] - expected = torch.cat([output.flatten(start_dim=1) for output in outputs], dim=0) - - actual = mm_projection_auto(_FlattenProjector(), outputs) - - torch.testing.assert_close(actual, expected) - assert actual.shape == (5, 12) - - -def test_mm_projection_auto_single_item_avoids_cat_copy(): - output = torch.randn(5, 4, 3) - - actual = mm_projection_auto(_FlattenProjector(), [output]) - - torch.testing.assert_close(actual, output.flatten(start_dim=1)) - assert actual.data_ptr() == output.data_ptr() - - -if __name__ == "__main__": - raise SystemExit(pytest.main([__file__, "-v"])) diff --git a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py b/test/registered/unit/server_args/test_model_config_reads_resolved_input.py deleted file mode 100644 index 7e988c91c..000000000 --- a/test/registered/unit/server_args/test_model_config_reads_resolved_input.py +++ /dev/null @@ -1,765 +0,0 @@ -"""`ModelConfig` is built from values resolution has already decided. - -Resolution builds a `ModelConfig` partway through and keys later decisions off -it, so the pipeline reads its own output through that object. The loop is only -benign while every field `ModelConfig.from_server_args` reads has been resolved -by the time it is built -- otherwise the model configuration describes a -half-resolved input, and every handler downstream of it inherits that. - -Nothing enforces the ordering today; it holds because the path and quantization -handlers happen to run early. So this derives both sides from the source -- the -fields the constructor reads, and the step each is declared at -- and pins the -one field that is deliberately read before resolution touches it. -""" - -import ast -import functools -import pathlib -import unittest - -import sglang -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=17, suite="base-a-test-cpu") - -_SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" - -# Read for what the caller asked for: the constructor passes it through and -# never stores it, while resolution declares the value the architecture implies. -# Two quantities sharing one name. -_READ_BEFORE_RESOLUTION = frozenset({"is_embedding"}) - -# Declared after the first `model_config_of()`, so the cached configuration -# holds the earlier value. Nothing reads the stale copy today (its one consumer -# is on the `is_draft_model` branch, built after resolution), and fixing it -# means moving the build or the hook. Pinned so a second field in this position -# has to be looked at. -_STALE_IN_THE_MODEL_CONFIG = frozenset({"speculative_algorithm"}) - -# Behind the expert-pack build. `expert_pack_hook.handle_expert_pack` builds a -# model configuration, and it always did -- the walk stopped at the record's -# file and never saw it, so these three read as decided before the first build. -# The call sits behind `load_format != "expert_pack": return`, so it is the -# first build only on an expert-pack launch. Pre-existing; named rather than -# fixed, because fixing it means moving the build or the hook. -_STALE_BEHIND_THE_EXPERT_PACK_BUILD = frozenset( - { - "_speculative_draft_quantization_explicitly_set", - "model_path", - "speculative_draft_model_quantization", - } -) - -# The same staleness through the registries: `_handle_model_specific_adjustments` -# builds the model configuration and *then* collects the override declarations, -# both inside one handler body. Named rather than fixed (that means moving the -# build or the collection), so a fifth field here has to be looked at -- and so -# does fixing the ordering. -_STALE_FROM_THE_REGISTRIES = frozenset( - { - "disable_hybrid_swa_memory", - "dtype", - "enable_multi_layer_eagle", - "quantization", - } -) - - -@functools.lru_cache(maxsize=None) -def _parsed(path): - return ast.parse(path.read_text(encoding="utf-8-sig")) - - -@functools.lru_cache(maxsize=None) -def _declared_resolution_fields(path): - fields = set() - for node in ast.walk(_parsed(path)): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "declare_resolution" - ): - fields |= {kw.arg for kw in node.keywords if kw.arg} - return frozenset(fields) - - -def _registry_declared_fields(): - """What the live registries and passes declare. - - Imported from the chain ratchet by path instead of re-derived: two - derivations of the same set drift, and the one that drifts narrower makes - this check quietly vacuous. Keying on `self._declare(...)` alone is what - hid these four -- 26 of the providers register through a helper call, and - none of them spell a keyword this file can see. - """ - import importlib.util - - ratchet = ( - pathlib.Path(__file__).resolve().parent.parent / "test_chain_read_ratchet.py" - ) - spec = importlib.util.spec_from_file_location("_chain_ratchet_for_pin", ratchet) - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module._declared_by_registry_and_passes() - - -def _registry_collection_is_after_the_build(): - """(collection line, first build line) inside the model-specific handler. - - Handler-local ordering only -- the caller still has to compare against the - pipeline-wide first build, which sits in an *earlier* step: hoisting the - collection above this handler's own `model_config_of()` call does not move - it above the configuration another handler already cached. - """ - handler = None - for source, wanted in ( - (_SRT / "server_args.py", "_handle_model_specific_adjustments"), - *( - (path, "handle_model_specific_adjustments") - for path in sorted((_SRT / "arg_groups").glob("*.py")) - ), - ): - for node in ast.walk(_parsed(source)): - if isinstance(node, ast.FunctionDef) and node.name == wanted: - if any( - isinstance(child, ast.Call) - and getattr(child.func, "attr", getattr(child.func, "id", None)) - == "collect_model_override_declarations" - for child in ast.walk(node) - ): - handler = node - break - if handler is not None: - break - assert handler is not None, "the model-specific handler was not found" - build = collect = None - for node in ast.walk(handler): - if not isinstance(node, ast.Call): - continue - # Both spellings: an Attribute call and a bare Name call. - func = node.func - if isinstance(func, ast.Attribute): - name = func.attr - elif isinstance(func, ast.Name): - name = func.id - else: - continue - if name == "model_config_of" and build is None: - build = node.lineno - if name == "collect_model_override_declarations" and collect is None: - collect = node.lineno - return collect, build - - -def _server_args_names(tree, path): - """Every local that names the record, including the read views over it. - - A resolution-time reader reads through `resolving_view(server_args)` (the - declaration stash over the fields): declaration-only resolvers write no - field, so a field read there answers with the raw input. `cfg.dtype` after `cfg = resolving_view(sa)` is - the same read this scan is looking for, so the local it binds counts. - """ - names = {"self"} if path.name == "server_args.py" else {"server_args"} - for node in ast.walk(tree): - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - args = node.args - for arg in args.posonlyargs + args.args + args.kwonlyargs: - annotation = arg.annotation - if isinstance(annotation, ast.Constant): - text = annotation.value - elif isinstance(annotation, ast.Name): - text = annotation.id - elif isinstance(annotation, ast.Attribute): - text = annotation.attr - else: - continue - if text == "ServerArgs": - names.add(arg.arg) - # `cfg = resolving_view(server_args)` / `resolved_view(server_args)` - for _ in range(2): # a view over a view-holding local is still one - for node in ast.walk(tree): - if not isinstance(node, ast.Assign): - continue - value = node.value - bare = ( - isinstance(value, ast.Call) - and isinstance(value.func, ast.Name) - and value.func.id in ("resolving_view", "resolved_view") - and value.args - and isinstance(value.args[0], ast.Name) - and value.args[0].id in names - ) - # `resolved = self._resolved()` is the same view, spelled as the - # resolution vocabulary. - member = ( - isinstance(value, ast.Call) - and isinstance(value.func, ast.Attribute) - and isinstance(value.func, ast.Name) - and value.func.id == "resolved_view" - and isinstance(value.func.value, ast.Name) - and value.func.value.id in names - ) - if not (bare or member): - continue - names |= {t.id for t in node.targets if isinstance(t, ast.Name)} - return names - - -def _constructor_reads(): - """Fields `ModelConfig.from_server_args` takes off the record.""" - path = _SRT / "configs/model_config.py" - tree = _parsed(path) - constructor = next( - node - for node in ast.walk(tree) - if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) - and node.name == "from_server_args" - ) - names = _server_args_names(tree, path) - reads = { - node.attr - for node in ast.walk(constructor) - if isinstance(node, ast.Attribute) - and isinstance(node.value, ast.Name) - and node.value.id in names - and isinstance(node.ctx, ast.Load) - } - # `getattr(server_args, "field", default)` is the normal spelling for an - # optional input and is a `Call`, not an `Attribute`. - for node in ast.walk(constructor): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "getattr" - and len(node.args) >= 2 - and isinstance(node.args[0], ast.Name) - and node.args[0].id in names - and isinstance(node.args[1], ast.Constant) - and isinstance(node.args[1].value, str) - ): - reads.add(node.args[1].value) - return reads - - -def _late_resolution_fields(): - """Fields written through `_late_resolution` / `declare_late_resolution`. - - All of them land after the model configuration is built: the launcher's - validation stage runs long after `__post_init__`. - """ - fields = set() - for name in ( - "server_args.py", - "arg_groups/overrides.py", - "parser/template_detection.py", - ): - path = _SRT / name - # A named file that moved away has to be loud; skipping it silently - # leaves the scan believing it read a module it never opened. - assert path.exists(), f"{name} is not where this scan looks for it" - tree = _parsed(path) - for node in ast.walk(tree): - if not isinstance(node, ast.Call): - continue - called = ( - node.func.attr - if isinstance(node.func, ast.Attribute) - else getattr(node.func, "id", "") - ) - if called == "declare_late_resolution": - fields |= {kw.arg for kw in node.keywords if kw.arg} - return fields - - -def _hook_declarations(dispatch, source_module): - """{field: dispatcher line} for hooks the dispatch calls on other objects. - - `handle_speculative_decoding(self)` is not a `self.()` call, so a - scan of the dispatcher's own method calls never reaches its - `declare_resolution` sites -- and the speculative hooks decide - `speculative_algorithm`, which the model configuration reads. - - The platform hook is *not* covered here: it reaches the pipeline as a - callback argument, so there is no call node to follow and its writes live - outside this tree. Its position is pinned instead -- - `test_every_opaque_callback_is_still_late`. - """ - imported = {} - for node in ast.walk(_parsed(source_module)): - if isinstance(node, ast.ImportFrom) and node.module: - for alias in node.names: - imported[alias.asname or alias.name] = node.module - - out = {} - for node in ast.walk(dispatch): - if not isinstance(node, ast.Call): - continue - name = ( - node.func.id - if isinstance(node.func, ast.Name) - else (node.func.attr if isinstance(node.func, ast.Attribute) else None) - ) - module = imported.get(name) - if not module or not module.startswith("sglang.srt."): - continue - path = _SRT / (module[len("sglang.srt.") :].replace(".", "/") + ".py") - if not path.exists(): - continue - for field in _declared_resolution_fields(path): - out[field] = max(out.get(field, 0), node.lineno) - return out - - -# The dispatcher's own file: its imports are what map a bare-name call in it -# to the family that defines the callable. -_DISPATCH_MODULE = _SRT / "arg_groups" / "pipeline.py" - - -def _hook_functions(): - """Module-level resolution functions under `arg_groups/`. - - A handler that moved out of the record leaves a slot behind that imports - one of these and calls it. Without following that hop the scan stops at - the slot and silently loses everything the handler does. - """ - functions = {} - for path in sorted((_SRT / "arg_groups").glob("*.py")): - for node in _parsed(path).body: - if isinstance(node, ast.FunctionDef): - functions.setdefault(node.name, node) - return functions - - -def _pipeline(): - """(ordered steps, {step: methods it reaches}) for the resolution dispatch.""" - tree = _parsed(_SRT / "server_args.py") - record = next( - node - for node in tree.body - if isinstance(node, ast.ClassDef) and node.name == "ServerArgs" - ) - methods = { - node.name: node for node in record.body if isinstance(node, ast.FunctionDef) - } - # The dispatcher calls its hooks by bare name, so the walk resolves those - # against `arg_groups/` alongside the record's own methods. - hooks = _hook_functions() - methods.update({name: node for name, node in hooks.items() if name not in methods}) - dispatch = methods["run_resolution_pipeline"] - # A step is either a record method (`self._x()`) or a bare-name hook call. - steps = [ - name - for _line, name in sorted( - ( - node.lineno, - ( - node.func.attr - if isinstance(node.func, ast.Attribute) - else node.func.id - ), - ) - for node in ast.walk(dispatch) - if isinstance(node, ast.Call) - and ( - ( - isinstance(node.func, ast.Attribute) - and isinstance(node.func.value, ast.Name) - and node.func.value.id == "self" - ) - or (isinstance(node.func, ast.Name) and node.func.id in hooks) - ) - ) - ] - - def reaches(name, seen=None): - seen = seen if seen is not None else set() - if name in seen or name not in methods: - return seen - seen.add(name) - for node in ast.walk(methods[name]): - if not isinstance(node, ast.Call): - continue - if ( - isinstance(node.func, ast.Attribute) - and isinstance(node.func.value, ast.Name) - and node.func.value.id == "self" - and node.func.attr in methods - ): - reaches(node.func.attr, seen) - elif isinstance(node.func, ast.Name) and node.func.id in hooks: - reaches(node.func.id, seen) - return seen - - step_lines = {} - for node in ast.walk(dispatch): - if not isinstance(node, ast.Call): - continue - if ( - isinstance(node.func, ast.Attribute) - and isinstance(node.func.value, ast.Name) - and node.func.value.id == "self" - ): - step_lines.setdefault(node.func.attr, node.lineno) - elif isinstance(node.func, ast.Name) and node.func.id in hooks: - step_lines.setdefault(node.func.id, node.lineno) - return steps, methods, {name: reaches(name) for name in steps}, step_lines - - -def _opaque_callback_positions(dispatch, source_module): - """{callback spelling: dispatcher line} for every resolver handed in. - - `declare_direct_writes(record, source, callback)` runs a callable instead of - code in this tree -- a platform plugin, a registered speculative algorithm. - Which fields such a callback writes is not a static question; only *when* it - runs is, so the position is what gets pinned. - - Two spellings reach the pipeline: the dispatcher wraps a callback itself, or - it calls a hook in this tree that wraps one. The line recorded is always the - dispatcher's, because that is where the ordering against the build is - decided -- a hook body sits further down its own file and says nothing about - it. - """ - imported = {} - for node in ast.walk(_parsed(source_module)): - if isinstance(node, ast.ImportFrom) and node.module: - for alias in node.names: - imported[alias.asname or alias.name] = node.module - - def callbacks_in(tree): - found = [] - for node in ast.walk(tree): - if not isinstance(node, ast.Call): - continue - name = ( - node.func.id - if isinstance(node.func, ast.Name) - else (node.func.attr if isinstance(node.func, ast.Attribute) else None) - ) - if name == "declare_direct_writes" and len(node.args) > 2: - found.append(ast.unparse(node.args[2])) - return found - - positions = {} - for spelling in callbacks_in(dispatch): - positions[spelling] = min( - positions.get(spelling, 10**9), - next( - node.lineno - for node in ast.walk(dispatch) - if isinstance(node, ast.Call) - and getattr(node.func, "id", None) == "declare_direct_writes" - ), - ) - for node in ast.walk(dispatch): - if not isinstance(node, ast.Call): - continue - name = ( - node.func.id - if isinstance(node.func, ast.Name) - else (node.func.attr if isinstance(node.func, ast.Attribute) else None) - ) - module = imported.get(name) - if not module or not module.startswith("sglang.srt."): - continue - path = _SRT / (module[len("sglang.srt.") :].replace(".", "/") + ".py") - if not path.exists(): - continue - for spelling in callbacks_in(_parsed(path)): - positions[spelling] = min(positions.get(spelling, 10**9), node.lineno) - return positions - - -def _declaration_positions(): - """({field: position}, first_build) over the fields the constructor reads. - - A position is `(step index, rank)`, and `rank` is 0 only for a declaration - that sits *above* the build in the very method that builds: a declaration - applies where it is written, so one statement earlier in the same body is - genuinely earlier. Everything else in the build's step gets rank 1 and - counts as late -- line numbers say nothing across two method bodies, since - a handler sits further down the file than the dispatcher that calls it. - - One derivation, two callers: the check below asks which fields land after - the build, and the pin check asks whether an exempted field is still one - of them. Two derivations of that answer drift apart. - """ - steps, methods, reached, step_lines = _pipeline() - wanted = _constructor_reads() - - def build_site(): - """(step index, method name, line) of the first `model_config_of()`.""" - for index, step in enumerate(steps): - for method in reached[step]: - for node in ast.walk(methods[method]): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "model_config_of" - ): - return index, step, method, node.lineno - return None - - site = build_site() - if site is None: - return {}, None - build_index, build_step, build_method, build_line_in_body = site - first_build = (build_index, build_step) - - source_module = _DISPATCH_MODULE - imported = {} - for node in ast.walk(_parsed(source_module)): - if isinstance(node, ast.ImportFrom) and node.module: - for alias in node.names: - imported[alias.asname or alias.name] = node.module - - def _hook_declared_fields(name): - """Fields a hook imported from `sglang.srt` declares, by callable name.""" - module = imported.get(name) - if not module or not module.startswith("sglang.srt."): - return frozenset() - path = _SRT / (module[len("sglang.srt.") :].replace(".", "/") + ".py") - if not path.exists(): - return frozenset() - return _declared_resolution_fields(path) - - declared_at = {} - for index, step in enumerate(steps): - for method in reached[step]: - for node in ast.walk(methods[method]): - if not isinstance(node, ast.Call): - continue - same_body = index == build_index and method == build_method - rank = 0 if same_body and node.lineno < build_line_in_body else 1 - if ( - isinstance(node.func, ast.Name) - and node.func.id == "declare_resolution" - ): - fields = {kw.arg for kw in node.keywords if kw.arg} - # A handler that calls an imported hook (the Kimi and DeepSeek - # defaults live in arg_groups modules) declares through it, and - # the hook can sit below the build inside the same handler. - elif isinstance(node.func, ast.Name): - fields = _hook_declared_fields(node.func.id) - else: - continue - for field in fields: - if field in wanted: - # The *last* declaration is the one that has to precede - # the build. - declared_at[field] = max( - declared_at.get(field, (index, rank)), (index, rank) - ) - - # Hooks the dispatch calls on other objects declare too, and a hook below - # the first build is late by definition. Both positions are read *inside - # the dispatcher*: a handler body sits further down the file than the - # dispatcher that calls it, so a line number taken from one scope says - # nothing about ordering against the other. - dispatch = methods["run_resolution_pipeline"] - build_line = step_lines[first_build[1]] - for field, line in _hook_declarations(dispatch, _DISPATCH_MODULE).items(): - if field in wanted and line > build_line: - declared_at[field] = max( - declared_at.get(field, (build_index, 1)), (10**6, 1) - ) - - # Late resolution is the other channel that can decide a field the - # constructor reads, and it runs after every build. - for field in _late_resolution_fields(): - if field in wanted: - declared_at[field] = (10**6, 1) - return declared_at, first_build - - -class TestModelConfigReadsResolvedInput(CustomTestCase): - def test_every_field_it_reads_is_resolved_before_it_is_built(self): - declared_at, first_build = _declaration_positions() - self.assertIsNotNone( - first_build, "no handler builds a ModelConfig; the scan broke" - ) - - known = ( - _READ_BEFORE_RESOLUTION - | _STALE_IN_THE_MODEL_CONFIG - | _STALE_FROM_THE_REGISTRIES - | _STALE_BEHIND_THE_EXPERT_PACK_BUILD - ) - late = sorted( - field - for field, position in declared_at.items() - if position >= (first_build[0], 1) and field not in known - ) - self.assertEqual( - late, - [], - "resolution decides these after it builds the ModelConfig that reads " - f"them, so the model configuration describes a half-resolved input " - f"(first build: step {first_build[0]}, {first_build[1]}): {late}", - ) - - def test_the_registry_stale_set_is_exactly_what_is_late(self): - """Equality, not membership. - - A fifth field the registries decide after the build fails here, and so - does fixing the ordering -- either way someone has to come back and - read this. The earlier version of this file derived declarations only - from `self._declare(...)` keywords, so it passed while these four were - already stale. - """ - collect_line, build_line = _registry_collection_is_after_the_build() - self.assertIsNotNone( - collect_line, "the handler no longer collects registry declarations" - ) - reads = _constructor_reads() - registry = _registry_declared_fields() - self.assertGreater( - len(registry), 20, "the registry-declared set collapsed; nothing to compare" - ) - # Late against the *pipeline-wide* first build, not only the build in - # the collection's own handler: `_handle_gpu_memory_settings` builds - # the configuration many steps earlier, so hoisting the collection - # above the local build still leaves that cache describing raw input. - steps, methods, reached, _step_lines = _pipeline() - _declared_at, first_build = _declaration_positions() - self.assertIsNotNone( - first_build, "no handler builds a ModelConfig; the scan broke" - ) - collecting_steps = [ - index - for index, step in enumerate(steps) - for method in reached[step] - if any( - isinstance(node, ast.Call) - and ( - node.func.attr - if isinstance(node.func, ast.Attribute) - else getattr(node.func, "id", None) - ) - == "collect_model_override_declarations" - for node in ast.walk(methods[method]) - ) - ] - self.assertTrue(collecting_steps, "no pipeline step collects the registry") - collection_is_late = min(collecting_steps) > first_build[0] or ( - build_line is not None and collect_line > build_line - ) - late = frozenset(reads & registry) if collection_is_late else frozenset() - self.assertEqual( - sorted(late), - sorted(_STALE_FROM_THE_REGISTRIES), - "the set of ModelConfig-read fields the registries decide after the " - f"build changed (collection at line {collect_line}, build at line " - f"{build_line}); read the comment on _STALE_FROM_THE_REGISTRIES " - "before editing it", - ) - - def test_the_pinned_stale_field_is_still_stale(self): - """If the ordering gets fixed, this pin has to be retired, not kept. - - A pin that outlives the defect it describes is worse than none: it - documents a hazard that no longer exists and hides the day one appears. - """ - steps, methods, reached, step_lines = _pipeline() - dispatch = methods["run_resolution_pipeline"] - hooks = _hook_declarations(dispatch, _DISPATCH_MODULE) - build_line = min( - step_lines[step] - for step in steps - for method in reached[step] - if any( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "model_config_of" - for node in ast.walk(methods[method]) - ) - ) - for field in _STALE_IN_THE_MODEL_CONFIG: - self.assertIn( - field, - hooks, - f"{field} is pinned as decided after the build, but no hook " - "declares it any more; retire the pin", - ) - self.assertGreater( - hooks[field], - build_line, - f"{field} is now decided before the model configuration is " - "built; retire the pin", - ) - - def test_every_opaque_callback_is_still_late(self): - """The opaque resolvers all run after the model configuration is built. - - A plugin that rewrites `dtype` or `model_path` in one of them is - invisible to the configuration already cached, and no scan can say - whether it does: the implementations are out of tree. So the positions - are the pin, and the set of callbacks is pinned with them -- a new one - has to be placed against the build by whoever adds it. Moving them all - above the build fixes the hazard and fails this test; retire the pin - then, rather than keeping a note about a hazard that is gone. - """ - steps, methods, reached, step_lines = _pipeline() - dispatch = methods["run_resolution_pipeline"] - positions = _opaque_callback_positions(dispatch, _DISPATCH_MODULE) - self.assertEqual( - sorted(positions), - [ - "algo.handle_server_args", - "algo.validate_server_args", - "current_platform.apply_server_args_defaults", - ], - "the set of resolvers handed to declare_direct_writes changed; each " - "one needs its position against the ModelConfig build looked at", - ) - _declared_at, first_build = _declaration_positions() - self.assertIsNotNone( - first_build, "no handler builds a ModelConfig; the scan broke" - ) - build_line = step_lines[first_build[1]] - for spelling, line in sorted(positions.items()): - self.assertGreater( - line, - build_line, - f"{spelling} now runs before the model configuration is built, " - "so a plugin's writes reach it; retire the pin", - ) - - def test_the_documented_exception_is_still_the_only_one(self): - """A field pinned as read-before-resolution has to still be all three. - - Read by the constructor, written by resolution, and written *after* the - build -- the last one is what makes the exemption load-bearing. Without - it, moving the declaration earlier leaves the name sitting in the - exempt set with nothing to exempt, and the next field that lands in - this position gets waved through by a pin nobody re-read. - """ - wanted = _constructor_reads() - declared_at, first_build = _declaration_positions() - self.assertIsNotNone( - first_build, "no handler builds a ModelConfig; the scan broke" - ) - for field in sorted(_READ_BEFORE_RESOLUTION): - self.assertIn( - field, - wanted, - f"{field} is pinned as read before resolution, but the " - "constructor no longer reads it; retire the pin", - ) - self.assertIn( - field, - declared_at, - f"{field} is pinned as read before resolution, but resolution " - "no longer writes it; retire the pin", - ) - self.assertGreaterEqual( - declared_at[field], - (first_build[0], 1), - f"{field} is now decided before the model configuration is " - "built, so the exemption covers nothing; retire the pin", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/server_args/test_no_public_non_field_slot.py b/test/registered/unit/server_args/test_no_public_non_field_slot.py deleted file mode 100644 index 8d4d9e33b..000000000 --- a/test/registered/unit/server_args/test_no_public_non_field_slot.py +++ /dev/null @@ -1,86 +0,0 @@ -"""The record grows no attribute the projection cannot see. - -A publicly-named attribute that is not a dataclass field is invisible to every -other guard here: the namespace coverage walks fields, the projection walks -fields, and the read ratchets watch field reads. Three of them accumulated that -way -- a `ModelConfig` cache, an `moe_ep_size` that only a log line read, and an -env-derived `grpc_worker_threads` that one entry point read across the boundary. - -Leading-underscore names are the record's own bookkeeping and stay: the -read-only guard classifies writability by that spelling, so a private name is -already outside the config tier by construction. -""" - -import ast -import dataclasses -import pathlib -import unittest - -import sglang -from sglang.srt.server_args import ServerArgs -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=10, suite="base-a-test-cpu") - - -def _self_written_attributes() -> set: - """Names `ServerArgs` writes on itself, by either spelling.""" - source = ( - pathlib.Path(next(iter(sglang.__path__))) / "srt" / "server_args.py" - ).read_text(encoding="utf-8-sig") - tree = ast.parse(source) - cls = next( - node - for node in tree.body - if isinstance(node, ast.ClassDef) and node.name == "ServerArgs" - ) - written = set() - for node in ast.walk(cls): - if isinstance(node, ast.Assign): - for target in node.targets: - if ( - isinstance(target, ast.Attribute) - and isinstance(target.value, ast.Name) - and target.value.id == "self" - ): - written.add(target.attr) - if ( - isinstance(node, ast.Call) - and getattr(node.func, "attr", None) == "__setattr__" - and getattr(getattr(node.func, "value", None), "id", None) == "object" - and len(node.args) >= 2 - and isinstance(node.args[1], ast.Constant) - ): - written.add(node.args[1].value) - return written - - -class TestNoPublicNonFieldSlot(CustomTestCase): - def test_every_public_attribute_is_a_field(self): - written = _self_written_attributes() - # Anchor on a name, not a count: the count falls every time a derived - # read leaves the record, so a floor erodes with what it measures. - self.assertIn( - "_resolution_finished", - written, - f"the scan did not find the resolution flag the record sets on " - f"itself, so it is the scan that is broken, not the record: " - f"{sorted(written)}", - ) - fields = {field.name for field in dataclasses.fields(ServerArgs)} - stray = sorted( - name for name in written if not name.startswith("_") and name not in fields - ) - self.assertEqual( - [], - stray, - "these are written on the record under a public name but are not " - "fields, so the projection cannot see them and no other guard " - "watches them: make each a field, or give it the leading underscore " - f"that says it is the record's own bookkeeping: {stray}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/server_args/test_record_member_calls_resolve.py b/test/registered/unit/server_args/test_record_member_calls_resolve.py deleted file mode 100644 index 238caa8c0..000000000 --- a/test/registered/unit/server_args/test_record_member_calls_resolve.py +++ /dev/null @@ -1,144 +0,0 @@ -"""Every `server_args.()` in the tree names something the record has. - -Removing a member from `ServerArgs` means rewriting its callers, and the ones -inside `server_args.py` are the ones you fix by reflex. The cross-file caller is -what bites: `ServerArgs.ssl_verify()` moved to `serving_hook.ssl_verify_of()` and -one call site kept the old spelling as `self.server_args.ssl_verify()` -- a grep -for `server_args.ssl_verify()` does not find that, and nothing else looks. Every -`HttpServerEngineAdapter` request raised `AttributeError` before sending. - -So this resolves the call sites instead of grepping for them: every attribute -*called* on something statically known to be a record has to exist on the record. -It is deliberately not limited to methods the refactor touched -- the next -removal gets the same check for free. - -`multimodal_gen` carries a different, same-named class outside this contract, as -the other record ratchets also record. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=19, suite="base-a-test-cpu") - -import ast -import dataclasses -import pathlib -import unittest - -import sglang -from sglang.srt.server_args import ServerArgs -from sglang.test.test_utils import CustomTestCase - -_ROOTS = ( - pathlib.Path(next(iter(sglang.__path__))) / "srt", - pathlib.Path(__file__).resolve().parents[3], # test/ -) -_EXCLUDED = ("multimodal_gen",) - -# Attribute names that hold a `ServerArgs`. `resolving_view` and `resolved_view` -# proxy the record but answer for names it does not carry, so they are not here. -_RECORD_NAMES = ("server_args", "_server_args") - - -def _is_record(node) -> bool: - """`server_args`, `self.server_args`, `self._server_args`, `cls.server_args`.""" - if isinstance(node, ast.Name): - return node.id in _RECORD_NAMES - if isinstance(node, ast.Attribute): - return node.attr in _RECORD_NAMES - return False - - -def _rebound_locally(tree) -> set: - """Names assigned something that is plainly not a record. - - `server_args` is also a natural name for a dict of CLI flags or a list of - argv strings in test helpers, and those legitimately answer `.update()` and - `.items()`. A function that assigns one of those to the name is not talking - about the record in that scope. - """ - literal = (ast.Dict, ast.List, ast.DictComp, ast.ListComp) - builders = {"dict", "list", "tuple", "set"} - - def _not_a_record(value) -> bool: - if isinstance(value, literal): - return True - # `dict(...)` / `list(...)`, and an annotated `server_args: list[str] = [...]` - return ( - isinstance(value, ast.Call) and getattr(value.func, "id", None) in builders - ) - - rebound = set() - for node in ast.walk(tree): - if isinstance(node, ast.AnnAssign): - targets, value = [node.target], node.value - elif isinstance(node, ast.Assign): - targets, value = node.targets, node.value - else: - continue - if value is None or not _not_a_record(value): - continue - for target in targets: - if isinstance(target, ast.Name) and target.id in _RECORD_NAMES: - rebound.add(target.id) - return rebound - - -def _called_members(): - """{name: [file:line]} for every `.(...)` in the tree.""" - found: dict[str, list[str]] = {} - for root in _ROOTS: - for path in sorted(root.rglob("*.py")): - text = path.as_posix() - if any(part in text for part in _EXCLUDED): - continue - source = path.read_text(encoding="utf-8-sig") - if "server_args" not in source: - continue - try: - tree = ast.parse(source) - except SyntaxError: - continue - rebound = _rebound_locally(tree) - for node in ast.walk(tree): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Attribute) - and _is_record(node.func.value) - and getattr(node.func.value, "id", None) not in rebound - ): - found.setdefault(node.func.attr, []).append( - f"{path.name}:{node.lineno}" - ) - return found - - -class TestRecordMemberCallsResolve(CustomTestCase): - def test_every_called_member_exists_on_the_record(self): - called = _called_members() - self.assertGreater( - len(called), - 5, - f"only {len(called)} members called on a record; the scan is broken, " - "not the tree", - ) - available = set(dir(ServerArgs)) | { - field.name for field in dataclasses.fields(ServerArgs) - } - missing = { - name: sites - for name, sites in sorted(called.items()) - if name not in available - } - self.assertEqual( - {}, - missing, - "these are called on a ServerArgs but the record has no such member -- " - "each one raises AttributeError at the call. A member that moved out of " - "the record has to be rewritten at every call site, including the ones " - f"reached through `self.server_args`: {missing}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/server_args/test_resolution_declarations.py b/test/registered/unit/server_args/test_resolution_declarations.py index 6262cb99b..51667b3e6 100644 --- a/test/registered/unit/server_args/test_resolution_declarations.py +++ b/test/registered/unit/server_args/test_resolution_declarations.py @@ -1,15 +1,10 @@ """Resolution writes are recorded, not just applied. The projection that replaces field materialization reads the declaration stash, -so a resolution write that only assigns the field is invisible to it. Every -resolver declares now -- the record's handlers through `self._declare`, the -hooks and hardware defaults through `declare_resolution` -- and that is pinned -two ways: no bare assignment to a field survives anywhere a ServerArgs instance -is in reach, and after resolution `resolution_result` answers for every declared -field with what the stash holds. The second check is what the stash is measured -against: the two can disagree only if something wrote behind the stash's back. A -third check runs the other way -- every field resolution moved has to be -explained by the stash, which covers the spellings a source scan cannot see. +so a resolution write that bypasses the stash is invisible to it. These tests +compare the raw input, resolved record, declaration result, and published bags +across representative configurations. A field that moves without a declaration +or is projected into the wrong namespace therefore fails on observed state. """ import ast @@ -122,114 +117,6 @@ _REACHED_BY_SHAPES = frozenset( ) -def _late_resolvers(): - """Callables that reach `declare_late_resolution`, derived per module.""" - found = set() - for relative in ("server_args.py", "parser/template_detection.py"): - tree = ast.parse((_SRT / relative).read_text(encoding="utf-8-sig")) - functions = { - node.name: node - for node in ast.walk(tree) - if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) - } - - def reaches(name, seen=None): - seen = seen if seen is not None else set() - if name in seen or name not in functions: - return False - seen.add(name) - for node in ast.walk(functions[name]): - if not isinstance(node, ast.Call): - continue - called = ( - node.func.attr - if isinstance(node.func, ast.Attribute) - else getattr(node.func, "id", None) - ) - if called == "declare_late_resolution": - return True - if called and reaches(called, seen): - return True - return False - - found |= {name for name in functions if reaches(name)} - return found - - -def _server_args_writers(tree, path): - """Assignment targets that land on a ServerArgs instance. - - Two mechanisms reach the same instance during resolution: a handler writing - `self.`, and a helper elsewhere in the tree writing through a - `ServerArgs`-annotated parameter -- `set_default_server_args(args)` is - called from the pipeline and writes `args.`. Both bypass the - declaration stash, so both have to be scanned; scanning only the handlers - would let a field look converted while a second writer still assigns it. - """ - names = {"self"} if path.name == "server_args.py" else set() - # A parameter *named* `server_args` counts with or without the annotation. - for node in ast.walk(tree): - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - args = node.args - for arg in args.posonlyargs + args.args + args.kwonlyargs: - annotation = arg.annotation - if isinstance(annotation, ast.Constant): - text = annotation.value - elif isinstance(annotation, ast.Name): - text = annotation.id - elif isinstance(annotation, ast.Attribute): - text = annotation.attr - else: - continue - if text == "ServerArgs": - names.add(arg.arg) - names |= { - arg.arg for arg in args.posonlyargs + args.args if arg.arg == "server_args" - } - return names - - -def _bare_assignments(): - """Assignments to a converted field that never reach the stash.""" - found = [] - for path in sorted(_SRT.rglob("*.py")): - try: - tree = ast.parse(path.read_text(encoding="utf-8-sig")) - except SyntaxError: - continue - names = _server_args_writers(tree, path) - if not names: - continue - for node in ast.walk(tree): - if isinstance(node, ast.Assign): - targets = node.targets - elif isinstance(node, (ast.AugAssign, ast.AnnAssign)): - targets = [node.target] - else: - continue - # Destructured targets count: `(sa.a, sa.b) = f()` writes two - # fields and is not an `ast.Attribute` at the top level. - flat = [] - for target in targets: - if isinstance(target, (ast.Tuple, ast.List)): - flat.extend(target.elts) - else: - flat.append(target) - for target in flat: - if ( - isinstance(target, ast.Attribute) - and isinstance(target.value, ast.Name) - and target.value.id in names - and target.attr in _RESOLVED_FIELDS - ): - found.append( - f"{path.relative_to(_SRT)}:{node.lineno} " - f"{target.value.id}.{target.attr}" - ) - return sorted(found) - - def shape_key(shape): """A shape rendered short enough for a failure message.""" return ",".join(f"{k}={v}" for k, v in sorted(shape.items())) or "defaults" @@ -299,15 +186,6 @@ class TestResolutionDeclarations(CustomTestCase): server_args.resolve_once() return server_args - def test_converted_fields_are_not_assigned_bare(self): - bare = _bare_assignments() - self.assertEqual( - bare, - [], - "a converted field is assigned directly, so the projection would " - "not see this write:\n " + "\n ".join(bare), - ) - def test_the_stash_accounts_for_every_change_resolution_made(self): """The other direction: a field resolution moved is in the stash. @@ -419,10 +297,7 @@ class TestResolutionDeclarations(CustomTestCase): the last hop: whether the leaf is reachable through the path the metadata declares, and whether it carries the resolved value once it is. Both sides here come from that metadata, so this cannot tell that - a field is assigned to the *wrong* group -- the readers are the - independent source for that, and - `test_server_args_namespaces.py::test_the_readers_agree_with_the_namespace_metadata` - is where the two are compared. + a field is assigned to the *wrong* group. """ import sglang.srt.runtime_context as runtime_context from sglang.srt.arg_groups.arg_utils import namespace_of @@ -643,54 +518,6 @@ class TestResolutionDeclarations(CustomTestCase): "normalization is a declaration; the record keeps what was passed", ) - def test_the_launcher_finishes_resolving_before_it_publishes(self): - """Every late resolver runs above the publish, in the source. - - A published record refuses to be written, so a late resolver below the - publish raises at startup rather than at test time -- and only for the - configuration that reaches it, which is why the LoRA path can break - while every other launch stays green. Both sides are derived: which - callables reach `declare_late_resolution`, and where the launcher calls - them. - """ - launcher = _SRT / "entrypoints/engine.py" - late = {"check_server_args", "resolve_auto_parsers"} | _late_resolvers() - tree = ast.parse(launcher.read_text(encoding="utf-8-sig")) - function = next( - node - for node in ast.walk(tree) - if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) - and node.name == "_launch_subprocesses" - ) - published_at = [ - node.lineno - for node in ast.walk(function) - if isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "publish" - ] - self.assertEqual(len(published_at), 1, "the launcher publishes once") - too_late = sorted( - f"{name}() at line {node.lineno}" - for node in ast.walk(function) - if isinstance(node, ast.Call) - for name in [ - ( - node.func.attr - if isinstance(node.func, ast.Attribute) - else getattr(node.func, "id", None) - ) - ] - if name in late and node.lineno > published_at[0] - ) - self.assertEqual( - too_late, - [], - f"these resolve after the launcher publishes at line " - f"{published_at[0]}, and a published record refuses to be " - f"written:\n " + "\n ".join(too_late), - ) - def test_an_undeclared_field_still_holds_the_raw_input(self): """Nothing writes a field behind the stash's back. @@ -832,48 +659,6 @@ class TestResolutionDeclarations(CustomTestCase): "so a decision made inside the declared object was dropped", ) - def test_every_platform_hook_that_takes_the_record_is_captured(self): - """A second out-of-tree config hook must not arrive uncaptured. - - `apply_server_args_defaults` is the one method on the platform - interface that is handed the record, and its implementations live in - other distributions -- no source scan of this tree can see what they - write, so the pipeline diffs the record across the call instead. A new - hook of the same shape would be invisible again, and this is what - notices. Derived from the interface rather than listed: a rename keeps - working, an addition fails. - """ - interface = _SRT / "platforms" / "interface.py" - tree = ast.parse(interface.read_text(encoding="utf-8-sig")) - taking_the_record = set() - for node in ast.walk(tree): - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - arguments = node.args - names = [ - arg.arg - for arg in arguments.posonlyargs + arguments.args + arguments.kwonlyargs - ] - if any(name == "server_args" or name.endswith("_args") for name in names): - taking_the_record.add(node.name) - self.assertEqual( - taking_the_record, - {"apply_server_args_defaults"}, - "the platform interface hands the startup record to a method this " - "test does not know about; either it only reads, or its writes need " - "capturing like apply_server_args_defaults", - ) - - pipeline = (_SRT / "arg_groups" / "pipeline.py").read_text(encoding="utf-8-sig") - for hook in sorted(taking_the_record): - self.assertIn( - f"current_platform.{hook},", - pipeline, - f"{hook} is called directly instead of through the write " - "capture, so an out-of-tree plugin's defaults would be dropped " - "by the projection", - ) - def test_the_shapes_reach_the_fields_they_are_meant_to(self): """A green agreement check over an empty stash would prove nothing.""" declared = set() diff --git a/test/registered/unit/server_args/test_resolution_is_reproducible.py b/test/registered/unit/server_args/test_resolution_is_reproducible.py index 2b791107c..998692fdf 100644 --- a/test/registered/unit/server_args/test_resolution_is_reproducible.py +++ b/test/registered/unit/server_args/test_resolution_is_reproducible.py @@ -35,7 +35,6 @@ import unittest.mock import torch -import sglang from sglang.srt.arg_groups.overrides import model_config_of, resolution_result from sglang.srt.environ import EnvField, envs from sglang.srt.server_args import ServerArgs @@ -481,274 +480,6 @@ class TestResolutionIsReproducible(_RestoresProcessState, CustomTestCase): self.assertEqual(getattr(first, "_resolved_overrides", None), first_provenance) -class TestProgramsResolveBeforeReadingResolution(CustomTestCase): - """A program that builds its own record resolves it before reading what - resolution decides. - - Construction is inert, so a program that builds a record and then reads a - resolution-written field reads the CLI default. Two of these shipped past - the earlier censuses because those are rooted at the `sglang` package: the - model gateway's launcher sized its worker plan from a raw `dp_size` - (`--dwdp-size 4` launched one server instead of four) and a speculative - benchmark forwarded `--mem-fraction-static None` to the server it spawns. - So the universe here is the *repository*, not the package. - """ - - # Entries that hand the record on instead of reading it. Reason required. - _EXEMPT: dict = {} - - def _repo_root(self): - # /python/sglang/__init__.py -> - root = pathlib.Path(next(iter(sglang.__path__))).resolve().parents[1] - if root.name == "python": - root = root.parent - return root - - def _written_fields(self): - """Fields resolution declares, read out of the pipeline's own source. - - Deliberately local: the chain ratchet has a wider derivation (it also - walks the model-override registries), but it arrives later in this - series, and a check that imports it would fail at this PR's boundary. - Coarser is fine here -- what this needs is the fields the entries below - actually read -- and the floor keeps it from drifting narrower. - """ - import ast - import dataclasses as _dataclasses - - from sglang.srt.server_args import ServerArgs as _ServerArgs - - srt = pathlib.Path(next(iter(sglang.__path__))).resolve() / "srt" - declarers = {"declare_resolution", "declare_late_resolution"} - fields = set() - field_names = {field.name for field in _dataclasses.fields(_ServerArgs)} - # The record plus every module under `arg_groups/`: a handler declares - # from whichever of the two it lives in. - sources = [srt / "server_args.py", *sorted((srt / "arg_groups").rglob("*.py"))] - for source in sources: - tree = ast.parse(source.read_text(encoding="utf-8-sig")) - for node in ast.walk(tree): - # Registry data: provider dict keys are field names as - # *data*, invisible to the keyword scan below. Filtered - # against the real field set. - if isinstance(node, ast.Dict): - fields |= { - key.value - for key in node.keys - if isinstance(key, ast.Constant) - and isinstance(key.value, str) - and key.value in field_names - } - if not isinstance(node, ast.Call): - continue - func = node.func - called = ( - func.attr - if isinstance(func, ast.Attribute) - else getattr(func, "id", "") - ) - if called in declarers or called == "update": - fields |= { - kw.arg - for kw in node.keywords - if kw.arg and (called != "update" or kw.arg in field_names) - } - return fields - - def _candidates(self, root): - """Source files that build a record, with the names they bind it to.""" - import ast - - skip = {".git", "build", "dist", "node_modules", ".venv", "target"} - found = {} - for path in sorted(root.rglob("*.py")): - parts = set(path.relative_to(root).parts) - if parts & skip: - continue - rel = path.relative_to(root).as_posix() - # Tests build raw records on purpose. - if rel.startswith("test/") or "/test/" in rel or "/tests/" in rel: - continue - try: - tree = ast.parse(path.read_text(encoding="utf-8-sig")) - except (SyntaxError, UnicodeDecodeError): - continue - # Which local names are *the srt record*, by import source: the - # diffusion runtime has a same-spelled `ServerArgs` with no - # resolution, so the spelling alone is not enough. - record_classes, record_helpers = set(), set() - for node in ast.walk(tree): - if not isinstance(node, ast.ImportFrom): - continue - for alias in node.names: - bound = alias.asname or alias.name - if node.module == "sglang" and alias.name == "ServerArgs": - record_classes.add(bound) - if node.module == "sglang.srt.server_args": - if alias.name == "ServerArgs": - record_classes.add(bound) - if alias.name == "prepare_server_args": - record_helpers.add(bound) - names = set() - for node in ast.walk(tree): - if isinstance(node, ast.Assign): - targets = node.targets - elif isinstance(node, (ast.AnnAssign, ast.NamedExpr)): - # `x: ServerArgs = ...` is an AnnAssign, not an Assign. - targets = [node.target] - else: - continue - call = getattr(node, "value", None) - if not isinstance(call, ast.Call): - continue - func = call.func - # `prepare_server_args(argv)` is the CLI launcher's way. - builds = ( - isinstance(func, ast.Name) - and func.id in (record_classes | record_helpers) - ) or ( - isinstance(func, ast.Attribute) - and func.attr == "from_cli_args" - and isinstance(func.value, ast.Name) - and func.value.id in record_classes - ) - if builds: - names |= {t.id for t in targets if isinstance(t, ast.Name)} - if names: - found[rel] = (tree, names, path) - return found - - def test_every_program_that_builds_a_record_resolves_it(self): - import ast - - root = self._repo_root() - candidates = self._candidates(root) - self.assertGreater( - len(candidates), - 10, - f"only {len(candidates)} files build a record under {root}; either " - "this is not a source checkout or the scan broke", - ) - written = self._written_fields() - self.assertGreater(len(written), 50, "the written-field set collapsed") - # What the escaped entries actually read: a narrower derivation goes - # quiet on exactly those. - for field in ("dp_size", "mem_fraction_static"): - self.assertIn(field, written) - - offenders = [] - for rel, (tree, names, path) in sorted(candidates.items()): - source = path.read_text(encoding="utf-8-sig") - if "resolve_once(" in source or "publish(" in source: - continue - reads = sorted( - { - f"{node.attr}:{node.lineno}" - for node in ast.walk(tree) - if isinstance(node, ast.Attribute) - and isinstance(node.value, ast.Name) - and node.value.id in names - and node.attr in written - } - ) - if reads and rel not in self._EXEMPT: - offenders.append(f"{rel} reads {', '.join(reads[:4])}") - self.assertEqual( - offenders, - [], - "a program builds its own record and reads what resolution decides " - "without resolving it, so it reads the CLI default:\n " - + "\n ".join(offenders), - ) - self.assertEqual( - sorted(set(self._EXEMPT) - set(candidates)), - [], - "an exemption names a file that no longer builds a record", - ) - - -class TestForksResolveFirst(CustomTestCase): - """A process that forks a child to run the record resolves it first. - - The pipeline probes the device (the default attention backend reads the CUDA - capability), and a forked child cannot initialize CUDA once its parent has. - Construction used to resolve, so the probe always happened in whoever built - the record; now it happens at the gate, and the gate must not be reached for - the first time inside a fork. - """ - - # Sites inside the launcher: `_launch_subprocesses` resolves at its top, so - # every fork below it already has a resolved record. - _AFTER_LAUNCHER_RESOLVE = { - "srt/entrypoints/engine.py", - "srt/managers/data_parallel_controller.py", - "srt/disaggregation/encoder/grpc_server.py", - "srt/disaggregation/encoder/runtime.py", - "srt/elastic_ep/expert_backup_manager.py", - } - - def test_every_fork_of_a_record_has_a_resolved_one(self): - import ast - - package_root = pathlib.Path(next(iter(sglang.__path__))).resolve() - offenders, examined = [], 0 - for path in sorted(package_root.rglob("*.py")): - rel = path.relative_to(package_root).as_posix() - if rel.startswith("test/") or "/test/" in rel: - continue - # The diffusion runtime has its own record with no gate. - if rel.startswith("multimodal_gen/"): - continue - try: - source = path.read_text(encoding="utf-8-sig") - if "Process" not in source: - continue - tree = ast.parse(source) - except (SyntaxError, UnicodeDecodeError): - continue - for node in ast.walk(tree): - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - forks = [ - call - for call in ast.walk(node) - if isinstance(call, ast.Call) - and ( - ( - isinstance(call.func, ast.Attribute) - and call.func.attr == "Process" - ) - or ( - isinstance(call.func, ast.Name) - and call.func.id == "Process" - ) - ) - and "server_args" in (ast.get_source_segment(source, call) or "") - ] - if not forks: - continue - body = ast.get_source_segment(source, node) or "" - examined += 1 - # `spawn` starts a fresh interpreter, so the child may probe. - if 'get_context("spawn")' in body or "'spawn'" in body: - continue - if "resolve_once(" in body or "publish(" in body: - continue - if rel in self._AFTER_LAUNCHER_RESOLVE: - continue - offenders.append(f"{rel}:{forks[0].lineno} {node.name}") - self.assertGreater( - examined, 5, f"only {examined} fork sites found; the scan broke" - ) - self.assertEqual( - offenders, - [], - "these fork a child that will resolve the record, without resolving " - "it first -- the child cannot initialize CUDA if this process " - f"already has:\n " + "\n ".join(offenders), - ) - - class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase): """A resolved record copied with `dataclasses.replace` loses what makes it resolved, and the next publish resolves it a second time -- over values it @@ -876,215 +607,6 @@ class TestACopyStaysResolved(_RestoresProcessState, CustomTestCase): "the parent's resolution decided", ) - def test_no_bare_replace_of_a_record_outside_the_helper(self): - """`dataclasses.replace` on a record is the helper's job now. - - Derived, not listed: any `dataclasses.replace` whose first argument is - named for a record. The helper's own call is the positive control -- if - the scan stops seeing it, the scan broke rather than the tree. - """ - import ast - - # The repository, not the package: the gateway is outside `sglang/`. - package_root = pathlib.Path(next(iter(sglang.__path__))).resolve().parents[1] - if package_root.name == "python": - package_root = package_root.parent - helper = "python/sglang/srt/server_args.py" - bare, inside_helper = [], 0 - - def replaces_a_record(node, record_names): - if not ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Attribute) - and node.func.attr == "replace" - and isinstance(node.func.value, ast.Name) - and node.func.value.id == "dataclasses" - and node.args - ): - return False - first = node.args[0] - name = ( - first.id if isinstance(first, ast.Name) else getattr(first, "attr", "") - ) - return name in record_names or "server_args" in name - - for path in sorted(package_root.rglob("*.py")): - try: - tree = ast.parse(path.read_text(encoding="utf-8-sig")) - except SyntaxError: - continue - rel = path.relative_to(package_root).as_posix() - # `self` is a record only inside the record's own class body. - in_record_class = [ - node - for node in ast.walk(tree) - if isinstance(node, ast.ClassDef) and node.name == "ServerArgs" - ] - for scope, record_names in [(tree, set())] + [ - (klass, {"self"}) for klass in in_record_class - ]: - for node in ast.walk(scope): - if not replaces_a_record(node, record_names): - continue - if rel == helper and record_names: - inside_helper += 1 - elif not record_names: - bare.append(f"{rel}:{node.lineno}") - self.assertEqual( - inside_helper, - 1, - "the scan no longer finds `replace_resolved`'s own call; it broke", - ) - self.assertEqual( - bare, - [], - "a record is copied with a bare `dataclasses.replace`, so the copy " - "loses the parent's resolution and the next publish resolves it " - "again: " + ", ".join(bare), - ) - - -class TestTheResolutionSeamHasOneCaller(CustomTestCase): - """The pipeline is entered from exactly one place, and that place decides - whether it runs at all. - - ``resolve_once`` is the gate: the handlers are not written to survive a - second pass over their own output, so a record must go through the pipeline - at most once. Keeping the pipeline itself down to a single caller is what - makes that gate impossible to bypass -- and what keeps the remaining move - (construction time to publish time) a matter of who calls the gate. - """ - - def test_only_the_gate_runs_the_pipeline(self): - import ast - from pathlib import Path - - import sglang - - package_root = Path(next(iter(sglang.__path__))) - callers = [] - for path in sorted(package_root.rglob("*.py")): - try: - source = path.read_text() - if "run_resolution_pipeline" not in source: - continue - tree = ast.parse(source) - except SyntaxError: - continue - # The full (class, function, ...) scope chain, so the assertion can - # say "the one caller is ServerArgs.__post_init__" -- not merely - # that nothing outside a function named __post_init__ calls it. - scopes = {} - for node in ast.walk(tree): - own = scopes.get(id(node), ()) - if isinstance( - node, (ast.ClassDef, ast.FunctionDef, ast.AsyncFunctionDef) - ): - own = own + (node.name,) - for child in ast.iter_child_nodes(node): - scopes[id(child)] = own - for node in ast.walk(tree): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "run_resolution_pipeline" - ): - rel = path.relative_to(package_root).as_posix() - callers.append((rel, ".".join(scopes.get(id(node), ())))) - # Every call, compared whole: a removed call, a duplicate inside - # __post_init__, or another class growing a same-named __post_init__ - # all show up here. - self.assertEqual( - [("srt/server_args.py", "ServerArgs.resolve_once")], - callers, - "the resolution pipeline must be entered exactly once, from " - f"ServerArgs.resolve_once; found: {callers}", - ) - - def test_the_gate_is_reached_from_the_launcher_and_from_publish(self): - """Both entries go through the gate, so neither can resolve twice. - - The launcher resolves the engine's record before reading any resolved - value from it; every publishing process asks the gate on the way in and - finds nothing left to do when the record arrived resolved. - """ - import ast - from pathlib import Path - - import sglang - - package_root = Path(next(iter(sglang.__path__))) - callers = [] - for path in sorted(package_root.rglob("*.py")): - try: - source = path.read_text() - if "resolve_once" not in source: - continue - tree = ast.parse(source) - except SyntaxError: - continue - for node in ast.walk(tree): - # `self.resolve_once()` at construction; publish looks the - # attribute up first, so it appears as a bare name call. - called = ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Attribute) - and node.func.attr == "resolve_once" - ) or ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "resolve_once" - ) - if called: - callers.append(path.relative_to(package_root).as_posix()) - machinery = {"srt/entrypoints/engine.py", "srt/runtime_context.py"} - self.assertEqual( - [ - # Program entries: each builds a record from its own - # arguments and then reads effective configuration, or hands it - # to a fork that must not be the first to probe the device. - "benchmark/endpoint.py", - "benchmark/offline_throughput.py", - "benchmark/one_batch.py", - "benchmark/one_batch_server.py", - "compile_deep_gemm.py", - "lang/backend/runtime_endpoint.py", - "launch_server.py", - # The mechanism. - "srt/entrypoints/engine.py", - "srt/entrypoints/http_server_engine.py", - "srt/runtime_context.py", - ], - sorted(set(callers)), - f"the resolution gate grew or lost a caller: {sorted(set(callers))}", - ) - # The rule the list stands for: a caller that is not the mechanism - # resolves a record it built itself from argv. Anything else was handed - # one someone already resolved, or should publish. - for caller in sorted(set(callers) - machinery): - source = (package_root / caller).read_text() - # Either the module turned argv into the record -- the dataclass, - # the CLI classmethod, or the argv helper `launch_server.py` uses - # -- or it hands the record to a fork, which has to resolve first: - # the pipeline probes the device and a forked child cannot - # re-initialize CUDA. A worker handed a resolved record is neither. - builds_its_own = any( - spelling in source - for spelling in ( - "ServerArgs(", - ".from_cli_args(", - "prepare_server_args(", - ) - ) or ("Process(" in source and "server_args" in source) - # `assertTrue`, not `assertIn`: the container is a whole module. - self.assertTrue( - builds_its_own, - f"{caller} calls the resolution gate but does not build the " - "record it resolves; a record it was handed is already " - "resolved by whoever built it, and publish resolves what it " - "is handed", - ) - class TestResolutionStaysLazy(CustomTestCase): """Resolving a dummy model must not load the families it never reaches. @@ -1096,92 +618,6 @@ class TestResolutionStaysLazy(CustomTestCase): `override_server_args` in the test suite. """ - def test_no_hook_module_imports_another_at_module_scope(self): - import ast - - import sglang - - srt = pathlib.Path(next(iter(sglang.__path__))).resolve() / "srt" - offenders = [] - for path in sorted((srt / "arg_groups").glob("*.py")): - for node in ast.parse(path.read_text(encoding="utf-8-sig")).body: - if ( - isinstance(node, ast.ImportFrom) - and node.module - and node.module.startswith("sglang.srt.arg_groups") - and node.module.endswith("_hook") - ): - offenders.append(f"{path.name}:{node.lineno} -> {node.module}") - self.assertEqual( - offenders, - [], - "a hook module imports another at module scope, so loading one " - "family drags in a family it may never call. Import it inside the " - "function that calls it:\n " + "\n ".join(offenders), - ) - - def test_no_family_is_imported_before_the_step_that_calls_it(self): - """Source-level, so it holds whatever else the process has imported. - - Every hook import inside the dispatcher must come after the imports of - the families reached earlier and before its own first call -- what an - eager block at the top of the function breaks, and what a `sys.modules` - diff cannot see once another test has loaded those modules. - """ - import ast - - import sglang - - srt = pathlib.Path(next(iter(sglang.__path__))).resolve() / "srt" - tree = ast.parse( - (srt / "arg_groups" / "pipeline.py").read_text(encoding="utf-8-sig") - ) - dispatch = next( - node - for node in ast.walk(tree) - if isinstance(node, ast.FunctionDef) - and node.name == "run_resolution_pipeline" - ) - early_return = min( - ( - n.lineno - for n in ast.walk(dispatch) - if isinstance(n, ast.Return) and n.value is None - ), - default=None, - ) - self.assertIsNotNone(early_return, "the dummy short circuit is gone") - - imported_early, called_early = set(), set() - for node in ast.walk(dispatch): - if ( - isinstance(node, ast.ImportFrom) - and node.module - and node.module.endswith("_hook") - and node.lineno < early_return - ): - imported_early.add(node.module.rsplit(".", 1)[-1]) - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.lineno < early_return - ): - called_early.add(node.func.id) - - hooks = {} - for path in sorted((srt / "arg_groups").glob("*_hook.py")): - for node in ast.parse(path.read_text(encoding="utf-8-sig")).body: - if isinstance(node, ast.FunctionDef): - hooks[node.name] = path.stem - needed_early = {hooks[name] for name in called_early if name in hooks} - self.assertEqual( - imported_early - needed_early, - set(), - "the dispatcher imports a hook family before the dummy short " - "circuit without calling it there, so every dummy resolution pays " - "for a family it never reaches", - ) - def test_a_dummy_resolution_loads_only_what_it_reaches(self): """The same claim measured, in an interpreter of its own. diff --git a/test/registered/unit/server_args/test_resolution_reads_no_bag.py b/test/registered/unit/server_args/test_resolution_reads_no_bag.py deleted file mode 100644 index 6929784d0..000000000 --- a/test/registered/unit/server_args/test_resolution_reads_no_bag.py +++ /dev/null @@ -1,333 +0,0 @@ -"""Resolution does not read the config bags, because they do not exist yet. - -The bags are projected from what resolution decides, so anything the pipeline -calls has to read the resolving state instead — `resolved_view(server_args)`, -or the view a handler already holds. A bag read reached from resolution raises -`config namespace ... not published`, and only on the branch that reaches it: -the diffusion-LM page-size pass needed one model family, the Marlin LoRA -validation needed one MoE runner backend. Both were written, merged into a -branch, and stayed green for everything except the configuration that triggers -them. - -`test_publish_precedes_bag_reads.py` is the same worry from the other side, but -it walks the *process entries* — it cannot see a helper the pipeline calls, and -neither of the two above appeared in it. - -The walk starts from three places: the symbols the pipeline imports, the live -resolution registries (every pass and override provider, taken from the -registries themselves rather than from the decorator that put it there -- most -providers register through a helper call), and the passes named at a -`run_post_process_pass(sa, fn)` call site. From there it follows calls -in-module, one hop out, and matches an accessor whether it is spelled bare or -through an object. - -What this still cannot see: a bag read reached through a method rather than a -module-level function, one behind an import the walk does not follow, and one -in a callable that reaches the pipeline through a variable no call site names. -It is a ratchet, not a proof. -""" - -import ast -import functools -import inspect -import pathlib -import unittest - -import sglang -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=22, suite="base-a-test-cpu") - -_SRT = pathlib.Path(sglang.__file__).resolve().parent / "srt" - - -def _accessor_names(): - """Every bag accessor `runtime_context` exports, read from the module. - - Listing them by hand is how this went stale once already: the list had - eighteen names while the module exported twenty-five, so a resolution-time - `get_flags().x` or `get_resources().y` would have walked straight past. - """ - tree = ast.parse((_SRT / "runtime_context.py").read_text(encoding="utf-8-sig")) - names = { - node.name - for node in tree.body - if isinstance(node, ast.FunctionDef) and node.name.startswith("get_") - } - # Two that are not bags: the context object itself, and the platform facts. - # Both answer before anything is published. - return frozenset(names - {"get_context", "get_platform"}) - - -_BAG_ACCESSORS = _accessor_names() - -# `get_device` also names the device-string utility and the platform method, -# so only the bare spelling is the accessor. -_ATTRIBUTE_SPELLED = _BAG_ACCESSORS - {"get_device"} - -# The pipeline itself and the mechanism it publishes through: `runtime_context` -# defines the accessors, and `arg_groups` is the pipeline's own code. -_OWN = ("server_args.py", "runtime_context.py") - - -def _pipeline_sources(): - """The record plus every module under `arg_groups/`. - - A handler that moved out of the record takes its imports with it, so - seeding the walk from two files would stop covering it. - """ - return [_SRT / "server_args.py", *sorted((_SRT / "arg_groups").rglob("*.py"))] - - -def _module_of(name): - """`sglang.srt.a.b` -> the file, if it is one of ours.""" - if not name or not name.startswith("sglang.srt."): - return None - rel = name[len("sglang.srt.") :].replace(".", "/") - for candidate in (_SRT / f"{rel}.py", _SRT / rel / "__init__.py"): - if candidate.exists(): - return candidate - return None - - -def _imported_symbols(paths): - """{module file: {symbol names imported from it}} across the given sources.""" - out = {} - for path in paths: - for node in ast.walk(ast.parse(path.read_text(encoding="utf-8-sig"))): - if not isinstance(node, ast.ImportFrom): - continue - target = _module_of(node.module) - if target is None or target.name in _OWN: - continue - out.setdefault(target, set()).update(alias.name for alias in node.names) - return out - - -def _registry_functions(): - """Every callable the resolution registries will call, from the registries. - - Not from decorator syntax: most model-override providers register through - a `_register_for(...)` helper rather than a decorator, so a scan keyed on - the decorator name walked past all of them -- 39 entries found where the - registries hold 65. However a provider registers, it is in the registry - once its module is imported, and `inspect` says where it came from. - """ - from sglang.srt.arg_groups import overrides - - functions = list(overrides.POST_PROCESS_PASSES) - functions += [fn for fns in overrides._MODEL_OVERRIDE_FNS.values() for fn in fns] - functions += [fn for _predicate, fn in overrides._PREDICATE_OVERRIDE_FNS] - return functions - - -@functools.lru_cache(maxsize=None) -def _registered_entries(): - """Entries the import map cannot reach: passes and override providers. - - A pass arrives at the pipeline as a value, and the registry calls its - providers by dictionary lookup. Both run during resolution, so a bag read - inside one raises exactly like a bag read in a handler -- and neither is - named by an import the walk can follow. - """ - entries = set() - for fn in _registry_functions(): - target = inspect.unwrap(fn) - name = getattr(target, "__name__", "") - if not name or name == "": - continue - source = inspect.getsourcefile(target) - if source is None: - continue - path = pathlib.Path(source).resolve() - if _SRT in path.parents: - entries.add((path, name)) - # A pass handed over by value is in no registry, so its call sites are read - # from the source. The entry carries the *defining* file: `_reaches_a_bag` - # walks functions in the entry's file, so a call-site key walks nothing. - by_value = set() - sources = { - path: path.read_text(encoding="utf-8-sig") - for path in sorted(_SRT.rglob("*.py")) - } - for path, source in sources.items(): - if "run_post_process_pass" not in source: - continue - try: - tree = ast.parse(source) - except SyntaxError: - continue - for node in ast.walk(tree): - if ( - isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.func.id == "run_post_process_pass" - and len(node.args) >= 2 - and isinstance(node.args[1], ast.Name) - ): - by_value.add(node.args[1].id) - trees = {} - for path, source in sources.items(): - if not any(name in source for name in by_value): - continue - try: - trees[path] = ast.parse(source) - except SyntaxError: - continue - for name in sorted(by_value): - defined_in = [ - path - for path, tree in trees.items() - if any( - isinstance(node, ast.FunctionDef) and node.name == name - for node in tree.body - ) - ] - if not defined_in: - raise AssertionError( - f"pass {name!r} is handed to run_post_process_pass by value " - "but defined in no scanned module; the walk cannot see it" - ) - for path in defined_in: - entries.add((path, name)) - return entries - - -@functools.lru_cache(maxsize=None) -def _functions_in(path): - tree = ast.parse(path.read_text(encoding="utf-8-sig")) - return { - node.name: node - for node in ast.walk(tree) - if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) - } - - -def _locally_shadowed_accessors(path): - """Accessor names this file imports from somewhere that is not the context. - - `get_device` is both the `device` bag accessor and the hardware probe in - `utils.common`. Matching the bare name would report the probe as a bag read, - so a name imported from elsewhere in this file is not the accessor. - """ - shadowed = set() - for node in ast.walk(ast.parse(path.read_text(encoding="utf-8-sig"))): - if isinstance(node, ast.ImportFrom) and node.module: - if node.module.endswith("runtime_context"): - continue - for alias in node.names: - name = alias.asname or alias.name - if name in _BAG_ACCESSORS: - shadowed.add(name) - return shadowed - - -def _reaches_a_bag(path, entry): - """Does `entry` in `path` reach a bag accessor, following calls in-module?""" - functions = _functions_in(path) - shadowed = _locally_shadowed_accessors(path) - seen = set() - - def walk(name): - if name in seen or name not in functions: - return None - seen.add(name) - for node in ast.walk(functions[name]): - if not isinstance(node, ast.Call): - continue - if isinstance(node.func, ast.Attribute): - # `rc.get_exec()`, `self.get_schedule()`: the same accessor - # reached through a module alias or an object. - if node.func.attr in _ATTRIBUTE_SPELLED: - return node.lineno - continue - if not isinstance(node.func, ast.Name): - continue - if node.func.id in _BAG_ACCESSORS and node.func.id not in shadowed: - return node.lineno - found = walk(node.func.id) - if found is not None: - return found - return None - - return walk(entry) - - -class TestResolutionReadsNoBag(CustomTestCase): - def test_the_accessor_set_is_derived_and_whole(self): - """A shrunken accessor set would make every other check pass quietly.""" - self.assertGreaterEqual( - len(_BAG_ACCESSORS), - 15, - f"only {len(_BAG_ACCESSORS)} accessors were derived from " - "runtime_context; the derivation broke", - ) - # Spelled out so a rename that drops one fails here. - for name in ("get_exec", "get_flags", "get_parallel", "get_resources"): - self.assertIn(name, _BAG_ACCESSORS) - - def test_the_walk_finds_something_to_walk(self): - """A collapsed import map would make the pin vacuous.""" - imported = _imported_symbols(_pipeline_sources()) - self.assertGreater( - len(imported), - 20, - f"the pipeline only imports from {len(imported)} of our modules; " - "the scan broke", - ) - - def test_the_registered_entries_are_found(self): - """The passes and providers are the half the import map cannot see.""" - entries = _registered_entries() - self.assertGreater( - len(entries), - 60, - f"only {len(entries)} passes and providers were found; the scan broke", - ) - # Every registered callable that lives in our tree has to appear: the - # derivation reads the registries through `inspect`, so an interpreter - # that imported a *different* checkout would resolve them outside - # `_SRT` and quietly leave the walk with nothing to walk. - missing = sorted( - name - for name in ( - getattr(inspect.unwrap(fn), "__name__", "") - for fn in _registry_functions() - ) - if name - and name != "" - and name not in {entry for _path, entry in entries} - ) - self.assertEqual( - missing, - [], - "a registered pass or provider did not resolve to a file under " - f"{_SRT}; the entry set is narrower than the registries:\n " - + "\n ".join(missing), - ) - - def test_nothing_the_pipeline_calls_reads_a_bag(self): - imported = _imported_symbols(_pipeline_sources()) - reachable = { - (path, symbol) for path, symbols in imported.items() for symbol in symbols - } | _registered_entries() - found = [] - for path, symbol in sorted(reachable): - line = _reaches_a_bag(path, symbol) - if line is not None: - found.append( - f"{path.relative_to(_SRT)}:{line} reached from " - f"{symbol}(), which resolution calls" - ) - self.assertEqual( - found, - [], - "resolution reaches a config-bag read, which raises on whichever " - "branch gets there first; read the resolving state instead:\n " - + "\n ".join(found), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_bench_long_context.py b/test/registered/unit/test_bench_long_context.py deleted file mode 100644 index 0af59ca67..000000000 --- a/test/registered/unit/test_bench_long_context.py +++ /dev/null @@ -1,128 +0,0 @@ -"""Unit test for benchmark/hicache/bench_long_context.py. - -Guards against the regression where ContextWorkloadGenerator.__init__ replaces -WorkloadGenerator.__init__ entirely but forgets to set attributes the inherited -request_sender/handle_request methods need (e.g. self.request_func). -""" - -import json -import sys -import tempfile -import unittest -from pathlib import Path -from types import SimpleNamespace -from unittest.mock import MagicMock, patch - -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=10, suite="base-a-test-cpu") - -REPO_ROOT = Path(__file__).resolve().parents[3] -HICACHE_DIR = REPO_ROOT / "benchmark" / "hicache" -if str(HICACHE_DIR) not in sys.path: - sys.path.insert(0, str(HICACHE_DIR)) - -import bench_long_context # noqa: E402 - -from sglang.test.kits.cache_hit_kit import async_request_sglang_generate # noqa: E402 - - -def _build_args(dataset_path: str) -> SimpleNamespace: - return SimpleNamespace( - host="localhost", - port=30000, - model_path="meta-llama/Llama-3.2-1B-Instruct", - distribution="poisson", - request_rate=1.0, - dataset_path=dataset_path, - num_clients=2, - max_parallel=2, - log_file="performance_metrics.jsonl", - tag="", - ) - - -def _fake_dataset() -> dict: - return { - "contexts": ["ctx-zero ", "ctx-one "], - "queries": [ - {"context": 0, "question": "q0", "reference_answer": "a0"}, - {"context": 1, "question": "q1", "reference_answer": "a1"}, - ], - } - - -class TestContextWorkloadGeneratorInit(CustomTestCase): - """Verify ContextWorkloadGenerator wires up everything its inherited - request_sender/handle_request/run methods rely on.""" - - def setUp(self): - self._tmp = tempfile.NamedTemporaryFile(mode="w", suffix=".json", delete=False) - json.dump(_fake_dataset(), self._tmp) - self._tmp.close() - self.dataset_path = self._tmp.name - - mock_tokenizer = MagicMock() - mock_tokenizer.encode.return_value = [1, 2, 3, 4] - mock_tokenizer.return_value = {"input_ids": [5, 6]} - - self._tok_patch = patch.object( - bench_long_context, "get_tokenizer", return_value=mock_tokenizer - ) - self._tok_patch.start() - - def tearDown(self): - self._tok_patch.stop() - Path(self.dataset_path).unlink(missing_ok=True) - - def test_request_func_is_set(self): - """The bug we're guarding against: request_func not being set caused - AttributeError as soon as the request_sender thread fired.""" - gen = bench_long_context.ContextWorkloadGenerator( - _build_args(self.dataset_path) - ) - self.assertTrue(callable(getattr(gen, "request_func", None))) - self.assertIs(gen.request_func, async_request_sglang_generate) - - def test_inherits_workload_generator_contract(self): - """All attributes WorkloadGenerator's run-time methods touch must exist.""" - gen = bench_long_context.ContextWorkloadGenerator( - _build_args(self.dataset_path) - ) - - # handle_request (bench_multiturn.py) reads these - for attr in ("request_func", "url", "pbar", "response_queue", "finished_time"): - self.assertTrue(hasattr(gen, attr), f"missing attribute: {attr}") - - # request_sender reads these - for attr in ( - "sent_requests", - "completed_requests", - "max_parallel", - "ready_queue", - "distribution", - "request_rate", - ): - self.assertTrue(hasattr(gen, attr), f"missing attribute: {attr}") - - # run() reads these - for attr in ("performance_metrics", "enable_round_barrier"): - self.assertTrue(hasattr(gen, attr), f"missing attribute: {attr}") - - def test_url_targets_sglang_generate_endpoint(self): - gen = bench_long_context.ContextWorkloadGenerator( - _build_args(self.dataset_path) - ) - self.assertEqual(gen.url, "http://localhost:30000/generate") - - def test_ready_queue_size_matches_dataset(self): - gen = bench_long_context.ContextWorkloadGenerator( - _build_args(self.dataset_path) - ) - # 2 queries in fake dataset, num_clients=2 → 2 init requests - self.assertEqual(len(gen.ready_queue.requests), 2) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_context_accessor_shadowing.py b/test/registered/unit/test_context_accessor_shadowing.py deleted file mode 100644 index af57f830e..000000000 --- a/test/registered/unit/test_context_accessor_shadowing.py +++ /dev/null @@ -1,191 +0,0 @@ -"""No local may shadow a ``runtime_context`` accessor it also calls. - -A mechanical sweep that rewrites ``self.server_args.mamba_cache_chunk_size`` -into ``mamba_cache_chunk_size()`` turns - - mamba_cache_chunk_size = self.server_args.mamba_cache_chunk_size - -into ``mamba_cache_chunk_size = mamba_cache_chunk_size()``, which is a -self-referential local: the name is local for the whole function, so the call -raises ``UnboundLocalError`` the first time that line runs. Five of these -shipped in one sweep and only one had unit coverage — a mamba model on the -radix-cache-v2 path found it at request time. - -This scans for the shape directly: a function-scope assignment whose target -name is an imported accessor. -""" - -import ast -import unittest -from pathlib import Path - -import sglang -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=24, suite="base-a-test-cpu") - -_PACKAGE_ROOT = Path(next(iter(sglang.__path__))) -_CONTEXT_MODULE = "sglang.srt.runtime_context" - - -def _module_level_accessor_imports(tree: ast.AST) -> set[str]: - """Accessors imported at module scope — visible in every function. - - A *function-local* import is visible only inside its own scope, so it is - collected per function in the scan below: charging it file-wide would flag - an unrelated sibling function that binds the same name, where no shadowing - can occur. - """ - names: set[str] = set() - stack = list(tree.body) - while stack: - stmt = stack.pop() - if isinstance( - stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.Lambda, ast.ClassDef) - ): - continue - if isinstance(stmt, ast.ImportFrom) and stmt.module == _CONTEXT_MODULE: - for alias in stmt.names: - names.add(alias.asname or alias.name) - stack.extend(ast.iter_child_nodes(stmt)) - return names - - -def _bound_names(target): - """Every name a binding target introduces, unpacking included. - - ``a, (b, c) = ...`` and ``for x, y in ...`` bind through Tuple/List/Starred - nodes, so a check that only accepts a bare ``ast.Name`` misses them. - """ - if isinstance(target, ast.Name): - yield target.id - elif isinstance(target, ast.Starred): - yield from _bound_names(target.value) - elif isinstance(target, (ast.Tuple, ast.List)): - for element in target.elts: - yield from _bound_names(element) - - -def _own_scope_statements(node) -> tuple: - """This function's OWN scope: its statements, plus the (name, lineno) of - each nested ``def``/``class`` — the definition's *name* is a binding in - this scope (an earlier accessor call raises UnboundLocalError just like an - assignment), while its *body* is the nested scope's own and descending into - it would misattribute bindings.""" - own_scope = [] - nested_def_bindings = [] - pending = list(node.body) - while pending: - stmt = pending.pop() - if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): - nested_def_bindings.append((stmt.name, stmt.lineno)) - continue - if isinstance(stmt, ast.Lambda): - continue - own_scope.append(stmt) - pending.extend(ast.iter_child_nodes(stmt)) - return own_scope, nested_def_bindings - - -def _child_functions(body) -> list: - """Function defs directly beneath this scope — descending through plain - statements and class bodies (a method closes over the enclosing function's - names, not the class's), but never into another function.""" - funcs = [] - pending = list(body) - while pending: - stmt = pending.pop() - if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): - funcs.append(stmt) - continue - if isinstance(stmt, ast.Lambda): - continue - pending.extend(ast.iter_child_nodes(stmt)) - return funcs - - -def _shadowing_assignments(tree: ast.AST, module_accessors: set[str]): - """Function-local bindings whose name shadows an accessor visible in that - scope -- every statement form that binds a local, not just ``=``. - - Python decides a name is local from *any* binding in the function, so a - loop variable, a ``with ... as``, a walrus, a comprehension target, or an - ``except ... as`` all shadow the accessor for the whole function body, - exactly like an assignment does. - - Visibility follows lexical scope: module-level imports reach every - function; a function-local import reaches its own scope and nested - functions (closure), but NOT an unrelated sibling — charging it file-wide - would flag bindings where no shadowing occurs. A function-scope *re-import* - of the accessor is itself fine: it binds the name to the same callable, so - calls after it behave identically (and the module is full of deliberate - local imports). - """ - stack = [(fn, module_accessors) for fn in _child_functions(tree.body)] - while stack: - node, inherited = stack.pop() - own_scope, nested_def_bindings = _own_scope_statements(node) - local_imports = { - alias.asname or alias.name - for stmt in own_scope - if isinstance(stmt, ast.ImportFrom) and stmt.module == _CONTEXT_MODULE - for alias in stmt.names - } - visible = inherited | local_imports - # ``def get_exec(): ...`` nested in the function binds the name in - # THIS scope, exactly like an assignment would. - for name, lineno in nested_def_bindings: - if name in visible: - yield node.name, name, lineno - for inner in own_scope: - targets = [] - if isinstance(inner, ast.Assign): - targets = inner.targets - elif isinstance(inner, (ast.AnnAssign, ast.AugAssign)): - targets = [inner.target] - elif isinstance(inner, (ast.For, ast.AsyncFor, ast.comprehension)): - targets = [inner.target] - elif isinstance(inner, ast.NamedExpr): - targets = [inner.target] - elif isinstance(inner, (ast.With, ast.AsyncWith)): - targets = [i.optional_vars for i in inner.items if i.optional_vars] - elif isinstance(inner, ast.ExceptHandler) and inner.name: - targets = [ast.Name(id=inner.name, ctx=ast.Store())] - for target in targets: - for name in _bound_names(target): - if name in visible: - yield node.name, name, getattr(inner, "lineno", node.lineno) - for nested in _child_functions(node.body): - stack.append((nested, visible)) - - -class TestNoAccessorShadowing(CustomTestCase): - def test_no_local_shadows_a_context_accessor(self): - offenders = [] - for path in sorted(_PACKAGE_ROOT.rglob("*.py")): - rel = path.relative_to(_PACKAGE_ROOT).as_posix() - if rel.startswith("srt/runtime_context.py"): - continue - source = path.read_text() - # A file that never names the module cannot import an accessor - # from it, at module scope or inside any function. - if _CONTEXT_MODULE not in source: - continue - try: - tree = ast.parse(source) - except SyntaxError: - continue - module_accessors = _module_level_accessor_imports(tree) - for func, name, lineno in _shadowing_assignments(tree, module_accessors): - offenders.append(f"{rel}:{lineno}: {func}() binds {name!r}") - self.assertFalse( - offenders, - "locals shadow a runtime_context accessor imported in the same " - "module; the name is local for the whole function, so any call to " - "the accessor there raises UnboundLocalError:\n" + "\n".join(offenders), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_dead_server_args_parameter_ratchet.py b/test/registered/unit/test_dead_server_args_parameter_ratchet.py deleted file mode 100644 index 87eb59827..000000000 --- a/test/registered/unit/test_dead_server_args_parameter_ratchet.py +++ /dev/null @@ -1,86 +0,0 @@ -"""A function does not take the record it never reads. - -A `server_args` parameter that the body never names keeps a reference to the -whole record alive across a call boundary, and it reads as an invitation: the -next person to need one value takes it off the parameter that is already there, -instead of deciding where that value should come from. Removing one usually -uncovers the next -- the caller that only had a record to pass it along. - -Class methods are exempt: a base class, an override, or one implementation of a -strategy carries the parameter for its contract, and the body of any single one -of them is not evidence. This walks module-level functions only. -""" - -import ast -import pathlib -import unittest - -import sglang -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=14, suite="base-a-test-cpu") - -_PACKAGE_ROOT = pathlib.Path(next(iter(sglang.__path__))) - -# The resolution pipeline builds the record, so a parameter there is the subject -# rather than a passenger. `multimodal_gen` has a different, same-named class -# outside this contract, as the mutation ratchet also records. -_EXCLUDED = ("srt/arg_groups", "srt/server_args.py", "multimodal_gen") - -_BASELINE = 0 - - -def _dead_parameters(): - found = [] - scanned = 0 - for path in sorted(_PACKAGE_ROOT.rglob("*.py")): - rel = path.relative_to(_PACKAGE_ROOT).as_posix() - if rel.startswith(_EXCLUDED): - continue - source = path.read_text(encoding="utf-8-sig") - if "server_args" not in source: - continue - scanned += 1 - try: - tree = ast.parse(source) - except SyntaxError: - continue - for node in tree.body: - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - taken = [a.arg for a in node.args.args] + [ - a.arg for a in node.args.kwonlyargs - ] - if "server_args" not in taken: - continue - named = any( - isinstance(inner, ast.Name) and inner.id == "server_args" - for inner in ast.walk(node) - if inner is not node - ) - if not named: - found.append(f"{rel}:{node.lineno} {node.name}") - return found, scanned - - -class TestNoDeadServerArgsParameter(CustomTestCase): - def test_no_module_level_function_takes_a_record_it_ignores(self): - found, scanned = _dead_parameters() - self.assertGreater( - scanned, - 50, - f"only {scanned} files mention server_args; the scan is broken, not " - "the tree", - ) - self.assertEqual( - _BASELINE, - len(found), - "these functions take `server_args` and never name it; drop the " - "parameter and the argument at every call site, then check whether " - f"the caller still needs its own: {found}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_legacy_global_ratchet.py b/test/registered/unit/test_legacy_global_ratchet.py deleted file mode 100644 index 5b4f2d96f..000000000 --- a/test/registered/unit/test_legacy_global_ratchet.py +++ /dev/null @@ -1,65 +0,0 @@ -"""Ratchet guard: legacy global-accessor call-sites may only decrease. - -The process-wide ``ServerArgs`` is owned by the runtime context; the legacy -``get_global_server_args`` / ``set_global_server_args_for_*`` names survive as -thin shims for the existing call-sites. New code should use the -``sglang.srt.runtime_context`` accessors (``get_server_args()`` / -``get_context().set_server_args()``), so the shim call-site counts below must -never grow. When your change removes call-sites, lower the matching baseline -to the new count. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=11, suite="base-a-test-cpu") - -import re -import unittest -from pathlib import Path - -import sglang.srt -from sglang.test.test_utils import CustomTestCase - -_SRT_ROOT = Path(next(iter(sglang.srt.__path__))) - -# Baselines counted over python/sglang/srt/**/*.py, including each function's -# own def line. Ratchet: decrease-only. -_RATCHETS = [ - # Down to the shim definition itself; every call-site now goes through - # runtime_context.get_server_args(). - ("get_global_server_args", r"\bget_global_server_args\s*\(", 1), - ( - "set_global_server_args_for_*", - r"\bset_global_server_args_for_(?:scheduler|tokenizer)\s*\(", - 2, - ), -] - - -class TestLegacyGlobalRatchet(CustomTestCase): - def test_legacy_accessor_call_sites_match_the_baselines(self): - # Exact pin, failing in BOTH directions: a grown count means new code - # bypassed the runtime_context accessors; a shrunk count means a - # removal forgot to lower the baseline, which would let later changes - # silently re-add call-sites up to the stale ceiling. - sources = [ - path.read_text(encoding="utf-8", errors="replace") - for path in sorted(_SRT_ROOT.rglob("*.py")) - ] - for name, pattern, baseline in _RATCHETS: - count = sum(len(re.findall(pattern, source)) for source in sources) - if count > baseline: - self.fail( - f"{name} call-sites grew: {count} > baseline {baseline}. " - "New code must use the sglang.srt.runtime_context accessors " - "(get_server_args() / get_context().set_server_args())." - ) - if count < baseline: - self.fail( - f"{name} call-sites shrank: {count} < baseline {baseline}. " - "Lower the baseline in this file to lock in the progress." - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_model_override_split.py b/test/registered/unit/test_model_override_split.py deleted file mode 100644 index 1e693e5b3..000000000 --- a/test/registered/unit/test_model_override_split.py +++ /dev/null @@ -1,120 +0,0 @@ -"""Two family modules must never declare the same field for the same architecture. - -An architecture claimed by two family modules is normal -- ``Qwen3NextForCausalLM`` -gets its attention shape from ``qwen3_5`` and its MoE runner from ``qwen3_moe``. -Two modules declaring the *same* field for it is not: nobody owns that value, -and which module supplies it is decided by nothing more deliberate than the -order the imports happen to be in. That is a defect in the declarations, so -this forbids it outright rather than choosing a winner. - -The ordering follows from the rule and is not itself pinned. ``__init__.py`` is -a list of imports, importing is what registers, and the gate applies matching -declarations in registration order with the last writer winning -- so an -overlap would make an import list into a behavioural statement, which tools -reorder freely. With no overlap the list can be sorted however anyone likes. - -The declared-field sets are read with the chain ratchet's own extractor rather -than a second implementation of the same scan, for the reason its docstring -gives: two censuses of one thing that disagree are worse than either alone. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=10, suite="base-a-test-cpu") - -import ast -import importlib.util -import pathlib -import sys -import unittest - -from sglang.srt.arg_groups import model_overrides -from sglang.srt.arg_groups.model_override_base import ( - _MODEL_OVERRIDE_FNS, - MODEL_OVERRIDES, -) -from sglang.test.test_utils import CustomTestCase - -_RATCHET = pathlib.Path(__file__).resolve().parent / "test_chain_read_ratchet.py" - - -def _returned_field_names(fn): - spec = importlib.util.spec_from_file_location("_chain_ratchet_for_split", _RATCHET) - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - source = pathlib.Path(sys.modules[fn.__module__].__file__).read_text( - encoding="utf-8-sig" - ) - body = next( - node - for node in ast.walk(ast.parse(source)) - if isinstance(node, ast.FunctionDef) and node.name == fn.__name__ - ) - return module._returned_field_names(body) - - -class TestModelOverrideSplit(CustomTestCase): - def test_no_field_is_declared_by_two_family_modules(self): - contested = { - arch: fns for arch, fns in _MODEL_OVERRIDE_FNS.items() if len(fns) > 1 - } - self.assertTrue(contested, "the scan found no architecture with two claimants") - for arch, fns in sorted(contested.items()): - with self.subTest(architecture=arch): - seen: dict[str, str] = {} - for fn in fns: - for field in _returned_field_names(fn): - earlier = seen.get(field) - self.assertIsNone( - earlier, - f"{arch}: {fn.__module__}.{fn.__name__} and {earlier} " - f"both declare {field!r}, so which one wins now depends " - f"on the order of the imports in " - f"arg_groups/model_overrides/__init__.py", - ) - seen[field] = f"{fn.__module__}.{fn.__name__}" - - def test_the_constant_table_does_not_contest_a_callable(self): - """``MODEL_OVERRIDES`` applies before the callables, so a field it and a - callable both name is decided by that ordering instead.""" - for arch, const in sorted(MODEL_OVERRIDES.items()): - for fn in _MODEL_OVERRIDE_FNS.get(arch, ()): - with self.subTest(architecture=arch, fn=fn.__name__): - self.assertFalse( - set(const) & _returned_field_names(fn), - f"{arch}: MODEL_OVERRIDES and {fn.__name__} both declare " - f"{sorted(set(const) & _returned_field_names(fn))}", - ) - - def test_the_import_list_names_every_family_module(self): - """Importing is what registers, so a module missing from the list is a - family that silently stops applying -- and the tests that import a - provider directly would not notice.""" - package = pathlib.Path(model_overrides.__file__).parent - on_disk = { - path.stem for path in package.glob("*.py") if path.stem != "__init__" - } - imported = { - alias.name - for node in ast.walk(ast.parse((package / "__init__.py").read_text())) - if isinstance(node, ast.ImportFrom) - and node.module == "sglang.srt.arg_groups.model_overrides" - for alias in node.names - } - self.assertEqual(on_disk, imported) - - def test_every_declaration_comes_from_its_own_family_module(self): - """The split itself: nothing was left behind in overrides.py.""" - for arch, fns in _MODEL_OVERRIDE_FNS.items(): - for fn in fns: - with self.subTest(architecture=arch, fn=fn.__name__): - self.assertTrue( - fn.__module__.startswith( - "sglang.srt.arg_groups.model_overrides." - ), - f"{fn.__name__} still lives in {fn.__module__}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_module_state_ratchet.py b/test/registered/unit/test_module_state_ratchet.py deleted file mode 100644 index 562df28fd..000000000 --- a/test/registered/unit/test_module_state_ratchet.py +++ /dev/null @@ -1,61 +0,0 @@ -"""Ratchet guard: module-level runtime state in the flag-owning layers may -only shrink. - -Runtime flags belong on ``get_flags()`` groups, which have lifecycle reset and -a scoped test-override primitive; a module-level ``global`` has neither and -leaks across test teardowns. ``_PINNED_GLOBALS`` names the survivors -- -migrating one must shrink it, adding one fails the ratchet. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=11, suite="base-a-test-cpu") - -import ast -import unittest -from pathlib import Path - -import sglang.srt -from sglang.test.test_utils import CustomTestCase - -_SRT_ROOT = Path(next(iter(sglang.srt.__path__))) - -_PINNED_GLOBALS = { - "layers/moe/utils.py": frozenset(), - "layers/dp_attention.py": frozenset( - { - # DP-attention topology (parallel vertical scope). - "_ATTN_DP_RANK", - "_ATTN_DP_SIZE", - } - ), -} - - -class TestModuleStateRatchet(CustomTestCase): - def test_global_statements_match_the_pins(self): - for rel, pinned in _PINNED_GLOBALS.items(): - tree = ast.parse((_SRT_ROOT / rel).read_text()) - declared = { - name - for node in ast.walk(tree) - if isinstance(node, ast.Global) - for name in node.names - } - grown = declared - pinned - self.assertFalse( - grown, - f"{rel} declares new module-level runtime state {sorted(grown)}; " - "put runtime flags on a get_flags() group instead " - "(see runtime_context.MoeFlags / DpFlags).", - ) - shrunk = pinned - declared - self.assertFalse( - shrunk, - f"{rel} no longer declares {sorted(shrunk)}; " - "shrink the pin in this file to lock in the progress.", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_parallel_adoption_ratchet.py b/test/registered/unit/test_parallel_adoption_ratchet.py deleted file mode 100644 index 73d601f16..000000000 --- a/test/registered/unit/test_parallel_adoption_ratchet.py +++ /dev/null @@ -1,73 +0,0 @@ -"""Ratchet guard: legacy parallel-getter calls in swept directories may only -shrink. - -``models/`` and ``layers/`` read parallel topology through -``get_parallel().`` (the read-through wrapper in ``runtime_context``), -which gives one import, one naming scheme, and the scoped ``override()`` test -primitive. Exemptions are pinned in ``_EXEMPT``, each with its reason; sweeping -one must remove it from there. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=11, suite="base-a-test-cpu") - -import re -import unittest -from pathlib import Path - -import sglang.srt -from sglang.test.test_utils import CustomTestCase - -_SRT_ROOT = Path(next(iter(sglang.srt.__path__))) - -_BANNED_CALLS = re.compile( - r"\b(?:dcp_enabled|get_(?:" - r"tensor_model_parallel_(?:world_size|rank)" - r"|pipeline_model_parallel_(?:world_size|rank)" - r"|moe_expert_parallel_(?:world_size|rank)" - r"|moe_tensor_parallel_(?:world_size|rank)" - r"|moe_data_parallel_(?:world_size|rank)" - r"|attn_tensor_model_parallel_(?:world_size|rank)" - r"|attn_context_model_parallel_(?:world_size|rank)" - r"|dcp_(?:world_size|rank)" - r"|dcp_group(?:_no_assert)?" - r"|attention_dcp_(?:world_size|rank)" - r"|attention_(?:tp|cp)_(?:group|rank|size)" - r"))\(\)" -) - -# The whole package is swept; the exemptions are the substrate itself. -_SWEPT_DIRS = ("",) - -_EXEMPT = ( - "distributed/", # parallel_state: defines the canonical getters - "runtime_context.py", # delegates DCP reads to canonical getters - "layers/dp_attention.py", # delegation substrate for the attn-DP dims - "layers/dcp/comm.py", # deprecated out-of-tree DCP compatibility shims - # The dumper's megatron plugin calls third-party getters that share the - # parallel_state names (self._mpu.get_tensor_model_parallel_rank()). - "debug_utils/dumper.py", -) - - -class TestParallelAdoptionRatchet(CustomTestCase): - def test_no_legacy_parallel_getters_in_swept_dirs(self): - offenders = [] - for top in _SWEPT_DIRS: - for path in sorted((_SRT_ROOT / top).rglob("*.py")): - rel = path.relative_to(_SRT_ROOT).as_posix() - if rel.startswith(_EXEMPT): - continue - for i, line in enumerate(path.read_text().split("\n"), 1): - if _BANNED_CALLS.search(line): - offenders.append(f"{rel}:{i}") - self.assertFalse( - offenders, - "legacy parallel-getter calls in swept directories (use " - f"get_parallel(). instead): {offenders}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_platform_address_not_frozen.py b/test/registered/unit/test_platform_address_not_frozen.py deleted file mode 100644 index faef851d2..000000000 --- a/test/registered/unit/test_platform_address_not_frozen.py +++ /dev/null @@ -1,81 +0,0 @@ -"""No module-scope name may freeze a platform fact. - -The address exists so `override_platform(...)` reaches every reader at once. -A module-level `_is_sm120 = get_platform().is_sm120` defeats that completely: -the value is read when the module is first imported and never again, so whether -an override is visible depends on import order -- and the line *looks* like it -went through the address, which is worse than the bare probe it replaced. - -Four of these were written during this refactor's own conversion (three in -`fp8_utils`, one in `deepseek_v4_backend`), by substituting the accessor into -lines that were already frozen. Substituting the call is not the conversion; the -conversion is the reader asking at the point of decision. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=12, suite="base-a-test-cpu") - -import ast -import pathlib -import unittest - -import sglang -from sglang.test.test_utils import CustomTestCase - -_ROOT = pathlib.Path(next(iter(sglang.__path__))) / "srt" - - -def _frozen_platform_reads(): - """(file, line, name) for each module-scope `x = get_platform().y`.""" - found = [] - for path in sorted(_ROOT.rglob("*.py")): - source = path.read_text(encoding="utf-8-sig") - if "get_platform" not in source: - continue - try: - tree = ast.parse(source) - except SyntaxError: - continue - # Module scope only: inside a function the call runs per invocation, - # which is the shape the address is for. - for node in tree.body: - if not isinstance(node, ast.Assign): - continue - value = node.value - if not ( - isinstance(value, ast.Attribute) - and isinstance(value.value, ast.Call) - and getattr(value.value.func, "id", None) == "get_platform" - ): - continue - for target in node.targets: - if isinstance(target, ast.Name): - rel = path.relative_to(_ROOT).as_posix() - found.append(f"{rel}:{node.lineno} {target.id}") - return found - - -class TestPlatformAddressNotFrozen(CustomTestCase): - def test_the_scan_reaches_the_address(self): - """The premise: `get_platform()` is used somewhere under srt/.""" - users = [ - path - for path in _ROOT.rglob("*.py") - if "get_platform()" in path.read_text(encoding="utf-8-sig") - ] - self.assertGreater(len(users), 20, "the scan found almost no readers") - - def test_no_module_scope_name_freezes_a_platform_fact(self): - frozen = _frozen_platform_reads() - self.assertEqual( - [], - frozen, - "these read a platform fact once at import and keep the answer, so " - "`override_platform(...)` cannot reach them and the result depends " - f"on import order. Ask at the point of decision instead: {frozen}", - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_pre_publish_readers.py b/test/registered/unit/test_pre_publish_readers.py index c87c419e8..4fb7c464d 100644 --- a/test/registered/unit/test_pre_publish_readers.py +++ b/test/registered/unit/test_pre_publish_readers.py @@ -8,25 +8,18 @@ normally free. The launcher's own reads are not: everything `multimodal_gen` calls into the same code with a `ServerArgs` of its own that never publishes them. -This is a class no other test in the tree catches: a converted reader is -exercised everywhere by tests that publish first, so it passes unit CI and then -takes the server down on the first real launch -- which is how `configure_logger` -shipped. So the protected set is *derived from the launch path* rather than -listed here: the callees named in `_launch_subprocesses` above its `publish` are -read out of the source, and each is called against a context where nothing has -been published. A conversion of any of them turns this red without a boot. +These tests call the launch helpers directly with no published context. A +conversion that asks a config bag before publication therefore fails without +requiring a full server boot. """ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=12, suite="base-a-test-cpu") -import ast import logging -import pathlib import unittest -import sglang from sglang.srt.entrypoints.engine import _set_envs_and_config from sglang.srt.runtime_context import get_observability, reset_context from sglang.srt.server_args import ServerArgs @@ -40,44 +33,6 @@ _EXERCISED = { "_set_envs_and_config": _set_envs_and_config, } -# Named by the launcher before `publish`, but none of them reads config out of a -# record. Listed so the set above stays a statement about all of the callees. -_NOT_EXERCISED = { - "load_plugins", - "resolve_auto_parsers", - "snapshot_context", - "resolving_view", -} - - -def _pre_publish_callees(): - """Functions `_launch_subprocesses` calls before it publishes. - - Read from the source so the set cannot go stale: a call added above the - `publish(...)` line joins the protected set on its own. - """ - source = ( - pathlib.Path(next(iter(sglang.__path__))) / "srt" / "entrypoints" / "engine.py" - ).read_text(encoding="utf-8-sig") - tree = ast.parse(source) - launcher = next( - node - for node in ast.walk(tree) - if isinstance(node, ast.FunctionDef) and node.name == "_launch_subprocesses" - ) - publish_line = min( - node.lineno - for node in ast.walk(launcher) - if isinstance(node, ast.Call) and getattr(node.func, "id", None) == "publish" - ) - return { - node.func.id - for node in ast.walk(launcher) - if isinstance(node, ast.Call) - and isinstance(node.func, ast.Name) - and node.lineno < publish_line - } - class TestPrePublishReaders(CustomTestCase): def setUp(self): @@ -99,11 +54,6 @@ class TestPrePublishReaders(CustomTestCase): get_observability() self.assertIn("not published", str(caught.exception)) - def test_the_protected_set_is_what_the_launcher_calls(self): - """If the launcher stops calling one of these, or starts calling - something new before publishing, this file has to be looked at.""" - self.assertEqual(set(_EXERCISED) | _NOT_EXERCISED, _pre_publish_callees()) - def test_none_of_them_asks_a_bag(self): """Each is called with nothing published. A bag read raises `config namespace ... not published` -- any other failure is the diff --git a/test/registered/unit/test_publish_precedes_bag_reads.py b/test/registered/unit/test_publish_precedes_bag_reads.py deleted file mode 100644 index 65695bc80..000000000 --- a/test/registered/unit/test_publish_precedes_bag_reads.py +++ /dev/null @@ -1,479 +0,0 @@ -"""A process entry publishes before it reads a config namespace. - -The functions checked here are found by walking the package for `publish` -calls, not by listing them: a hand-kept list can name a function that no longer -exists and still pass, which is how the Ray actor entry went unchecked. - -Every such function starts a process (or is the first thing a spawned worker -runs), so the runtime context it inherits is empty. A bag read placed above the -publish raises `config namespace ... not published` -- in a spawned worker, -which no unit test starts, so the failure only shows up as a server that never -comes up. - -A process entry reaches its bag reads through what it calls -- `Scheduler(...)`, -`configure_scheduler_process(...)`, `self.init_tokenizer_and_processor()` -- not -by naming an accessor itself, so a scan of the entry's own body sees nothing to -order and passes whatever the code does. The read line is therefore taken over -what the entry calls: a call resolved inside the module, through a parameter's -default (`detokenizer_manager_class=DetokenizerManager`), or one hop out through -that module's import table, followed the same way at every depth. Following the -import table only out of the entry's own body would stop one call short of the -expert-backup read, which is reached as `ExpertBackupManager(...)` -> -`backup_weights_from_disk` -> imported loader code -> `get_model()`. - -Reaching no read is not a pass. The walk is a static one, and every call it -cannot resolve -- a callable handed in as a parameter, an attribute off -something other than `self` -- turns into "reaches no accessor", which is also -what a defect looks like. So every publisher that reaches none is named in -`_UNREAD_ENTRIES` with the reason, and that map is asserted against the whole -set the walk finds: a publisher that stops reaching a read, or a new one that -never reached any, fails here instead of becoming an entry with nothing to -check. Restricting the comparison to `_KNOWN_ENTRIES` would exempt exactly the -newly discovered entry the derivation exists to catch. - -What this cannot pin: a publish moving across code that reads only the handed -`server_args` instance. Such code names no accessor, so there is no read for the -walk to order it against. -""" - -import ast -import pathlib -import unittest - -import sglang -from sglang.test.ci.ci_register import register_cpu_ci -from sglang.test.config_publishers import publisher_names -from sglang.test.test_utils import CustomTestCase - -register_cpu_ci(est_time=24, suite="base-a-test-cpu") - -_PACKAGE_ROOT = pathlib.Path(sglang.__file__).resolve().parent - -_ACCESSORS = frozenset( - { - "get_exec", - "get_memory", - "get_schedule", - "get_model", - "get_spec", - "get_serving", - "get_observability", - "get_disagg", - "get_lora", - "get_mm", - "get_device", - "get_parallel", - } -) - -# The process entries that must be found by the walk. A derivation that stops -# matching -- an import rewritten, a publish moved behind a helper -- would -# otherwise leave this test green over an empty set. -_KNOWN_ENTRIES = frozenset( - { - ("srt/managers/scheduler.py", "run_scheduler_process"), - ("srt/managers/detokenizer_manager.py", "run_detokenizer_process"), - ( - "srt/managers/data_parallel_controller.py", - "run_data_parallel_controller_process", - ), - ("srt/ray/scheduler_actor.py", "__init__"), - ("srt/disaggregation/encoder/http_server.py", "launch_server"), - ("srt/entrypoints/engine.py", "_launch_subprocesses"), - ( - "srt/elastic_ep/expert_backup_manager.py", - "run_expert_backup_manager_process", - ), - ("srt/weight_cache/daemon.py", "load"), - # The multi-tokenizer worker, the benchmark work functions (run - # inline or spawned per rank), and the encoder's gRPC / spawned-TP / - # spawned-DP entries. - ("srt/entrypoints/http_server.py", "init_multi_tokenizer"), - ("benchmark/one_batch.py", "latency_test"), - ("benchmark/one_batch.py", "correctness_test"), - ("srt/disaggregation/encoder/grpc_server.py", "serve_grpc_encoder"), - ("srt/disaggregation/encoder/server.py", "launch_encoder"), - ("srt/disaggregation/encoder/runtime.py", "launch_dp_worker"), - } -) - -# Every publisher the walk finds whose callees reach no bag accessor at this -# revision, and why. Asserted exactly against what the walk finds -- not -# intersected with `_KNOWN_ENTRIES`, which would drop a newly discovered entry -# and check no ordering for the one case the derivation exists to catch. -# "Reaches none" is also what the walk answers when it cannot resolve a call, so -# every one of them is named. An entry leaves this map in the commit that gives -# it a bag read. -_UNREAD_ENTRIES: dict = { - # Not process entries: these publish to set up a context for themselves. - ("kernels/aot/tests/test_fused_qk_norm_rope.py", "test_fused_qk_norm_rope"): ( - "a kernel test publishing its own context" - ), - ("multimodal_gen/test/unit/test_disagg_trace.py", "_srt_trace_server_args"): ( - "a trace fixture publishing its own context" - ), - ( - "multimodal_gen/runtime/managers/gpu_worker.py", - "init_device_and_model", - ): ( - "a worker installing a placeholder when its process has nothing " - "published; it reads its own config, not the srt bags" - ), -} - -# `publish` itself and its named wrappers live here; a call inside them is the -# definition, not a process entry. -_PUBLISH_HOMES = frozenset({"srt/runtime_context.py", "srt/server_args.py"}) - -_CONFIG_MODULES = frozenset({"sglang.srt.runtime_context", "sglang.srt.server_args"}) - - -def _module_path(dotted: str): - """The package-relative file a `sglang.` import names, if it is one.""" - if not dotted.startswith("sglang."): - return None - parts = dotted.split(".")[1:] - for candidate in ("/".join(parts) + ".py", "/".join(parts) + "/__init__.py"): - if (_PACKAGE_ROOT / candidate).exists(): - return candidate - return None - - -_CALLS = {} - - -def _calls(fn): - """(callee key, line) per call in this body. - - `self.f()` / `cls.f()` are keyed to the owning class; a bare name is a - module-level def, an imported name, or a class -- for a class the call runs - its `__init__`. `x.f()` for any other `x` yields two keys: the bare `x`, - which resolves when `x` is a class (`PortArgs.init_new()`), and a - module-qualified one that keeps `f`, which resolves when `x` is a module - the file imported. Keeping only the bare name loses `f` entirely, so a - helper called as `foo.initialize()` contributes nothing to the walk. - Anything deeper (`a.b.c()`, a callable off an attribute) is not resolved. - """ - if id(fn) in _CALLS: - return _CALLS[id(fn)] - out = [] - for node in ast.walk(fn): - if not isinstance(node, ast.Call): - continue - func = node.func - if isinstance(func, ast.Name): - out.append((("name", func.id), node.lineno)) - elif isinstance(func, ast.Attribute) and isinstance(func.value, ast.Name): - if func.value.id in ("self", "cls"): - out.append((("self", func.attr), node.lineno)) - else: - out.append((("name", func.value.id), node.lineno)) - out.append(((f"module:{func.value.id}", func.attr), node.lineno)) - _CALLS[id(fn)] = out - return out - - -_PUBLISH_NAMES = publisher_names(_PACKAGE_ROOT / "srt") - - -class _Module: - """One parsed module: what it calls the config API, and what it defines. - - The publisher and accessor names are resolved from the imports rather than - matched by name: a model's ``index_topk_share.publish()`` and a platform's - ``get_device()`` are unrelated methods that a name-only match reports as - config calls. - """ - - def __init__(self, rel: str, tree): - self.rel, self.tree = rel, tree - self.publishers, self.accessors, self.qualified = set(), set(), set() - self.imported = {} - # Names bound to another sglang module rather than to a symbol in one: - # `import sglang.srt.foo as foo` / `from sglang.srt import foo`. Without - # these, `foo.initialize()` reaches nothing. - self.modules = {} - self.functions = {} - self.classes = {} - self.owner = {} - for node in ast.walk(tree): - if isinstance(node, ast.ImportFrom) and node.module: - if node.module in _CONFIG_MODULES: - for alias in node.names: - local = alias.asname or alias.name - if alias.name in _PUBLISH_NAMES: - self.publishers.add(local) - elif alias.name in _ACCESSORS: - self.accessors.add(local) - target = _module_path(node.module) - if target is not None and target != rel: - for alias in node.names: - self.imported[alias.asname or alias.name] = (target, alias.name) - for alias in node.names: - # `from sglang.srt import runtime_context` binds the module - # itself; `runtime_context.get_serving()` is the same read. - dotted = f"{node.module}.{alias.name}" - if dotted in _CONFIG_MODULES: - self.qualified.add(alias.asname or alias.name) - bound = _module_path(dotted) - if bound is not None and bound != rel: - self.modules[alias.asname or alias.name] = bound - elif isinstance(node, ast.Import): - for alias in node.names: - if alias.name in _CONFIG_MODULES: - self.qualified.add(alias.asname or alias.name.split(".")[0]) - # Only `import a.b.c as name` binds a usable name; without - # the alias the call reads `a.b.c.f()`, which is deeper than - # this walk resolves. - bound = _module_path(alias.name) if alias.asname else None - if bound is not None and bound != rel: - self.modules[alias.asname] = bound - elif isinstance(node, ast.ClassDef): - methods = self.classes.setdefault(node.name, {}) - for stmt in node.body: - if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)): - methods[stmt.name] = stmt - self.owner[id(stmt)] = node.name - for node in ast.walk(tree): - if ( - isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) - and id(node) not in self.owner - ): - self.functions.setdefault(node.name, node) - self._direct = {} - - def is_publish(self, node) -> bool: - if not isinstance(node, ast.Call): - return False - if isinstance(node.func, ast.Name) and node.func.id in self.publishers: - return True - return ( - isinstance(node.func, ast.Attribute) - and isinstance(node.func.value, ast.Name) - and node.func.value.id in self.qualified - ) - - def is_read(self, node) -> bool: - if not isinstance(node, ast.Call): - return False - if isinstance(node.func, ast.Name) and node.func.id in self.accessors: - return True - return ( - isinstance(node.func, ast.Attribute) - and isinstance(node.func.value, ast.Name) - and node.func.value.id in self.qualified - and node.func.attr in _ACCESSORS - ) - - def resolve(self, key, cls): - """The def in this module a callee key names, if it names one. - - A module-qualified key names nothing here: its def lives in the module - the alias is bound to, which `_targets` follows. - """ - kind, name = key - if kind == "self": - return self.classes.get(cls, {}).get(name) - if kind.startswith("module:"): - return None - if name in self.functions: - return self.functions[name] - return self.classes.get(name, {}).get("__init__") - - def direct_read(self, fn) -> bool: - """Whether this body names an accessor itself.""" - if id(fn) not in self._direct: - self._direct[id(fn)] = any(self.is_read(n) for n in ast.walk(fn)) - return self._direct[id(fn)] - - -_MODULES = {} - - -def _module(rel: str): - # A module this walk cannot parse would resolve to "reaches no config", - # which is the answer that hides a defect. utf-8-sig because a file in the - # package carries a BOM; anything still unparsable fails the test. - if rel not in _MODULES: - _MODULES[rel] = _Module( - rel, ast.parse((_PACKAGE_ROOT / rel).read_text(encoding="utf-8-sig")) - ) - return _MODULES[rel] - - -def _defaulted_parameters(fn): - """`{parameter: default name}` for the parameters this def gives a plain - name as default. `run_detokenizer_process` reaches `DetokenizerManager` - only this way -- the body calls the parameter, and the class it is really - handed is written once, as that parameter's default.""" - if id(fn) in _DEFAULTS: - return _DEFAULTS[id(fn)] - arguments = fn.args - positional = arguments.posonlyargs + arguments.args - pairs = list( - zip(positional[len(positional) - len(arguments.defaults) :], arguments.defaults) - ) - pairs += zip(arguments.kwonlyargs, arguments.kw_defaults) - _DEFAULTS[id(fn)] = { - parameter.arg: default.id - for parameter, default in pairs - if isinstance(default, ast.Name) - } - return _DEFAULTS[id(fn)] - - -_DEFAULTS = {} - - -def _targets(mod, key, cls, fn=None): - """The (module, def, owning class) a callee key names, here and one hop out. - - A bare name that is one of `fn`'s parameters resolves through that - parameter's default as well, which is how a factory handed in as an - argument is followed to the class the entry actually constructs. - - A module-qualified key (`foo.initialize()`) resolves in the module `foo` is - bound to, so a helper reached that way joins the walk instead of dropping - out of it. - """ - keys = [key] - if fn is not None and key[0] == "name": - default = _defaulted_parameters(fn).get(key[1]) - if default is not None: - keys.append(("name", default)) - for key in keys: - target = mod.resolve(key, cls) - if target is not None: - yield mod, target, mod.owner.get(id(target), cls) - if key[0].startswith("module:"): - home = mod.modules.get(key[0][len("module:") :]) - if home is None: - continue - other = _module(home) - target = other.resolve(("name", key[1]), None) - if target is not None: - yield other, target, other.owner.get(id(target)) - continue - hop = mod.imported.get(key[1]) if key[0] == "name" else None - if hop is None: - continue - other = _module(hop[0]) - target = other.resolve(("name", hop[1]), None) - if target is not None: - yield other, target, other.owner.get(id(target)) - - -# The callee names from a def down to the read it reaches, for the defs a -# witness has been found for. Only "reaches a read" is carried between -# questions: a witness path stays one, while "reaches none" can be the answer a -# recursive call gets when it re-enters a def the walk is still inside, which -# is that call's answer and not the def's own. -_WITNESS = {} - - -def _witness(mod, fn, cls, seen): - """The callees from this def down to a bag read, ending in the file that - reads, or None for a def that reaches no read. - - `seen` belongs to one question. Skipping a def already on this walk keeps - the answer for the def the question was asked about -- that def opened the - skipped one, so a read below it still comes back up the path that opened - it -- and bounds the walk on a call cycle. - """ - key = (mod.rel, id(fn)) - if key in _WITNESS: - return _WITNESS[key] - if key in seen: - return None - seen.add(key) - if mod.direct_read(fn): - _WITNESS[key] = [mod.rel] - return _WITNESS[key] - for call, _ in _calls(fn): - for target in _targets(mod, call, cls, fn): - below = _witness(*target, seen) - if below is not None: - _WITNESS[key] = [call[1]] + below - return _WITNESS[key] - return None - - -def _publishing_functions(): - """(relative path, function node, module) per publisher.""" - for path in sorted(_PACKAGE_ROOT.rglob("*.py")): - rel = path.relative_to(_PACKAGE_ROOT).as_posix() - if rel in _PUBLISH_HOMES: - continue - source = path.read_text(encoding="utf-8-sig") - if not any(name in source for name in _PUBLISH_NAMES): - continue - mod = _module(rel) - if mod is None or not mod.publishers: - continue - for fn in ast.walk(mod.tree): - if not isinstance(fn, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - if any(mod.is_publish(n) for n in ast.walk(fn)): - yield rel, fn, mod - - -def _first_read(fn, mod): - """(line, what read it) of the earliest config read this entry reaches.""" - cls = mod.owner.get(id(fn)) - marks = [(n.lineno, "a bag accessor") for n in ast.walk(fn) if mod.is_read(n)] - for call, line in _calls(fn): - for target in _targets(mod, call, cls, fn): - if target[1] is fn: - continue - below = _witness(*target, set()) - if below is not None: - chain = [call[1]] + below - marks.append( - (line, " -> ".join(chain[:-1]) + f", which reads in {chain[-1]}") - ) - break - return min(marks) if marks else None - - -class TestPublishPrecedesBagReads(CustomTestCase): - def test_every_publishing_entry_publishes_first(self): - offenders = [] - found = set() - unread = set() - for rel, fn, mod in _publishing_functions(): - found.add((rel, fn.name)) - # ast.walk yields breadth-first, so the first match is not the - # earliest line; take the minimum. - publish_line = min(n.lineno for n in ast.walk(fn) if mod.is_publish(n)) - read = _first_read(fn, mod) - if read is None: - unread.add((rel, fn.name)) - elif read[0] < publish_line: - offenders.append( - f"{rel}:{fn.name} reaches a config namespace at line " - f"{read[0]} through {read[1]}, before its publish at " - f"{publish_line}" - ) - self.assertEqual( - sorted(_KNOWN_ENTRIES - found), - [], - "the walk stopped finding known process entries; the derivation " - "is broken, not the tree", - ) - self.assertEqual( - sorted(unread), - sorted(_UNREAD_ENTRIES), - "a publisher reaching no bag accessor checks nothing; either the " - "walk stopped resolving a call, or a publisher appeared or moved " - "its reads and _UNREAD_ENTRIES has to say so", - ) - self.assertEqual( - offenders, - [], - "a spawned worker starts with an empty context:\n " - + "\n ".join(offenders), - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_ray_driver_reads_the_bags.py b/test/registered/unit/test_ray_driver_reads_the_bags.py index fa6c18979..3f3fb61ce 100644 --- a/test/registered/unit/test_ray_driver_reads_the_bags.py +++ b/test/registered/unit/test_ray_driver_reads_the_bags.py @@ -22,8 +22,7 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=10, suite="base-a-test-cpu") # `sglang.srt.ray.engine` imports `ray` at module scope, and the CPU runner has -# no ray wheel. The file-scoped source scan below is the part that has to run -# everywhere; the three arithmetic cases need the import. +# no ray wheel. _HAS_RAY = importlib.util.find_spec("ray") is not None _needs_ray = unittest.skipUnless(_HAS_RAY, "ray is not installed") @@ -64,57 +63,6 @@ class TestRayDriverReadsTheBags(CustomTestCase): self.assertEqual(get_parallel().tp_size, 8) self.assertEqual(_compute_world_size(), 8) - def test_the_driver_modules_read_no_field_off_a_record(self): - """File-scoped: neither Ray driver module reads a config field off an - instance any more. - - The Ray path has no CI coverage, so this is what keeps a new - `server_args.tp_size` from appearing in it -- the placement arithmetic - runs after the publish, and the bags are the surface that carries what - resolution decided. - """ - import ast - import dataclasses - import pathlib - - import sglang - from sglang.srt.server_args import ServerArgs - - fields = {field.name for field in dataclasses.fields(ServerArgs)} - srt = pathlib.Path(sglang.__file__).resolve().parent / "srt" - offenders = [] - for rel in ("ray/engine.py", "ray/data_parallel_controller.py"): - tree = ast.parse((srt / rel).read_text(encoding="utf-8-sig")) - holders = {"server_args", "sa"} - for node in ast.walk(tree): - if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): - continue - for arg in list(node.args.args) + list(node.args.kwonlyargs): - if arg.annotation is not None and "ServerArgs" in ast.dump( - arg.annotation - ): - holders.add(arg.arg) - for node in ast.walk(tree): - if ( - isinstance(node, ast.Attribute) - and node.attr in fields - and isinstance(node.ctx, ast.Load) - and ( - (isinstance(node.value, ast.Name) and node.value.id in holders) - or ( - isinstance(node.value, ast.Attribute) - and node.value.attr == "server_args" - ) - ) - ): - offenders.append(f"{rel}:{node.lineno} reads .{node.attr}") - self.assertEqual( - offenders, - [], - "the Ray driver reads a config field off a record; the driver runs " - "after the publish, so read `get_parallel()`:\n " + "\n ".join(offenders), - ) - if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/test_runtime_context.py b/test/registered/unit/test_runtime_context.py index eaddcd9a9..37160bd5a 100644 --- a/test/registered/unit/test_runtime_context.py +++ b/test/registered/unit/test_runtime_context.py @@ -346,39 +346,6 @@ class TestAssertPublished(_IsolatedServerArgs): self.assertEqual(publish_role(), "tokenizer") - def test_no_constructor_publishes_outside_the_two_entries(self): - """Publishing from an `__init__` is an entry's job or a bug. - - It is right when the constructor *is* the entry -- an `Engine` being - (re)built, the Ray actor that stands in for `run_scheduler_process`, - where resetting the bags is the point. It is wrong anywhere else, - because the process is already live with a record and re-projecting - drops its overrides. The census is pinned, so a new constructor publish - fails here until it is one of the two. - - Both the publisher set and "which `__init__` reaches one" come from - `sglang.test.config_publishers`, which derives them from the code -- - a hand-written spelling list here missed a constructor that publishes - one hop away through a helper. The derivation follows helpers defined - in the same module; a constructor that publishes through a helper in - *another* module is not seen, which is the one hole left here. - """ - import pathlib - - import sglang - from sglang.test.config_publishers import constructor_publishers - - srt = pathlib.Path(sglang.__file__).resolve().parent / "srt" - self.assertEqual( - constructor_publishers(srt), - { - ("entrypoints/engine.py", "Engine", "publish"), - ("ray/scheduler_actor.py", "SchedulerActor", "publish"), - }, - "a constructor publishes and it is not one of the two entries; " - "publish at the process entry and let the constructor assert", - ) - class TestServerArgsScopedOverride(_IsolatedServerArgs): """ctx.override_server_args: the config tier's scoped test override — @@ -482,11 +449,6 @@ class TestServerArgsScopedOverride(_IsolatedServerArgs): with self.assertRaises(AssertionError): override.install() - def test_module_global_removed(self): - # The legacy storage must not survive: a stale _global_server_args would - # silently fork the config into two objects. - self.assertFalse(hasattr(server_args_module, "_global_server_args")) - @dataclasses.dataclass class _FakeCaptureGroup(_FlagGroupBase): @@ -1364,64 +1326,6 @@ class TestAdaptiveDraftBoundLifecycle(_IsolatedServerArgs): self.assertEqual(max_speculative_num_draft_tokens(), 7) -class TestNamedAccessorsCallWhatTheyWrap(CustomTestCase): - """A named accessor must *call* a member that is a method. - - `return get_server_args().x` hands back a bound method when `x` is defined - with `def`; the failure then lands far away, in whatever arithmetic the - caller does with it. Checked statically so accessors that need a real model - config are covered too. - """ - - def test_accessors_that_wrap_methods_call_them(self): - import ast - import functools - import inspect - - import sglang.srt.runtime_context as rc - from sglang.srt.server_args import ServerArgs - - tree = ast.parse(inspect.getsource(rc)) - wrong = [] - for node in tree.body: - if not isinstance(node, ast.FunctionDef): - continue - for inner in ast.walk(node): - if not (isinstance(inner, ast.Return) and inner.value is not None): - continue - value = inner.value - called = isinstance(value, ast.Call) - target = value.func if called else value - if not ( - isinstance(target, ast.Attribute) - and isinstance(target.value, ast.Call) - and isinstance(target.value.func, ast.Name) - and target.value.func.id == "get_server_args" - ): - continue - member = getattr(ServerArgs, target.attr, None) - # A `property` / `functools.cached_property` member is already - # evaluated by the attribute access, so it is named here to keep - # the failure message from calling it "not a method" -- the fix - # for those is the opposite one. - kind = ( - "a property" - if isinstance(member, (property, functools.cached_property)) - else "not a method" - ) - if inspect.isfunction(member) and not called: - wrong.append( - f"{node.name}(): returns ServerArgs.{target.attr} without " - "calling it, so callers get a bound method" - ) - if not inspect.isfunction(member) and called: - wrong.append( - f"{node.name}(): calls ServerArgs.{target.attr}, which is " - f"{kind} -- the attribute access already produced the value" - ) - self.assertEqual([], wrong, "\n".join(wrong)) - - class TestParallelLeafReads(_IsolatedServerArgs): """The contract ``ParallelContext.__getattr__`` answers a parallel leaf on.""" @@ -1640,21 +1544,6 @@ class TestDerivedWidths(_IsolatedOverrides): publish(ServerArgs(model_path="dummy", tp_size=1), role="test") self.assertEqual(get_parallel().attn_tp_size, 1) - def test_the_arithmetic_has_one_home(self): - """`parallel_state` builds its groups from the same dict it stamps, and - `dp_attention` derives the pair it needs for the ranks, so a second copy - of a quotient would let two answers to one width drift apart.""" - for rel, spelling in ( - ("distributed/parallel_state.py", "derive_parallel_widths("), - ("layers/dp_attention.py", "derive_attention_widths("), - ): - source = (_SRT / rel).read_text(encoding="utf-8-sig") - self.assertNotIn("// attn_dp_size // attn_cp_size", source, rel) - self.assertNotIn("// attn_cp_size // attn_dp_size", source, rel) - self.assertNotIn("// moe_ep_size // moe_dp_size", source, rel) - self.assertNotIn("if enable_dp_attention else 1", source, rel) - self.assertIn(spelling, source, rel) - def test_the_rank_helper_agrees_with_the_stamp(self): """`compute_dp_attention_world_info` keeps the ranks and takes the widths from the same derivation the stamp uses.""" diff --git a/test/registered/unit/test_runtime_context_config_bags.py b/test/registered/unit/test_runtime_context_config_bags.py index 890d3c034..ac0d39fd9 100644 --- a/test/registered/unit/test_runtime_context_config_bags.py +++ b/test/registered/unit/test_runtime_context_config_bags.py @@ -32,33 +32,6 @@ class _CollisionFake: x: A[int, NS("exec.moe.topk")] = 1 -_TOP = ( - rc.get_device, - rc.get_model, - rc.get_exec, - rc.get_schedule, - rc.get_memory, - rc.get_spec, - rc.get_lora, - rc.get_mm, - rc.get_disagg, - rc.get_serving, - rc.get_observability, -) -_EXEC_SUBS = ( - "kernel", - "moe", - "graph", - "comm", - "mamba", - "overlap", - "offload", - "dllm", - "deterministic", - "features", -) - - class TestConfigBags(CustomTestCase): def _callTestMethod(self, method): # No retry: CustomTestCase retries once in CI, but `addCleanup` runs @@ -217,14 +190,6 @@ class TestConfigBags(CustomTestCase): restore_process_state() return sa, resolve() - def test_all_accessors_and_exec_subgroups(self): - self._publish() - for acc in _TOP: - self.assertIsNotNone(acc()) - exec_cfg = rc.get_exec() - for sub in _EXEC_SUBS: - self.assertTrue(hasattr(exec_cfg, sub), f"exec.{sub} missing") - def test_read_only_by_bare_assignment(self): self._publish() with self.assertRaises(AttributeError): diff --git a/test/registered/unit/test_server_args_cli_metadata.py b/test/registered/unit/test_server_args_cli_metadata.py index 7d5abd984..8e34487d9 100644 --- a/test/registered/unit/test_server_args_cli_metadata.py +++ b/test/registered/unit/test_server_args_cli_metadata.py @@ -1,9 +1,6 @@ """Unit tests for migrated ServerArgs CLI metadata.""" import argparse -import ast -import inspect -import textwrap import unittest from sglang.srt.server_args import ServerArgs @@ -14,53 +11,6 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=11, suite="base-a-test-cpu") -MIGRATED_OPTIONS = frozenset( - { - "--dtype", - "--quantization", - "--quantization-param-path", - "--kv-cache-dtype", - "--enable-fp32-lm-head", - "--modelopt-quant", - "--modelopt-checkpoint-restore-path", - "--modelopt-checkpoint-save-path", - "--modelopt-export-path", - "--quantize-and-serve", - "--rl-quant-profile", - "--mem-fraction-static", - "--max-running-requests", - "--max-queued-requests", - "--max-total-tokens", - "--chunked-prefill-size", - "--prefill-max-requests", - "--enable-dynamic-chunking", - "--max-prefill-tokens", - "--schedule-policy", - "--enable-priority-scheduling", - "--disable-priority-preemption", - "--default-priority-value", - "--abort-on-priority-when-disabled", - "--schedule-low-priority-values-first", - "--priority-scheduling-preemption-threshold", - "--schedule-conservativeness", - "--page-size", - "--swa-full-tokens-ratio", - "--disable-hybrid-swa-memory", - "--radix-eviction-policy", - "--enable-prefill-delayer", - "--prefill-delayer-max-delay-passes", - "--prefill-delayer-token-usage-low-watermark", - "--prefill-delayer-forward-passes-buckets", - "--prefill-delayer-wait-seconds-buckets", - "--prefill-delayer-queue-min-ratio", - "--prefill-delayer-max-delay-ms", - "--data-parallel-size", - "--dp-size", - "--load-balance-method", - } -) - - class TestServerArgsMigratedCliMetadata(CustomTestCase): @classmethod def setUpClass(cls): @@ -72,22 +22,6 @@ class TestServerArgsMigratedCliMetadata(CustomTestCase): for option in action.option_strings } - def test_migrated_options_are_registered_by_dataclass_metadata(self): - add_cli_args_source = textwrap.dedent( - inspect.getsource(ServerArgs.add_cli_args) - ) - add_cli_args_tree = ast.parse(add_cli_args_source) - manual_options = { - node.value - for node in ast.walk(add_cli_args_tree) - if isinstance(node, ast.Constant) - and isinstance(node.value, str) - and node.value.startswith("--") - } - - self.assertFalse(MIGRATED_OPTIONS & manual_options) - self.assertIn("--prefill-round-robin-balance", manual_options) - def test_argparse_shape_is_preserved_for_representative_migrated_options(self): self.assertEqual(self.actions_by_option["--dtype"].default, ServerArgs.dtype) self.assertEqual( diff --git a/test/registered/unit/test_server_args_mutation_ratchet.py b/test/registered/unit/test_server_args_mutation_ratchet.py deleted file mode 100644 index 06828d9e2..000000000 --- a/test/registered/unit/test_server_args_mutation_ratchet.py +++ /dev/null @@ -1,81 +0,0 @@ -"""Ratchet guard: server_args mutations outside the resolution pipeline may -only decrease. - -After ``ServerArgs.__post_init__`` returns, the instance carries the resolved -configuration and the resolution pipeline (``server_args.py`` and -``arg_groups/``) is the only place that computes it: resolved config changes go -to the context bags via ``get_context().override(source, **fields)``, and a -value one runner or worker owns travels as a constructor argument. The baseline -is therefore an exact pin at zero -- new mutations must not appear, and removals -must lower it. - -``ServerArgs.__setattr__`` already raises on a bare assignment after -resolution; this textual scan is what reaches the sites tests never execute. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=12, suite="base-a-test-cpu") - -import re -import unittest -from pathlib import Path - -import sglang -from sglang.test.test_utils import CustomTestCase - -_SGLANG_ROOT = Path(next(iter(sglang.__path__))) - -# Assignments to a server_args attribute (``server_args.x = ...``, -# ``self.server_args.x = ...``, and the ``sa`` alias used by a few helpers). -# ``==`` comparisons are excluded by the negative lookahead. -_MUTATION_PATTERNS = [ - # (?![=}]) skips ``==`` comparisons and f-string ``{x=}`` debug specs. - re.compile(r"\bserver_args\.[a-z0-9_]+\s*=(?![=}])"), - re.compile(r"\bsa\.[a-z0-9_]+\s*=(?![=}])"), - re.compile(r"get_(?:global_)?server_args\(\)\.[a-z0-9_]+\s*=(?![=}])"), - # setattr is the same write with the attribute name behind a variable. - re.compile( - r"setattr\(\s*(?:[\w.]+\.)?(?:server_args|sa|get_(?:global_)?server_args\(\))\s*," - ), -] - -# The resolution pipeline itself (mutation is its job) and multimodal_gen, -# whose ServerArgs is a different class outside this contract. -_EXCLUDED = ( - "srt/server_args.py", - "srt/arg_groups", - "multimodal_gen", -) - -_BASELINE = 0 - - -class TestServerArgsMutationRatchet(CustomTestCase): - def test_out_of_pipeline_mutations_match_the_baseline(self): - count = 0 - for path in sorted(_SGLANG_ROOT.rglob("*.py")): - rel = path.relative_to(_SGLANG_ROOT).as_posix() - if rel.startswith(_EXCLUDED): - continue - source = path.read_text() - count += sum(len(p.findall(source)) for p in _MUTATION_PATTERNS) - if count > _BASELINE: - self.fail( - f"server_args mutations outside the resolution pipeline grew: " - f"{count} > baseline {_BASELINE}. Configuration is resolved in " - "ServerArgs.__post_init__; declare through the pipeline " - "(passes / declare_late_resolution), change resolved config " - "with get_context().override(source, ...), or hand the value " - "to its runner as a constructor argument — do not assign fields." - ) - if count < _BASELINE: - self.fail( - f"server_args mutations outside the resolution pipeline " - f"shrank: {count} < baseline {_BASELINE}. Lower the baseline " - "in this file to lock in the progress." - ) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_server_args_no_instance_mutation_entry.py b/test/registered/unit/test_server_args_no_instance_mutation_entry.py deleted file mode 100644 index 15153c77a..000000000 --- a/test/registered/unit/test_server_args_no_instance_mutation_entry.py +++ /dev/null @@ -1,119 +0,0 @@ -"""``ServerArgs`` has no in-place mutation entry, and nothing calls one. - -``ServerArgs.override(source, **fields)`` used to mutate a resolved instance, -and ``ServerArgs.derive(source, **fields)`` used to copy-and-edit one; after -resolution the fields are the record the config bags were projected from, so a -write desyncs every namespace reader, and a copy invites publishing stale -variants. Both are gone: post-publish changes go to the bags -(``get_context().override``), a value one runner or worker owns travels as a -constructor argument, and late launcher-stage resolution declares through -``arg_groups.overrides.declare_late_resolution``, which writes no field and -refuses the published instance. - -The textual half of this guard matters because the resolution pipeline's own file -is exempt from the mutation ratchet: a ``self.override(...)`` there — exactly -what the LoRA normalization used — is invisible to it. -""" - -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=15, suite="base-a-test-cpu") - -import re -import unittest -from pathlib import Path - -import sglang -from sglang.srt.server_args import ServerArgs -from sglang.test.test_utils import CustomTestCase - -_SGLANG_ROOT = Path(next(iter(sglang.__path__))) - -# ``x.override(`` on anything that is a ServerArgs by name, including the -# pipeline's own ``self.override(`` inside server_args.py. -_PATTERNS = [ - re.compile(r"\bself\.override\("), - re.compile(r"\bserver_args\.override\("), - re.compile(r"\bsa\.override\("), - re.compile(r"\bargs\.override\("), -] - -_EXCLUDED = ("multimodal_gen",) - - -class TestNoServerArgsMutationEntry(CustomTestCase): - def test_the_methods_are_gone(self): - self.assertFalse( - hasattr(ServerArgs, "override"), - "ServerArgs.override is back; post-publish changes belong on the bags " - "(get_context().override); a value one runner owns travels as a " - "constructor argument.", - ) - self.assertFalse( - hasattr(ServerArgs, "derive"), - "ServerArgs.derive is back; a value one runner or worker owns travels " - "as a constructor argument (draft_attention_backend, MMEncoder " - "gpu_id), and test doubles copy via " - "sglang.test.test_utils.server_args_variant.", - ) - - def test_nothing_calls_an_instance_override(self): - offenders = [] - for path in sorted(_SGLANG_ROOT.rglob("*.py")): - rel = path.relative_to(_SGLANG_ROOT).as_posix() - if rel.startswith(_EXCLUDED): - continue - source = path.read_text() - for pattern in _PATTERNS: - for match in pattern.finditer(source): - line = source.count("\n", 0, match.start()) + 1 - offenders.append(f"{rel}:{line}: {match.group(0)}") - if offenders: - self.fail( - "in-place ServerArgs mutation call-sites:\n" - + "\n".join(offenders) - + "\n\nUse get_context().override(source, ...) for resolved config " - "or declare_late_resolution(...) for pre-publish launcher " - "resolution." - ) - - def test_nothing_derives(self): - """A config is never copied-and-edited in the package: a value one - runner consumes travels as a constructor argument, and test doubles - copy via ``server_args_variant`` (test_utils).""" - derive_pattern = re.compile(r"\.derive\(") - offenders = [] - for path in sorted(_SGLANG_ROOT.rglob("*.py")): - rel = path.relative_to(_SGLANG_ROOT).as_posix() - if rel.startswith(_EXCLUDED): - continue - source = path.read_text() - for match in derive_pattern.finditer(source): - line = source.count("\n", 0, match.start()) + 1 - offenders.append(f"{rel}:{line}") - self.assertFalse( - offenders, - ".derive( call-sites in the package (the method no longer exists):\n" - + "\n".join(offenders), - ) - - def test_late_resolution_refuses_the_published_config(self): - from sglang.srt.arg_groups.overrides import ( - declare_late_resolution, - resolution_result, - ) - from sglang.srt.runtime_context import get_context - - override = get_context().override_server_args(tp_size=2) - published = override.install() - self.addCleanup(override.restore) - - with self.assertRaises(ValueError): - declare_late_resolution(published, "test", tp_size=4) - # The refusal left the resolution alone: the hook's declaration stands, - # and the record still carries the operator's input. - self.assertEqual(resolution_result(published, "tp_size"), 2) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/registered/unit/test_split_attention_backend_decisions.py b/test/registered/unit/test_split_attention_backend_decisions.py index 446432928..164910c2c 100644 --- a/test/registered/unit/test_split_attention_backend_decisions.py +++ b/test/registered/unit/test_split_attention_backend_decisions.py @@ -17,14 +17,10 @@ cost, before the sweep these cases guard: - a deterministic-inference knob left unset (prefill truncation align). `attention_backends()` is the shared answer: the pair with the base-field -fallback applied. The callable decisions are checked by calling them; the rest -are pinned statically, since reproducing them means building a model or a -scheduler. +fallback applied. The tests below exercise the callable decisions directly. """ -import ast import unittest -from pathlib import Path from sglang.srt.runtime_context import attention_backends, get_context from sglang.srt.server_args import ServerArgs @@ -33,26 +29,8 @@ from sglang.test.test_utils import CustomTestCase register_cpu_ci(est_time=12, suite="base-a-test-cpu") -import sglang from sglang.srt.arg_groups.overrides import attention_backends_of, resolved_view -_PACKAGE_ROOT = Path(next(iter(sglang.__path__))) / "srt" - -# The decisions this file is about, and which half of the pair each one needs. -# A base-only read here is the regression; the resolution pipeline and the two -# modules that own the config are exempt because "did the operator pin the base -# field?" is a real question *there*. -_PAIR_READERS = { - "models/inkling_common/attn.py": "the half serving the forward (mirrors hybrid dispatch)", - "model_executor/model_runner_components/misc_utils.py": "prefill (chunked prefix cache)", - "layers/rotary_embedding/mrope.py": "both (triton availability)", - "mem_cache/allocation.py": "prefill (req-to-token writer); both (get_last_loc)", - "batch_overlap/two_batch_overlap.py": "prefill (extend positions)", - "managers/scheduler.py": "prefill (truncation align knobs)", - "entrypoints/engine.py": "either half (flashinfer version floor)", - "models/sarvam_moe.py": "the half serving the forward (attn dispatch)", -} - class TestSplitBackendsReachTheDecisions(CustomTestCase): def setUp(self): @@ -191,25 +169,6 @@ class TestSplitBackendsReachTheDecisions(CustomTestCase): # reads as "supported". self.assertTrue(support_triton(None)) - def test_no_listed_decision_reads_the_base_field_alone(self): - offenders = [] - for rel, why in _PAIR_READERS.items(): - tree = ast.parse((_PACKAGE_ROOT / rel).read_text()) - for node in ast.walk(tree): - # Any attribute read named `attention_backend` is the base - # field, whatever the base expression is spelled as -- a bag - # chain, a record, or a local alias of either - # (`k = get_exec().kernel; k.attention_backend`). The pair - # helpers are calls, not attributes, so they never match. - if isinstance(node, ast.Attribute) and node.attr == "attention_backend": - offenders.append(f"{rel}:{node.lineno}: base-only read ({why})") - self.assertEqual( - [], - offenders, - "these decisions must read attention_backends() (the pair with the " - "base-field fallback), not the base field:\n" + "\n".join(offenders), - ) - class TestDraftFactoryStamping(CustomTestCase): """The factory's products carry the stamp `serving_attention_backend` @@ -316,48 +275,6 @@ class TestDraftFactoryStamping(CustomTestCase): product.attn_backends[0].decode_attention_backend_str, "triton" ) - def test_the_real_map_never_stamps_an_alias(self): - # The factory's real constructors each answer their effective name; - # this pins that no map key with an aliased or host-dependent - # constructor ("nsa", "cutedsl_mla", "hybrid_linear_attn") can leak - # its request name into a stamp: whatever the leaf built, the name it - # answered is a concrete kernel, never one of the alias keys. - import ast as _ast - import inspect - - from sglang.srt.speculative import draft_utils - - tree = _ast.parse(inspect.getsource(draft_utils)) - offenders = [] - for node in _ast.walk(tree): - if not isinstance(node, _ast.FunctionDef): - continue - if not node.name.startswith("_create_") or "_backend" not in node.name: - continue - for ret in _ast.walk(node): - if not isinstance(ret, _ast.Return) or ret.value is None: - continue - # Leaf returns are ("name", ctor(...)); delegations return the - # inner call. A bare backend return would silently miss the - # stamp contract. - if isinstance(ret.value, _ast.Tuple): - name = ret.value.elts[0] - if isinstance(name, _ast.Constant) and name.value in ( - "nsa", - "hybrid_linear_attn", - ): - offenders.append(f"{node.name}: stamps alias {name.value!r}") - elif isinstance(ret.value, _ast.Call): - fn = ret.value.func - is_delegation = isinstance( - fn, _ast.Attribute - ) and fn.attr.startswith("_create_") - if not is_delegation: - offenders.append( - f"{node.name}: returns a bare backend (no effective name)" - ) - self.assertEqual([], offenders, "\n".join(offenders)) - if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/utils/test_torch_npu_patch_utils.py b/test/registered/unit/utils/test_torch_npu_patch_utils.py deleted file mode 100644 index 1f177a7ff..000000000 --- a/test/registered/unit/utils/test_torch_npu_patch_utils.py +++ /dev/null @@ -1,41 +0,0 @@ -import types -import unittest - -from sglang.srt.utils.torch_npu_patch_utils import apply_torch_npu_patches -from sglang.test.ci.ci_register import register_cpu_ci - -register_cpu_ci(est_time=6, suite="base-a-test-cpu") - - -class TestTorchNpuPatchUtils(unittest.TestCase): - def test_apply_torch_npu_patches_uses_targeted_api_when_available(self): - calls = [] - torch_npu = types.SimpleNamespace( - _apply_patches=lambda patches: calls.append(("_apply_patches", patches)), - _apply_all_patches=lambda: calls.append(("_apply_all_patches", None)), - ) - patches = [["profiler.profile", object()]] - - apply_torch_npu_patches(torch_npu, patches) - - self.assertEqual(calls, [("_apply_patches", patches)]) - - def test_apply_torch_npu_patches_uses_all_patches_api_when_targeted_api_missing( - self, - ): - calls = [] - torch_npu = types.SimpleNamespace( - _apply_all_patches=lambda: calls.append("_apply_all_patches") - ) - - apply_torch_npu_patches(torch_npu, [["profiler.profile", object()]]) - - self.assertEqual(calls, ["_apply_all_patches"]) - - def test_apply_torch_npu_patches_requires_supported_api(self): - with self.assertRaises(AttributeError): - apply_torch_npu_patches(types.SimpleNamespace(), []) - - -if __name__ == "__main__": - unittest.main() diff --git a/test/run_suite.py b/test/run_suite.py index 665de6773..3edbff392 100644 --- a/test/run_suite.py +++ b/test/run_suite.py @@ -77,6 +77,15 @@ PER_COMMIT_SUITES = { "base-b-kernel-unit-test-4-gpu-b200", "base-b-kernel-unit-test-8-gpu-h200", "base-b-kernel-benchmark-test-1-gpu-large", + # Diffusion keeps case-level pytest partitioning behind registered + # bridge files while sharing this discovery and dispatch entry point. + "base-b-test-diffusion-1-gpu-h100", + "base-b-test-diffusion-1-gpu-5090", + "base-b-test-diffusion-1-gpu-b200", + "base-b-test-diffusion-2-gpu-h100", + "base-b-test-diffusion-bcg-1-gpu-h100", + "base-b-test-diffusion-component-2-gpu-h100", + "base-b-test-diffusion-unit-1-gpu-h100", "base-c-test-4-gpu-h100", "base-c-test-4-gpu-b200", "base-c-test-4-gpu-gb300",