mirror of
https://github.com/wassname/pytorch-lightning.git
synced 2026-09-09 11:32:07 +08:00
Deployed f0af138 with MkDocs version: 1.0.4
This commit is contained in:
@@ -564,26 +564,39 @@
|
||||
<hr />
|
||||
<h3 id="freeze">freeze</h3>
|
||||
<p>Freeze all params for inference</p>
|
||||
<pre><code class="python">model = MyLightningModule(...)
|
||||
model.freeze()
|
||||
</code></pre>
|
||||
<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="n">model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">freeze</span><span class="p">()</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<hr />
|
||||
<h3 id="load_from_metrics">load_from_metrics</h3>
|
||||
<p>This is the easiest/fastest way which uses the meta_tags.csv file from test-tube to rebuild the model.
|
||||
The meta_tags.csv file can be found in the test-tube experiment save_dir. </p>
|
||||
<pre><code class="python">pretrained_model = MyLightningModule.load_from_metrics(
|
||||
weights_path='/path/to/pytorch_checkpoint.ckpt',
|
||||
tags_csv='/path/to/test_tube/experiment/version/meta_tags.csv',
|
||||
on_gpu=True,
|
||||
map_location=None
|
||||
)
|
||||
<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="n">pretrained_model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="o">.</span><span class="n">load_from_metrics</span><span class="p">(</span>
|
||||
<span class="n">weights_path</span><span class="o">=</span><span class="s1">'/path/to/pytorch_checkpoint.ckpt'</span><span class="p">,</span>
|
||||
<span class="n">tags_csv</span><span class="o">=</span><span class="s1">'/path/to/test_tube/experiment/version/meta_tags.csv'</span><span class="p">,</span>
|
||||
<span class="n">on_gpu</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span>
|
||||
<span class="n">map_location</span><span class="o">=</span><span class="bp">None</span>
|
||||
<span class="p">)</span>
|
||||
|
||||
# predict
|
||||
pretrained_model.eval()
|
||||
pretrained_model.freeze()
|
||||
y_hat = pretrained_model(x)
|
||||
</code></pre>
|
||||
<span class="c1"># predict</span>
|
||||
<span class="n">pretrained_model</span><span class="o">.</span><span class="n">eval</span><span class="p">()</span>
|
||||
<span class="n">pretrained_model</span><span class="o">.</span><span class="n">freeze</span><span class="p">()</span>
|
||||
<span class="n">y_hat</span> <span class="o">=</span> <span class="n">pretrained_model</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
<p><strong>Params</strong> </p>
|
||||
<table>
|
||||
@@ -617,9 +630,11 @@ y_hat = pretrained_model(x)
|
||||
<hr />
|
||||
<h3 id="unfreeze">unfreeze</h3>
|
||||
<p>Unfreeze all params for inference</p>
|
||||
<pre><code class="python">model = MyLightningModule(...)
|
||||
model.unfreeze()
|
||||
</code></pre>
|
||||
<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="n">model</span> <span class="o">=</span> <span class="n">MyLightningModule</span><span class="p">(</span><span class="o">...</span><span class="p">)</span>
|
||||
<span class="n">model</span><span class="o">.</span><span class="n">unfreeze</span><span class="p">()</span>
|
||||
</pre></div>
|
||||
</td></tr></table>
|
||||
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user