move dead sglang.test files to test/manual (#25316)
This commit is contained in:
@@ -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
|
||||
"""
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user