已合并
【bugfix】ascendv1_save_module_preprocess 允许保存新prefix #106
【bugfix】ascendv1_save_module_preprocess 允许保存新prefix #106
已合并
zhangz200102创建于 2月3日
共 4 个文件变更+38-64
@@ -552,5 +552,5 @@ class AscendV1Saver(AutoSaverProcessor):
552 def _process_module(self, prefix: str, module: nn.Module):552 def _process_module(self, prefix: str, module: nn.Module):
553 if isinstance(self.adapter, AscendV1SaveInterface):553 if isinstance(self.adapter, AscendV1SaveInterface):
554 self.processed_modules.add(module)554 self.processed_modules.add(module)
555- module = self.adapter.ascendv1_save_module_preprocess(prefix, module, self.model) or module555+ prefix, module = self.adapter.ascendv1_save_module_preprocess(prefix, module, self.model)
556 super()._process_module(prefix=prefix, module=module)556 super()._process_module(prefix=prefix, module=module)
@@ -20,7 +20,7 @@ See the Mulan PSL v2 for more details.
20"""20"""
21 21 
22from abc import ABC, abstractmethod22from abc import ABC, abstractmethod
23-from typing import Optional23+from typing import Tuple
24import torch.nn as nn24import torch.nn as nn
25 25 
26 26 
@@ -33,6 +33,12 @@ class AscendV1SaveInterface(ABC):
33 """33 """
34 pass34 pass
35 35 
36- def ascendv1_save_module_preprocess(self, prefix: str, module: nn.Module, model: nn.Module) -> Optional[nn.Module]:36+ def ascendv1_save_module_preprocess(self, prefix: str, module: nn.Module, model: nn.Module) -> Tuple[str, nn.Module]:
37- pass37+ """
38- 38+ 在保存模块前,对模块进行预处理,返回新的前缀和模块
39+ @param prefix: 模块的前缀路径
40+ @param module: 待处理的模块
41+ @param model: 模型
42+ @return: 返回(prefix, module)
43+ """
44+ pass
@@ -113,10 +113,10 @@ class Qwen3NextModelAdapter(TransformersModel,
113 ])113 ])
114 return adapter_config114 return adapter_config
115 115 
116- def ascendv1_save_module_preprocess(self, prefix: str, module: nn.Module, model: nn.Module) -> Optional[nn.Module]:116+ def ascendv1_save_module_preprocess(self, prefix: str, module: nn.Module, model: nn.Module) -> Tuple[str, nn.Module]:
117 if 'input_layernorm' in prefix and module.__class__.__name__ == 'Qwen3RMSNorm':117 if 'input_layernorm' in prefix and module.__class__.__name__ == 'Qwen3RMSNorm':
118 new_module = Qwen3NextRMSNorm(module.weight.shape[0], module.variance_epsilon)118 new_module = Qwen3NextRMSNorm(module.weight.shape[0], module.variance_epsilon)
119 new_module.weight.data = module.weight.data - 1119 new_module.weight.data = module.weight.data - 1
120 model.set_submodule(prefix, new_module)120 model.set_submodule(prefix, new_module)
121- return new_module121+ return prefix, new_module
122- return None122+ return prefix, module
@@ -22,6 +22,7 @@ import unittest
22from pathlib import Path22from pathlib import Path
23from unittest.mock import MagicMock, patch23from unittest.mock import MagicMock, patch
24 24 
25+import torch
25import torch.nn as nn26import torch.nn as nn
26 27 
27from msmodelslim.core.base.protocol import ProcessRequest28from msmodelslim.core.base.protocol import ProcessRequest
@@ -298,66 +299,33 @@ class TestQwen3NextModelAdapter(unittest.TestCase):
298 )299 )
299 300 
300 # 创建模拟的Qwen3RMSNorm模块301 # 创建模拟的Qwen3RMSNorm模块
302+ test_prefix = "model.layers.0.input_layernorm"
303+ original_weight_data = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0])
304+ expected_weight_data = original_weight_data - 1 # 期望的weight值:原始值 - 1
305+
301 mock_module = MagicMock()306 mock_module = MagicMock()
302 mock_module.__class__.__name__ = 'Qwen3RMSNorm'307 mock_module.__class__.__name__ = 'Qwen3RMSNorm'
303 mock_module.weight = MagicMock()308 mock_module.weight = MagicMock()
304- mock_module.weight.shape = [100]309+ mock_module.weight.shape = [5]
305 mock_module.variance_epsilon = 1e-6310 mock_module.variance_epsilon = 1e-6
306- mock_module.weight.data = MagicMock()311+ mock_module.weight.data = original_weight_data.clone()
307- mock_module.weight.data.__sub__ = MagicMock(return_value=MagicMock())
308- mock_module.weight.data.__radd__ = MagicMock(return_value=MagicMock())
309 312 
310 mock_model = MagicMock()313 mock_model = MagicMock()
314+
315+ # Mock Qwen3NextRMSNorm 的创建,使用真实的nn.Module来存储weight
316+ with patch('msmodelslim.model.qwen3_next.model_adapter.Qwen3NextRMSNorm') as mock_qwen3_next_rms_norm:
317+ # 创建一个真实的新模块来存储weight
318+ class MockNewModule(nn.Module):
319+ def __init__(self):
320+ super().__init__()
321+ self.weight = nn.Parameter(torch.zeros(5))
322+
323+ mock_new_module = MockNewModule()
324+ mock_qwen3_next_rms_norm.return_value = mock_new_module
325+
326+ new_prefix, new_module = adapter.ascendv1_save_module_preprocess(test_prefix, mock_module, mock_model)
311 327 
312- result = adapter.ascendv1_save_module_preprocess("model.layers.0.input_layernorm", mock_module, mock_model)328+ # 验证prefix没有变化
313- 329+ self.assertEqual(new_prefix, test_prefix)
314- # 验证返回值不为None330+ # 验证新模块的weight.data是原始weight.data - 1
315- self.assertIsNotNone(result)331+ self.assertTrue(torch.allclose(new_module.weight.data, expected_weight_data))
316- # 验证model.set_submodule被调用
317- mock_model.set_submodule.assert_called_once()
318- 
319- @patch.dict('sys.modules', {
320- 'transformers.models.qwen3_next.modeling_qwen3_next': MagicMock(),
321- })
322- def test_ascendv1_save_module_preprocess_without_input_layernorm(self):
323- """测试ascendv1_save_module_preprocess方法当prefix不包含input_layernorm时"""
324- from msmodelslim.model.qwen3_next.model_adapter import Qwen3NextModelAdapter
325- with patch('msmodelslim.model.qwen3_next.model_adapter.TransformersModel.__init__', return_value=None):
326- adapter = Qwen3NextModelAdapter(
327- model_type=self.model_type,
328- model_path=self.model_path
329- )
330- 
331- mock_module = MagicMock()
332- mock_model = MagicMock()
333- 
334- result = adapter.ascendv1_save_module_preprocess("model.layers.0.some_other_module", mock_module,
335- mock_model)
336- 
337- # 验证返回值为None
338- self.assertIsNone(result)
339- # 验证model.set_submodule没有被调用
340- mock_model.set_submodule.assert_not_called()
341- 
342- @patch.dict('sys.modules', {
343- 'transformers.models.qwen3_next.modeling_qwen3_next': MagicMock(),
344- })
345- def test_ascendv1_save_module_preprocess_with_wrong_module_type(self):
346- """测试ascendv1_save_module_preprocess方法当module不是Qwen3RMSNorm时"""
347- from msmodelslim.model.qwen3_next.model_adapter import Qwen3NextModelAdapter
348- with patch('msmodelslim.model.qwen3_next.model_adapter.TransformersModel.__init__', return_value=None):
349- adapter = Qwen3NextModelAdapter(
350- model_type=self.model_type,
351- model_path=self.model_path
352- )
353- 
354- mock_module = MagicMock()
355- mock_module.__class__.__name__ = 'SomeOtherModule'
356- mock_model = MagicMock()
357- 
358- result = adapter.ascendv1_save_module_preprocess("model.layers.0.input_layernorm", mock_module, mock_model)
359- 
360- # 验证返回值为None
361- self.assertIsNone(result)
362- # 验证model.set_submodule没有被调用
363- mock_model.set_submodule.assert_not_called()