已合并
【bugfix】ascendv1_save_module_preprocess 允许保存新prefix #106
zhangz200102创建于 2月3日
【bugfix】ascendv1_save_module_preprocess 允许保存新prefix #106
已合并
共 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 module | 555 | + 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 | ||
| 22 | from abc import ABC, abstractmethod | 22 | from abc import ABC, abstractmethod |
| 23 | -from typing import Optional | 23 | +from typing import Tuple |
| 24 | import torch.nn as nn | 24 | import torch.nn as nn |
| 25 | 25 | ||
| 26 | 26 | ||
| @@ -33,6 +33,12 @@ class AscendV1SaveInterface(ABC): | |||
| 33 | """ | 33 | """ |
| 34 | pass | 34 | 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 | - pass | 37 | + """ |
| 38 | - | 38 | + 在保存模块前,对模块进行预处理,返回新的前缀和模块 |
| 39 | + | ||
| 40 | + | ||
| 41 | + | ||
| 42 | + | ||
| 43 | + """ | ||
| 44 | + pass | ||
| @@ -113,10 +113,10 @@ class Qwen3NextModelAdapter(TransformersModel, | |||
| 113 | ]) | 113 | ]) |
| 114 | return adapter_config | 114 | 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 - 1 | 119 | 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_module | 121 | + return prefix, new_module |
| 122 | - return None | 122 | + return prefix, module |
| @@ -22,6 +22,7 @@ import unittest | |||
| 22 | from pathlib import Path | 22 | from pathlib import Path |
| 23 | from unittest.mock import MagicMock, patch | 23 | from unittest.mock import MagicMock, patch |
| 24 | 24 | ||
| 25 | +import torch | ||
| 25 | import torch.nn as nn | 26 | import torch.nn as nn |
| 26 | 27 | ||
| 27 | from msmodelslim.core.base.protocol import ProcessRequest | 28 | from 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-6 | 310 | 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 | - # 验证返回值不为None | 330 | + # 验证新模块的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 | - | ||
| 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 | - | ||
| 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() | ||