已合并
【bugfix】修复使用GLM-5-FP8权重量化时保存后处理报错 #226
【bugfix】修复使用GLM-5-FP8权重量化时保存后处理报错 #226
已合并
zhangz200102创建于 3月17日
共 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, None416 global_rotation, norm_weight = None, None
415 417
416 # catch the global rotation418 # 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.weight426 # 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.json461 # 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"] = True464 description_data["is_rot_used"] = True