已合并
[pytorch][bugfix]fix lora_target_modules in ckpt save_lora_to_hf #4016
温一盏创建于 1月6日
[pytorch][bugfix]fix lora_target_modules in ckpt save_lora_to_hf #4016
已合并
共 5 个文件变更+61-39
| @@ -79,7 +79,9 @@ def main(): | |||
| 79 | parser.add_argument('--ckpt-format', default='torch', | 79 | parser.add_argument('--ckpt-format', default='torch', |
| 80 | choices=['torch', 'torch_dist', 'zarr'], | 80 | choices=['torch', 'torch_dist', 'zarr'], |
| 81 | help='Checkpoint format to use.') | 81 | help='Checkpoint format to use.') |
| 82 | - | 82 | + parser.add_argument('--lora-target-modules', nargs='+', type=str, default=[], |
| 83 | + help='Lora target modules.') | ||
| 84 | + | ||
| 83 | known_args, _ = parser.parse_known_args() | 85 | known_args, _ = parser.parse_known_args() |
| 84 | 86 | ||
| 85 | use_saver = known_args.load_model_type is None | 87 | use_saver = known_args.load_model_type is None |
| @@ -178,8 +178,9 @@ def get_message_layer_attn(message, model, md=None, **kwargs): | |||
| 178 | qb_weight.append(model.get_layers_self_attention_linear_q_up_proj_weight(**kwargs)) | 178 | qb_weight.append(model.get_layers_self_attention_linear_q_up_proj_weight(**kwargs)) |
| 179 | kvb_weight.append(model.get_layers_self_attention_linear_kv_up_proj_weight(**kwargs)) | 179 | kvb_weight.append(model.get_layers_self_attention_linear_kv_up_proj_weight(**kwargs)) |
| 180 | 180 | ||
| 181 | - if margs.save_lora_to_hf: | 181 | + if margs.save_lora_to_hf and 'linear_proj' in margs.lora_target_modules: |
| 182 | proj_lora_A_weight.append(model.get_layers_self_attention_linear_proj_lora_A_default_weight(**kwargs)) | 182 | proj_lora_A_weight.append(model.get_layers_self_attention_linear_proj_lora_A_default_weight(**kwargs)) |
| 183 | + if margs.save_lora_to_hf and 'linear_qkv' in margs.lora_target_modules: | ||
| 183 | qkv_lora_B_weight.append(model.get_layers_self_attention_linear_qkv_lora_B_default_weight(**kwargs)) | 184 | qkv_lora_B_weight.append(model.get_layers_self_attention_linear_qkv_lora_B_default_weight(**kwargs)) |
| 184 | 185 | ||
| 185 | # Handle gated linear units | 186 | # Handle gated linear units |
| @@ -204,9 +205,10 @@ def get_message_layer_attn(message, model, md=None, **kwargs): | |||
| 204 | if md.linear_bias or margs.add_dense_bias: | 205 | if md.linear_bias or margs.add_dense_bias: |
| 205 | message["dense bias"] = model.get_layers_self_attention_linear_proj_bias(**kwargs) | 206 | message["dense bias"] = model.get_layers_self_attention_linear_proj_bias(**kwargs) |
| 206 | 207 | ||
| 207 | - if margs.save_lora_to_hf: | 208 | + if margs.save_lora_to_hf and 'linear_proj' in margs.lora_target_modules: |
| 208 | message["proj lora A"] = torch.cat(proj_lora_A_weight, dim=1) | 209 | message["proj lora A"] = torch.cat(proj_lora_A_weight, dim=1) |
| 209 | message["proj lora B"] = model.get_layers_self_attention_linear_proj_lora_B_default_weight(**kwargs) | 210 | message["proj lora B"] = model.get_layers_self_attention_linear_proj_lora_B_default_weight(**kwargs) |
| 211 | + if margs.save_lora_to_hf and 'linear_qkv' in margs.lora_target_modules: | ||
| 210 | message["qkv lora A"] = model.get_layers_self_attention_linear_qkv_lora_A_default_weight(**kwargs) | 212 | message["qkv lora A"] = model.get_layers_self_attention_linear_qkv_lora_A_default_weight(**kwargs) |
| 211 | message["qkv lora B"] = torch.cat(qkv_lora_B_weight, dim=0) | 213 | message["qkv lora B"] = torch.cat(qkv_lora_B_weight, dim=0) |
| 212 | 214 | ||
| @@ -225,17 +227,19 @@ def _get_message_layer_mlp(message, model, md=None, is_moe_mlp=False, **kwargs): | |||
| 225 | if is_moe_mlp: | 227 | if is_moe_mlp: |
| 226 | mlp_l0_weight.append(model.get_layers_mlp_experts_linear_fc1_weight(**kwargs)) | 228 | mlp_l0_weight.append(model.get_layers_mlp_experts_linear_fc1_weight(**kwargs)) |
| 227 | mlp_l1_weight.append(model.get_layers_mlp_experts_linear_fc2_weight(**kwargs)) | 229 | mlp_l1_weight.append(model.get_layers_mlp_experts_linear_fc2_weight(**kwargs)) |
| 228 | - if margs.save_lora_to_hf: | 230 | + if margs.save_lora_to_hf and 'linear_fc1' in margs.lora_target_modules: |
| 229 | fc1_lora_A = model.get_layers_mlp_experts_linear_fc1_lora_A_default_weight(**kwargs) | 231 | fc1_lora_A = model.get_layers_mlp_experts_linear_fc1_lora_A_default_weight(**kwargs) |
| 230 | fc1_lora_B.append(model.get_layers_mlp_experts_linear_fc1_lora_B_default_weight(**kwargs)) | 232 | fc1_lora_B.append(model.get_layers_mlp_experts_linear_fc1_lora_B_default_weight(**kwargs)) |
| 233 | + if margs.save_lora_to_hf and 'linear_fc2' in margs.lora_target_modules: | ||
| 231 | fc2_lora_A.append(model.get_layers_mlp_experts_linear_fc2_lora_A_default_weight(**kwargs)) | 234 | fc2_lora_A.append(model.get_layers_mlp_experts_linear_fc2_lora_A_default_weight(**kwargs)) |
| 232 | fc2_lora_B = model.get_layers_mlp_experts_linear_fc2_lora_B_default_weight(**kwargs) | 235 | fc2_lora_B = model.get_layers_mlp_experts_linear_fc2_lora_B_default_weight(**kwargs) |
| 233 | else: | 236 | else: |
| 234 | mlp_l0_weight.append(model.get_layers_mlp_linear_fc1_weight(**kwargs)) | 237 | mlp_l0_weight.append(model.get_layers_mlp_linear_fc1_weight(**kwargs)) |
| 235 | mlp_l1_weight.append(model.get_layers_mlp_linear_fc2_weight(**kwargs)) | 238 | mlp_l1_weight.append(model.get_layers_mlp_linear_fc2_weight(**kwargs)) |
| 236 | - if margs.save_lora_to_hf: | 239 | + if margs.save_lora_to_hf and 'linear_fc1' in margs.lora_target_modules: |
| 237 | fc1_lora_A = model.get_layers_mlp_linear_fc1_lora_A_default_weight(**kwargs) | 240 | fc1_lora_A = model.get_layers_mlp_linear_fc1_lora_A_default_weight(**kwargs) |
| 238 | fc1_lora_B.append(model.get_layers_mlp_linear_fc1_lora_B_default_weight(**kwargs)) | 241 | fc1_lora_B.append(model.get_layers_mlp_linear_fc1_lora_B_default_weight(**kwargs)) |
| 242 | + if margs.save_lora_to_hf and 'linear_fc2' in margs.lora_target_modules: | ||
| 239 | fc2_lora_A.append(model.get_layers_mlp_linear_fc2_lora_A_default_weight(**kwargs)) | 243 | fc2_lora_A.append(model.get_layers_mlp_linear_fc2_lora_A_default_weight(**kwargs)) |
| 240 | fc2_lora_B = model.get_layers_mlp_linear_fc2_lora_B_default_weight(**kwargs) | 244 | fc2_lora_B = model.get_layers_mlp_linear_fc2_lora_B_default_weight(**kwargs) |
| 241 | if md.linear_bias: | 245 | if md.linear_bias: |
| @@ -265,9 +269,10 @@ def _get_message_layer_mlp(message, model, md=None, is_moe_mlp=False, **kwargs): | |||
| 265 | else: | 269 | else: |
| 266 | message[f"mlp l0 bias"] = torch.cat(mlp_l0_bias, dim=0) | 270 | message[f"mlp l0 bias"] = torch.cat(mlp_l0_bias, dim=0) |
| 267 | message[f"mlp l1 bias"] = model.get_layers_mlp_linear_fc2_bias(**kwargs) | 271 | message[f"mlp l1 bias"] = model.get_layers_mlp_linear_fc2_bias(**kwargs) |
| 268 | - if margs.save_lora_to_hf: | 272 | + if margs.save_lora_to_hf and 'linear_fc1' in margs.lora_target_modules: |
| 269 | message[f"fc1 lora A"] = fc1_lora_A | 273 | message[f"fc1 lora A"] = fc1_lora_A |
| 270 | message[f"fc1 lora B"] = torch.cat(fc1_lora_B, dim=0) | 274 | message[f"fc1 lora B"] = torch.cat(fc1_lora_B, dim=0) |
| 275 | + if margs.save_lora_to_hf and 'linear_fc2' in margs.lora_target_modules: | ||
| 271 | message[f"fc2 lora A"] = torch.cat(fc2_lora_A, dim=1) | 276 | message[f"fc2 lora A"] = torch.cat(fc2_lora_A, dim=1) |
| 272 | message[f"fc2 lora B"] = fc2_lora_B | 277 | message[f"fc2 lora B"] = fc2_lora_B |
| 273 | 278 | ||
| @@ -257,14 +257,16 @@ class ModelBase(abc.ABC): | |||
| 257 | def set_attn_state(self, src_layer_idx, dst_layer_idx, src_model): | 257 | def set_attn_state(self, src_layer_idx, dst_layer_idx, src_model): |
| 258 | """Set self-attention params.""" | 258 | """Set self-attention params.""" |
| 259 | if self.args.save_lora_to_hf: | 259 | if self.args.save_lora_to_hf: |
| 260 | - qkv_lora_A_weight = src_model.get_layers_self_attention_linear_qkv_lora_A_default_weight(layer_idx=src_layer_idx) | 260 | + if 'linear_qkv' in self.args.lora_target_modules: |
| 261 | - qkv_lora_B_weight = src_model.get_layers_self_attention_linear_qkv_lora_B_default_weight(layer_idx=src_layer_idx) | 261 | + qkv_lora_A_weight = src_model.get_layers_self_attention_linear_qkv_lora_A_default_weight(layer_idx=src_layer_idx) |
| 262 | - proj_lora_A_weight = src_model.get_layers_self_attention_linear_proj_lora_A_default_weight(layer_idx=src_layer_idx) | 262 | + qkv_lora_B_weight = src_model.get_layers_self_attention_linear_qkv_lora_B_default_weight(layer_idx=src_layer_idx) |
| 263 | - proj_lora_B_weight = src_model.get_layers_self_attention_linear_proj_lora_B_default_weight(layer_idx=src_layer_idx) | 263 | + self.set_layers_self_attention_linear_qkv_lora_A_default_weight(layer_idx=dst_layer_idx, data=qkv_lora_A_weight) |
| 264 | - self.set_layers_self_attention_linear_qkv_lora_A_default_weight(layer_idx=dst_layer_idx, data=qkv_lora_A_weight) | 264 | + self.set_layers_self_attention_linear_qkv_lora_B_default_weight(layer_idx=dst_layer_idx, data=qkv_lora_B_weight) |
| 265 | - self.set_layers_self_attention_linear_qkv_lora_B_default_weight(layer_idx=dst_layer_idx, data=qkv_lora_B_weight) | 265 | + if 'linear_proj' in self.args.lora_target_modules: |
| 266 | - self.set_layers_self_attention_linear_proj_lora_A_default_weight(layer_idx=dst_layer_idx, data=proj_lora_A_weight) | 266 | + proj_lora_A_weight = src_model.get_layers_self_attention_linear_proj_lora_A_default_weight(layer_idx=src_layer_idx) |
| 267 | - self.set_layers_self_attention_linear_proj_lora_B_default_weight(layer_idx=dst_layer_idx, data=proj_lora_B_weight) | 267 | + proj_lora_B_weight = src_model.get_layers_self_attention_linear_proj_lora_B_default_weight(layer_idx=src_layer_idx) |
| 268 | + self.set_layers_self_attention_linear_proj_lora_A_default_weight(layer_idx=dst_layer_idx, data=proj_lora_A_weight) | ||
| 269 | + self.set_layers_self_attention_linear_proj_lora_B_default_weight(layer_idx=dst_layer_idx, data=proj_lora_B_weight) | ||
| 268 | else: | 270 | else: |
| 269 | if getattr(src_model.get_args(), "qk_layernorm", False): | 271 | if getattr(src_model.get_args(), "qk_layernorm", False): |
| 270 | if getattr(src_model.get_args(), "multi_latent_attention", False): | 272 | if getattr(src_model.get_args(), "multi_latent_attention", False): |
| @@ -300,14 +302,16 @@ class ModelBase(abc.ABC): | |||
| 300 | def _set_mlp_state(self, src_model, **kwargs): | 302 | def _set_mlp_state(self, src_model, **kwargs): |
| 301 | """Set MLP params.""" | 303 | """Set MLP params.""" |
| 302 | if self.args.save_lora_to_hf: | 304 | if self.args.save_lora_to_hf: |
| 303 | - fc1_lora_A_weight = src_model.get_layers_mlp_linear_fc1_lora_A_default_weight(**kwargs) | 305 | + if 'linear_fc1' in self.args.lora_target_modules: |
| 304 | - fc1_lora_B_weight = src_model.get_layers_mlp_linear_fc1_lora_B_default_weight(**kwargs) | 306 | + fc1_lora_A_weight = src_model.get_layers_mlp_linear_fc1_lora_A_default_weight(**kwargs) |
| 305 | - fc2_lora_A_weight = src_model.get_layers_mlp_linear_fc2_lora_A_default_weight(**kwargs) | 307 | + fc1_lora_B_weight = src_model.get_layers_mlp_linear_fc1_lora_B_default_weight(**kwargs) |
| 306 | - fc2_lora_B_weight = src_model.get_layers_mlp_linear_fc2_lora_B_default_weight(**kwargs) | 308 | + self.set_layers_mlp_linear_fc1_lora_A_default_weight(data=fc1_lora_A_weight, **kwargs) |
| 307 | - self.set_layers_mlp_linear_fc1_lora_A_default_weight(data=fc1_lora_A_weight, **kwargs) | 309 | + self.set_layers_mlp_linear_fc1_lora_B_default_weight(data=fc1_lora_B_weight, **kwargs) |
| 308 | - self.set_layers_mlp_linear_fc1_lora_B_default_weight(data=fc1_lora_B_weight, **kwargs) | 310 | + if 'linear_fc2' in self.args.lora_target_modules: |
| 309 | - self.set_layers_mlp_linear_fc2_lora_A_default_weight(data=fc2_lora_A_weight, **kwargs) | 311 | + fc2_lora_A_weight = src_model.get_layers_mlp_linear_fc2_lora_A_default_weight(**kwargs) |
| 310 | - self.set_layers_mlp_linear_fc2_lora_B_default_weight(data=fc2_lora_B_weight, **kwargs) | 312 | + fc2_lora_B_weight = src_model.get_layers_mlp_linear_fc2_lora_B_default_weight(**kwargs) |
| 313 | + self.set_layers_mlp_linear_fc2_lora_A_default_weight(data=fc2_lora_A_weight, **kwargs) | ||
| 314 | + self.set_layers_mlp_linear_fc2_lora_B_default_weight(data=fc2_lora_B_weight, **kwargs) | ||
| 311 | else: | 315 | else: |
| 312 | fc1_weight = src_model.get_layers_mlp_linear_fc1_weight(**kwargs) | 316 | fc1_weight = src_model.get_layers_mlp_linear_fc1_weight(**kwargs) |
| 313 | fc2_weight = src_model.get_layers_mlp_linear_fc2_weight(**kwargs) | 317 | fc2_weight = src_model.get_layers_mlp_linear_fc2_weight(**kwargs) |
| @@ -328,14 +332,16 @@ class ModelBase(abc.ABC): | |||
| 328 | def _set_mlp_experts_state(self, src_model, **kwargs): | 332 | def _set_mlp_experts_state(self, src_model, **kwargs): |
| 329 | """Set MLP experts params.""" | 333 | """Set MLP experts params.""" |
| 330 | if self.args.save_lora_to_hf: | 334 | if self.args.save_lora_to_hf: |
| 331 | - fc1_lora_A_weight = src_model.get_layers_mlp_experts_linear_fc1_lora_A_default_weight(**kwargs) | 335 | + if 'linear_fc1' in self.args.lora_target_modules: |
| 332 | - fc1_lora_B_weight = src_model.get_layers_mlp_experts_linear_fc1_lora_B_default_weight(**kwargs) | 336 | + fc1_lora_A_weight = src_model.get_layers_mlp_experts_linear_fc1_lora_A_default_weight(**kwargs) |
| 333 | - fc2_lora_A_weight = src_model.get_layers_mlp_experts_linear_fc2_lora_A_default_weight(**kwargs) | 337 | + fc1_lora_B_weight = src_model.get_layers_mlp_experts_linear_fc1_lora_B_default_weight(**kwargs) |
| 334 | - fc2_lora_B_weight = src_model.get_layers_mlp_experts_linear_fc2_lora_B_default_weight(**kwargs) | 338 | + self.set_layers_mlp_experts_linear_fc1_lora_A_default_weight(data=fc1_lora_A_weight, **kwargs) |
| 335 | - self.set_layers_mlp_experts_linear_fc1_lora_A_default_weight(data=fc1_lora_A_weight, **kwargs) | 339 | + self.set_layers_mlp_experts_linear_fc1_lora_B_default_weight(data=fc1_lora_B_weight, **kwargs) |
| 336 | - self.set_layers_mlp_experts_linear_fc1_lora_B_default_weight(data=fc1_lora_B_weight, **kwargs) | 340 | + if 'linear_fc2' in self.args.lora_target_modules: |
| 337 | - self.set_layers_mlp_experts_linear_fc2_lora_A_default_weight(data=fc2_lora_A_weight, **kwargs) | 341 | + fc2_lora_A_weight = src_model.get_layers_mlp_experts_linear_fc2_lora_A_default_weight(**kwargs) |
| 338 | - self.set_layers_mlp_experts_linear_fc2_lora_B_default_weight(data=fc2_lora_B_weight, **kwargs) | 342 | + fc2_lora_B_weight = src_model.get_layers_mlp_experts_linear_fc2_lora_B_default_weight(**kwargs) |
| 343 | + self.set_layers_mlp_experts_linear_fc2_lora_A_default_weight(data=fc2_lora_A_weight, **kwargs) | ||
| 344 | + self.set_layers_mlp_experts_linear_fc2_lora_B_default_weight(data=fc2_lora_B_weight, **kwargs) | ||
| 339 | else: | 345 | else: |
| 340 | fc1_weight = src_model.get_layers_mlp_experts_linear_fc1_weight(**kwargs) | 346 | fc1_weight = src_model.get_layers_mlp_experts_linear_fc1_weight(**kwargs) |
| 341 | fc2_weight = src_model.get_layers_mlp_experts_linear_fc2_weight(**kwargs) | 347 | fc2_weight = src_model.get_layers_mlp_experts_linear_fc2_weight(**kwargs) |
| @@ -504,6 +510,7 @@ class HuggingfaceModel(ModelBase): | |||
| 504 | self.args.add_dense_bias = self.args_cmd.add_dense_bias | 510 | self.args.add_dense_bias = self.args_cmd.add_dense_bias |
| 505 | self.args.post_norm = self.args_cmd.post_norm | 511 | self.args.post_norm = self.args_cmd.post_norm |
| 506 | self.args.save_lora_to_hf = self.args_cmd.save_lora_to_hf | 512 | self.args.save_lora_to_hf = self.args_cmd.save_lora_to_hf |
| 513 | + self.args.lora_target_modules = self.args_cmd.lora_target_modules | ||
| 507 | self.args.noop_layers = self.args_cmd.noop_layers | 514 | self.args.noop_layers = self.args_cmd.noop_layers |
| 508 | if self.args.noop_layers is not None and self.args_cmd.save_model_type == 'hf': | 515 | if self.args.noop_layers is not None and self.args_cmd.save_model_type == 'hf': |
| 509 | mg_num_layers = self.args.num_layers + len(self.args.noop_layers) | 516 | mg_num_layers = self.args.num_layers + len(self.args.noop_layers) |
| @@ -209,9 +209,10 @@ def set_model_layer_attn(model_mg, msg, md, **kwargs): | |||
| 209 | if md.linear_bias or margs.add_qkv_bias: | 209 | if md.linear_bias or margs.add_qkv_bias: |
| 210 | qkv_bias = torch.chunk(msg.pop("qkv bias"), margs.tensor_model_parallel_size, dim=0) | 210 | qkv_bias = torch.chunk(msg.pop("qkv bias"), margs.tensor_model_parallel_size, dim=0) |
| 211 | 211 | ||
| 212 | - if margs.save_lora_to_hf: | 212 | + if margs.save_lora_to_hf and 'linear_qkv' in margs.lora_target_modules: |
| 213 | qkv_lora_A = msg.pop("qkv lora A") | 213 | qkv_lora_A = msg.pop("qkv lora A") |
| 214 | qkv_lora_B = msg.pop("qkv lora B") | 214 | qkv_lora_B = msg.pop("qkv lora B") |
| 215 | + if margs.save_lora_to_hf and 'linear_proj' in margs.lora_target_modules: | ||
| 215 | proj_lora_A = msg.pop("proj lora A") | 216 | proj_lora_A = msg.pop("proj lora A") |
| 216 | proj_lora_B = msg.pop("proj lora B") | 217 | proj_lora_B = msg.pop("proj lora B") |
| 217 | 218 | ||
| @@ -266,10 +267,12 @@ def set_model_layer_attn(model_mg, msg, md, **kwargs): | |||
| 266 | if margs.add_dense_bias: | 267 | if margs.add_dense_bias: |
| 267 | model_mg.set_layers_self_attention_linear_proj_bias(**kwargs, data=dense_bias) | 268 | model_mg.set_layers_self_attention_linear_proj_bias(**kwargs, data=dense_bias) |
| 268 | 269 | ||
| 269 | - if margs.save_lora_to_hf: | 270 | + if margs.save_lora_to_hf and 'linear_proj' in margs.lora_target_modules: |
| 270 | - logger.info(f"begin to convert attn of lora.") | 271 | + logger.info(f"begin to convert attn linear_proj of lora.") |
| 271 | model_mg.set_layers_self_attention_linear_proj_lora_A_default_weight(**kwargs, data=proj_lora_A) | 272 | model_mg.set_layers_self_attention_linear_proj_lora_A_default_weight(**kwargs, data=proj_lora_A) |
| 272 | model_mg.set_layers_self_attention_linear_proj_lora_B_default_weight(**kwargs, data=proj_lora_B) | 273 | model_mg.set_layers_self_attention_linear_proj_lora_B_default_weight(**kwargs, data=proj_lora_B) |
| 274 | + if margs.save_lora_to_hf and 'linear_qkv' in margs.lora_target_modules: | ||
| 275 | + logger.info(f"begin to convert attn linear_qkv of lora.") | ||
| 273 | model_mg.set_layers_self_attention_linear_qkv_lora_A_default_weight(**kwargs, data=qkv_lora_A) | 276 | model_mg.set_layers_self_attention_linear_qkv_lora_A_default_weight(**kwargs, data=qkv_lora_A) |
| 274 | model_mg.set_layers_self_attention_linear_qkv_lora_B_default_weight(**kwargs, data=qkv_lora_B) | 277 | model_mg.set_layers_self_attention_linear_qkv_lora_B_default_weight(**kwargs, data=qkv_lora_B) |
| 275 | 278 | ||
| @@ -282,9 +285,10 @@ def _set_set_model_layer_mlp(model_mg, msg, md, pop_flag=True, is_moe_mlp=False, | |||
| 282 | num_experts_local = margs.num_experts // margs.expert_model_parallel_size | 285 | num_experts_local = margs.num_experts // margs.expert_model_parallel_size |
| 283 | # Save them to the model | 286 | # Save them to the model |
| 284 | 287 | ||
| 285 | - if margs.save_lora_to_hf: | 288 | + if margs.save_lora_to_hf and 'linear_fc1' in margs.lora_target_modules: |
| 286 | fc1_lora_A = func(f"fc1 lora A") | 289 | fc1_lora_A = func(f"fc1 lora A") |
| 287 | fc1_lora_B = func(f"fc1 lora B") | 290 | fc1_lora_B = func(f"fc1 lora B") |
| 291 | + if margs.save_lora_to_hf and 'linear_fc2' in margs.lora_target_modules: | ||
| 288 | fc2_lora_A = func(f"fc2 lora A") | 292 | fc2_lora_A = func(f"fc2 lora A") |
| 289 | fc2_lora_B = func(f"fc2 lora B") | 293 | fc2_lora_B = func(f"fc2 lora B") |
| 290 | if md.linear_bias: | 294 | if md.linear_bias: |
| @@ -313,19 +317,23 @@ def _set_set_model_layer_mlp(model_mg, msg, md, pop_flag=True, is_moe_mlp=False, | |||
| 313 | if is_moe_mlp: | 317 | if is_moe_mlp: |
| 314 | model_mg.set_layers_mlp_experts_linear_fc1_weight(**kwargs, data=mlp_l0_weight[tp_rank]) | 318 | model_mg.set_layers_mlp_experts_linear_fc1_weight(**kwargs, data=mlp_l0_weight[tp_rank]) |
| 315 | model_mg.set_layers_mlp_experts_linear_fc2_weight(**kwargs, data=mlp_l1_weight[tp_rank]) | 319 | model_mg.set_layers_mlp_experts_linear_fc2_weight(**kwargs, data=mlp_l1_weight[tp_rank]) |
| 316 | - if margs.save_lora_to_hf: | 320 | + if margs.save_lora_to_hf and 'linear_fc1' in margs.lora_target_modules: |
| 317 | - logger.info(f"begin to convert mlp experts of lora.") | 321 | + logger.info(f"begin to convert mlp experts linear_fc1 of lora.") |
| 318 | model_mg.set_layers_mlp_experts_linear_fc1_lora_A_default_weight(**kwargs, data=fc1_lora_A) | 322 | model_mg.set_layers_mlp_experts_linear_fc1_lora_A_default_weight(**kwargs, data=fc1_lora_A) |
| 319 | model_mg.set_layers_mlp_experts_linear_fc1_lora_B_default_weight(**kwargs, data=fc1_lora_B) | 323 | model_mg.set_layers_mlp_experts_linear_fc1_lora_B_default_weight(**kwargs, data=fc1_lora_B) |
| 324 | + if margs.save_lora_to_hf and 'linear_fc2' in margs.lora_target_modules: | ||
| 325 | + logger.info(f"begin to convert mlp experts linear_fc2 of lora.") | ||
| 320 | model_mg.set_layers_mlp_experts_linear_fc2_lora_A_default_weight(**kwargs, data=fc2_lora_A) | 326 | model_mg.set_layers_mlp_experts_linear_fc2_lora_A_default_weight(**kwargs, data=fc2_lora_A) |
| 321 | model_mg.set_layers_mlp_experts_linear_fc2_lora_B_default_weight(**kwargs, data=fc2_lora_B) | 327 | model_mg.set_layers_mlp_experts_linear_fc2_lora_B_default_weight(**kwargs, data=fc2_lora_B) |
| 322 | else: | 328 | else: |
| 323 | model_mg.set_layers_mlp_linear_fc1_weight(**kwargs, data=mlp_l0_weight[tp_rank]) | 329 | model_mg.set_layers_mlp_linear_fc1_weight(**kwargs, data=mlp_l0_weight[tp_rank]) |
| 324 | model_mg.set_layers_mlp_linear_fc2_weight(**kwargs, data=mlp_l1_weight[tp_rank]) | 330 | model_mg.set_layers_mlp_linear_fc2_weight(**kwargs, data=mlp_l1_weight[tp_rank]) |
| 325 | - if margs.save_lora_to_hf: | 331 | + if margs.save_lora_to_hf and 'linear_fc1' in margs.lora_target_modules: |
| 326 | - logger.info(f"begin to convert mlp of lora.") | 332 | + logger.info(f"begin to convert mlp linear_fc1 of lora.") |
| 327 | model_mg.set_layers_mlp_linear_fc1_lora_A_default_weight(**kwargs, data=fc1_lora_A) | 333 | model_mg.set_layers_mlp_linear_fc1_lora_A_default_weight(**kwargs, data=fc1_lora_A) |
| 328 | model_mg.set_layers_mlp_linear_fc1_lora_B_default_weight(**kwargs, data=fc1_lora_B) | 334 | model_mg.set_layers_mlp_linear_fc1_lora_B_default_weight(**kwargs, data=fc1_lora_B) |
| 335 | + if margs.save_lora_to_hf and 'linear_fc2' in margs.lora_target_modules: | ||
| 336 | + logger.info(f"begin to convert mlp linear_fc2 of lora.") | ||
| 329 | model_mg.set_layers_mlp_linear_fc2_lora_A_default_weight(**kwargs, data=fc2_lora_A) | 337 | model_mg.set_layers_mlp_linear_fc2_lora_A_default_weight(**kwargs, data=fc2_lora_A) |
| 330 | model_mg.set_layers_mlp_linear_fc2_lora_B_default_weight(**kwargs, data=fc2_lora_B) | 338 | model_mg.set_layers_mlp_linear_fc2_lora_B_default_weight(**kwargs, data=fc2_lora_B) |
| 331 | 339 | ||
| @@ -20,7 +20,7 @@ def _load_from_state_dict_wrapper(fn): | |||
| 20 | fn(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) | 20 | fn(self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs) |
| 21 | from megatron.training import get_args | 21 | from megatron.training import get_args |
| 22 | args = get_args() | 22 | args = get_args() |
| 23 | - if hasattr(args, 'lora_target_modules') and args.lora_target_modules: | 23 | + if not getattr(args, 'save_lora_to_hf', False) and hasattr(args, 'lora_target_modules') and args.lora_target_modules: |
| 24 | if not any(('lora_a' in key.lower() or 'lora_b' in key.lower()) and key.endswith('weight') for key in state_dict): | 24 | if not any(('lora_a' in key.lower() or 'lora_b' in key.lower()) and key.endswith('weight') for key in state_dict): |
| 25 | import warnings | 25 | import warnings |
| 26 | warnings.warn("The lora weights is missing from the checkpoint and will be randomly initialized.", RuntimeWarning) | 26 | warnings.warn("The lora weights is missing from the checkpoint and will be randomly initialized.", RuntimeWarning) |