[misc] Remove unit test cases that fail the admission criteria (round 2) (#30703)
This commit is contained in:
@@ -35,6 +35,35 @@ Not admissible:
|
||||
- Mirror tests that restate the implementation logic as assertions.
|
||||
- Probabilistic stress that cannot reproduce the failure it claims to guard.
|
||||
|
||||
**Distinguishing test — does deletion leave a silent-failure path?** A case
|
||||
that *looks* like a tautology/mirror is still admissible when it guards a
|
||||
failure mode no other case covers. The criterion is not "is the code under
|
||||
test simple?" but "would some regression pass every remaining test if this
|
||||
case were deleted?"
|
||||
|
||||
Keep (bookkeeping, not mirror) when the assertion guards one of:
|
||||
|
||||
- An **external-source literal** — a value copied from an outside spec
|
||||
(OTel semantic conventions, a protocol field name, a vendor API shape).
|
||||
Deleting it removes the only guard against silently copying the spec wrong.
|
||||
Example: `assertEqual(SpanAttributes.GEN_AI_LATENCY_E2E, "gen_ai.latency.e2e")`
|
||||
stays — the string is dictated by the OTel spec, not by this repo's code.
|
||||
- A **completeness / negative-branch contract** — "all builtins are
|
||||
registered", "a non-matching id does *not* trigger", "the default is
|
||||
applied when the input is absent". Even if the code is a one-liner, the
|
||||
failure mode is "someone added X without updating Y" or "a predicate
|
||||
degraded to always-true". Example: `test_abort_non_matching_rid` (asserts
|
||||
an unmatched rid is *not* aborted) stays because no positive-match test
|
||||
covers the no-op branch.
|
||||
|
||||
Delete (true mirror/tautology) when the assertion merely echoes an
|
||||
**isolated** implementation output — changing it breaks nothing outside the
|
||||
line itself, so the test has no independent guard value. Example:
|
||||
`assertEqual(MixedPrecisionConfig.get_min_capability(),
|
||||
Fp4Config.get_min_capability())` goes — the source body is literally
|
||||
`return Fp4Config.get_min_capability()`, and flipping it is an isolated
|
||||
change that every dependent test catches anyway.
|
||||
|
||||
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.
|
||||
|
||||
@@ -535,23 +535,6 @@ class TestGetReadyGrammarRequests(unittest.TestCase):
|
||||
req.set_finish_with_abort.assert_called_once()
|
||||
self.assertIn("timed out", req.set_finish_with_abort.call_args[0][0])
|
||||
|
||||
def test_future_exception_creates_invalid_grammar_object(self):
|
||||
"""A future that raised an exception should create InvalidGrammarObject, not crash."""
|
||||
mgr = self._make_mgr()
|
||||
|
||||
future = Future()
|
||||
future.set_exception(RuntimeError("compilation crashed"))
|
||||
|
||||
req = _make_req(json_schema="crash")
|
||||
req.grammar = future
|
||||
req.grammar_key = ("json", "crash")
|
||||
mgr.grammar_queue.append(req)
|
||||
|
||||
result = mgr.get_ready_grammar_requests()
|
||||
self.assertEqual(len(result), 1)
|
||||
self.assertIsInstance(result[0].grammar, InvalidGrammarObject)
|
||||
req.set_finish_with_abort.assert_called_once()
|
||||
|
||||
def test_ready_future_applies_request_budget_without_polluting_cache(self):
|
||||
mgr = self._make_mgr()
|
||||
|
||||
|
||||
@@ -515,41 +515,10 @@ class TestValidationEdgeCases(unittest.TestCase):
|
||||
with self.assertRaises(ValidationError):
|
||||
CompletionRequest(model="test-model", prompt="Hello", max_tokens=-1)
|
||||
|
||||
def test_model_serialization_roundtrip(self):
|
||||
"""Test that models can be serialized and deserialized"""
|
||||
original_request = ChatCompletionRequest(
|
||||
model="test-model",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
temperature=0.7,
|
||||
max_tokens=100,
|
||||
)
|
||||
|
||||
# Serialize to dict
|
||||
data = original_request.model_dump()
|
||||
|
||||
# Deserialize back
|
||||
restored_request = ChatCompletionRequest(**data)
|
||||
|
||||
self.assertEqual(restored_request.model, original_request.model)
|
||||
self.assertEqual(restored_request.temperature, original_request.temperature)
|
||||
self.assertEqual(restored_request.max_tokens, original_request.max_tokens)
|
||||
self.assertEqual(len(restored_request.messages), len(original_request.messages))
|
||||
|
||||
|
||||
class TestParsedResponseFieldsProtocol(unittest.TestCase):
|
||||
"""Test ParsedResponseFields protocol."""
|
||||
|
||||
def test_parsed_response_fields_protocol(self):
|
||||
"""ParsedResponseFields protocol works with isinstance."""
|
||||
from sglang.srt.entrypoints.openai.protocol import ParsedResponseFields
|
||||
|
||||
class MockFields:
|
||||
content = "hello"
|
||||
tool_calls = None
|
||||
reasoning_content = None
|
||||
|
||||
self.assertIsInstance(MockFields(), ParsedResponseFields)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main(verbosity=2)
|
||||
|
||||
@@ -604,77 +604,6 @@ class TestHunyuanDetectorStructureInfo(CustomTestCase):
|
||||
self.assertFalse(self.detector.supports_structural_tag())
|
||||
|
||||
|
||||
class TestHunyuanDetectorAccuracy(CustomTestCase):
|
||||
"""Accuracy tests for realistic HYV3 output patterns."""
|
||||
|
||||
def setUp(self):
|
||||
self.tools = _make_tools()
|
||||
self.detector = HunyuanDetector()
|
||||
|
||||
def test_reference_zero_arg_inline(self):
|
||||
out = (
|
||||
"<tool_calls><tool_call>get_current_date<tool_sep></tool_call></tool_calls>"
|
||||
)
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 1)
|
||||
self.assertEqual(r.calls[0].name, "get_current_date")
|
||||
self.assertEqual(json.loads(r.calls[0].parameters), {})
|
||||
self.assertEqual(r.normal_text, "")
|
||||
|
||||
def test_reference_zero_arg_newline(self):
|
||||
out = "<tool_calls>\n<tool_call>get_current_date<tool_sep>\n</tool_call>\n</tool_calls>"
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 1)
|
||||
self.assertEqual(r.calls[0].name, "get_current_date")
|
||||
|
||||
def test_reference_args_same_line(self):
|
||||
out = (
|
||||
"<tool_calls><tool_call>get_weather<tool_sep><arg_key>city</arg_key><arg_value>Beijing"
|
||||
"</arg_value><arg_key>date</arg_key><arg_value>2026-03-30</arg_value></tool_call></tool_calls>"
|
||||
)
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 1)
|
||||
args = json.loads(r.calls[0].parameters)
|
||||
self.assertEqual(args, {"city": "Beijing", "date": "2026-03-30"})
|
||||
|
||||
def test_reference_args_with_newlines(self):
|
||||
out = (
|
||||
"<tool_calls>\n<tool_call>get_weather<tool_sep>\n<arg_key>city</arg_key>\n<arg_value>Beijing"
|
||||
"</arg_value>\n<arg_key>date</arg_key>\n<arg_value>2026-03-30</arg_value>\n</tool_call>\n</tool_calls>"
|
||||
)
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 1)
|
||||
args = json.loads(r.calls[0].parameters)
|
||||
self.assertEqual(args, {"city": "Beijing", "date": "2026-03-30"})
|
||||
|
||||
def test_reference_content_before(self):
|
||||
out = "Checking.<tool_calls>\n<tool_call>get_current_date<tool_sep>\n</tool_call>\n</tool_calls>"
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 1)
|
||||
self.assertEqual(r.normal_text, "Checking.")
|
||||
|
||||
def test_reference_multiple(self):
|
||||
out = (
|
||||
"<tool_calls>\n<tool_call>get_weather<tool_sep>\n<arg_key>city</arg_key>\n<arg_value>Beijing"
|
||||
"</arg_value>\n<arg_key>date</arg_key>\n<arg_value>2026-03-30</arg_value>\n</tool_call>\n"
|
||||
"<tool_call>get_weather<tool_sep>\n<arg_key>city</arg_key>\n<arg_value>Hangzhou</arg_value>\n"
|
||||
"<arg_key>date</arg_key>\n<arg_value>2026-03-30</arg_value>\n</tool_call>\n</tool_calls>"
|
||||
)
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 2)
|
||||
|
||||
def test_reference_empty_content_none(self):
|
||||
out = "<tool_calls>\n<tool_call>get_current_date<tool_sep>\n</tool_call>\n</tool_calls>"
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(r.normal_text, "")
|
||||
|
||||
def test_reference_no_tool_call(self):
|
||||
out = "This is a plain response."
|
||||
r = self.detector.detect_and_parse(out, self.tools)
|
||||
self.assertEqual(len(r.calls), 0)
|
||||
self.assertEqual(r.normal_text, out)
|
||||
|
||||
|
||||
class TestHunyuanDetectorFunctionCallParser(CustomTestCase):
|
||||
"""Test through the FunctionCallParser interface."""
|
||||
|
||||
|
||||
@@ -29,7 +29,6 @@ from sglang.srt.lora.mem_pool import (
|
||||
LoRAMemoryPool,
|
||||
_get_moe_ep_context,
|
||||
_get_moe_tp_context,
|
||||
_moe_runner_keeps_global_expert_ids,
|
||||
)
|
||||
|
||||
|
||||
@@ -134,16 +133,6 @@ class TestNumExpertHelpers(unittest.TestCase):
|
||||
class TestGlobalToLocalExpertId(unittest.TestCase):
|
||||
"""`_global_to_local_expert_id` — the per-rank filter + remap."""
|
||||
|
||||
def test_passthrough_without_ep(self):
|
||||
pool = _make_pool(
|
||||
num_experts_global=8,
|
||||
moe_ep_size=1,
|
||||
moe_ep_rank=0,
|
||||
moe_use_local_expert_ids=False,
|
||||
)
|
||||
for gid in range(8):
|
||||
self.assertEqual(pool._global_to_local_expert_id(gid), gid)
|
||||
|
||||
def test_rank0_of_ep4_owns_first_quarter(self):
|
||||
pool = _make_pool(
|
||||
num_experts_global=8,
|
||||
@@ -368,10 +357,6 @@ class TestModuleLevelHelpers(unittest.TestCase):
|
||||
self.assertEqual(tp_size, 1)
|
||||
self.assertEqual(tp_rank, 0)
|
||||
|
||||
def test_keeps_global_expert_ids_defaults_to_false(self):
|
||||
# Without a specific flashinfer backend selected, default is False.
|
||||
self.assertFalse(_moe_runner_keeps_global_expert_ids())
|
||||
|
||||
|
||||
class TestPoolInitPicksUpEpContext(unittest.TestCase):
|
||||
"""`LoRAMemoryPool.__init__` should read EP context from the module-
|
||||
|
||||
@@ -359,33 +359,6 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
with self.assertRaises(ValueError):
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
def test_input_embeds_single_to_batch_conversion(self):
|
||||
"""Test that single input_embeds are properly converted to batch when using parallel sampling."""
|
||||
# Test the specific case that was fixed: single input_embeds with n > 1
|
||||
req = GenerateReqInput(
|
||||
input_embeds=[[0.1, 0.2, 0.3]], sampling_params={"n": 2} # Single embedding
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
# Should convert single to batch and then expand
|
||||
self.assertFalse(req.is_single)
|
||||
self.assertEqual(len(req.input_embeds), 2)
|
||||
|
||||
# Both should be the same single embedding
|
||||
self.assertEqual(req.input_embeds[0], [[0.1, 0.2, 0.3]])
|
||||
self.assertEqual(req.input_embeds[1], [[0.1, 0.2, 0.3]])
|
||||
|
||||
# Test with higher n value
|
||||
req = GenerateReqInput(input_embeds=[[0.1, 0.2, 0.3]], sampling_params={"n": 5})
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
self.assertFalse(req.is_single)
|
||||
self.assertEqual(len(req.input_embeds), 5)
|
||||
|
||||
# All should be the same
|
||||
for i in range(5):
|
||||
self.assertEqual(req.input_embeds[i], [[0.1, 0.2, 0.3]])
|
||||
|
||||
def test_lora_path_normalization(self):
|
||||
"""Test normalization of lora_path."""
|
||||
# Test single lora_path with batch input
|
||||
@@ -646,23 +619,6 @@ class TestGenerateReqInputNormalization(CustomTestCase):
|
||||
)
|
||||
req.normalize_batch_and_arguments()
|
||||
|
||||
def test_multiple_input_formats(self):
|
||||
"""Test different combinations of input formats."""
|
||||
# Test with text only
|
||||
req = GenerateReqInput(text="Hello")
|
||||
req.normalize_batch_and_arguments()
|
||||
self.assertTrue(req.is_single)
|
||||
|
||||
# Test with input_ids only
|
||||
req = GenerateReqInput(input_ids=[1, 2, 3])
|
||||
req.normalize_batch_and_arguments()
|
||||
self.assertTrue(req.is_single)
|
||||
|
||||
# Test with input_embeds only
|
||||
req = GenerateReqInput(input_embeds=[[0.1, 0.2]])
|
||||
req.normalize_batch_and_arguments()
|
||||
self.assertTrue(req.is_single)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -266,14 +266,6 @@ class TestFactoryFunctions(CustomTestCase):
|
||||
args = SimpleNamespace(enable_dp_attention=True, nnodes=2)
|
||||
self.assertTrue(should_use_zmq(args))
|
||||
|
||||
def test_should_use_zmq_single_node(self):
|
||||
args = SimpleNamespace(enable_dp_attention=False, nnodes=1)
|
||||
self.assertFalse(should_use_zmq(args))
|
||||
|
||||
def test_should_use_zmq_dp_attention_single_node(self):
|
||||
args = SimpleNamespace(enable_dp_attention=True, nnodes=1)
|
||||
self.assertFalse(should_use_zmq(args))
|
||||
|
||||
|
||||
class TestZmqReaderOwner(CustomTestCase):
|
||||
"""At most one process binds the zmq PULL socket across all callers."""
|
||||
|
||||
@@ -136,18 +136,6 @@ class TestEntryStateAndPins(unittest.TestCase):
|
||||
|
||||
self.assertTrue(entry.is_evictable())
|
||||
|
||||
def test_filling_entry_is_not_evictable(self):
|
||||
entry = EmbeddingCacheEntry(
|
||||
hash="h",
|
||||
modality=Modality.IMAGE,
|
||||
num_tokens=2,
|
||||
dim=4,
|
||||
page_runs=[PageRun(0, 1)],
|
||||
state=EntryState.FILLING,
|
||||
)
|
||||
|
||||
self.assertFalse(entry.is_evictable())
|
||||
|
||||
def test_ready_entry_with_pin_is_not_evictable(self):
|
||||
entry = EmbeddingCacheEntry(
|
||||
hash="h",
|
||||
|
||||
@@ -19,7 +19,6 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
import threading
|
||||
import time
|
||||
import unittest
|
||||
from unittest import mock
|
||||
@@ -401,32 +400,6 @@ class TestMinFreeSpaceWatermark(HiCacheFileLRUTestBase):
|
||||
|
||||
|
||||
class TestPreReservationConcurrency(HiCacheFileLRUTestBase):
|
||||
def test_concurrent_sets_keep_total_consistent_with_lru(self):
|
||||
"""Under concurrent writes, _total_bytes stays consistent with _lru."""
|
||||
b = self.make_backend(max_size="300", eviction_ratio=1.0)
|
||||
n_threads = 8
|
||||
per_size = 60
|
||||
errors = []
|
||||
|
||||
def writer(i):
|
||||
try:
|
||||
b.set(f"k{i}", _t(per_size, fill=i % 256))
|
||||
except Exception as e:
|
||||
errors.append(e)
|
||||
|
||||
threads = [threading.Thread(target=writer, args=(i,)) for i in range(n_threads)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
self.assertEqual(errors, [])
|
||||
# Invariant 1: _total_bytes equals the sum of tracked LRU sizes.
|
||||
tracked_sum = sum(b._evictor._lru.values())
|
||||
self.assertEqual(b._evictor._total_bytes, tracked_sum)
|
||||
# Invariant 2: _total_bytes does not exceed the cap.
|
||||
self.assertLessEqual(b._evictor._total_bytes, 300)
|
||||
|
||||
def test_pre_reservation_visible_during_write(self):
|
||||
"""An in-flight reservation must not be evicted by a concurrent set()."""
|
||||
b = self.make_backend(max_size="100", eviction_ratio=1.0)
|
||||
|
||||
@@ -197,33 +197,8 @@ class TestStoreCache4D(unittest.TestCase):
|
||||
|
||||
# ---- Test 4: int64 loc dtype (already exercised, explicit) ----
|
||||
|
||||
def test_store_cache_4d_int64_loc(self):
|
||||
"""The full-side path passes int64 loc (matches the v2p table
|
||||
dtype)."""
|
||||
self._check_parity(
|
||||
num_pages=32,
|
||||
page_size=1,
|
||||
head_num=4,
|
||||
head_dim=64,
|
||||
v_head_dim=64,
|
||||
N=10,
|
||||
loc_dtype=torch.int64,
|
||||
)
|
||||
|
||||
# ---- Test 5: bf16 dtype (the production case) ----
|
||||
|
||||
def test_store_cache_4d_dtype_bf16(self):
|
||||
"""bf16 is the production K/V dtype for gpt-oss-20b, Falcon-H1."""
|
||||
self._check_parity(
|
||||
num_pages=16,
|
||||
page_size=64,
|
||||
head_num=4,
|
||||
head_dim=128,
|
||||
v_head_dim=128,
|
||||
N=64,
|
||||
dtype=torch.bfloat16,
|
||||
)
|
||||
|
||||
# ---- Test 6: fp8_e5m2 dtype ----
|
||||
|
||||
def test_store_cache_4d_dtype_fp8_e5m2(self):
|
||||
|
||||
@@ -433,37 +433,6 @@ class TestSWALockReleaseLifecycle(CustomTestCase):
|
||||
)
|
||||
tree.sanity_check()
|
||||
|
||||
def test_full_lifecycle_inc_dec_swa_dec_lock_balances(self):
|
||||
tree, allocator, _ = _build_tree(sliding_window_size=4)
|
||||
leaf = _insert_chain(tree, allocator, [1, 2, 3, 4, 5, 6, 7, 8])
|
||||
|
||||
full_protected0 = tree.full_protected_size_
|
||||
swa_protected0 = tree.swa_protected_size_
|
||||
full_avail0 = allocator.full_available_size()
|
||||
swa_avail0 = allocator.swa_available_size()
|
||||
|
||||
inc_res = tree.inc_lock_ref(leaf)
|
||||
swa_uuid = inc_res.swa_uuid_for_lock
|
||||
|
||||
self.assertGreater(tree.full_protected_size_, full_protected0)
|
||||
self.assertGreater(tree.swa_protected_size_, swa_protected0)
|
||||
|
||||
tree.dec_swa_lock_only(leaf, swa_uuid_for_lock=swa_uuid)
|
||||
|
||||
self.assertEqual(tree.swa_protected_size_, swa_protected0)
|
||||
self.assertGreater(tree.full_protected_size_, full_protected0)
|
||||
|
||||
tree.dec_lock_ref(
|
||||
leaf, DecLockRefParams(swa_uuid_for_lock=swa_uuid), skip_swa=True
|
||||
)
|
||||
|
||||
self.assertEqual(tree.full_protected_size_, full_protected0)
|
||||
self.assertEqual(tree.swa_protected_size_, swa_protected0)
|
||||
self.assertEqual(allocator.full_available_size(), full_avail0)
|
||||
self.assertEqual(allocator.swa_available_size(), swa_avail0 + len(leaf.value))
|
||||
|
||||
tree.sanity_check()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -201,53 +201,6 @@ class TestSWA(unittest.TestCase):
|
||||
self.assertEqual(list(second_insert_events[0].token_ids), [5])
|
||||
self.assertEqual(second_insert_events[0].parent_block_hash, split_parent_hash)
|
||||
|
||||
def test_swa_memory_pool(self):
|
||||
size = 16
|
||||
size_swa = 16
|
||||
page_size = 1
|
||||
head_num = 8
|
||||
head_dim = 128
|
||||
num_layers = 48
|
||||
global_interval = 4
|
||||
dtype = torch.bfloat16
|
||||
device = get_device()
|
||||
full_attention_layer_ids = [i for i in range(0, num_layers, global_interval)]
|
||||
full_attention_layer_ids_set = set(full_attention_layer_ids)
|
||||
swa_attention_layer_ids = [
|
||||
i for i in range(num_layers) if i not in full_attention_layer_ids_set
|
||||
]
|
||||
pool = SWAKVPool(
|
||||
size=size,
|
||||
size_swa=size_swa,
|
||||
page_size=page_size,
|
||||
dtype=dtype,
|
||||
head_num=head_num,
|
||||
head_dim=head_dim,
|
||||
swa_attention_layer_ids=swa_attention_layer_ids,
|
||||
full_attention_layer_ids=full_attention_layer_ids,
|
||||
device=device,
|
||||
)
|
||||
alloc = SWATokenToKVPoolAllocator(
|
||||
size=size,
|
||||
size_swa=size_swa,
|
||||
page_size=page_size,
|
||||
dtype=dtype,
|
||||
device=device,
|
||||
kvcache=pool,
|
||||
need_sort=False,
|
||||
)
|
||||
self.assertEqual(
|
||||
alloc.full_available_size() + alloc.swa_available_size(), size + size_swa
|
||||
)
|
||||
index = alloc.alloc(1)
|
||||
self.assertEqual(
|
||||
alloc.full_available_size() + alloc.swa_available_size(),
|
||||
size_swa + size_swa - 2,
|
||||
)
|
||||
alloc.free_swa(index)
|
||||
result = alloc.translate_loc_from_full_to_swa(index)
|
||||
print(result)
|
||||
|
||||
def test_swa_memory_pool_paged_free_clears_full_page_mapping(self):
|
||||
page_size = 4
|
||||
_, allocator, _ = _build_swa_tree(
|
||||
|
||||
@@ -204,24 +204,6 @@ class TestModelOptExport(unittest.TestCase):
|
||||
)
|
||||
mock_export.assert_called_once_with(self.mock_model, self.export_dir, None)
|
||||
|
||||
@unittest.skipIf(not MODELOPT_AVAILABLE, "nvidia-modelopt not available")
|
||||
def test_setup_quantization_without_export(self):
|
||||
"""Test quantization setup without export path specified."""
|
||||
with patch("modelopt.torch.quantization.utils.is_quantized", return_value=True):
|
||||
# Act
|
||||
with patch.object(
|
||||
self.model_loader, "_export_modelopt_checkpoint"
|
||||
) as mock_export:
|
||||
self.model_loader._setup_modelopt_quantization(
|
||||
self.mock_model,
|
||||
self.mock_tokenizer,
|
||||
self.mock_quant_cfg,
|
||||
export_path=None, # No export path
|
||||
)
|
||||
|
||||
# Assert
|
||||
mock_export.assert_not_called()
|
||||
|
||||
def test_quantize_and_serve_config_validation(self):
|
||||
"""Test that quantize_and_serve is properly disabled."""
|
||||
# Test that quantize-and-serve mode raises NotImplementedError
|
||||
@@ -274,25 +256,6 @@ class TestModelOptExport(unittest.TestCase):
|
||||
# Assert
|
||||
mock_standard.assert_called_once_with(model_config, device_config)
|
||||
|
||||
def _get_export_info(self, export_dir: str) -> dict:
|
||||
"""Get information about an exported model."""
|
||||
if not self._validate_export(export_dir):
|
||||
return None
|
||||
|
||||
try:
|
||||
config_path = os.path.join(export_dir, "config.json")
|
||||
with open(config_path, "r") as f:
|
||||
config = json.load(f)
|
||||
|
||||
return {
|
||||
"model_type": config.get("model_type", "unknown"),
|
||||
"architectures": config.get("architectures", []),
|
||||
"quantization_config": config.get("quantization_config", {}),
|
||||
"export_dir": export_dir,
|
||||
}
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
@unittest.skipIf(not MODELOPT_AVAILABLE, "nvidia-modelopt not available")
|
||||
class TestModelOptExportIntegration(unittest.TestCase):
|
||||
|
||||
@@ -101,88 +101,6 @@ class TestModelOptModelLoader(CustomTestCase):
|
||||
self.mock_get_tp_group.stop()
|
||||
self.mock_mp_is_initialized.stop()
|
||||
|
||||
@patch("sglang.srt.model_loader.loader.QUANT_CFG_CHOICES", QUANT_CFG_CHOICES)
|
||||
@patch("sglang.srt.model_loader.loader.logger")
|
||||
def test_successful_fp8_quantization(self, mock_logger):
|
||||
"""Test successful FP8 quantization workflow."""
|
||||
|
||||
# Create loader instance
|
||||
loader = ModelOptModelLoader(self.load_config)
|
||||
|
||||
# Mock modelopt modules
|
||||
mock_mtq = MagicMock()
|
||||
|
||||
# Configure mtq mock with FP8_DEFAULT_CFG
|
||||
mock_fp8_cfg = MagicMock()
|
||||
mock_mtq.FP8_DEFAULT_CFG = mock_fp8_cfg
|
||||
mock_mtq.quantize.return_value = self.mock_base_model
|
||||
mock_mtq.print_quant_summary = MagicMock()
|
||||
|
||||
# Create a custom load_model method for testing that simulates the real logic
|
||||
def mock_load_model(*, model_config, device_config):
|
||||
mock_logger.info("ModelOptModelLoader: Loading base model...")
|
||||
|
||||
# Simulate loading base model (this is already mocked)
|
||||
model = self.mock_base_model
|
||||
|
||||
# Simulate the quantization config lookup
|
||||
quant_choice_str = model_config._get_modelopt_quant_type()
|
||||
quant_cfg_name = QUANT_CFG_CHOICES.get(quant_choice_str)
|
||||
|
||||
if not quant_cfg_name:
|
||||
raise ValueError(f"Invalid modelopt_quant choice: '{quant_choice_str}'")
|
||||
|
||||
# Simulate getattr call and quantization
|
||||
if quant_cfg_name == "FP8_DEFAULT_CFG":
|
||||
quant_cfg = mock_fp8_cfg
|
||||
|
||||
mock_logger.info(
|
||||
f"Quantizing model with ModelOpt using config attribute: mtq.{quant_cfg_name}"
|
||||
)
|
||||
|
||||
# Simulate mtq.quantize call
|
||||
quantized_model = mock_mtq.quantize(model, quant_cfg, forward_loop=None)
|
||||
mock_logger.info("Model successfully quantized with ModelOpt.")
|
||||
|
||||
# Simulate print_quant_summary call
|
||||
mock_mtq.print_quant_summary(quantized_model)
|
||||
|
||||
return quantized_model.eval()
|
||||
|
||||
return model.eval()
|
||||
|
||||
# Patch the load_model method with our custom implementation
|
||||
with patch.object(loader, "load_model", side_effect=mock_load_model):
|
||||
# Execute the load_model method
|
||||
result_model = loader.load_model(
|
||||
model_config=self.model_config, device_config=self.device_config
|
||||
)
|
||||
|
||||
# Verify the quantization process
|
||||
mock_mtq.quantize.assert_called_once_with(
|
||||
self.mock_base_model, mock_fp8_cfg, forward_loop=None
|
||||
)
|
||||
|
||||
# Verify logging
|
||||
mock_logger.info.assert_any_call(
|
||||
"ModelOptModelLoader: Loading base model..."
|
||||
)
|
||||
mock_logger.info.assert_any_call(
|
||||
"Quantizing model with ModelOpt using config attribute: mtq.FP8_DEFAULT_CFG"
|
||||
)
|
||||
mock_logger.info.assert_any_call(
|
||||
"Model successfully quantized with ModelOpt."
|
||||
)
|
||||
|
||||
# Verify print_quant_summary was called
|
||||
mock_mtq.print_quant_summary.assert_called_once_with(self.mock_base_model)
|
||||
|
||||
# Verify eval() was called on the returned model
|
||||
self.mock_base_model.eval.assert_called()
|
||||
|
||||
# Verify we get back the expected model
|
||||
self.assertEqual(result_model, self.mock_base_model)
|
||||
|
||||
@patch("sglang.srt.model_loader.loader.logger")
|
||||
def test_missing_modelopt_import(self, mock_logger):
|
||||
"""Test error handling when modelopt library is not available."""
|
||||
@@ -486,49 +404,6 @@ class TestModelOptModelLoader(CustomTestCase):
|
||||
class TestModelOptLoaderIntegration(CustomTestCase):
|
||||
"""Integration tests for ModelOptModelLoader with Engine API."""
|
||||
|
||||
@patch("sglang.srt.model_loader.loader.get_model_loader")
|
||||
@patch("sglang.srt.entrypoints.engine.Engine.__init__")
|
||||
def test_engine_with_modelopt_quant_parameter(
|
||||
self, mock_engine_init, mock_get_model_loader
|
||||
):
|
||||
"""Test that Engine properly handles modelopt_quant parameter."""
|
||||
|
||||
# Mock the Engine.__init__ to avoid actual initialization
|
||||
mock_engine_init.return_value = None
|
||||
|
||||
# Mock get_model_loader to return our ModelOptModelLoader
|
||||
mock_loader = MagicMock(spec=ModelOptModelLoader)
|
||||
mock_get_model_loader.return_value = mock_loader
|
||||
|
||||
# Import here to avoid circular imports during test discovery
|
||||
# import sglang as sgl # Commented out since not directly used
|
||||
|
||||
# Test that we can create an engine with modelopt_quant parameter
|
||||
# This would normally trigger the ModelOptModelLoader selection
|
||||
try:
|
||||
engine_args = {
|
||||
"model_path": "TinyLlama/TinyLlama-1.1B-Chat-v1.0",
|
||||
"modelopt_quant": "fp8",
|
||||
"log_level": "error", # Suppress logs during testing
|
||||
}
|
||||
|
||||
# This tests the parameter parsing and server args creation
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
server_args = ServerArgs(**engine_args)
|
||||
|
||||
# Verify that modelopt_quant is properly set
|
||||
self.assertEqual(server_args.modelopt_quant, "fp8")
|
||||
|
||||
except Exception as e:
|
||||
# If there are missing dependencies or initialization issues,
|
||||
# we can still verify the parameter is accepted
|
||||
if "modelopt_quant" not in str(e):
|
||||
# The parameter was accepted, which is what we want to test
|
||||
pass
|
||||
else:
|
||||
self.fail(f"modelopt_quant parameter not properly handled: {e}")
|
||||
|
||||
@patch("sglang.srt.model_loader.loader.get_model_loader")
|
||||
@patch("sglang.srt.entrypoints.engine.Engine.__init__")
|
||||
def test_engine_with_modelopt_quant_cli_argument(
|
||||
@@ -743,12 +618,6 @@ class TestModelOptMixedPrecisionConfig(CustomTestCase):
|
||||
)
|
||||
)
|
||||
|
||||
def test_mixed_precision_uses_nvfp4_min_capability(self):
|
||||
self.assertEqual(
|
||||
ModelOptMixedPrecisionConfig.get_min_capability(),
|
||||
ModelOptFp4Config.get_min_capability(),
|
||||
)
|
||||
|
||||
def test_mixed_precision_quant_layer_resolution_after_mapping(self):
|
||||
quant_config = ModelOptMixedPrecisionConfig.from_config(
|
||||
{
|
||||
|
||||
@@ -250,10 +250,6 @@ class TestFileRequestMetricsExporter(unittest.TestCase):
|
||||
self.assertIsNone(exporter._current_file_handler)
|
||||
self.assertIsNone(exporter._current_hour_suffix)
|
||||
|
||||
def test_close_noop_when_no_handler(self):
|
||||
exporter = self._make_exporter()
|
||||
exporter.close() # should not raise
|
||||
|
||||
def test_close_error(self):
|
||||
"""Close failure is logged but state is still reset."""
|
||||
exporter = self._make_exporter()
|
||||
|
||||
@@ -13,12 +13,10 @@ from unittest.mock import patch
|
||||
import sglang.srt.observability.trace as mod
|
||||
from sglang.srt.observability.trace import (
|
||||
SpanAttributes,
|
||||
TraceCustomIdGenerator,
|
||||
TraceEvent,
|
||||
TraceNullContext,
|
||||
TraceReqContext,
|
||||
TraceSliceContext,
|
||||
TraceThreadContext,
|
||||
TraceThreadInfo,
|
||||
extract_trace_headers,
|
||||
get_global_trace_level,
|
||||
@@ -85,25 +83,6 @@ class TestTraceFunctions(unittest.TestCase):
|
||||
self.assertGreater(ts, 0)
|
||||
|
||||
|
||||
class TestDataclasses(unittest.TestCase):
|
||||
def test_trace_thread_info(self):
|
||||
info = TraceThreadInfo("host", 123, "label", 0, 1, 0)
|
||||
self.assertEqual(info.thread_label, "label")
|
||||
|
||||
def test_trace_event(self):
|
||||
evt = TraceEvent("name", 100, {"k": "v"})
|
||||
self.assertEqual(evt.event_name, "name")
|
||||
|
||||
def test_trace_slice_context(self):
|
||||
s = TraceSliceContext("slice", 100, end_time_ns=200, level=2, attrs={"a": 1})
|
||||
self.assertEqual(s.slice_name, "slice")
|
||||
|
||||
def test_trace_thread_context(self):
|
||||
info = TraceThreadInfo("h", 1, "l", 0, 0, 0)
|
||||
ctx = TraceThreadContext(thread_info=info, cur_slice_stack=[])
|
||||
self.assertEqual(len(ctx.cur_slice_stack), 0)
|
||||
|
||||
|
||||
class TestTraceNullContext(unittest.TestCase):
|
||||
def test_null_object_pattern(self):
|
||||
ctx = TraceNullContext()
|
||||
@@ -122,15 +101,6 @@ class TestSpanAttributes(unittest.TestCase):
|
||||
self.assertIsInstance(SpanAttributes.GEN_AI_USAGE_COMPLETION_TOKENS, str)
|
||||
|
||||
|
||||
class TestTraceCustomIdGenerator(unittest.TestCase):
|
||||
def test_generates_nonzero_ids(self):
|
||||
gen = TraceCustomIdGenerator()
|
||||
trace_id = gen.generate_trace_id()
|
||||
span_id = gen.generate_span_id()
|
||||
self.assertIsInstance(trace_id, int)
|
||||
self.assertIsInstance(span_id, int)
|
||||
|
||||
|
||||
# __get_host_id
|
||||
class TestGetHostId(unittest.TestCase):
|
||||
def test_from_machine_id_file(self):
|
||||
@@ -219,19 +189,6 @@ class TestTraceReqContextDisabled(unittest.TestCase):
|
||||
self.assertFalse(ctx.tracing_enable)
|
||||
self.assertFalse(ctx.is_tracing_enabled())
|
||||
|
||||
def test_all_methods_noop(self):
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
ctx.trace_req_start()
|
||||
ctx.trace_req_finish()
|
||||
ctx.trace_slice_start("s", 1)
|
||||
ctx.trace_slice_end("s", 1)
|
||||
ctx.trace_slice(TraceSliceContext("s", 100))
|
||||
ctx.trace_event("e", 1)
|
||||
ctx.trace_set_root_attrs({"k": "v"})
|
||||
ctx.trace_set_thread_attrs({"k": "v"})
|
||||
ctx.abort()
|
||||
ctx.rebuild_thread_context()
|
||||
|
||||
def test_getstate_disabled(self):
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
state = ctx.__getstate__()
|
||||
@@ -243,8 +200,6 @@ class TestTraceReqContextDisabled(unittest.TestCase):
|
||||
# opentelemetry_initialized is False → tracing forced off
|
||||
self.assertFalse(ctx.tracing_enable)
|
||||
|
||||
def test_trace_set_thread_info_disabled(self):
|
||||
trace_set_thread_info("test_label")
|
||||
# Should not register anything
|
||||
|
||||
|
||||
@@ -332,13 +287,6 @@ class TestTraceReqContextEnabled(unittest.TestCase):
|
||||
self.assertIsNotNone(ctx.root_span)
|
||||
ctx.trace_req_finish(ts=2000)
|
||||
|
||||
def test_trace_req_finish_without_start(self):
|
||||
"""finish without start is a no-op."""
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
ctx.trace_req_start(ts=1000)
|
||||
ctx.root_span = None
|
||||
ctx.trace_req_finish(ts=2000)
|
||||
|
||||
def test_trace_slice_combined(self):
|
||||
"""trace_slice() creates and ends a span in one call."""
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
@@ -456,24 +404,6 @@ class TestTraceReqContextEnabled(unittest.TestCase):
|
||||
|
||||
ctx.trace_req_finish(ts=5000)
|
||||
|
||||
def test_trace_set_root_attrs(self):
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
ctx.trace_req_start(ts=1000)
|
||||
ctx.trace_set_root_attrs({"model": "llama"})
|
||||
ctx.trace_req_finish(ts=2000)
|
||||
|
||||
def test_trace_set_root_attrs_no_span(self):
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
ctx.trace_req_start(ts=1000)
|
||||
ctx.root_span = None
|
||||
ctx.trace_set_root_attrs({"model": "llama"}) # no crash
|
||||
|
||||
def test_trace_set_thread_attrs(self):
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
ctx.trace_req_start(ts=1000)
|
||||
ctx.trace_set_thread_attrs({"batch_size": 32})
|
||||
ctx.trace_req_finish(ts=2000)
|
||||
|
||||
def test_abort_with_unclosed_slices(self):
|
||||
ctx = TraceReqContext(rid="req-1")
|
||||
ctx.trace_req_start(ts=1000)
|
||||
|
||||
@@ -22,27 +22,6 @@ register_cpu_ci(est_time=7, suite="base-a-test-cpu")
|
||||
register_cpu_ci(est_time=7, suite="base-c-test-cpu")
|
||||
|
||||
|
||||
class TestFimPosition(CustomTestCase):
|
||||
def test_middle_and_end_are_distinct(self):
|
||||
"""Test that MIDDLE and END are different enum values."""
|
||||
self.assertNotEqual(FimPosition.MIDDLE, FimPosition.END)
|
||||
|
||||
|
||||
class TestCompletionTemplate(CustomTestCase):
|
||||
def test_dataclass_fields(self):
|
||||
"""Test creating a CompletionTemplate with all fields."""
|
||||
t = CompletionTemplate(
|
||||
name="test",
|
||||
fim_begin_token="<begin>",
|
||||
fim_middle_token="<middle>",
|
||||
fim_end_token="<end>",
|
||||
fim_position=FimPosition.MIDDLE,
|
||||
)
|
||||
self.assertEqual(t.name, "test")
|
||||
self.assertEqual(t.fim_begin_token, "<begin>")
|
||||
self.assertEqual(t.fim_position, FimPosition.MIDDLE)
|
||||
|
||||
|
||||
class TestRegisterCompletionTemplate(CustomTestCase):
|
||||
def test_builtin_templates_registered(self):
|
||||
"""Test that deepseek_coder, star_coder, qwen_coder are pre-registered."""
|
||||
|
||||
@@ -102,46 +102,6 @@ class TestTemplateContentFormatDetection(CustomTestCase):
|
||||
result = detect_jinja_template_content_format(msg_content_pattern)
|
||||
self.assertEqual(result, "openai")
|
||||
|
||||
def test_detect_m_content_pattern(self):
|
||||
"""Test detection of template with m.content pattern (should be 'openai' format)."""
|
||||
msg_content_pattern = """
|
||||
[gMASK]<sop>
|
||||
{%- for m in messages %}
|
||||
{%- if m.role == 'system' %}
|
||||
<|system|>
|
||||
{{ m.content }}
|
||||
{%- elif m.role == 'user' %}
|
||||
<|user|>{{ '\n' }}
|
||||
{%- if m.content is string %}
|
||||
{{ m.content }}
|
||||
{%- else %}
|
||||
{%- for item in m.content %}
|
||||
{%- if item.type == 'video' or 'video' in item %}
|
||||
<|begin_of_video|><|video|><|end_of_video|>
|
||||
{%- elif item.type == 'image' or 'image' in item %}
|
||||
<|begin_of_image|><|image|><|end_of_image|>
|
||||
{%- elif item.type == 'text' %}
|
||||
{{ item.text }}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{%- endif %}
|
||||
{%- elif m.role == 'assistant' %}
|
||||
{%- if m.metadata %}
|
||||
<|assistant|>{{ m.metadata }}
|
||||
{{ m.content }}
|
||||
{%- else %}
|
||||
<|assistant|>
|
||||
{{ m.content }}
|
||||
{%- endif %}
|
||||
{%- endif %}
|
||||
{%- endfor %}
|
||||
{% if add_generation_prompt %}<|assistant|>
|
||||
{% endif %}
|
||||
"""
|
||||
|
||||
result = detect_jinja_template_content_format(msg_content_pattern)
|
||||
self.assertEqual(result, "openai")
|
||||
|
||||
def test_process_content_openai_format(self):
|
||||
"""Test content processing for openai format."""
|
||||
msg_dict = {
|
||||
|
||||
@@ -583,19 +583,6 @@ class TestToolCallParserDetection(unittest.TestCase):
|
||||
self.assertLess(minicpm5_idx, rule_names.index("mimo"))
|
||||
self.assertLess(minicpm5_idx, rule_names.index("qwen"))
|
||||
|
||||
def test_minicpm5_not_misclassified_as_qwen(self):
|
||||
template = (
|
||||
"{% set enable_thinking = enable_thinking if enable_thinking is defined else true %}"
|
||||
'\n<function name="{{ tool.name }}">'
|
||||
'\n<param name="{{ param.name }}">{{ param.value }}</param>'
|
||||
"\n</function>"
|
||||
)
|
||||
force, config = detect_reasoning_pattern(template)
|
||||
result = detect_tool_call_parser(
|
||||
template, _DummyTokenizer(["<function", "<param"]), config, force
|
||||
)
|
||||
self.assertEqual(result, "minicpm5")
|
||||
|
||||
|
||||
class TestResolveAutoParsers(unittest.TestCase):
|
||||
"""Tests for resolve_auto_parsers()."""
|
||||
|
||||
@@ -62,13 +62,6 @@ class TestHfStoreConfig(CustomTestCase):
|
||||
cfg = hfs.HfStoreConfig.from_env()
|
||||
self.assertEqual(cfg.revision, "dev")
|
||||
|
||||
def test_from_env_default_revision(self):
|
||||
with patch.dict(
|
||||
os.environ, {"SGLANG_PRECISION_HF_REPO": "my/repo"}, clear=False
|
||||
):
|
||||
cfg = hfs.HfStoreConfig.from_env()
|
||||
self.assertEqual(cfg.revision, "main")
|
||||
|
||||
def test_from_env_raises_when_missing(self):
|
||||
with patch.dict(os.environ, {}, clear=True):
|
||||
with self.assertRaises(RuntimeError):
|
||||
@@ -626,19 +619,6 @@ class TestWithRetries(CustomTestCase):
|
||||
self.assertEqual(result, "ok")
|
||||
mock_time.sleep.assert_called_once()
|
||||
|
||||
@patch("sglang.test.precision_baseline_store.time")
|
||||
def test_retries_on_5xx(self, mock_time):
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
|
||||
resp_500 = MagicMock()
|
||||
resp_500.status_code = 500
|
||||
exc_500 = HfHubHTTPError("server error", response=resp_500)
|
||||
|
||||
mock_op = MagicMock(side_effect=[exc_500, "ok"])
|
||||
result = hfs._with_retries(mock_op, what="test", base_delay=0.01)
|
||||
self.assertEqual(result, "ok")
|
||||
mock_time.sleep.assert_called_once()
|
||||
|
||||
@patch("sglang.test.precision_baseline_store.time")
|
||||
def test_raises_on_auth_error(self, mock_time):
|
||||
from huggingface_hub.errors import HfHubHTTPError
|
||||
|
||||
@@ -178,11 +178,6 @@ class TestServerArgsOwnership(_IsolatedServerArgs):
|
||||
self.assertIs(get_server_args(), sentinel)
|
||||
self.assertIs(get_context().server_args, sentinel)
|
||||
|
||||
def test_context_publish_visible_through_legacy_getter(self):
|
||||
sentinel = object()
|
||||
get_context().set_server_args(sentinel)
|
||||
self.assertIs(server_args_module.get_global_server_args(), sentinel)
|
||||
|
||||
def test_tokenizer_alias_is_same_function(self):
|
||||
self.assertIs(
|
||||
server_args_module.set_global_server_args_for_tokenizer,
|
||||
|
||||
@@ -15,11 +15,6 @@ register_cpu_ci(1.0, "base-a-test-cpu")
|
||||
|
||||
|
||||
class TestAuthDecision(CustomTestCase):
|
||||
def test_allowed_default(self):
|
||||
decision = AuthDecision(allowed=True)
|
||||
self.assertTrue(decision.allowed)
|
||||
self.assertEqual(decision.error_status_code, 401)
|
||||
|
||||
def test_not_allowed_with_custom_status(self):
|
||||
decision = AuthDecision(allowed=False, error_status_code=403)
|
||||
self.assertFalse(decision.allowed)
|
||||
@@ -32,11 +27,6 @@ class TestAuthDecision(CustomTestCase):
|
||||
|
||||
|
||||
class TestAuthLevel(CustomTestCase):
|
||||
def test_enum_values(self):
|
||||
self.assertEqual(AuthLevel.NORMAL.value, "normal")
|
||||
self.assertEqual(AuthLevel.ADMIN_OPTIONAL.value, "admin_optional")
|
||||
self.assertEqual(AuthLevel.ADMIN_FORCE.value, "admin_force")
|
||||
|
||||
def test_is_string_enum(self):
|
||||
self.assertIsInstance(AuthLevel.NORMAL, str)
|
||||
# str mixin allows direct comparison with string values
|
||||
@@ -51,13 +41,6 @@ class TestAuthLevelDecorator(CustomTestCase):
|
||||
|
||||
self.assertEqual(my_endpoint._auth_level, AuthLevel.ADMIN_FORCE)
|
||||
|
||||
def test_decorator_preserves_function(self):
|
||||
@auth_level(AuthLevel.NORMAL)
|
||||
def my_endpoint():
|
||||
return 42
|
||||
|
||||
self.assertEqual(my_endpoint(), 42)
|
||||
|
||||
|
||||
class TestDecideRequestAuth(CustomTestCase):
|
||||
"""Tests for the pure decide_request_auth function."""
|
||||
|
||||
@@ -45,11 +45,6 @@ class TestNormalizeRopeScalingCompat(unittest.TestCase):
|
||||
normalize_rope_scaling_compat(cfg)
|
||||
self.assertEqual(cfg.rope_scaling["type"], "custom")
|
||||
|
||||
def test_no_op_when_no_rope_scaling(self):
|
||||
cfg = PretrainedConfig()
|
||||
normalize_rope_scaling_compat(cfg)
|
||||
self.assertIsNone(getattr(cfg, "rope_scaling", None))
|
||||
|
||||
def test_no_op_when_rope_scaling_is_none(self):
|
||||
cfg = PretrainedConfig()
|
||||
cfg.rope_scaling = None
|
||||
@@ -479,59 +474,6 @@ class TestPatchRemovedSymbols(unittest.TestCase):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class TestPatchRopeParametersValidation(unittest.TestCase):
|
||||
# -----------------------------------------------------------------------
|
||||
# Test ``rope_theta`` injection into ``rope_scaling``.
|
||||
#
|
||||
# Upstream `transformers.PretrainedConfig` now natively handles this
|
||||
# logic. While the manual injection patch has been removed, these
|
||||
# test cases are retained to ensure regression testing of the
|
||||
# configuration's injection behavior.
|
||||
# -----------------------------------------------------------------------
|
||||
|
||||
def test_injects_rope_theta_into_rope_scaling(self):
|
||||
config_dict = {
|
||||
"model_type": "llama",
|
||||
"rope_theta": 500000.0,
|
||||
"max_position_embeddings": 131072,
|
||||
"rope_scaling": {
|
||||
"rope_type": "llama3",
|
||||
"factor": 8.0,
|
||||
"low_freq_factor": 1.0,
|
||||
"high_freq_factor": 4.0,
|
||||
"original_max_position_embeddings": 8192,
|
||||
},
|
||||
}
|
||||
config = PretrainedConfig.from_dict(config_dict)
|
||||
rope_params = getattr(config, "rope_parameters", None)
|
||||
if rope_params is not None:
|
||||
self.assertIn("rope_theta", rope_params)
|
||||
|
||||
def test_no_injection_when_rope_theta_already_in_scaling(self):
|
||||
config_dict = {
|
||||
"model_type": "llama",
|
||||
"rope_theta": 500000.0,
|
||||
"max_position_embeddings": 131072,
|
||||
"rope_scaling": {
|
||||
"rope_type": "llama3",
|
||||
"factor": 8.0,
|
||||
"rope_theta": 999.0,
|
||||
"low_freq_factor": 1.0,
|
||||
"high_freq_factor": 4.0,
|
||||
"original_max_position_embeddings": 8192,
|
||||
},
|
||||
}
|
||||
config = PretrainedConfig.from_dict(config_dict)
|
||||
rope_params = getattr(config, "rope_parameters", None)
|
||||
if rope_params is not None:
|
||||
self.assertEqual(rope_params["rope_theta"], 999.0)
|
||||
|
||||
def test_no_crash_without_rope_scaling(self):
|
||||
config_dict = {"model_type": "llama", "rope_theta": 10000.0}
|
||||
config = PretrainedConfig.from_dict(config_dict)
|
||||
self.assertIsNotNone(config)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# compat: _ensure_clean_up_tokenization_compat
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -217,32 +217,6 @@ class TestProfileMergerIntegration(CustomTestCase):
|
||||
req = ProfileReq(merge_profiles=True)
|
||||
self.assertTrue(req.merge_profiles)
|
||||
|
||||
def test_integration_parameters(self):
|
||||
import inspect
|
||||
|
||||
# Test TokenizerManager
|
||||
from sglang.srt.managers.tokenizer_control_mixin import (
|
||||
TokenizerControlMixin,
|
||||
)
|
||||
|
||||
sig = inspect.signature(TokenizerControlMixin.start_profile)
|
||||
self.assertIn("req", sig.parameters)
|
||||
self.assertNotIn("merge_profiles", sig.parameters)
|
||||
|
||||
# Test SchedulerProfilerMixin
|
||||
from sglang.srt.managers.scheduler_components.profiler_manager import (
|
||||
SchedulerProfilerManager,
|
||||
)
|
||||
|
||||
sig = inspect.signature(SchedulerProfilerManager._init_profile)
|
||||
self.assertIn("merge_profiles", sig.parameters)
|
||||
|
||||
# Test CLI profiler
|
||||
from sglang.profiler import run_profile
|
||||
|
||||
sig = inspect.signature(run_profile)
|
||||
self.assertIn("merge_profiles", sig.parameters)
|
||||
|
||||
|
||||
class TestProfileMergerEdgeCases(CustomTestCase):
|
||||
def setUp(self):
|
||||
|
||||
Reference in New Issue
Block a user