move dead sglang.test files to test/manual (#25316)

This commit is contained in:
Liangsheng Yin
2026-05-14 20:02:44 -07:00
committed by GitHub
parent 8d5b347edd
commit d89b678d69
29 changed files with 0 additions and 1 deletions
+57
View File
@@ -0,0 +1,57 @@
import torch
import torch.nn as nn
class DummyModel(nn.Module):
def __init__(self, d_in=2048, n_heads=128, softmax_scale=0.5):
super().__init__()
self.weights_proj = nn.Linear(d_in, 1024)
self.n_heads = n_heads
self.softmax_scale = softmax_scale
def _get_logits_head_gate_orig(self, x: torch.Tensor, q_scale: torch.Tensor):
weights = self.weights_proj(x)
weights = weights * self.n_heads**-0.5
q_scale = q_scale.unsqueeze(1) # (B,1,1)
weights = weights.unsqueeze(-1) * q_scale * self.softmax_scale
return weights
def _get_logits_head_gate_opt(self, x: torch.Tensor, q_scale: torch.Tensor):
weights = self.weights_proj(x)
q_scale = q_scale.unsqueeze(1) # (B,1,1)
scale_const = self.n_heads**-0.5 * q_scale * self.softmax_scale # (B,1,1)
weights = weights.unsqueeze(-1) * scale_const # (B,1024,1)
return weights
def main():
torch.manual_seed(0)
model = DummyModel(d_in=2048, n_heads=128, softmax_scale=0.5)
x = torch.randn(128, 2048) # batch=128, d_in=2048
q_scale = torch.randn(128, 1)
import time
start = time.time()
for _ in range(1000):
out_orig = model._get_logits_head_gate_orig(x, q_scale)
print("Original version time:", time.time() - start)
start = time.time()
for _ in range(1000):
out_opt = model._get_logits_head_gate_opt(x, q_scale)
print("Optimized version time:", time.time() - start)
print("Difference:", (out_orig - out_opt).abs().max().item())
assert torch.allclose(out_orig, out_opt), "Mismatch between original and optimized"
if __name__ == "__main__":
main()
"""
Original version time: 0.49235057830810547
Optimized version time: 0.4087331295013428
Difference: 1.4901161193847656e-08
"""
+53
View File
@@ -0,0 +1,53 @@
"""
Simple wrapper to run a test file with retry logic.
Usage:
python3 -m sglang.test.ci.run_with_retry test_file.py [--max-attempts 2] [--retry-wait 60]
"""
import argparse
import sys
from sglang.test.ci.ci_utils import TestFile, run_unittest_files
def main():
parser = argparse.ArgumentParser(description="Run a test file with retry logic")
parser.add_argument("test_file", help="The test file to run")
parser.add_argument(
"--max-attempts",
type=int,
default=2,
help="Maximum number of attempts (default: 2)",
)
parser.add_argument(
"--retry-wait",
type=int,
default=60,
help="Seconds to wait between retries (default: 60)",
)
parser.add_argument(
"--timeout",
type=int,
default=1200,
help="Timeout per attempt in seconds (default: 1200)",
)
args = parser.parse_args()
# Create a TestFile with a reasonable estimated time
test_file = TestFile(name=args.test_file, estimated_time=args.timeout)
exit_code = run_unittest_files(
files=[test_file],
timeout_per_file=args.timeout,
continue_on_error=False,
enable_retry=True,
max_attempts=args.max_attempts,
retry_wait_seconds=args.retry_wait,
)
sys.exit(exit_code)
if __name__ == "__main__":
main()
+132
View File
@@ -0,0 +1,132 @@
"""Unit tests for dump_metric() function."""
import json
import os
import tempfile
import unittest
from pathlib import Path
from sglang.test.test_utils import dump_metric
class TestDumpMetric(unittest.TestCase):
"""Test suite for dump_metric() function."""
_ENV_KEYS_TO_CLEAN = ["SGLANG_TEST_METRICS_OUTPUT", "PYTEST_CURRENT_TEST"]
def setUp(self):
"""Clean up env vars before each test."""
for key in self._ENV_KEYS_TO_CLEAN:
os.environ.pop(key, None)
def tearDown(self):
"""Clean up env vars after each test."""
for key in self._ENV_KEYS_TO_CLEAN:
os.environ.pop(key, None)
def test_writes_valid_jsonl(self):
"""Test that dump_metric writes one valid JSON line when env is set."""
with tempfile.TemporaryDirectory() as tmpdir:
base_path = os.path.join(tmpdir, "metrics")
os.environ["SGLANG_TEST_METRICS_OUTPUT"] = base_path
dump_metric("test_accuracy", 0.95, labels={"model": "llama"})
# Check file exists with PID suffix
pid = os.getpid()
jsonl_path = f"{base_path}.{pid}.jsonl"
self.assertTrue(os.path.exists(jsonl_path))
# Read and validate
with open(jsonl_path, encoding="utf-8") as f:
lines = f.readlines()
self.assertEqual(len(lines), 1)
record = json.loads(lines[0])
# Validate required fields
self.assertIn("filename", record)
self.assertIn("test_case", record)
self.assertEqual(record["metric_name"], "test_accuracy")
self.assertEqual(record["value"], 0.95)
# Validate optional fields
self.assertIn("ts", record)
self.assertIsInstance(record["ts"], (int, float))
self.assertEqual(record["labels"], {"model": "llama"})
def test_no_env_no_file(self):
"""Test that dump_metric doesn't create file when env var not set."""
with tempfile.TemporaryDirectory() as tmpdir:
# Don't set env var
dump_metric("test_metric", 42)
# Verify no files created
files = list(Path(tmpdir).glob("*.jsonl"))
self.assertEqual(len(files), 0)
def test_labels_not_serializable_stringified(self):
"""Test that non-serializable labels are stringified."""
with tempfile.TemporaryDirectory() as tmpdir:
base_path = os.path.join(tmpdir, "metrics")
os.environ["SGLANG_TEST_METRICS_OUTPUT"] = base_path
# Non-serializable label
class NonSerializable:
pass
dump_metric("test_metric", 100, labels={"obj": NonSerializable()})
pid = os.getpid()
jsonl_path = f"{base_path}.{pid}.jsonl"
with open(jsonl_path, encoding="utf-8") as f:
lines = f.readlines()
record = json.loads(lines[0])
self.assertIn("labels", record)
self.assertIsInstance(record["labels"], str)
def test_bool_to_int(self):
"""Test that bool values are converted to int."""
with tempfile.TemporaryDirectory() as tmpdir:
base_path = os.path.join(tmpdir, "metrics")
os.environ["SGLANG_TEST_METRICS_OUTPUT"] = base_path
dump_metric("bool_true", True)
dump_metric("bool_false", False)
pid = os.getpid()
jsonl_path = f"{base_path}.{pid}.jsonl"
with open(jsonl_path, encoding="utf-8") as f:
lines = f.readlines()
self.assertEqual(len(lines), 2)
record1 = json.loads(lines[0])
record2 = json.loads(lines[1])
self.assertEqual(record1["value"], 1) # True -> 1
self.assertEqual(record2["value"], 0) # False -> 0
def test_pytest_current_test_parsing(self):
"""Test PYTEST_CURRENT_TEST parsing for test_case."""
with tempfile.TemporaryDirectory() as tmpdir:
base_path = os.path.join(tmpdir, "metrics")
os.environ["SGLANG_TEST_METRICS_OUTPUT"] = base_path
os.environ["PYTEST_CURRENT_TEST"] = (
"test/srt/test_example.py::TestClass::test_method (call)"
)
dump_metric("pytest_metric", 123)
pid = os.getpid()
jsonl_path = f"{base_path}.{pid}.jsonl"
with open(jsonl_path, encoding="utf-8") as f:
lines = f.readlines()
record = json.loads(lines[0])
# Only assert test_case parsing, not filename
self.assertEqual(record["test_case"], "TestClass.test_method")
if __name__ == "__main__":
unittest.main()