已合并
【bugfix】修复使用GLM-5-FP8权重量化时保存后处理报错 #226
zhangz200102创建于 3月17日
【bugfix】修复使用GLM-5-FP8权重量化时保存后处理报错 #226
已合并
共 1 个文件变更+6-2
| @@ -411,6 +411,8 @@ class GLM5ModelAdapter(TransformersModel, | |||
| 411 | return [pre_run], [pair for pair in rot_pairs.values()] | 411 | return [pre_run], [pair for pair in rot_pairs.values()] |
| 412 | 412 | ||
| 413 | def ascendv1_save_postprocess(self, model: nn.Module, save_directory: str) -> None: | 413 | def ascendv1_save_postprocess(self, model: nn.Module, save_directory: str) -> None: |
| 414 | + from msmodelslim.utils.security import json_safe_dump | ||
| 415 | + | ||
| 414 | global_rotation, norm_weight = None, None | 416 | global_rotation, norm_weight = None, None |
| 415 | 417 | ||
| 416 | # catch the global rotation | 418 | # catch the global rotation |
| @@ -422,7 +424,10 @@ class GLM5ModelAdapter(TransformersModel, | |||
| 422 | raise UnsupportedError("Global rotation is not found.") | 424 | raise UnsupportedError("Global rotation is not found.") |
| 423 | 425 | ||
| 424 | # catch the original model.norm.weight | 426 | # catch the original model.norm.weight |
| 425 | - weight_path = os.path.join(self.model_path, "model-00282-of-00282.safetensors") | 427 | + origin_index_path = os.path.join(self.model_path, "model.safetensors.index.json") |
| 428 | + origin_index_data = json_safe_load(origin_index_path) | ||
| 429 | + | ||
| 430 | + weight_path = os.path.join(self.model_path, origin_index_data["weight_map"]["model.norm.weight"]) | ||
| 426 | with safe_open(weight_path, framework='pt', device='cpu') as f: | 431 | with safe_open(weight_path, framework='pt', device='cpu') as f: |
| 427 | norm_weight = f.get_tensor("model.norm.weight") | 432 | norm_weight = f.get_tensor("model.norm.weight") |
| 428 | if norm_weight is None: | 433 | if norm_weight is None: |
| @@ -454,7 +459,6 @@ class GLM5ModelAdapter(TransformersModel, | |||
| 454 | save_file({"rot.weight": rot_weight}, os.path.join(save_directory, "rot.safetensors")) | 459 | save_file({"rot.weight": rot_weight}, os.path.join(save_directory, "rot.safetensors")) |
| 455 | 460 | ||
| 456 | # update quant_model_description.json | 461 | # update quant_model_description.json |
| 457 | - from msmodelslim.utils.security import json_safe_dump | ||
| 458 | description_path = os.path.join(save_directory, "quant_model_description.json") | 462 | description_path = os.path.join(save_directory, "quant_model_description.json") |
| 459 | description_data = json_safe_load(description_path) | 463 | description_data = json_safe_load(description_path) |
| 460 | description_data["is_rot_used"] = True | 464 | description_data["is_rot_used"] = True |