MCPcopy Create free account
hub / github.com/OpenGVLab/InternVL / save_checkpoint

Function save_checkpoint

classification/main_accelerate.py:166–177  ·  view source on GitHub ↗
(save_dir, accelerator, epoch, max_acc, config, lr_scheduler=None)

Source from the content-addressed store, hash-verified

164
165
166def save_checkpoint(save_dir, accelerator, epoch, max_acc, config, lr_scheduler=None):
167 # let accelerator handle the model and optimizer state for ddp and deepspeed.
168 accelerator.save_state(save_dir)
169
170 if accelerator.is_main_process:
171 save_state = {
172 'lr_scheduler': lr_scheduler.state_dict(),
173 'max_acc': max_acc,
174 'epoch': epoch,
175 'config': config
176 }
177 torch.save(save_state, os.path.join(save_dir, 'additional_state.pth'))
178
179
180def load_checkpoint_if_needed(accelerator, config, lr_scheduler=None):

Callers 1

trainFunction · 0.70

Calls 1

state_dictMethod · 0.80

Tested by

no test coverage detected