已合并
[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
已合并
温一盏创建于 1月6日
共 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 None87 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 units186 # 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_A273 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_B277 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_bias510 self.args.add_dense_bias = self.args_cmd.add_dense_bias
505 self.args.post_norm = self.args_cmd.post_norm511 self.args.post_norm = self.args_cmd.post_norm
506 self.args.save_lora_to_hf = self.args_cmd.save_lora_to_hf512 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_layers514 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_size285 num_experts_local = margs.num_experts // margs.expert_model_parallel_size
283 # Save them to the model286 # 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_args21 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 warnings25 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)