{"cells":[{"metadata":{},"cell_type":"markdown","source":"Keras has a `save_best_only` option for its [ModelCheckpoint](https://keras.io/callbacks/#modelcheckpoint). Ignite [doesn't](https://pytorch.org/ignite/_modules/ignite/handlers/checkpoint.html) - it only has a `save_interval`.   \n   \nBelow is a very simple code snippet you can use with Ignite if you want to have this feature."},{"metadata":{},"cell_type":"markdown","source":"## Dummy Ignite trainer/engine preparation"},{"metadata":{"trusted":true},"cell_type":"code","source":"import os\nimport torch\nfrom torchvision import models\n\nfrom ignite.engine import Events, create_supervised_trainer, create_supervised_evaluator\nfrom ignite.metrics import Loss, Accuracy\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\nmodel = models.resnet18(pretrained=True)\n\ncriterion = torch.nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=0.0003)\n\nmetrics = {\n    'loss': Loss(criterion),\n    'accuracy': Accuracy(),\n}\n\ntrainer = create_supervised_trainer(model, optimizer, criterion, device=device)\nval_evaluator = create_supervised_evaluator(model, metrics=metrics, device=device)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Below is the reusable `save_best_only` code snippet"},{"metadata":{"trusted":true},"cell_type":"code","source":"# create models directory if it doesn't exist\n!mkdir -p models","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"def get_saved_model_path(epoch):\n    return f'models/Model_{model_name}_{epoch}.pth'\n\nbest_acc = 0.\nbest_epoch = 1\nbest_epoch_file = ''\n\n@trainer.on(Events.EPOCH_COMPLETED)\ndef save_best_epoch_only(engine):\n    epoch = engine.state.epoch\n\n    global best_acc\n    global best_epoch\n    global best_epoch_file\n    best_acc = 0. if epoch == 1 else best_acc\n    best_epoch = 1 if epoch == 1 else best_epoch\n    best_epoch_file = '' if epoch == 1 else best_epoch_file\n\n    metrics = val_evaluator.run(val_loader).metrics\n\n    if metrics['accuracy'] > best_acc:\n        prev_best_epoch_file = get_saved_model_path(best_epoch)\n        if os.path.exists(prev_best_epoch_file):\n            os.remove(prev_best_epoch_file)\n            \n        best_acc = metrics['accuracy']\n        best_epoch = epoch\n        best_epoch_file = get_saved_model_path(best_epoch)\n        print(f'\\nEpoch: {best_epoch} - New best accuracy! Accuracy: {best_acc}\\n\\n\\n')\n        torch.save(model.state_dict(), best_epoch_file)","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":1}