MCPcopy Create free account
hub / github.com/clinicalml/TabLLM / save_model

Method save_model

t-few/src/models/EncoderDecoder.py:340–364  ·  view source on GitHub ↗
(self, finish=False)

Source from the content-addressed store, hash-verified

338 ), f"Load model failed, unexpected keys {load_result.unexpected_keys.__str__()}"
339
340 def save_model(self, finish=False):
341 if self.config.save_model and (finish or self._last_global_step_saved != self.global_step):
342 if finish:
343 model_fname = os.path.join(self.config.exp_dir, "finish.pt")
344 else:
345 model_fname = os.path.join(self.config.exp_dir, f"global_step{self.global_step}.pt")
346
347 if self.use_deepspeed or self.use_ddp:
348 distributed_save_path = os.path.join(self.config.exp_dir, "saved_model")
349 self.trainer.model.save_checkpoint(distributed_save_path)
350 torch.distributed.barrier()
351 if dist.get_rank() == 0:
352 trainable_states = zero_to_fp32.get_fp32_state_dict_from_zero_checkpoint(distributed_save_path)
353 prefix_length = len("module.model.")
354 trainable_states = {k[prefix_length:]: v for k, v in trainable_states.items()}
355 torch.save(trainable_states, model_fname)
356 else:
357 trainable_states = {
358 param_name: param_weight.cpu()
359 for param_name, param_weight in self.model.state_dict().items()
360 if param_name in self.trainable_param_names
361 }
362 torch.save(trainable_states, model_fname)
363
364 self._last_global_step_saved = self.global_step
365
366 def on_before_optimizer_step(self, optimizer, optimizer_idx):
367 if self.config.fishmask_mode is not None:

Callers 3

training_stepMethod · 0.95
validation_epoch_endMethod · 0.95
on_train_endMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected