mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
2295 lines
103 KiB
HTML
2295 lines
103 KiB
HTML
|
|
|
|
|
|
|
|
<!doctype html>
|
|
<html lang="en" class="no-js">
|
|
<head>
|
|
|
|
<meta charset="utf-8">
|
|
<meta name="viewport" content="width=device-width,initial-scale=1">
|
|
<meta http-equiv="x-ua-compatible" content="ie=edge">
|
|
|
|
<meta name="description" content="Documentation for PyTorch LightningModule, the researcher version of keras.">
|
|
|
|
|
|
|
|
|
|
<meta name="lang:clipboard.copy" content="Copy to clipboard">
|
|
|
|
<meta name="lang:clipboard.copied" content="Copied to clipboard">
|
|
|
|
<meta name="lang:search.language" content="en">
|
|
|
|
<meta name="lang:search.pipeline.stopwords" content="True">
|
|
|
|
<meta name="lang:search.pipeline.trimmer" content="True">
|
|
|
|
<meta name="lang:search.result.none" content="No matching documents">
|
|
|
|
<meta name="lang:search.result.one" content="1 matching document">
|
|
|
|
<meta name="lang:search.result.other" content="# matching documents">
|
|
|
|
<meta name="lang:search.tokenizer" content="[\s\-]+">
|
|
|
|
<link rel="shortcut icon" href="../../assets/images/favicon.png">
|
|
<meta name="generator" content="mkdocs-1.0.4, mkdocs-material-4.4.0">
|
|
|
|
|
|
|
|
<title>Lightning Module interface - PyTorch lightning Documentation</title>
|
|
|
|
|
|
|
|
<link rel="stylesheet" href="../../assets/stylesheets/application.0284f74d.css">
|
|
|
|
|
|
|
|
|
|
<script src="../../assets/javascripts/modernizr.74668098.js"></script>
|
|
|
|
|
|
|
|
<link href="https://fonts.gstatic.com" rel="preconnect" crossorigin>
|
|
<link rel="stylesheet" href="https://fonts.googleapis.com/css?family=Roboto:300,400,400i,700|Roboto+Mono&display=fallback">
|
|
<style>body,input{font-family:"Roboto","Helvetica Neue",Helvetica,Arial,sans-serif}code,kbd,pre{font-family:"Roboto Mono","Courier New",Courier,monospace}</style>
|
|
|
|
|
|
<link rel="stylesheet" href="../../assets/fonts/material-icons.css">
|
|
|
|
|
|
|
|
|
|
|
|
|
|
</head>
|
|
|
|
<body dir="ltr">
|
|
|
|
<svg class="md-svg">
|
|
<defs>
|
|
|
|
|
|
<svg xmlns="http://www.w3.org/2000/svg" width="416" height="448" viewBox="0 0 416 448" id="__github"><path fill="currentColor" d="M160 304q0 10-3.125 20.5t-10.75 19T128 352t-18.125-8.5-10.75-19T96 304t3.125-20.5 10.75-19T128 256t18.125 8.5 10.75 19T160 304zm160 0q0 10-3.125 20.5t-10.75 19T288 352t-18.125-8.5-10.75-19T256 304t3.125-20.5 10.75-19T288 256t18.125 8.5 10.75 19T320 304zm40 0q0-30-17.25-51T296 232q-10.25 0-48.75 5.25Q229.5 240 208 240t-39.25-2.75Q130.75 232 120 232q-29.5 0-46.75 21T56 304q0 22 8 38.375t20.25 25.75 30.5 15 35 7.375 37.25 1.75h42q20.5 0 37.25-1.75t35-7.375 30.5-15 20.25-25.75T360 304zm56-44q0 51.75-15.25 82.75-9.5 19.25-26.375 33.25t-35.25 21.5-42.5 11.875-42.875 5.5T212 416q-19.5 0-35.5-.75t-36.875-3.125-38.125-7.5-34.25-12.875T37 371.5t-21.5-28.75Q0 312 0 260q0-59.25 34-99-6.75-20.5-6.75-42.5 0-29 12.75-54.5 27 0 47.5 9.875t47.25 30.875Q171.5 96 212 96q37 0 70 8 26.25-20.5 46.75-30.25T376 64q12.75 25.5 12.75 54.5 0 21.75-6.75 42 34 40 34 99.5z"/></svg>
|
|
|
|
</defs>
|
|
</svg>
|
|
<input class="md-toggle" data-md-toggle="drawer" type="checkbox" id="__drawer" autocomplete="off">
|
|
<input class="md-toggle" data-md-toggle="search" type="checkbox" id="__search" autocomplete="off">
|
|
<label class="md-overlay" data-md-component="overlay" for="__drawer"></label>
|
|
|
|
<a href="#lightning-module-interface" tabindex="1" class="md-skip">
|
|
Skip to content
|
|
</a>
|
|
|
|
|
|
<header class="md-header" data-md-component="header">
|
|
<nav class="md-header-nav md-grid">
|
|
<div class="md-flex">
|
|
<div class="md-flex__cell md-flex__cell--shrink">
|
|
<a href="../.." title="PyTorch lightning Documentation" class="md-header-nav__button md-logo">
|
|
|
|
<i class="md-icon"></i>
|
|
|
|
</a>
|
|
</div>
|
|
<div class="md-flex__cell md-flex__cell--shrink">
|
|
<label class="md-icon md-icon--menu md-header-nav__button" for="__drawer"></label>
|
|
</div>
|
|
<div class="md-flex__cell md-flex__cell--stretch">
|
|
<div class="md-flex__ellipsis md-header-nav__title" data-md-component="title">
|
|
|
|
<span class="md-header-nav__topic">
|
|
PyTorch lightning Documentation
|
|
</span>
|
|
<span class="md-header-nav__topic">
|
|
|
|
Lightning Module interface
|
|
|
|
</span>
|
|
|
|
</div>
|
|
</div>
|
|
<div class="md-flex__cell md-flex__cell--shrink">
|
|
|
|
<label class="md-icon md-icon--search md-header-nav__button" for="__search"></label>
|
|
|
|
<div class="md-search" data-md-component="search" role="dialog">
|
|
<label class="md-search__overlay" for="__search"></label>
|
|
<div class="md-search__inner" role="search">
|
|
<form class="md-search__form" name="search">
|
|
<input type="text" class="md-search__input" name="query" placeholder="Search" autocapitalize="off" autocorrect="off" autocomplete="off" spellcheck="false" data-md-component="query" data-md-state="active">
|
|
<label class="md-icon md-search__icon" for="__search"></label>
|
|
<button type="reset" class="md-icon md-search__icon" data-md-component="reset" tabindex="-1">
|
|

|
|
</button>
|
|
</form>
|
|
<div class="md-search__output">
|
|
<div class="md-search__scrollwrap" data-md-scrollfix>
|
|
<div class="md-search-result" data-md-component="result">
|
|
<div class="md-search-result__meta">
|
|
Type to start searching
|
|
</div>
|
|
<ol class="md-search-result__list"></ol>
|
|
</div>
|
|
</div>
|
|
</div>
|
|
</div>
|
|
</div>
|
|
|
|
</div>
|
|
|
|
<div class="md-flex__cell md-flex__cell--shrink">
|
|
<div class="md-header-nav__source">
|
|
|
|
|
|
|
|
|
|
|
|
<a href="https://github.com/williamFalcon/pytorch-lightning/" title="Go to repository" class="md-source" data-md-source="github">
|
|
|
|
<div class="md-source__icon">
|
|
<svg viewBox="0 0 24 24" width="24" height="24">
|
|
<use xlink:href="#__github" width="24" height="24"></use>
|
|
</svg>
|
|
</div>
|
|
|
|
<div class="md-source__repository">
|
|
williamFalcon/pytorch-lightning
|
|
</div>
|
|
</a>
|
|
</div>
|
|
</div>
|
|
|
|
</div>
|
|
</nav>
|
|
</header>
|
|
|
|
<div class="md-container">
|
|
|
|
|
|
|
|
|
|
<main class="md-main">
|
|
<div class="md-main__inner md-grid" data-md-component="container">
|
|
|
|
|
|
<div class="md-sidebar md-sidebar--primary" data-md-component="navigation">
|
|
<div class="md-sidebar__scrollwrap">
|
|
<div class="md-sidebar__inner">
|
|
<nav class="md-nav md-nav--primary" data-md-level="0">
|
|
<label class="md-nav__title md-nav__title--site" for="__drawer">
|
|
<a href="../.." title="PyTorch lightning Documentation" class="md-nav__button md-logo">
|
|
|
|
<i class="md-icon"></i>
|
|
|
|
</a>
|
|
PyTorch lightning Documentation
|
|
</label>
|
|
|
|
<div class="md-nav__source">
|
|
|
|
|
|
|
|
|
|
|
|
<a href="https://github.com/williamFalcon/pytorch-lightning/" title="Go to repository" class="md-source" data-md-source="github">
|
|
|
|
<div class="md-source__icon">
|
|
<svg viewBox="0 0 24 24" width="24" height="24">
|
|
<use xlink:href="#__github" width="24" height="24"></use>
|
|
</svg>
|
|
</div>
|
|
|
|
<div class="md-source__repository">
|
|
williamFalcon/pytorch-lightning
|
|
</div>
|
|
</a>
|
|
</div>
|
|
|
|
<ul class="md-nav__list" data-md-scrollfix>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../.." title="Home" class="md-nav__link">
|
|
Home
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item md-nav__item--active md-nav__item--nested">
|
|
|
|
<input class="md-toggle md-nav__toggle" data-md-toggle="nav-2" type="checkbox" id="nav-2" checked>
|
|
|
|
<label class="md-nav__link" for="nav-2">
|
|
LightningModule
|
|
</label>
|
|
<nav class="md-nav" data-md-component="collapsible" data-md-level="1">
|
|
<label class="md-nav__title" for="nav-2">
|
|
LightningModule
|
|
</label>
|
|
<ul class="md-nav__list" data-md-scrollfix>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item md-nav__item--active">
|
|
|
|
<input class="md-toggle md-nav__toggle" data-md-toggle="toc" type="checkbox" id="__toc">
|
|
|
|
|
|
|
|
|
|
<label class="md-nav__link md-nav__link--active" for="__toc">
|
|
Lightning Module interface
|
|
</label>
|
|
|
|
<a href="./" title="Lightning Module interface" class="md-nav__link md-nav__link--active">
|
|
Lightning Module interface
|
|
</a>
|
|
|
|
|
|
<nav class="md-nav md-nav--secondary">
|
|
|
|
|
|
|
|
|
|
|
|
<label class="md-nav__title" for="__toc">Table of contents</label>
|
|
<ul class="md-nav__list" data-md-scrollfix>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#minimal-example" title="Minimal example" class="md-nav__link">
|
|
Minimal example
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#how-do-these-methods-fit-into-the-broader-training" title="How do these methods fit into the broader training?" class="md-nav__link">
|
|
How do these methods fit into the broader training?
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#required-methods" title="Required Methods" class="md-nav__link">
|
|
Required Methods
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#training_step" title="training_step" class="md-nav__link">
|
|
training_step
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#training_end" title="training_end" class="md-nav__link">
|
|
training_end
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#train_dataloader" title="train_dataloader" class="md-nav__link">
|
|
train_dataloader
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#configure_optimizers" title="configure_optimizers" class="md-nav__link">
|
|
configure_optimizers
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_1" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#optional-methods" title="Optional Methods" class="md-nav__link">
|
|
Optional Methods
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#validation_step" title="validation_step" class="md-nav__link">
|
|
validation_step
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#validation_end" title="validation_end" class="md-nav__link">
|
|
validation_end
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#test_step" title="test_step" class="md-nav__link">
|
|
test_step
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#test_end" title="test_end" class="md-nav__link">
|
|
test_end
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#on_save_checkpoint" title="on_save_checkpoint" class="md-nav__link">
|
|
on_save_checkpoint
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_2" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#on_load_checkpoint" title="on_load_checkpoint" class="md-nav__link">
|
|
on_load_checkpoint
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_3" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#val_dataloader" title="val_dataloader" class="md-nav__link">
|
|
val_dataloader
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_4" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#test_dataloader" title="test_dataloader" class="md-nav__link">
|
|
test_dataloader
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_5" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#add_model_specific_args" title="add_model_specific_args" class="md-nav__link">
|
|
add_model_specific_args
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_6" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
</ul>
|
|
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../methods/" title="Methods" class="md-nav__link">
|
|
Methods
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../properties/" title="Properties" class="md-nav__link">
|
|
Properties
|
|
</a>
|
|
</li>
|
|
|
|
|
|
</ul>
|
|
</nav>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item md-nav__item--nested">
|
|
|
|
<input class="md-toggle md-nav__toggle" data-md-toggle="nav-3" type="checkbox" id="nav-3">
|
|
|
|
<label class="md-nav__link" for="nav-3">
|
|
Trainer
|
|
</label>
|
|
<nav class="md-nav" data-md-component="collapsible" data-md-level="1">
|
|
<label class="md-nav__title" for="nav-3">
|
|
Trainer
|
|
</label>
|
|
<ul class="md-nav__list" data-md-scrollfix>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/" title="Trainer" class="md-nav__link">
|
|
Trainer
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/Checkpointing/" title="Checkpointing" class="md-nav__link">
|
|
Checkpointing
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/Distributed training/" title="Distributed training" class="md-nav__link">
|
|
Distributed training
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/Logging/" title="Logging" class="md-nav__link">
|
|
Logging
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/SLURM Managed Cluster/" title="SLURM Managed Cluster" class="md-nav__link">
|
|
SLURM Managed Cluster
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/Testing loop/" title="Testing loop" class="md-nav__link">
|
|
Testing loop
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/Training Loop/" title="Training Loop" class="md-nav__link">
|
|
Training Loop
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/Validation loop/" title="Validation loop" class="md-nav__link">
|
|
Validation loop
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/debugging/" title="Debugging" class="md-nav__link">
|
|
Debugging
|
|
</a>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../Trainer/hooks/" title="Hooks" class="md-nav__link">
|
|
Hooks
|
|
</a>
|
|
</li>
|
|
|
|
|
|
</ul>
|
|
</nav>
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item md-nav__item--nested">
|
|
|
|
<input class="md-toggle md-nav__toggle" data-md-toggle="nav-4" type="checkbox" id="nav-4">
|
|
|
|
<label class="md-nav__link" for="nav-4">
|
|
Examples
|
|
</label>
|
|
<nav class="md-nav" data-md-component="collapsible" data-md-level="1">
|
|
<label class="md-nav__title" for="nav-4">
|
|
Examples
|
|
</label>
|
|
<ul class="md-nav__list" data-md-scrollfix>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
<li class="md-nav__item">
|
|
<a href="../../examples/Examples/" title="Examples" class="md-nav__link">
|
|
Examples
|
|
</a>
|
|
</li>
|
|
|
|
|
|
</ul>
|
|
</nav>
|
|
</li>
|
|
|
|
|
|
</ul>
|
|
</nav>
|
|
</div>
|
|
</div>
|
|
</div>
|
|
|
|
|
|
<div class="md-sidebar md-sidebar--secondary" data-md-component="toc">
|
|
<div class="md-sidebar__scrollwrap">
|
|
<div class="md-sidebar__inner">
|
|
|
|
<nav class="md-nav md-nav--secondary">
|
|
|
|
|
|
|
|
|
|
|
|
<label class="md-nav__title" for="__toc">Table of contents</label>
|
|
<ul class="md-nav__list" data-md-scrollfix>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#minimal-example" title="Minimal example" class="md-nav__link">
|
|
Minimal example
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#how-do-these-methods-fit-into-the-broader-training" title="How do these methods fit into the broader training?" class="md-nav__link">
|
|
How do these methods fit into the broader training?
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#required-methods" title="Required Methods" class="md-nav__link">
|
|
Required Methods
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#training_step" title="training_step" class="md-nav__link">
|
|
training_step
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#training_end" title="training_end" class="md-nav__link">
|
|
training_end
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#train_dataloader" title="train_dataloader" class="md-nav__link">
|
|
train_dataloader
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#configure_optimizers" title="configure_optimizers" class="md-nav__link">
|
|
configure_optimizers
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_1" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#optional-methods" title="Optional Methods" class="md-nav__link">
|
|
Optional Methods
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#validation_step" title="validation_step" class="md-nav__link">
|
|
validation_step
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#validation_end" title="validation_end" class="md-nav__link">
|
|
validation_end
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#test_step" title="test_step" class="md-nav__link">
|
|
test_step
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#test_end" title="test_end" class="md-nav__link">
|
|
test_end
|
|
</a>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#on_save_checkpoint" title="on_save_checkpoint" class="md-nav__link">
|
|
on_save_checkpoint
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_2" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#on_load_checkpoint" title="on_load_checkpoint" class="md-nav__link">
|
|
on_load_checkpoint
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_3" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#val_dataloader" title="val_dataloader" class="md-nav__link">
|
|
val_dataloader
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_4" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#test_dataloader" title="test_dataloader" class="md-nav__link">
|
|
test_dataloader
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_5" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#add_model_specific_args" title="add_model_specific_args" class="md-nav__link">
|
|
add_model_specific_args
|
|
</a>
|
|
|
|
<nav class="md-nav">
|
|
<ul class="md-nav__list">
|
|
|
|
<li class="md-nav__item">
|
|
<a href="#return_6" title="Return" class="md-nav__link">
|
|
Return
|
|
</a>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
</ul>
|
|
</nav>
|
|
|
|
</li>
|
|
|
|
|
|
|
|
|
|
|
|
</ul>
|
|
|
|
</nav>
|
|
</div>
|
|
</div>
|
|
</div>
|
|
|
|
|
|
<div class="md-content">
|
|
<article class="md-content__inner md-typeset">
|
|
|
|
|
|
<a href="https://github.com/williamFalcon/pytorch-lightning/edit/master/docs/LightningModule/RequiredTrainerInterface.md" title="Edit this page" class="md-icon md-content__icon"></a>
|
|
|
|
|
|
<h1 id="lightning-module-interface">Lightning Module interface</h1>
|
|
<p>[<a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/root_module/root_module.py">Github Code</a>]</p>
|
|
<p>A lightning module is a strict superclass of nn.Module, it provides a standard interface for the trainer to interact with the model.</p>
|
|
<p>The easiest thing to do is copy the <a href="https://williamfalcon.github.io/pytorch-lightning/LightningModule/RequiredTrainerInterface/#minimal-example">minimal example</a> below and modify accordingly. </p>
|
|
<p>Otherwise, to Define a Lightning Module, implement the following methods:</p>
|
|
<p><strong>Required</strong>: </p>
|
|
<ul>
|
|
<li><a href="./#training_step">training_step</a> </li>
|
|
<li><a href="./#train_dataloader">train_dataloader</a> </li>
|
|
<li><a href="./#configure_optimizers">configure_optimizers</a> </li>
|
|
</ul>
|
|
<p><strong>Optional</strong>: </p>
|
|
<ul>
|
|
<li><a href="./#training_end">training_end</a> </li>
|
|
<li><a href="./#validation_step">validation_step</a> </li>
|
|
<li><a href="./#validation_end">validation_end</a> </li>
|
|
<li><a href="./#test_step">test_step</a> </li>
|
|
<li><a href="./#test_end">test_end</a> </li>
|
|
<li><a href="./#val_dataloader">val_dataloader</a> </li>
|
|
<li><a href="./#test_dataloader">test_dataloader</a> </li>
|
|
<li><a href="./#on_save_checkpoint">on_save_checkpoint</a> </li>
|
|
<li><a href="./#on_load_checkpoint">on_load_checkpoint</a> </li>
|
|
<li><a href="./#add_model_specific_args">add_model_specific_args</a> </li>
|
|
</ul>
|
|
<hr />
|
|
<h3 id="minimal-example">Minimal example</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22
|
|
23
|
|
24
|
|
25
|
|
26
|
|
27
|
|
28
|
|
29
|
|
30
|
|
31
|
|
32
|
|
33
|
|
34
|
|
35
|
|
36
|
|
37
|
|
38
|
|
39
|
|
40
|
|
41
|
|
42
|
|
43
|
|
44
|
|
45
|
|
46
|
|
47
|
|
48
|
|
49
|
|
50
|
|
51
|
|
52
|
|
53
|
|
54
|
|
55
|
|
56
|
|
57
|
|
58
|
|
59
|
|
60
|
|
61
|
|
62
|
|
63
|
|
64
|
|
65
|
|
66</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="kn">import</span> <span class="nn">os</span>
|
|
<span class="kn">import</span> <span class="nn">torch</span>
|
|
<span class="kn">from</span> <span class="nn">torch.nn</span> <span class="kn">import</span> <span class="n">functional</span> <span class="k">as</span> <span class="n">F</span>
|
|
<span class="kn">from</span> <span class="nn">torch.utils.data</span> <span class="kn">import</span> <span class="n">DataLoader</span>
|
|
<span class="kn">from</span> <span class="nn">torchvision.datasets</span> <span class="kn">import</span> <span class="n">MNIST</span>
|
|
<span class="kn">import</span> <span class="nn">torchvision.transforms</span> <span class="kn">as</span> <span class="nn">transforms</span>
|
|
|
|
<span class="kn">import</span> <span class="nn">pytorch_lightning</span> <span class="kn">as</span> <span class="nn">pl</span>
|
|
|
|
<span class="k">class</span> <span class="nc">CoolModel</span><span class="p">(</span><span class="n">pl</span><span class="o">.</span><span class="n">LightningModule</span><span class="p">):</span>
|
|
|
|
<span class="k">def</span> <span class="fm">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="nb">super</span><span class="p">(</span><span class="n">CoolModel</span><span class="p">,</span> <span class="bp">self</span><span class="p">)</span><span class="o">.</span><span class="fm">__init__</span><span class="p">()</span>
|
|
<span class="c1"># not the best model...</span>
|
|
<span class="bp">self</span><span class="o">.</span><span class="n">l1</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">nn</span><span class="o">.</span><span class="n">Linear</span><span class="p">(</span><span class="mi">28</span> <span class="o">*</span> <span class="mi">28</span><span class="p">,</span> <span class="mi">10</span><span class="p">)</span>
|
|
|
|
<span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
|
|
<span class="k">return</span> <span class="n">torch</span><span class="o">.</span><span class="n">relu</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">l1</span><span class="p">(</span><span class="n">x</span><span class="o">.</span><span class="n">view</span><span class="p">(</span><span class="n">x</span><span class="o">.</span><span class="n">size</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="o">-</span><span class="mi">1</span><span class="p">)))</span>
|
|
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="c1"># REQUIRED</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
<span class="n">y_hat</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'loss'</span><span class="p">:</span> <span class="n">F</span><span class="o">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)}</span>
|
|
|
|
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="c1"># OPTIONAL</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
<span class="n">y_hat</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'val_loss'</span><span class="p">:</span> <span class="n">F</span><span class="o">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)}</span>
|
|
|
|
<span class="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="c1"># OPTIONAL</span>
|
|
<span class="n">avg_loss</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">stack</span><span class="p">([</span><span class="n">x</span><span class="p">[</span><span class="s1">'val_loss'</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">])</span><span class="o">.</span><span class="n">mean</span><span class="p">()</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'avg_val_loss'</span><span class="p">:</span> <span class="n">avg_loss</span><span class="p">}</span>
|
|
|
|
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="c1"># OPTIONAL</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
<span class="n">y_hat</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'test_loss'</span><span class="p">:</span> <span class="n">F</span><span class="o">.</span><span class="n">cross_entropy</span><span class="p">(</span><span class="n">y_hat</span><span class="p">,</span> <span class="n">y</span><span class="p">)}</span>
|
|
|
|
<span class="k">def</span> <span class="nf">test_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="c1"># OPTIONAL</span>
|
|
<span class="n">avg_loss</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">stack</span><span class="p">([</span><span class="n">x</span><span class="p">[</span><span class="s1">'test_loss'</span><span class="p">]</span> <span class="k">for</span> <span class="n">x</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">])</span><span class="o">.</span><span class="n">mean</span><span class="p">()</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'avg_test_loss'</span><span class="p">:</span> <span class="n">avg_loss</span><span class="p">}</span>
|
|
|
|
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="c1"># REQUIRED</span>
|
|
<span class="k">return</span> <span class="n">torch</span><span class="o">.</span><span class="n">optim</span><span class="o">.</span><span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.02</span><span class="p">)</span>
|
|
|
|
<span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">train_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
|
|
|
|
<span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="c1"># OPTIONAL</span>
|
|
<span class="c1"># can also return a list of val dataloaders</span>
|
|
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
|
|
|
|
<span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="c1"># OPTIONAL</span>
|
|
<span class="c1"># can also return a list of test dataloaders</span>
|
|
<span class="k">return</span> <span class="n">DataLoader</span><span class="p">(</span><span class="n">MNIST</span><span class="p">(</span><span class="n">os</span><span class="o">.</span><span class="n">getcwd</span><span class="p">(),</span> <span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">()),</span> <span class="n">batch_size</span><span class="o">=</span><span class="mi">32</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<hr />
|
|
<h3 id="how-do-these-methods-fit-into-the-broader-training">How do these methods fit into the broader training?</h3>
|
|
<p>The LightningModule interface is on the right. Each method corresponds to a part of a research project. Lightning automates everything not in blue. </p>
|
|
<p align="center">
|
|
<a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/docs/source/_static/overview_flat.jpg">
|
|
<img alt="" src="https://github.com/williamFalcon/pytorch-lightning/blob/master/docs/source/_static/overview_flat.jpg" height="900px">
|
|
</a>
|
|
</p>
|
|
|
|
<h2 id="required-methods">Required Methods</h2>
|
|
<h3 id="training_step">training_step</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>In this step you'd normally do the forward pass and calculate the loss for a batch. You can also do fancier things like multiple forward passes or something specific to your model.</p>
|
|
<p><strong>Params</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Param</th>
|
|
<th>description</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>batch</td>
|
|
<td>The output of your dataloader. A tensor, tuple or list</td>
|
|
</tr>
|
|
<tr>
|
|
<td>batch_nb</td>
|
|
<td>Integer displaying which batch this is</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Return</strong> </p>
|
|
<p>Dictionary or OrderedDict </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>key</th>
|
|
<th>value</th>
|
|
<th>is required</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>loss</td>
|
|
<td>tensor scalar</td>
|
|
<td>Y</td>
|
|
</tr>
|
|
<tr>
|
|
<td>progress_bar</td>
|
|
<td>Dict for progress bar display. Must have only tensors</td>
|
|
<td>N</td>
|
|
</tr>
|
|
<tr>
|
|
<td>log</td>
|
|
<td>Dict of metrics to add to logger. Must have only tensors (no images, etc)</td>
|
|
<td>N</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">z</span> <span class="o">=</span> <span class="n">batch</span>
|
|
|
|
<span class="c1"># implement your own</span>
|
|
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
|
|
|
|
<span class="n">logger_logs</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'training_loss'</span><span class="p">:</span> <span class="n">loss</span><span class="p">}</span> <span class="c1"># optional (MUST ALL BE TENSORS)</span>
|
|
|
|
<span class="c1"># if using TestTubeLogger or TensorboardLogger you can nest scalars</span>
|
|
<span class="n">logger_logs</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'losses'</span><span class="p">:</span> <span class="n">logger_logs</span><span class="p">}</span> <span class="c1"># optional (MUST ALL BE TENSORS)</span>
|
|
|
|
<span class="n">output</span> <span class="o">=</span> <span class="p">{</span>
|
|
<span class="s1">'loss'</span><span class="p">:</span> <span class="n">loss</span><span class="p">,</span> <span class="c1"># required</span>
|
|
<span class="s1">'progress_bar'</span><span class="p">:</span> <span class="p">{</span><span class="s1">'training_loss'</span><span class="p">:</span> <span class="n">loss</span><span class="p">},</span> <span class="c1"># optional (MUST ALL BE TENSORS)</span>
|
|
<span class="s1">'log'</span><span class="p">:</span> <span class="n">logger_logs</span>
|
|
<span class="p">}</span>
|
|
|
|
<span class="c1"># return a dict</span>
|
|
<span class="k">return</span> <span class="n">output</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you define multiple optimizers, this step will also be called with an additional <code>optimizer_idx</code> param. </p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># Multiple optimizers (ie: GANs) </span>
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer_idx</span><span class="p">):</span>
|
|
<span class="k">if</span> <span class="n">optimizer_idx</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
|
|
<span class="c1"># do training_step with encoder</span>
|
|
<span class="k">if</span> <span class="n">optimizer_idx</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
|
|
<span class="c1"># do training_step with decoder </span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you add truncated back propagation through time you will also get an additional argument with the hidden states of the previous step. </p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># Truncated back-propagation through time </span>
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">hiddens</span><span class="p">):</span>
|
|
<span class="c1"># hiddens are the hiddens from the previous truncated backprop step</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>You can also return a -1 instead of a dict to stop the current loop. This is useful if you want to
|
|
break out of the current training epoch early.</p>
|
|
<hr />
|
|
<h3 id="training_end">training_end</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">training_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">train_step_outputs</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>In certain cases (dp, ddp2), you might want to use all outputs of every process to do something.
|
|
For instance, if using negative samples, you could run a batch via dp and use ALL the outputs
|
|
for a single softmax across the full batch (ie: the denominator would use the full batch).</p>
|
|
<p>In this case you should define training_end to perform those calculations.</p>
|
|
<p><strong>Params</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Param</th>
|
|
<th>description</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>outputs</td>
|
|
<td>What you return in training_step.</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Return</strong> </p>
|
|
<p>Dictionary or OrderedDict </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>key</th>
|
|
<th>value</th>
|
|
<th>is required</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>loss</td>
|
|
<td>tensor scalar</td>
|
|
<td>Y</td>
|
|
</tr>
|
|
<tr>
|
|
<td>progress_bar</td>
|
|
<td>Dict for progress bar display. Must have only tensors</td>
|
|
<td>N</td>
|
|
</tr>
|
|
<tr>
|
|
<td>log</td>
|
|
<td>Dict of metrics to add to logger. Must have only tensors (no images, etc)</td>
|
|
<td>N</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22
|
|
23
|
|
24
|
|
25
|
|
26
|
|
27
|
|
28</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># WITHOUT training_end</span>
|
|
<span class="c1"># if used in DP or DDP2, this batch is 1/nb_gpus large</span>
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="c1"># batch is 1/nb_gpus big</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
|
|
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="n">nce_loss</span><span class="p">(</span><span class="n">loss</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'loss'</span><span class="p">:</span> <span class="n">loss</span><span class="p">}</span>
|
|
|
|
<span class="c1"># --------------</span>
|
|
<span class="c1"># with training_end to do softmax over the full batch</span>
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="c1"># batch is 1/nb_gpus big</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
|
|
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'out'</span><span class="p">:</span> <span class="n">out</span><span class="p">}</span>
|
|
|
|
<span class="k">def</span> <span class="nf">training_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="c1"># this out is now the full size of the batch</span>
|
|
<span class="n">out</span> <span class="o">=</span> <span class="n">outputs</span><span class="p">[</span><span class="s1">'out'</span><span class="p">]</span>
|
|
|
|
<span class="c1"># this softmax now uses the full batch size</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">softmax</span><span class="p">(</span><span class="n">out</span><span class="p">)</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="n">nce_loss</span><span class="p">(</span><span class="n">loss</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">{</span><span class="s1">'loss'</span><span class="p">:</span> <span class="n">loss</span><span class="p">}</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you define multiple optimizers, this step will also be called with an additional <code>optimizer_idx</code> param. </p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># Multiple optimizers (ie: GANs) </span>
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer_idx</span><span class="p">):</span>
|
|
<span class="k">if</span> <span class="n">optimizer_idx</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
|
|
<span class="c1"># do training_step with encoder</span>
|
|
<span class="k">if</span> <span class="n">optimizer_idx</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
|
|
<span class="c1"># do training_step with decoder </span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you add truncated back propagation through time you will also get an additional argument with the hidden states of the previous step. </p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># Truncated back-propagation through time </span>
|
|
<span class="k">def</span> <span class="nf">training_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">hiddens</span><span class="p">):</span>
|
|
<span class="c1"># hiddens are the hiddens from the previous truncated backprop step</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>You can also return a -1 instead of a dict to stop the current loop. This is useful if you want to
|
|
break out of the current training epoch early.</p>
|
|
<hr />
|
|
<h3 id="train_dataloader">train_dataloader</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">train_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>Called by lightning during training loop. Make sure to use the @pl.data_loader decorator, this ensures not calling this function until the data are needed. <br />
|
|
If you want to change the data during every epoch DON'T use the data_loader decorator.</p>
|
|
<h5 id="return">Return</h5>
|
|
<p>PyTorch DataLoader</p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">train_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="n">transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Normalize</span><span class="p">((</span><span class="mf">0.5</span><span class="p">,),</span> <span class="p">(</span><span class="mf">1.0</span><span class="p">,))])</span>
|
|
<span class="n">dataset</span> <span class="o">=</span> <span class="n">MNIST</span><span class="p">(</span><span class="n">root</span><span class="o">=</span><span class="s1">'/path/to/mnist/'</span><span class="p">,</span> <span class="n">train</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transform</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
|
|
<span class="n">loader</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">DataLoader</span><span class="p">(</span>
|
|
<span class="n">dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
|
|
<span class="n">batch_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">hparams</span><span class="o">.</span><span class="n">batch_size</span><span class="p">,</span>
|
|
<span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
|
|
<span class="p">)</span>
|
|
<span class="k">return</span> <span class="n">loader</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<hr />
|
|
<h3 id="configure_optimizers">configure_optimizers</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>Set up as many optimizers and (optionally) learning rate schedulers as you need. Normally you'd need one. But in the case of GANs or something more esoteric you might have multiple.
|
|
Lightning will call .backward() and .step() on each one in every epoch. If you use 16 bit precision it will also handle that.</p>
|
|
<p><strong>Note:</strong> If you use multiple optimizers, training_step will have an additional <code>optimizer_idx</code> parameter. <br />
|
|
<strong>Note 2:</strong> If you use LBFGS lightning handles the closure function automatically for you.</p>
|
|
<h5 id="return_1">Return</h5>
|
|
<p>Return any of these 3 options: <br />
|
|
Single optimizer <br />
|
|
List or Tuple - List of optimizers <br />
|
|
Two lists - The first list has multiple optimizers, the second a list of learning-rate schedulers</p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># most cases</span>
|
|
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="n">opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.01</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="n">opt</span>
|
|
|
|
<span class="c1"># multiple optimizer case (eg: GAN)</span>
|
|
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="n">generator_opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model_gen</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.01</span><span class="p">)</span>
|
|
<span class="n">disriminator_opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model_disc</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.02</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="n">generator_opt</span><span class="p">,</span> <span class="n">disriminator_opt</span>
|
|
|
|
<span class="c1"># example with learning_rate schedulers </span>
|
|
<span class="k">def</span> <span class="nf">configure_optimizers</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="n">generator_opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model_gen</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.01</span><span class="p">)</span>
|
|
<span class="n">disriminator_opt</span> <span class="o">=</span> <span class="n">Adam</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">model_disc</span><span class="o">.</span><span class="n">parameters</span><span class="p">(),</span> <span class="n">lr</span><span class="o">=</span><span class="mf">0.02</span><span class="p">)</span>
|
|
<span class="n">discriminator_sched</span> <span class="o">=</span> <span class="n">CosineAnnealing</span><span class="p">(</span><span class="n">discriminator_opt</span><span class="p">,</span> <span class="n">T_max</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="p">[</span><span class="n">generator_opt</span><span class="p">,</span> <span class="n">disriminator_opt</span><span class="p">],</span> <span class="p">[</span><span class="n">discriminator_sched</span><span class="p">]</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you need to control how often those optimizers step or override the default .step() schedule, override
|
|
the <a href="https://williamfalcon.github.io/pytorch-lightning/Trainer/hooks/#optimizer_step">optimizer_step</a> hook. </p>
|
|
<h2 id="optional-methods">Optional Methods</h2>
|
|
<h3 id="validation_step">validation_step</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3
|
|
4
|
|
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># if you have one val dataloader:</span>
|
|
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">)</span>
|
|
|
|
<span class="c1"># if you have multiple val dataloaders: </span>
|
|
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataloader_idxdx</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p><strong>OPTIONAL</strong> <br />
|
|
If you don't need to validate you don't need to implement this method. In this step you'd normally generate examples or calculate anything of interest such as accuracy. </p>
|
|
<p>When the validation_step is called, the model has been put in eval mode and PyTorch gradients have been disabled. At the end of validation, model goes back to training mode and gradients are enabled.</p>
|
|
<p>The dict you return here will be available in the <code>validation_end</code> method. </p>
|
|
<p><strong>Params</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Param</th>
|
|
<th>description</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>batch</td>
|
|
<td>The output of your dataloader. A tensor, tuple or list</td>
|
|
</tr>
|
|
<tr>
|
|
<td>batch_nb</td>
|
|
<td>Integer displaying which batch this is</td>
|
|
</tr>
|
|
<tr>
|
|
<td>dataloader_idx</td>
|
|
<td>Integer displaying which dataloader this is (only if multiple val datasets used)</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Return</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Return</th>
|
|
<th>description</th>
|
|
<th>optional</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>dict</td>
|
|
<td>Dict or OrderedDict - passed to the validation_end step</td>
|
|
<td>N</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22
|
|
23
|
|
24
|
|
25
|
|
26
|
|
27</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 1: A single validation dataset</span>
|
|
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
|
|
<span class="c1"># implement your own</span>
|
|
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
|
|
|
<span class="c1"># log 6 example images</span>
|
|
<span class="c1"># or generated text... or whatever</span>
|
|
<span class="n">sample_imgs</span> <span class="o">=</span> <span class="n">x</span><span class="p">[:</span><span class="mi">6</span><span class="p">]</span>
|
|
<span class="n">grid</span> <span class="o">=</span> <span class="n">torchvision</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">make_grid</span><span class="p">(</span><span class="n">sample_imgs</span><span class="p">)</span>
|
|
<span class="bp">self</span><span class="o">.</span><span class="n">logger</span><span class="o">.</span><span class="n">experiment</span><span class="o">.</span><span class="n">add_image</span><span class="p">(</span><span class="s1">'example_images'</span><span class="p">,</span> <span class="n">grid</span><span class="p">,</span> <span class="mi">0</span><span class="p">)</span>
|
|
|
|
<span class="c1"># calculate acc</span>
|
|
<span class="n">labels_hat</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
|
|
<span class="n">val_acc</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">y</span> <span class="o">==</span> <span class="n">labels_hat</span><span class="p">)</span><span class="o">.</span><span class="n">item</span><span class="p">()</span> <span class="o">/</span> <span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">y</span><span class="p">)</span> <span class="o">*</span> <span class="mf">1.0</span><span class="p">)</span>
|
|
|
|
<span class="c1"># all optional...</span>
|
|
<span class="c1"># return whatever you need for the collation function validation_end</span>
|
|
<span class="n">output</span> <span class="o">=</span> <span class="n">OrderedDict</span><span class="p">({</span>
|
|
<span class="s1">'val_loss'</span><span class="p">:</span> <span class="n">loss_val</span><span class="p">,</span>
|
|
<span class="s1">'val_acc'</span><span class="p">:</span> <span class="n">torch</span><span class="o">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">val_acc</span><span class="p">),</span> <span class="c1"># everything must be a tensor</span>
|
|
<span class="p">})</span>
|
|
|
|
<span class="c1"># return an optional dict</span>
|
|
<span class="k">return</span> <span class="n">output</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you pass in multiple validation datasets, validation_step will have an additional argument.</p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 2: multiple validation datasets</span>
|
|
<span class="k">def</span> <span class="nf">validation_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataset_idx</span><span class="p">):</span>
|
|
<span class="c1"># dataset_idx tells you which dataset this is. </span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>The <code>dataset_idx</code> corresponds to the order of datasets returned in <code>val_dataloader</code>. </p>
|
|
<hr />
|
|
<h3 id="validation_end">validation_end</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you didn't define a validation_step, this won't be called. Called at the end of the validation loop with the outputs of validation_step.</p>
|
|
<p>The outputs here are strictly for the progress bar. If you don't need to display anything, don't return anything. <br />
|
|
Any keys present in 'log', 'progress_bar' or the rest of the dictionary are available for callbacks to access.
|
|
<strong>Params</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Param</th>
|
|
<th>description</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>outputs</td>
|
|
<td>List of outputs you defined in validation_step, or if there are multiple dataloaders, a list containing a list of outputs for each dataloader</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Return</strong> </p>
|
|
<p>Dictionary or OrderedDict </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>key</th>
|
|
<th>value</th>
|
|
<th>is required</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>progress_bar</td>
|
|
<td>Dict for progress bar display. Must have only tensors</td>
|
|
<td>N</td>
|
|
</tr>
|
|
<tr>
|
|
<td>log</td>
|
|
<td>Dict of metrics to add to logger. Must have only tensors (no images, etc)</td>
|
|
<td>N</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Example</strong></p>
|
|
<p>With a single dataloader</p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="sd">"""</span>
|
|
<span class="sd"> Called at the end of validation to aggregate outputs</span>
|
|
<span class="sd"> :param outputs: list of individual outputs of each validation step</span>
|
|
<span class="sd"> :return:</span>
|
|
<span class="sd"> """</span>
|
|
<span class="n">val_loss_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="n">val_acc_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="k">for</span> <span class="n">output</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">:</span>
|
|
<span class="n">val_loss_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'val_loss'</span><span class="p">]</span>
|
|
<span class="n">val_acc_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'val_acc'</span><span class="p">]</span>
|
|
|
|
<span class="n">val_loss_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
|
|
<span class="n">val_acc_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
|
|
<span class="n">tqdm_dict</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'val_loss'</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">(),</span> <span class="s1">'val_acc'</span><span class="p">:</span> <span class="n">val_acc_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
|
|
<span class="c1"># show val_loss and val_acc in progress bar but only log val_loss</span>
|
|
<span class="n">results</span> <span class="o">=</span> <span class="p">{</span>
|
|
<span class="s1">'progress_bar'</span><span class="p">:</span> <span class="n">tqdm_dict</span><span class="p">,</span>
|
|
<span class="s1">'log'</span><span class="p">:</span> <span class="p">{</span><span class="s1">'val_loss'</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
<span class="p">}</span>
|
|
<span class="k">return</span> <span class="n">results</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>With multiple dataloaders, <code>outputs</code> will be a list of lists. The outer list contains
|
|
one entry per dataloader, while the inner list contains the individual outputs of
|
|
each validation step for that dataloader.</p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22
|
|
23
|
|
24
|
|
25</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">validation_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="sd">"""</span>
|
|
<span class="sd"> Called at the end of validation to aggregate outputs</span>
|
|
<span class="sd"> :param outputs: list of list of individual outputs of each validation step</span>
|
|
<span class="sd"> :return:</span>
|
|
<span class="sd"> """</span>
|
|
<span class="n">val_loss_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="n">val_acc_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="n">i</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="k">for</span> <span class="n">dataloader_outputs</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">:</span>
|
|
<span class="k">for</span> <span class="n">output</span> <span class="ow">in</span> <span class="n">dataloader_outputs</span><span class="p">:</span>
|
|
<span class="n">val_loss_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'val_loss'</span><span class="p">]</span>
|
|
<span class="n">val_acc_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'val_acc'</span><span class="p">]</span>
|
|
<span class="n">i</span> <span class="o">+=</span> <span class="mi">1</span>
|
|
|
|
<span class="n">val_loss_mean</span> <span class="o">/=</span> <span class="n">i</span>
|
|
<span class="n">val_acc_mean</span> <span class="o">/=</span> <span class="n">i</span>
|
|
<span class="n">tqdm_dict</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'val_loss'</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">(),</span> <span class="s1">'val_acc'</span><span class="p">:</span> <span class="n">val_acc_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
|
|
<span class="c1"># show val_loss and val_acc in progress bar but only log val_loss</span>
|
|
<span class="n">results</span> <span class="o">=</span> <span class="p">{</span>
|
|
<span class="s1">'progress_bar'</span><span class="p">:</span> <span class="n">tqdm_dict</span><span class="p">,</span>
|
|
<span class="s1">'log'</span><span class="p">:</span> <span class="p">{</span><span class="s1">'val_loss'</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
<span class="p">}</span>
|
|
<span class="k">return</span> <span class="n">results</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<h3 id="test_step">test_step</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3
|
|
4
|
|
5</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># if you have one test dataloader:</span>
|
|
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">)</span>
|
|
|
|
<span class="c1"># if you have multiple test dataloaders: </span>
|
|
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataloader_idxdx</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p><strong>OPTIONAL</strong> <br />
|
|
If you don't need to test you don't need to implement this method. In this step you'd normally generate examples or calculate anything of interest such as accuracy. </p>
|
|
<p>When the validation_step is called, the model has been put in eval mode and PyTorch gradients have been disabled. At the end of validation, model goes back to training mode and gradients are enabled.</p>
|
|
<p>The dict you return here will be available in the <code>test_end</code> method. </p>
|
|
<p>This function is used when you execute <code>trainer.test()</code>.</p>
|
|
<p><strong>Params</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Param</th>
|
|
<th>description</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>batch</td>
|
|
<td>The output of your dataloader. A tensor, tuple or list</td>
|
|
</tr>
|
|
<tr>
|
|
<td>batch_nb</td>
|
|
<td>Integer displaying which batch this is</td>
|
|
</tr>
|
|
<tr>
|
|
<td>dataloader_idx</td>
|
|
<td>Integer displaying which dataloader this is (only if multiple test datasets used)</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Return</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Return</th>
|
|
<th>description</th>
|
|
<th>optional</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>dict</td>
|
|
<td>Dict or OrderedDict with metrics to display in progress bar. All keys must be tensors.</td>
|
|
<td>Y</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 1: A single test dataset</span>
|
|
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">):</span>
|
|
<span class="n">x</span><span class="p">,</span> <span class="n">y</span> <span class="o">=</span> <span class="n">batch</span>
|
|
|
|
<span class="c1"># implement your own</span>
|
|
<span class="n">out</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
|
<span class="n">loss</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">loss</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
|
|
|
|
<span class="c1"># calculate acc</span>
|
|
<span class="n">labels_hat</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">argmax</span><span class="p">(</span><span class="n">out</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
|
|
<span class="n">test_acc</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">y</span> <span class="o">==</span> <span class="n">labels_hat</span><span class="p">)</span><span class="o">.</span><span class="n">item</span><span class="p">()</span> <span class="o">/</span> <span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">y</span><span class="p">)</span> <span class="o">*</span> <span class="mf">1.0</span><span class="p">)</span>
|
|
|
|
<span class="c1"># all optional...</span>
|
|
<span class="c1"># return whatever you need for the collation function test_end</span>
|
|
<span class="n">output</span> <span class="o">=</span> <span class="n">OrderedDict</span><span class="p">({</span>
|
|
<span class="s1">'test_loss'</span><span class="p">:</span> <span class="n">loss_test</span><span class="p">,</span>
|
|
<span class="s1">'test_acc'</span><span class="p">:</span> <span class="n">torch</span><span class="o">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">test_acc</span><span class="p">),</span> <span class="c1"># everything must be a tensor</span>
|
|
<span class="p">})</span>
|
|
|
|
<span class="c1"># return an optional dict</span>
|
|
<span class="k">return</span> <span class="n">output</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you pass in multiple test datasets, test_step will have an additional argument.</p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="c1"># CASE 2: multiple test datasets</span>
|
|
<span class="k">def</span> <span class="nf">test_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">batch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">dataset_idx</span><span class="p">):</span>
|
|
<span class="c1"># dataset_idx tells you which dataset this is. </span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>The <code>dataset_idx</code> corresponds to the order of datasets returned in <code>test_dataloader</code>. </p>
|
|
<hr />
|
|
<h3 id="test_end">test_end</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">test_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>If you didn't define a test_step, this won't be called. </p>
|
|
<p>Called at the end of the test step with the output of each test_step.</p>
|
|
<p>The outputs here are strictly for the progress bar. If you don't need to display anything, don't return anything. </p>
|
|
<p><strong>Params</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Param</th>
|
|
<th>description</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>outputs</td>
|
|
<td>List of outputs you defined in test_step, or if there are multiple dataloaders, a list containing a list of outputs for each dataloader</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Return</strong> </p>
|
|
<table>
|
|
<thead>
|
|
<tr>
|
|
<th>Return</th>
|
|
<th>description</th>
|
|
<th>optional</th>
|
|
</tr>
|
|
</thead>
|
|
<tbody>
|
|
<tr>
|
|
<td>dict</td>
|
|
<td>Dict of OrderedDict with metrics to display in progress bar</td>
|
|
<td>Y</td>
|
|
</tr>
|
|
</tbody>
|
|
</table>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">test_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="sd">"""</span>
|
|
<span class="sd"> Called at the end of test to aggregate outputs</span>
|
|
<span class="sd"> :param outputs: list of individual outputs of each test step</span>
|
|
<span class="sd"> :return:</span>
|
|
<span class="sd"> """</span>
|
|
<span class="n">test_loss_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="n">test_acc_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="k">for</span> <span class="n">output</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">:</span>
|
|
<span class="n">test_loss_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'test_loss'</span><span class="p">]</span>
|
|
<span class="n">test_acc_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'test_acc'</span><span class="p">]</span>
|
|
|
|
<span class="n">test_loss_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
|
|
<span class="n">test_acc_mean</span> <span class="o">/=</span> <span class="nb">len</span><span class="p">(</span><span class="n">outputs</span><span class="p">)</span>
|
|
<span class="n">tqdm_dict</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'test_loss'</span><span class="p">:</span> <span class="n">test_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">(),</span> <span class="s1">'test_acc'</span><span class="p">:</span> <span class="n">test_acc_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
|
|
<span class="c1"># show test_loss and test_acc in progress bar but only log test_loss</span>
|
|
<span class="n">results</span> <span class="o">=</span> <span class="p">{</span>
|
|
<span class="s1">'progress_bar'</span><span class="p">:</span> <span class="n">tqdm_dict</span><span class="p">,</span>
|
|
<span class="s1">'log'</span><span class="p">:</span> <span class="p">{</span><span class="s1">'test_loss'</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
<span class="p">}</span>
|
|
<span class="k">return</span> <span class="n">results</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>With multiple dataloaders, <code>outputs</code> will be a list of lists. The outer list contains
|
|
one entry per dataloader, while the inner list contains the individual outputs of
|
|
each validation step for that dataloader.</p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22
|
|
23
|
|
24
|
|
25</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">test_end</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">outputs</span><span class="p">):</span>
|
|
<span class="sd">"""</span>
|
|
<span class="sd"> Called at the end of test to aggregate outputs</span>
|
|
<span class="sd"> :param outputs: list of individual outputs of each test step</span>
|
|
<span class="sd"> :return:</span>
|
|
<span class="sd"> """</span>
|
|
<span class="n">test_loss_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="n">test_acc_mean</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="n">i</span> <span class="o">=</span> <span class="mi">0</span>
|
|
<span class="k">for</span> <span class="n">dataloader_outputs</span> <span class="ow">in</span> <span class="n">outputs</span><span class="p">:</span>
|
|
<span class="k">for</span> <span class="n">output</span> <span class="ow">in</span> <span class="n">dataloader_outputs</span><span class="p">:</span>
|
|
<span class="n">test_loss_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'test_loss'</span><span class="p">]</span>
|
|
<span class="n">test_acc_mean</span> <span class="o">+=</span> <span class="n">output</span><span class="p">[</span><span class="s1">'test_acc'</span><span class="p">]</span>
|
|
<span class="n">i</span> <span class="o">+=</span> <span class="mi">1</span>
|
|
|
|
<span class="n">test_loss_mean</span> <span class="o">/=</span> <span class="n">i</span>
|
|
<span class="n">test_acc_mean</span> <span class="o">/=</span> <span class="n">i</span>
|
|
<span class="n">tqdm_dict</span> <span class="o">=</span> <span class="p">{</span><span class="s1">'test_loss'</span><span class="p">:</span> <span class="n">test_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">(),</span> <span class="s1">'test_acc'</span><span class="p">:</span> <span class="n">test_acc_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
|
|
<span class="c1"># show test_loss and test_acc in progress bar but only log test_loss</span>
|
|
<span class="n">results</span> <span class="o">=</span> <span class="p">{</span>
|
|
<span class="s1">'progress_bar'</span><span class="p">:</span> <span class="n">tqdm_dict</span><span class="p">,</span>
|
|
<span class="s1">'log'</span><span class="p">:</span> <span class="p">{</span><span class="s1">'test_loss'</span><span class="p">:</span> <span class="n">val_loss_mean</span><span class="o">.</span><span class="n">item</span><span class="p">()}</span>
|
|
<span class="p">}</span>
|
|
<span class="k">return</span> <span class="n">results</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<hr />
|
|
<h3 id="on_save_checkpoint">on_save_checkpoint</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_save_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>Called by lightning to checkpoint your model. Lightning saves the training state (current epoch, global_step, etc)
|
|
and also saves the model state_dict. If you want to save anything else, use this method to add your own
|
|
key-value pair.</p>
|
|
<h5 id="return_2">Return</h5>
|
|
<p>Nothing</p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_save_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">):</span>
|
|
<span class="c1"># 99% of use cases you don't need to implement this method </span>
|
|
<span class="n">checkpoint</span><span class="p">[</span><span class="s1">'something_cool_i_want_to_save'</span><span class="p">]</span> <span class="o">=</span> <span class="n">my_cool_pickable_object</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<hr />
|
|
<h3 id="on_load_checkpoint">on_load_checkpoint</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_load_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>Called by lightning to restore your model. Lighting auto-restores global step, epoch, etc...
|
|
It also restores the model state_dict.
|
|
If you saved something with <strong>on_save_checkpoint</strong> this is your chance to restore this.</p>
|
|
<h5 id="return_3">Return</h5>
|
|
<p>Nothing </p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2
|
|
3</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_load_checkpoint</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">checkpoint</span><span class="p">):</span>
|
|
<span class="c1"># 99% of the time you don't need to implement this method</span>
|
|
<span class="bp">self</span><span class="o">.</span><span class="n">something_cool_i_want_to_save</span> <span class="o">=</span> <span class="n">checkpoint</span><span class="p">[</span><span class="s1">'something_cool_i_want_to_save'</span><span class="p">]</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<hr />
|
|
<h3 id="val_dataloader">val_dataloader</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p><strong>OPTIONAL</strong> <br />
|
|
If you don't need a validation dataset and a validation_step, you don't need to implement this method. </p>
|
|
<p>Called by lightning during validation loop. Make sure to use the @pl.data_loader decorator, this ensures not calling this function until the data are needed. <br />
|
|
If you want to change the data during every epoch DON'T use the data_loader decorator. </p>
|
|
<h5 id="return_4">Return</h5>
|
|
<p>PyTorch DataLoader or list of PyTorch Dataloaders. </p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="n">transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Normalize</span><span class="p">((</span><span class="mf">0.5</span><span class="p">,),</span> <span class="p">(</span><span class="mf">1.0</span><span class="p">,))])</span>
|
|
<span class="n">dataset</span> <span class="o">=</span> <span class="n">MNIST</span><span class="p">(</span><span class="n">root</span><span class="o">=</span><span class="s1">'/path/to/mnist/'</span><span class="p">,</span> <span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transform</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
|
|
<span class="n">loader</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">DataLoader</span><span class="p">(</span>
|
|
<span class="n">dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
|
|
<span class="n">batch_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">hparams</span><span class="o">.</span><span class="n">batch_size</span><span class="p">,</span>
|
|
<span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
|
|
<span class="p">)</span>
|
|
|
|
<span class="k">return</span> <span class="n">loader</span>
|
|
|
|
<span class="c1"># can also return multiple dataloaders </span>
|
|
<span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">val_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="k">return</span> <span class="p">[</span><span class="n">loader_a</span><span class="p">,</span> <span class="n">loader_b</span><span class="p">,</span> <span class="o">...</span><span class="p">,</span> <span class="n">loader_n</span><span class="p">]</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>In the case where you return multiple val_dataloaders, the validation_step will have an arguement <code>dataset_idx</code>
|
|
which matches the order here. </p>
|
|
<hr />
|
|
<h3 id="test_dataloader">test_dataloader</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p><strong>OPTIONAL</strong> <br />
|
|
If you don't need a test dataset and a test_step, you don't need to implement this method. </p>
|
|
<p>Called by lightning during test loop. Make sure to use the @pl.data_loader decorator, this ensures not calling this function until the data are needed.
|
|
If you want to change the data during every epoch DON'T use the data_loader decorator. </p>
|
|
<h5 id="return_5">Return</h5>
|
|
<p>PyTorch DataLoader</p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@pl.data_loader</span>
|
|
<span class="k">def</span> <span class="nf">test_dataloader</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
|
|
<span class="n">transform</span> <span class="o">=</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Compose</span><span class="p">([</span><span class="n">transforms</span><span class="o">.</span><span class="n">ToTensor</span><span class="p">(),</span> <span class="n">transforms</span><span class="o">.</span><span class="n">Normalize</span><span class="p">((</span><span class="mf">0.5</span><span class="p">,),</span> <span class="p">(</span><span class="mf">1.0</span><span class="p">,))])</span>
|
|
<span class="n">dataset</span> <span class="o">=</span> <span class="n">MNIST</span><span class="p">(</span><span class="n">root</span><span class="o">=</span><span class="s1">'/path/to/mnist/'</span><span class="p">,</span> <span class="n">train</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">transform</span><span class="p">,</span> <span class="n">download</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
|
|
<span class="n">loader</span> <span class="o">=</span> <span class="n">torch</span><span class="o">.</span><span class="n">utils</span><span class="o">.</span><span class="n">data</span><span class="o">.</span><span class="n">DataLoader</span><span class="p">(</span>
|
|
<span class="n">dataset</span><span class="o">=</span><span class="n">dataset</span><span class="p">,</span>
|
|
<span class="n">batch_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">hparams</span><span class="o">.</span><span class="n">batch_size</span><span class="p">,</span>
|
|
<span class="n">shuffle</span><span class="o">=</span><span class="bp">True</span>
|
|
<span class="p">)</span>
|
|
|
|
<span class="k">return</span> <span class="n">loader</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<hr />
|
|
<h3 id="add_model_specific_args">add_model_specific_args</h3>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
|
|
2</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@staticmethod</span>
|
|
<span class="k">def</span> <span class="nf">add_model_specific_args</span><span class="p">(</span><span class="n">parent_parser</span><span class="p">,</span> <span class="n">root_dir</span><span class="p">)</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
<p>Lightning has a list of default argparse commands.
|
|
This method is your chance to add or modify commands specific to your model.
|
|
The <a href="https://williamfalcon.github.io/test-tube/hyperparameter_optimization/HyperOptArgumentParser/">hyperparameter argument parser</a> is available anywhere in your model by calling self.hparams.</p>
|
|
<h5 id="return_6">Return</h5>
|
|
<p>An argument parser</p>
|
|
<p><strong>Example</strong></p>
|
|
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre> 1
|
|
2
|
|
3
|
|
4
|
|
5
|
|
6
|
|
7
|
|
8
|
|
9
|
|
10
|
|
11
|
|
12
|
|
13
|
|
14
|
|
15
|
|
16
|
|
17
|
|
18
|
|
19
|
|
20
|
|
21
|
|
22</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="nd">@staticmethod</span>
|
|
<span class="k">def</span> <span class="nf">add_model_specific_args</span><span class="p">(</span><span class="n">parent_parser</span><span class="p">,</span> <span class="n">root_dir</span><span class="p">):</span>
|
|
<span class="n">parser</span> <span class="o">=</span> <span class="n">HyperOptArgumentParser</span><span class="p">(</span><span class="n">strategy</span><span class="o">=</span><span class="n">parent_parser</span><span class="o">.</span><span class="n">strategy</span><span class="p">,</span> <span class="n">parents</span><span class="o">=</span><span class="p">[</span><span class="n">parent_parser</span><span class="p">])</span>
|
|
|
|
<span class="c1"># param overwrites</span>
|
|
<span class="c1"># parser.set_defaults(gradient_clip_val=5.0)</span>
|
|
|
|
<span class="c1"># network params</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">'--drop_prob'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mf">0.2</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">],</span> <span class="nb">type</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">'--in_features'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">28</span><span class="o">*</span><span class="mi">28</span><span class="p">)</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">'--out_features'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">'--hidden_dim'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">50000</span><span class="p">)</span> <span class="c1"># use 500 for CPU, 50000 for GPU to see speed difference</span>
|
|
|
|
<span class="c1"># data</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">add_argument</span><span class="p">(</span><span class="s1">'--data_root'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="n">os</span><span class="o">.</span><span class="n">path</span><span class="o">.</span><span class="n">join</span><span class="p">(</span><span class="n">root_dir</span><span class="p">,</span> <span class="s1">'mnist'</span><span class="p">),</span> <span class="nb">type</span><span class="o">=</span><span class="nb">str</span><span class="p">)</span>
|
|
|
|
<span class="c1"># training params (opt)</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">'--learning_rate'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mf">0.0001</span><span class="p">,</span> <span class="mf">0.0005</span><span class="p">,</span> <span class="mf">0.001</span><span class="p">,</span> <span class="mf">0.005</span><span class="p">],</span>
|
|
<span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">'--batch_size'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="mi">256</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">int</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="mi">32</span><span class="p">,</span> <span class="mi">64</span><span class="p">,</span> <span class="mi">128</span><span class="p">,</span> <span class="mi">256</span><span class="p">],</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
|
|
<span class="n">parser</span><span class="o">.</span><span class="n">opt_list</span><span class="p">(</span><span class="s1">'--optimizer_name'</span><span class="p">,</span> <span class="n">default</span><span class="o">=</span><span class="s1">'adam'</span><span class="p">,</span> <span class="nb">type</span><span class="o">=</span><span class="nb">str</span><span class="p">,</span> <span class="n">options</span><span class="o">=</span><span class="p">[</span><span class="s1">'adam'</span><span class="p">],</span> <span class="n">tunable</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
|
|
<span class="k">return</span> <span class="n">parser</span>
|
|
</pre></div>
|
|
</td></tr></table>
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
</article>
|
|
</div>
|
|
</div>
|
|
</main>
|
|
|
|
|
|
<footer class="md-footer">
|
|
|
|
<div class="md-footer-nav">
|
|
<nav class="md-footer-nav__inner md-grid">
|
|
|
|
<a href="../.." title="Home" class="md-flex md-footer-nav__link md-footer-nav__link--prev" rel="prev">
|
|
<div class="md-flex__cell md-flex__cell--shrink">
|
|
<i class="md-icon md-icon--arrow-back md-footer-nav__button"></i>
|
|
</div>
|
|
<div class="md-flex__cell md-flex__cell--stretch md-footer-nav__title">
|
|
<span class="md-flex__ellipsis">
|
|
<span class="md-footer-nav__direction">
|
|
Previous
|
|
</span>
|
|
Home
|
|
</span>
|
|
</div>
|
|
</a>
|
|
|
|
|
|
<a href="../methods/" title="Methods" class="md-flex md-footer-nav__link md-footer-nav__link--next" rel="next">
|
|
<div class="md-flex__cell md-flex__cell--stretch md-footer-nav__title">
|
|
<span class="md-flex__ellipsis">
|
|
<span class="md-footer-nav__direction">
|
|
Next
|
|
</span>
|
|
Methods
|
|
</span>
|
|
</div>
|
|
<div class="md-flex__cell md-flex__cell--shrink">
|
|
<i class="md-icon md-icon--arrow-forward md-footer-nav__button"></i>
|
|
</div>
|
|
</a>
|
|
|
|
</nav>
|
|
</div>
|
|
|
|
<div class="md-footer-meta md-typeset">
|
|
<div class="md-footer-meta__inner md-grid">
|
|
<div class="md-footer-copyright">
|
|
|
|
powered by
|
|
<a href="https://www.mkdocs.org">MkDocs</a>
|
|
and
|
|
<a href="https://squidfunk.github.io/mkdocs-material/">
|
|
Material for MkDocs</a>
|
|
</div>
|
|
|
|
</div>
|
|
</div>
|
|
</footer>
|
|
|
|
</div>
|
|
|
|
<script src="../../assets/javascripts/application.245445c6.js"></script>
|
|
|
|
<script>app.initialize({version:"1.0.4",url:{base:"../.."}})</script>
|
|
|
|
|
|
</body>
|
|
</html> |