Skip to content

Commit 2eaff29

Browse files
authored
[Misc] Fix UT error. (vllm-project#11581)
### What this PR does / why we need it? vllm-project#11053 introduce a UT error. This PR fix it. - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@1f486d9 Signed-off-by: menogrey <1299267905@qq.com>
1 parent e0d0502 commit 2eaff29

1 file changed

Lines changed: 104 additions & 11 deletions

File tree

tests/ut/quantization/test_modelslim_config.py

Lines changed: 104 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -353,22 +353,115 @@ def test_step3p5_mtp_maps_direct_and_step3p7_wrapped_quant_keys(self):
353353
self.assertEqual(prefix, expected)
354354

355355

356-
class TestGetCacheScale(TestBase):
357-
def test_c8_kv_cache_type_k_proj_scale(self):
356+
class TestGetCacheScaleMapper(TestBase):
357+
def test_c8_kv_cache_type_returns_mapper(self):
358358
config = AscendModelSlimConfig({"kv_cache_type": "C8"})
359-
result = config.get_cache_scale("model.layers.0.k_proj.kv_cache_scale")
360-
self.assertEqual(result, "model.layers.0.attn.k_cache_scale")
361-
result = config.get_cache_scale("model.layers.0.v_proj.kv_cache_offset")
362-
self.assertEqual(result, "model.layers.0.attn.v_cache_offset")
359+
mapper = config.get_cache_scale_mapper()
360+
self.assertIsNotNone(mapper)
361+
# C8 mappings: k_proj → attn
362+
self.assertEqual(
363+
mapper._map_name("model.layers.0.k_proj.kv_cache_scale"),
364+
"model.layers.0.attn.k_cache_scale",
365+
)
366+
self.assertEqual(
367+
mapper._map_name("model.layers.0.k_proj.kv_cache_offset"),
368+
"model.layers.0.attn.k_cache_offset",
369+
)
370+
self.assertEqual(
371+
mapper._map_name("model.layers.0.v_proj.kv_cache_scale"),
372+
"model.layers.0.attn.v_cache_scale",
373+
)
374+
self.assertEqual(
375+
mapper._map_name("model.layers.0.v_proj.kv_cache_offset"),
376+
"model.layers.0.attn.v_cache_offset",
377+
)
363378

364-
def test_no_match(self):
379+
def test_no_match_returns_none(self):
365380
config = AscendModelSlimConfig({"kv_cache_type": "FLOAT"})
366-
result = config.get_cache_scale("model.layers.0.k_proj.kv_cache_scale")
367-
self.assertIsNone(result)
381+
mapper = config.get_cache_scale_mapper()
382+
self.assertIsNone(mapper)
383+
384+
config = AscendModelSlimConfig({})
385+
mapper = config.get_cache_scale_mapper()
386+
self.assertIsNone(mapper)
387+
388+
def test_fa_quant_returns_mapper(self):
389+
config = AscendModelSlimConfig(
390+
{
391+
"fa_quant_type": "C8",
392+
"layers.1.fa_k.scale": "C8",
393+
}
394+
)
395+
mapper = config.get_cache_scale_mapper()
396+
self.assertIsNotNone(mapper)
397+
self.assertEqual(
398+
mapper._map_name("model.layers.1.fa_k.scale"),
399+
"model.layers.1.mla_attn.mla_attn.fa_k.scale",
400+
)
401+
self.assertEqual(
402+
mapper._map_name("model.layers.1.fa_q.scale"),
403+
"model.layers.1.mla_attn.mla_attn.fa_q.scale",
404+
)
405+
self.assertEqual(
406+
mapper._map_name("model.layers.1.fa_v.offset"),
407+
"model.layers.1.mla_attn.mla_attn.fa_v.offset",
408+
)
409+
410+
def test_indexer_quant_returns_mapper(self):
411+
config = AscendModelSlimConfig(
412+
{
413+
"indexer_quant_type": "INT8",
414+
"layers.1.indexer.quant_type": "INT8",
415+
}
416+
)
417+
mapper = config.get_cache_scale_mapper()
418+
self.assertIsNotNone(mapper)
419+
self.assertEqual(
420+
mapper._map_name("model.layers.1.indexer.q_rot"),
421+
"model.layers.1.mla_attn.mla_attn.indexer.q_rot",
422+
)
423+
self.assertEqual(
424+
mapper._map_name("model.layers.1.indexer.k_rot"),
425+
"model.layers.1.mla_attn.mla_attn.indexer.k_rot",
426+
)
427+
428+
def test_combined_quant_types_returns_mapper_with_all_suffixes(self):
429+
config = AscendModelSlimConfig(
430+
{
431+
"kv_cache_type": "C8",
432+
"fa_quant_type": "C8",
433+
"layers.1.fa_k.scale": "C8",
434+
"indexer_quant_type": "INT8",
435+
"layers.2.indexer.quant_type": "INT8",
436+
}
437+
)
438+
mapper = config.get_cache_scale_mapper()
439+
self.assertIsNotNone(mapper)
440+
# C8 mapping
441+
self.assertEqual(
442+
mapper._map_name("model.layers.0.k_proj.kv_cache_scale"),
443+
"model.layers.0.attn.k_cache_scale",
444+
)
445+
# FA quant mapping
446+
self.assertEqual(
447+
mapper._map_name("model.layers.1.fa_k.scale"),
448+
"model.layers.1.mla_attn.mla_attn.fa_k.scale",
449+
)
450+
# Indexer quant mapping
451+
self.assertEqual(
452+
mapper._map_name("model.layers.2.indexer.k_rot"),
453+
"model.layers.2.mla_attn.mla_attn.indexer.k_rot",
454+
)
368455

456+
def test_unmapped_key_returns_unchanged(self):
369457
config = AscendModelSlimConfig({"kv_cache_type": "C8"})
370-
result = config.get_cache_scale("model.layers.0.other_key")
371-
self.assertIsNone(result)
458+
mapper = config.get_cache_scale_mapper()
459+
self.assertIsNotNone(mapper)
460+
# C8 mapper doesn't cover FA quant keys
461+
self.assertEqual(
462+
mapper._map_name("model.layers.0.other_key"),
463+
"model.layers.0.other_key",
464+
)
372465

373466

374467
class TestGetKvQuantDtype(TestBase):

0 commit comments

Comments
 (0)