Files
pytorch-lightning/Trainer/hooks/index.html
T

919 lines
31 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">
GitHub
</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">
GitHub
</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="../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="#on_tng_metrics" title="on_tng_metrics" class="md-nav__link">
on_tng_metrics
</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="#on_after_backward" title="on_after_backward" class="md-nav__link">
on_after_backward
</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="#on_tng_metrics" title="on_tng_metrics" class="md-nav__link">
on_tng_metrics
</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="#on_after_backward" title="on_after_backward" class="md-nav__link">
on_after_backward
</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="on_tng_metrics">on_tng_metrics</h4>
<p>Called in the training loop, right before metrics are logged.
Although you can log at any time by using self.experiment, you can use
this callback to modify what will be logged.</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_tng_metrics</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">metrics</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">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="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="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="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">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>
</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>