已合并
[Bugfix] fix convert part_file_size not read from YAML config #737
[Bugfix] fix convert part_file_size not read from YAML config #737
已合并
caishengcheng创建于 7月15日
5 个文件变更+91-4
Mmsmodelslim/core/convert/config.py+9-0
@@ -121,6 +121,8 @@ class ConvertConfig(BaseModel):
121 save_path: str121 save_path: str
122 model_family: Optional[str] = None122 model_family: Optional[str] = None
123 dst_format: str = "ascendv1"123 dst_format: str = "ascendv1"
124+ # 分片文件大小(GB);0 表示不分片。来自 YAML ``spec.save[].part_file_size``。
125+ part_file_size: int = 4
124 defaults: ConvertDefaults = Field(default_factory=ConvertDefaults)126 defaults: ConvertDefaults = Field(default_factory=ConvertDefaults)
125 preprocess_rules: List[WeightMappingRule] = Field(default_factory=list)127 preprocess_rules: List[WeightMappingRule] = Field(default_factory=list)
126 module_rules: List[ModuleRule] = Field(default_factory=list)128 module_rules: List[ModuleRule] = Field(default_factory=list)
@@ -134,6 +136,13 @@ class ConvertConfig(BaseModel):
134 raise ValueError("path must be non-empty")136 raise ValueError("path must be non-empty")
135 return v137 return v
136 138 
139+ @field_validator("part_file_size", mode="after")
140+ @classmethod
141+ def _non_negative_part_file_size(cls, v: int) -> int:
142+ if v < 0:
143+ raise ValueError("part_file_size must be >= 0 (0 means no sharding)")
144+ return v
145+ 
137 @model_validator(mode="after")146 @model_validator(mode="after")
138 def _mxfp8_requires_ascendv1(self) -> ConvertConfig:147 def _mxfp8_requires_ascendv1(self) -> ConvertConfig:
139 """任意 convert_rule 目标为 W8A8_MXFP8 时,dst_format 必须为 ascendv1。"""148 """任意 convert_rule 目标为 W8A8_MXFP8 时,dst_format 必须为 ascendv1。"""
Mmsmodelslim/core/quant_service/modelslim_convert/config_mapper.py+8-0
@@ -252,6 +252,13 @@ def _resolve_dst_format(save: List[SaveConfig], defaults: ConvertDefaults) -> st
252 return defaults.dst_format252 return defaults.dst_format
253 253 
254 254 
255+def _resolve_part_file_size(save: List[SaveConfig]) -> int:
256+ """从 YAML ``spec.save[0].part_file_size`` 读取分片大小;未配置时默认 4GB。"""
257+ if save:
258+ return save[0].part_file_size
259+ return 4
260+ 
261+ 
255def spec_to_convert_config(262def spec_to_convert_config(
256 spec: Union[ModelslimConvertServiceConfig, Dict[str, Any]],263 spec: Union[ModelslimConvertServiceConfig, Dict[str, Any]],
257 model_path: str,264 model_path: str,
@@ -279,6 +286,7 @@ def spec_to_convert_config(
279 save_path=save_path,286 save_path=save_path,
280 model_family=model_family,287 model_family=model_family,
281 dst_format=_resolve_dst_format(spec.save, spec.defaults),288 dst_format=_resolve_dst_format(spec.save, spec.defaults),
289+ part_file_size=_resolve_part_file_size(spec.save),
282 defaults=spec.defaults,290 defaults=spec.defaults,
283 preprocess_rules=_preprocess_to_rules(spec),291 preprocess_rules=_preprocess_to_rules(spec),
284 module_rules=module_rules,292 module_rules=module_rules,
Mmsmodelslim/core/quant_service/modelslim_convert/impl/save_adapter.py+14-4
@@ -83,7 +83,8 @@ class SaveProcessorAdapter(ISaveProcessorAdapter):
83 save_dir: str,83 save_dir: str,
84 adapter: IModel,84 adapter: IModel,
85 ) -> None:85 ) -> None:
86- format_cfg = parse_format_config({"type": "compressed_tensors", "part_file_size": 4})86+ part_file_size = context.config.part_file_size
87+ format_cfg = parse_format_config({"type": "compressed_tensors", "part_file_size": part_file_size})
87 cfg = QuantSaveProcessorConfig(type="saver", format=format_cfg)88 cfg = QuantSaveProcessorConfig(type="saver", format=format_cfg)
88 cfg.set_save_directory(save_dir)89 cfg.set_save_directory(save_dir)
89 _lazy_init_unsaved_modules(context, tree)90 _lazy_init_unsaved_modules(context, tree)
@@ -95,7 +96,11 @@ class SaveProcessorAdapter(ISaveProcessorAdapter):
95 if name:96 if name:
96 saver.postprocess(BatchProcessRequest(name=name, module=module, datas=None, outputs=None))97 saver.postprocess(BatchProcessRequest(name=name, module=module, datas=None, outputs=None))
97 saver.post_run()98 saver.post_run()
98- logger.info("Saved HF/compressed_tensors checkpoint to %s", save_dir)99+ logger.info(
100+ "Saved HF/compressed_tensors checkpoint to %s (part_file_size=%s)",
101+ save_dir,
102+ part_file_size,
103+ )
99 104 
100 @staticmethod105 @staticmethod
101 def _save_ascendv1(106 def _save_ascendv1(
@@ -106,9 +111,14 @@ class SaveProcessorAdapter(ISaveProcessorAdapter):
106 ) -> None:111 ) -> None:
107 from msmodelslim.core.quant_service.modelslim_v1.save.ascendv1 import AscendV1Config, AscendV1Saver112 from msmodelslim.core.quant_service.modelslim_v1.save.ascendv1 import AscendV1Config, AscendV1Saver
108 113 
114+ part_file_size = context.config.part_file_size
109 _lazy_init_unsaved_modules(context, tree)115 _lazy_init_unsaved_modules(context, tree)
110- cfg = AscendV1Config(save_directory=save_dir, part_file_size=4)116+ cfg = AscendV1Config(save_directory=save_dir, part_file_size=part_file_size)
111 saver = AscendV1Saver(model=tree, config=cfg, adapter=adapter)117 saver = AscendV1Saver(model=tree, config=cfg, adapter=adapter)
112 saver.pre_run()118 saver.pre_run()
113 saver.post_run()119 saver.post_run()
114- logger.info("Saved AscendV1 checkpoint to %s", save_dir)120+ logger.info(
121+ "Saved AscendV1 checkpoint to %s (part_file_size=%s)",
122+ save_dir,
123+ part_file_size,
124+ )
Mtest/cases/core/quant_service/modelslim_convert/impl/test_save_adapter.py+43-0
@@ -72,3 +72,46 @@ class TestSaveProcessorAdapter:
72 cfg = mock_cls.call_args[0][1]72 cfg = mock_cls.call_args[0][1]
73 assert isinstance(cfg.format, CompressedTensorsQuantFormatConfig)73 assert isinstance(cfg.format, CompressedTensorsQuantFormatConfig)
74 assert cfg.format.part_file_size == 474 assert cfg.format.part_file_size == 4
75+ 
76+ def test_save_compressed_tensors_use_part_file_size_when_config_set(self):
77+ from msmodelslim.format.compressed_tensors_format.compressed_tensors import (
78+ CompressedTensorsQuantFormatConfig,
79+ )
80+ 
81+ tree = nn.Module()
82+ config = ConvertConfig(model_path="/m", save_path="/out", dst_format="huggingface", part_file_size=0)
83+ context = ConvertContext(config=config)
84+ context.reader = MagicMock()
85+ 
86+ with (
87+ patch.object(QuantSaveProcessor, "pre_run"),
88+ patch.object(QuantSaveProcessor, "postprocess"),
89+ patch.object(QuantSaveProcessor, "post_run"),
90+ patch(
91+ "msmodelslim.core.quant_service.modelslim_convert.impl.save_adapter.QuantSaveProcessor",
92+ ) as mock_cls,
93+ ):
94+ SaveProcessorAdapter().save(context, tree)
95+ cfg = mock_cls.call_args[0][1]
96+ assert isinstance(cfg.format, CompressedTensorsQuantFormatConfig)
97+ assert cfg.format.part_file_size == 0
98+ 
99+ def test_save_ascendv1_use_part_file_size_when_config_set(self):
100+ tree = nn.Module()
101+ config = ConvertConfig(model_path="/m", save_path="/out", dst_format="ascendv1", part_file_size=8)
102+ context = ConvertContext(config=config)
103+ context.reader = MagicMock()
104+ 
105+ with (
106+ patch("msmodelslim.core.quant_service.modelslim_v1.save.ascendv1.AscendV1Saver") as mock_saver_cls,
107+ patch(
108+ "msmodelslim.core.quant_service.modelslim_convert.impl.save_adapter._lazy_init_unsaved_modules",
109+ ),
110+ ):
111+ mock_saver = MagicMock()
112+ mock_saver_cls.return_value = mock_saver
113+ SaveProcessorAdapter().save(context, tree)
114+ saver_cfg = mock_saver_cls.call_args.kwargs["config"]
115+ assert saver_cfg.part_file_size == 8
116+ mock_saver.pre_run.assert_called_once()
117+ mock_saver.post_run.assert_called_once()
Mtest/cases/core/quant_service/modelslim_convert/test_config_mapper.py+17-0
@@ -74,9 +74,26 @@ class TestSpecToConvertConfig:
74 assert len(cfg.convert_rules) == 274 assert len(cfg.convert_rules) == 2
75 assert cfg.convert_rules[0].target_ir == IRKind.W8A8_MXFP875 assert cfg.convert_rules[0].target_ir == IRKind.W8A8_MXFP8
76 assert cfg.dst_format == "ascendv1"76 assert cfg.dst_format == "ascendv1"
77+ assert cfg.part_file_size == 4
77 assert cfg.parallel.max_workers == 878 assert cfg.parallel.max_workers == 8
78 assert cfg.model_family == "qwen3_5_moe"79 assert cfg.model_family == "qwen3_5_moe"
79 80 
81+ def test_spec_to_convert_config_map_part_file_size_when_save_given(self):
82+ spec = ModelslimConvertServiceConfig.model_validate(
83+ {
84+ "linears": [],
85+ "save": [{"type": "huggingface", "part_file_size": 0}],
86+ }
87+ )
88+ cfg = spec_to_convert_config(spec, model_path="/m", save_path="/o")
89+ assert cfg.dst_format == "huggingface"
90+ assert cfg.part_file_size == 0
91+ 
92+ def test_spec_to_convert_config_default_part_file_size_when_save_empty(self):
93+ spec = ModelslimConvertServiceConfig.model_validate({"linears": []})
94+ cfg = spec_to_convert_config(spec, model_path="/m", save_path="/o")
95+ assert cfg.part_file_size == 4
96+ 
80 def test_spec_to_convert_config_auto_route_infer_source_ir_from_catalog_later(self):97 def test_spec_to_convert_config_auto_route_infer_source_ir_from_catalog_later(self):
81 spec = ModelslimConvertServiceConfig.model_validate(98 spec = ModelslimConvertServiceConfig.model_validate(
82 {99 {