@@ -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
374467class TestGetKvQuantDtype (TestBase ):
0 commit comments