Files

1195 lines
45 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>Hooks - 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="#hooks" 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">
Hooks
</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">
&#xE5CD;
</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--nested">
<input class="md-toggle md-nav__toggle" data-md-toggle="nav-2" type="checkbox" id="nav-2">
<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">
<a href="../../LightningModule/RequiredTrainerInterface/" title="Lightning Module interface" class="md-nav__link">
Lightning Module interface
</a>
</li>
<li class="md-nav__item">
<a href="../../LightningModule/methods/" title="Methods" class="md-nav__link">
Methods
</a>
</li>
<li class="md-nav__item">
<a href="../../LightningModule/properties/" title="Properties" class="md-nav__link">
Properties
</a>
</li>
</ul>
</nav>
</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-3" type="checkbox" id="nav-3" checked>
<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="../" title="Trainer" class="md-nav__link">
Trainer
</a>
</li>
<li class="md-nav__item">
<a href="../Checkpointing/" title="Checkpointing" class="md-nav__link">
Checkpointing
</a>
</li>
<li class="md-nav__item">
<a href="../Distributed training/" title="Distributed training" class="md-nav__link">
Distributed training
</a>
</li>
<li class="md-nav__item">
<a href="../Logging/" title="Logging" class="md-nav__link">
Logging
</a>
</li>
<li class="md-nav__item">
<a href="../SLURM Managed Cluster/" title="SLURM Managed Cluster" class="md-nav__link">
SLURM Managed Cluster
</a>
</li>
<li class="md-nav__item">
<a href="../Testing loop/" title="Testing loop" class="md-nav__link">
Testing loop
</a>
</li>
<li class="md-nav__item">
<a href="../Training Loop/" title="Training Loop" class="md-nav__link">
Training Loop
</a>
</li>
<li class="md-nav__item">
<a href="../Validation loop/" title="Validation loop" class="md-nav__link">
Validation loop
</a>
</li>
<li class="md-nav__item">
<a href="../debugging/" title="Debugging" class="md-nav__link">
Debugging
</a>
</li>
<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">
Hooks
</label>
<a href="./" title="Hooks" class="md-nav__link md-nav__link--active">
Hooks
</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="#on_epoch_start" title="on_epoch_start" class="md-nav__link">
on_epoch_start
</a>
</li>
<li class="md-nav__item">
<a href="#on_epoch_end" title="on_epoch_end" class="md-nav__link">
on_epoch_end
</a>
</li>
<li class="md-nav__item">
<a href="#on_batch_start" title="on_batch_start" class="md-nav__link">
on_batch_start
</a>
</li>
<li class="md-nav__item">
<a href="#on_batch_end" title="on_batch_end" class="md-nav__link">
on_batch_end
</a>
</li>
<li class="md-nav__item">
<a href="#on_pre_performance_check" title="on_pre_performance_check" class="md-nav__link">
on_pre_performance_check
</a>
</li>
<li class="md-nav__item">
<a href="#on_post_performance_check" title="on_post_performance_check" class="md-nav__link">
on_post_performance_check
</a>
</li>
<li class="md-nav__item">
<a href="#optimizer_step" title="optimizer_step" class="md-nav__link">
optimizer_step
</a>
</li>
<li class="md-nav__item">
<a href="#on_before_zero_grad" title="on_before_zero_grad" class="md-nav__link">
on_before_zero_grad
</a>
</li>
<li class="md-nav__item">
<a href="#backward" title="backward" class="md-nav__link">
backward
</a>
</li>
<li class="md-nav__item">
<a href="#on_after_backward" title="on_after_backward" class="md-nav__link">
on_after_backward
</a>
</li>
<li class="md-nav__item">
<a href="#tbptt_split_batch" title="tbptt_split_batch" class="md-nav__link">
tbptt_split_batch
</a>
</li>
<li class="md-nav__item">
<a href="#configure_apex" title="configure_apex" class="md-nav__link">
configure_apex
</a>
</li>
<li class="md-nav__item">
<a href="#configure_ddp" title="configure_ddp" class="md-nav__link">
configure_ddp
</a>
</li>
<li class="md-nav__item">
<a href="#init_ddp_connection" title="init_ddp_connection" class="md-nav__link">
init_ddp_connection
</a>
</li>
</ul>
</nav>
</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="#on_epoch_start" title="on_epoch_start" class="md-nav__link">
on_epoch_start
</a>
</li>
<li class="md-nav__item">
<a href="#on_epoch_end" title="on_epoch_end" class="md-nav__link">
on_epoch_end
</a>
</li>
<li class="md-nav__item">
<a href="#on_batch_start" title="on_batch_start" class="md-nav__link">
on_batch_start
</a>
</li>
<li class="md-nav__item">
<a href="#on_batch_end" title="on_batch_end" class="md-nav__link">
on_batch_end
</a>
</li>
<li class="md-nav__item">
<a href="#on_pre_performance_check" title="on_pre_performance_check" class="md-nav__link">
on_pre_performance_check
</a>
</li>
<li class="md-nav__item">
<a href="#on_post_performance_check" title="on_post_performance_check" class="md-nav__link">
on_post_performance_check
</a>
</li>
<li class="md-nav__item">
<a href="#optimizer_step" title="optimizer_step" class="md-nav__link">
optimizer_step
</a>
</li>
<li class="md-nav__item">
<a href="#on_before_zero_grad" title="on_before_zero_grad" class="md-nav__link">
on_before_zero_grad
</a>
</li>
<li class="md-nav__item">
<a href="#backward" title="backward" class="md-nav__link">
backward
</a>
</li>
<li class="md-nav__item">
<a href="#on_after_backward" title="on_after_backward" class="md-nav__link">
on_after_backward
</a>
</li>
<li class="md-nav__item">
<a href="#tbptt_split_batch" title="tbptt_split_batch" class="md-nav__link">
tbptt_split_batch
</a>
</li>
<li class="md-nav__item">
<a href="#configure_apex" title="configure_apex" class="md-nav__link">
configure_apex
</a>
</li>
<li class="md-nav__item">
<a href="#configure_ddp" title="configure_ddp" class="md-nav__link">
configure_ddp
</a>
</li>
<li class="md-nav__item">
<a href="#init_ddp_connection" title="init_ddp_connection" class="md-nav__link">
init_ddp_connection
</a>
</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/Trainer/hooks.md" title="Edit this page" class="md-icon md-content__icon">&#xE3C9;</a>
<h1 id="hooks">Hooks</h1>
<p>[<a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/root_module/hooks.py">Github Code</a>] </p>
<p>There are cases when you might want to do something different at different parts of the training/validation loop.
To enable a hook, simply override the method in your LightningModule and the trainer will call it at the correct time.</p>
<p><strong>Contributing</strong> If there's a hook you'd like to add, simply: <br />
1. Fork PyTorchLightning. <br />
2. Add the hook <a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/root_module/hooks.py">here</a>. <br />
3. Add the correct place in the <a href="https://github.com/williamFalcon/pytorch-lightning/blob/master/pytorch_lightning/models/trainer.py">Trainer</a> where it should be called. </p>
<hr />
<h4 id="on_epoch_start">on_epoch_start</h4>
<p>Called in the training loop at the very beginning of the epoch. </p>
<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="k">def</span> <span class="nf">on_epoch_start</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the epoch starts</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_epoch_end">on_epoch_end</h4>
<p>Called in the training loop at the very end of the epoch. </p>
<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="k">def</span> <span class="nf">on_epoch_end</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the epoch ends </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_batch_start">on_batch_start</h4>
<p>Called in the training loop before anything happens for that batch. </p>
<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="k">def</span> <span class="nf">on_batch_start</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the batch starts</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_batch_end">on_batch_end</h4>
<p>Called in the training loop after the batch. </p>
<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="k">def</span> <span class="nf">on_batch_end</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something when the batch ends </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_pre_performance_check">on_pre_performance_check</h4>
<p>Called at the very beginning of the validation loop. </p>
<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="k">def</span> <span class="nf">on_pre_performance_check</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something before validation starts </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_post_performance_check">on_post_performance_check</h4>
<p>Called at the very end of the validation loop. </p>
<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="k">def</span> <span class="nf">on_post_performance_check</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># do something before validation end</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="optimizer_step">optimizer_step</h4>
<p>Calls .step() and .zero_grad for each optimizer.<br />
You can override this method to adjust how you do the optimizer step for each optimizer</p>
<p>Called once per optimizer</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"># DEFAULT</span>
<span class="k">def</span> <span class="nf">optimizer_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">current_epoch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">optimizer_i</span><span class="p">,</span> <span class="n">second_order_closure</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
<span class="c1"># Alternating schedule for optimizer steps (ie: GANs) </span>
<span class="k">def</span> <span class="nf">optimizer_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">current_epoch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">optimizer_i</span><span class="p">,</span> <span class="n">second_order_closure</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
<span class="c1"># update generator opt every 2 steps</span>
<span class="k">if</span> <span class="n">optimizer_i</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
<span class="k">if</span> <span class="n">batch_nb</span> <span class="o">%</span> <span class="mi">2</span> <span class="o">==</span> <span class="mi">0</span> <span class="p">:</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
<span class="c1"># update discriminator opt every 4 steps</span>
<span class="k">if</span> <span class="n">optimizer_i</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
<span class="k">if</span> <span class="n">batch_nb</span> <span class="o">%</span> <span class="mi">4</span> <span class="o">==</span> <span class="mi">0</span> <span class="p">:</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
<span class="c1"># ...</span>
<span class="c1"># add as many optimizers as you want </span>
</pre></div>
</td></tr></table>
<p>This step allows you to do a lot of non-standard training tricks such as learning-rate warm-up: </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="c1"># learning rate warm-up</span>
<span class="k">def</span> <span class="nf">optimizer_step</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">current_epoch</span><span class="p">,</span> <span class="n">batch_nb</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">,</span> <span class="n">optimizer_i</span><span class="p">,</span> <span class="n">second_order_closure</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
<span class="c1"># warm up lr</span>
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">global_step</span> <span class="o">&lt;</span> <span class="mi">500</span><span class="p">:</span>
<span class="n">lr_scale</span> <span class="o">=</span> <span class="nb">min</span><span class="p">(</span><span class="mf">1.</span><span class="p">,</span> <span class="nb">float</span><span class="p">(</span><span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">global_step</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span> <span class="o">/</span> <span class="mf">500.</span><span class="p">)</span>
<span class="k">for</span> <span class="n">pg</span> <span class="ow">in</span> <span class="n">optimizer</span><span class="o">.</span><span class="n">param_groups</span><span class="p">:</span>
<span class="n">pg</span><span class="p">[</span><span class="s1">&#39;lr&#39;</span><span class="p">]</span> <span class="o">=</span> <span class="n">lr_scale</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">learning_rate</span>
<span class="c1"># update params</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">step</span><span class="p">()</span>
<span class="n">optimizer</span><span class="o">.</span><span class="n">zero_grad</span><span class="p">()</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_before_zero_grad">on_before_zero_grad</h4>
<p>Called in the training loop after taking an optimizer step and before zeroing grads.
Good place to inspect weight information with weights updated.</p>
<p>Called once per optimizer</p>
<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="k">def</span> <span class="nf">on_before_zero_grad</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">):</span>
<span class="c1"># do something with the optimizer or inspect it. </span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="backward">backward</h4>
<p>Called to perform backward step.
Feel free to override as needed.</p>
<p>The loss passed in has already been scaled for accumulated gradients if requested.</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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">backward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">use_amp</span><span class="p">,</span> <span class="n">loss</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">):</span>
<span class="sd">&quot;&quot;&quot;</span>
<span class="sd"> Override backward with your own implementation if you need to</span>
<span class="sd"> :param use_amp: Whether amp was requested or not</span>
<span class="sd"> :param loss: Loss is already scaled by accumulated grads</span>
<span class="sd"> :param optimizer: Current optimizer being used</span>
<span class="sd"> :return:</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="k">if</span> <span class="n">use_amp</span><span class="p">:</span>
<span class="k">with</span> <span class="n">amp</span><span class="o">.</span><span class="n">scale_loss</span><span class="p">(</span><span class="n">loss</span><span class="p">,</span> <span class="n">optimizer</span><span class="p">)</span> <span class="k">as</span> <span class="n">scaled_loss</span><span class="p">:</span>
<span class="n">scaled_loss</span><span class="o">.</span><span class="n">backward</span><span class="p">()</span>
<span class="k">else</span><span class="p">:</span>
<span class="n">loss</span><span class="o">.</span><span class="n">backward</span><span class="p">()</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="on_after_backward">on_after_backward</h4>
<p>Called in the training loop after model.backward()
This is the ideal place to inspect or log gradient information </p>
<table class="codehilitetable"><tr><td class="linenos"><div class="linenodiv"><pre>1
2
3
4
5
6
7
8</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">on_after_backward</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="c1"># example to inspect gradient information in tensorboard</span>
<span class="k">if</span> <span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">global_step</span> <span class="o">%</span> <span class="mi">25</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="c1"># don&#39;t make the tf file huge</span>
<span class="n">params</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">state_dict</span><span class="p">()</span>
<span class="k">for</span> <span class="n">k</span><span class="p">,</span> <span class="n">v</span> <span class="ow">in</span> <span class="n">params</span><span class="o">.</span><span class="n">items</span><span class="p">():</span>
<span class="n">grads</span> <span class="o">=</span> <span class="n">v</span>
<span class="n">name</span> <span class="o">=</span> <span class="n">k</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_histogram</span><span class="p">(</span><span class="n">tag</span><span class="o">=</span><span class="n">name</span><span class="p">,</span> <span class="n">values</span><span class="o">=</span><span class="n">grads</span><span class="p">,</span> <span class="n">global_step</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">global_step</span><span class="p">)</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="tbptt_split_batch">tbptt_split_batch</h4>
<p>Called in the training loop after on_batch_start if <code>truncated_bptt_steps &gt; 0</code>. Each returned batch split is passed separately to training_step(...).</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="k">def</span> <span class="nf">tbptt_split_batch</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">split_size</span><span class="p">):</span>
<span class="n">splits</span> <span class="o">=</span> <span class="p">[]</span>
<span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">time_dims</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">split_size</span><span class="p">):</span>
<span class="n">batch_split</span> <span class="o">=</span> <span class="p">[]</span>
<span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">x</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">batch</span><span class="p">):</span>
<span class="k">if</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">x</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">split_x</span> <span class="o">=</span> <span class="n">x</span><span class="p">[:,</span> <span class="n">t</span><span class="p">:</span><span class="n">t</span> <span class="o">+</span> <span class="n">split_size</span><span class="p">]</span>
<span class="k">elif</span> <span class="nb">isinstance</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">collections</span><span class="o">.</span><span class="n">Sequence</span><span class="p">):</span>
<span class="n">split_x</span> <span class="o">=</span> <span class="p">[</span><span class="bp">None</span><span class="p">]</span> <span class="o">*</span> <span class="nb">len</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="k">for</span> <span class="n">batch_idx</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">x</span><span class="p">)):</span>
<span class="n">split_x</span><span class="p">[</span><span class="n">batch_idx</span><span class="p">]</span> <span class="o">=</span> <span class="n">x</span><span class="p">[</span><span class="n">batch_idx</span><span class="p">][</span><span class="n">t</span><span class="p">:</span><span class="n">t</span> <span class="o">+</span> <span class="n">split_size</span><span class="p">]</span>
<span class="n">batch_split</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">split_x</span><span class="p">)</span>
<span class="n">splits</span><span class="o">.</span><span class="n">append</span><span class="p">(</span><span class="n">batch_split</span><span class="p">)</span>
<span class="k">return</span> <span class="n">splits</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="configure_apex">configure_apex</h4>
<p>Overwrite to define your own Apex implementation init.</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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">configure_apex</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">amp</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">optimizers</span><span class="p">,</span> <span class="n">amp_level</span><span class="p">):</span>
<span class="sd">&quot;&quot;&quot;</span>
<span class="sd"> Override to init AMP your own way</span>
<span class="sd"> Must return a model and list of optimizers</span>
<span class="sd"> :param amp:</span>
<span class="sd"> :param model:</span>
<span class="sd"> :param optimizers:</span>
<span class="sd"> :param amp_level:</span>
<span class="sd"> :return: Apex wrapped model and optimizers</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="n">model</span><span class="p">,</span> <span class="n">optimizers</span> <span class="o">=</span> <span class="n">amp</span><span class="o">.</span><span class="n">initialize</span><span class="p">(</span>
<span class="n">model</span><span class="p">,</span> <span class="n">optimizers</span><span class="p">,</span> <span class="n">opt_level</span><span class="o">=</span><span class="n">amp_level</span><span class="p">,</span>
<span class="p">)</span>
<span class="k">return</span> <span class="n">model</span><span class="p">,</span> <span class="n">optimizers</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="configure_ddp">configure_ddp</h4>
<p>Overwrite to define your own DDP implementation init.
The only requirement is that:
1. On a validation batch the call goes to model.validation_step. <br />
2. On a training batch the call goes to model.training_step. <br />
3. On a testing batch, the call goes to model.test_step</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</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">configure_ddp</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">model</span><span class="p">,</span> <span class="n">device_ids</span><span class="p">):</span>
<span class="sd">&quot;&quot;&quot;</span>
<span class="sd"> Override to init DDP in a different way or use your own wrapper.</span>
<span class="sd"> Must return model.</span>
<span class="sd"> :param model:</span>
<span class="sd"> :param device_ids:</span>
<span class="sd"> :return: DDP wrapped model</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="c1"># Lightning DDP simply routes to test_step, val_step, etc...</span>
<span class="n">model</span> <span class="o">=</span> <span class="n">LightningDistributedDataParallel</span><span class="p">(</span>
<span class="n">model</span><span class="p">,</span>
<span class="n">device_ids</span><span class="o">=</span><span class="n">device_ids</span><span class="p">,</span>
<span class="n">find_unused_parameters</span><span class="o">=</span><span class="bp">True</span>
<span class="p">)</span>
<span class="k">return</span> <span class="n">model</span>
</pre></div>
</td></tr></table>
<hr />
<h4 id="init_ddp_connection">init_ddp_connection</h4>
<p>Override to init DDP in your own way. </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
29
30
31
32
33
34</pre></div></td><td class="code"><div class="codehilite"><pre><span></span><span class="k">def</span> <span class="nf">init_ddp_connection</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
<span class="sd">&quot;&quot;&quot;</span>
<span class="sd"> Connect all procs in the world using the env:// init</span>
<span class="sd"> Use the first node as the root address</span>
<span class="sd"> &quot;&quot;&quot;</span>
<span class="c1"># use slurm job id for the port number</span>
<span class="c1"># guarantees unique ports across jobs from same grid search</span>
<span class="k">try</span><span class="p">:</span>
<span class="c1"># use the last 4 numbers in the job id as the id</span>
<span class="n">default_port</span> <span class="o">=</span> <span class="n">os</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">&#39;SLURM_JOB_ID&#39;</span><span class="p">]</span>
<span class="n">default_port</span> <span class="o">=</span> <span class="n">default_port</span><span class="p">[</span><span class="o">-</span><span class="mi">4</span><span class="p">:]</span>
<span class="c1"># all ports should be in the 10k+ range</span>
<span class="n">default_port</span> <span class="o">=</span> <span class="nb">int</span><span class="p">(</span><span class="n">default_port</span><span class="p">)</span> <span class="o">+</span> <span class="mi">15000</span>
<span class="k">except</span> <span class="ne">Exception</span> <span class="k">as</span> <span class="n">e</span><span class="p">:</span>
<span class="n">default_port</span> <span class="o">=</span> <span class="mi">12910</span>
<span class="c1"># if user gave a port number, use that one instead</span>
<span class="k">try</span><span class="p">:</span>
<span class="n">default_port</span> <span class="o">=</span> <span class="n">os</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">&#39;MASTER_PORT&#39;</span><span class="p">]</span>
<span class="k">except</span> <span class="ne">Exception</span><span class="p">:</span>
<span class="n">os</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">&#39;MASTER_PORT&#39;</span><span class="p">]</span> <span class="o">=</span> <span class="nb">str</span><span class="p">(</span><span class="n">default_port</span><span class="p">)</span>
<span class="c1"># figure out the root node addr</span>
<span class="k">try</span><span class="p">:</span>
<span class="n">root_node</span> <span class="o">=</span> <span class="n">os</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">&#39;SLURM_NODELIST&#39;</span><span class="p">]</span><span class="o">.</span><span class="n">split</span><span class="p">(</span><span class="s1">&#39; &#39;</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span>
<span class="k">except</span> <span class="ne">Exception</span><span class="p">:</span>
<span class="n">root_node</span> <span class="o">=</span> <span class="s1">&#39;127.0.0.2&#39;</span>
<span class="n">root_node</span> <span class="o">=</span> <span class="bp">self</span><span class="o">.</span><span class="n">trainer</span><span class="o">.</span><span class="n">resolve_root_node_address</span><span class="p">(</span><span class="n">root_node</span><span class="p">)</span>
<span class="n">os</span><span class="o">.</span><span class="n">environ</span><span class="p">[</span><span class="s1">&#39;MASTER_ADDR&#39;</span><span class="p">]</span> <span class="o">=</span> <span class="n">root_node</span>
<span class="n">dist</span><span class="o">.</span><span class="n">init_process_group</span><span class="p">(</span><span class="s1">&#39;nccl&#39;</span><span class="p">,</span> <span class="n">rank</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">proc_rank</span><span class="p">,</span> <span class="n">world_size</span><span class="o">=</span><span class="bp">self</span><span class="o">.</span><span class="n">world_size</span><span class="p">)</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="../debugging/" title="Debugging" 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>
Debugging
</span>
</div>
</a>
<a href="../../examples/Examples/" title="Examples" 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>
Examples
</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>