last update, perhaps

This commit is contained in:
Morten Hjorth-Jensen
2024-10-14 13:54:09 +02:00
parent 7d4725e646
commit 0bc2e08294
52 changed files with 2146 additions and 1626 deletions
+65 -57
View File
@@ -2709,63 +2709,50 @@ Using TensorFlow results in a much better execution time. Try it!</p>
<span class="g g-Whitespace"> </span><span class="mi">19</span> <span class="n">x</span> <span class="o">=</span> <span class="nb">tuple</span><span class="p">(</span><span class="n">args</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="n">argnum</span><span class="p">)</span>
<span class="ne">---&gt; </span><span class="mi">20</span> <span class="k">return</span> <span class="n">unary_operator</span><span class="p">(</span><span class="n">unary_f</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="o">*</span><span class="n">nary_op_args</span><span class="p">,</span> <span class="o">**</span><span class="n">nary_op_kwargs</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/differential_operators.py:64,</span> in <span class="ni">jacobian</span><span class="nt">(fun, x)</span>
<span class="g g-Whitespace"> </span><span class="mi">62</span> <span class="n">jacobian_shape</span> <span class="o">=</span> <span class="n">ans_vspace</span><span class="o">.</span><span class="n">shape</span> <span class="o">+</span> <span class="n">vspace</span><span class="p">(</span><span class="n">x</span><span class="p">)</span><span class="o">.</span><span class="n">shape</span>
<span class="g g-Whitespace"> </span><span class="mi">63</span> <span class="n">grads</span> <span class="o">=</span> <span class="nb">map</span><span class="p">(</span><span class="n">vjp</span><span class="p">,</span> <span class="n">ans_vspace</span><span class="o">.</span><span class="n">standard_basis</span><span class="p">())</span>
<span class="ne">---&gt; </span><span class="mi">64</span> <span class="k">return</span> <span class="n">np</span><span class="o">.</span><span class="n">reshape</span><span class="p">(</span><span class="n">np</span><span class="o">.</span><span class="n">stack</span><span class="p">(</span><span class="n">grads</span><span class="p">),</span> <span class="n">jacobian_shape</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/differential_operators.py:60,</span> in <span class="ni">jacobian</span><span class="nt">(fun, x)</span>
<span class="g g-Whitespace"> </span><span class="mi">50</span> <span class="nd">@unary_to_nary</span>
<span class="g g-Whitespace"> </span><span class="mi">51</span> <span class="k">def</span> <span class="nf">jacobian</span><span class="p">(</span><span class="n">fun</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">52</span><span class="w"> </span><span class="sd">&quot;&quot;&quot;</span>
<span class="g g-Whitespace"> </span><span class="mi">53</span><span class="sd"> Returns a function which computes the Jacobian of `fun` with respect to</span>
<span class="g g-Whitespace"> </span><span class="mi">54</span><span class="sd"> positional argument number `argnum`, which must be a scalar or array. Unlike</span>
<span class="sd"> (...)</span>
<span class="g g-Whitespace"> </span><span class="mi">58</span><span class="sd"> (out1, out2, ...) then the Jacobian has shape (out1, out2, ..., in1, in2, ...).</span>
<span class="g g-Whitespace"> </span><span class="mi">59</span><span class="sd"> &quot;&quot;&quot;</span>
<span class="ne">---&gt; </span><span class="mi">60</span> <span class="n">vjp</span><span class="p">,</span> <span class="n">ans</span> <span class="o">=</span> <span class="n">_make_vjp</span><span class="p">(</span><span class="n">fun</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">61</span> <span class="n">ans_vspace</span> <span class="o">=</span> <span class="n">vspace</span><span class="p">(</span><span class="n">ans</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">62</span> <span class="n">jacobian_shape</span> <span class="o">=</span> <span class="n">ans_vspace</span><span class="o">.</span><span class="n">shape</span> <span class="o">+</span> <span class="n">vspace</span><span class="p">(</span><span class="n">x</span><span class="p">)</span><span class="o">.</span><span class="n">shape</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_wrapper.py:88,</span> in <span class="ni">stack</span><span class="nt">(arrays, axis)</span>
<span class="g g-Whitespace"> </span><span class="mi">83</span> <span class="k">def</span> <span class="nf">stack</span><span class="p">(</span><span class="n">arrays</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">84</span> <span class="c1"># this code is basically copied from numpy/core/shape_base.py&#39;s stack</span>
<span class="g g-Whitespace"> </span><span class="mi">85</span> <span class="c1"># we need it here because we want to re-implement stack in terms of the</span>
<span class="g g-Whitespace"> </span><span class="mi">86</span> <span class="c1"># primitives defined in this file</span>
<span class="ne">---&gt; </span><span class="mi">88</span> <span class="n">arrays</span> <span class="o">=</span> <span class="p">[</span><span class="n">array</span><span class="p">(</span><span class="n">arr</span><span class="p">)</span> <span class="k">for</span> <span class="n">arr</span> <span class="ow">in</span> <span class="n">arrays</span><span class="p">]</span>
<span class="g g-Whitespace"> </span><span class="mi">89</span> <span class="k">if</span> <span class="ow">not</span> <span class="n">arrays</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">90</span> <span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">&#39;need at least one array to stack&#39;</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/core.py:10,</span> in <span class="ni">make_vjp</span><span class="nt">(fun, x)</span>
<span class="g g-Whitespace"> </span><span class="mi">8</span> <span class="k">def</span> <span class="nf">make_vjp</span><span class="p">(</span><span class="n">fun</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">9</span> <span class="n">start_node</span> <span class="o">=</span> <span class="n">VJPNode</span><span class="o">.</span><span class="n">new_root</span><span class="p">()</span>
<span class="ne">---&gt; </span><span class="mi">10</span> <span class="n">end_value</span><span class="p">,</span> <span class="n">end_node</span> <span class="o">=</span> <span class="n">trace</span><span class="p">(</span><span class="n">start_node</span><span class="p">,</span> <span class="n">fun</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">11</span> <span class="k">if</span> <span class="n">end_node</span> <span class="ow">is</span> <span class="kc">None</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">12</span> <span class="k">def</span> <span class="nf">vjp</span><span class="p">(</span><span class="n">g</span><span class="p">):</span> <span class="k">return</span> <span class="n">vspace</span><span class="p">(</span><span class="n">x</span><span class="p">)</span><span class="o">.</span><span class="n">zeros</span><span class="p">()</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_wrapper.py:88,</span> in <span class="ni">&lt;listcomp&gt;</span><span class="nt">(.0)</span>
<span class="g g-Whitespace"> </span><span class="mi">83</span> <span class="k">def</span> <span class="nf">stack</span><span class="p">(</span><span class="n">arrays</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">84</span> <span class="c1"># this code is basically copied from numpy/core/shape_base.py&#39;s stack</span>
<span class="g g-Whitespace"> </span><span class="mi">85</span> <span class="c1"># we need it here because we want to re-implement stack in terms of the</span>
<span class="g g-Whitespace"> </span><span class="mi">86</span> <span class="c1"># primitives defined in this file</span>
<span class="ne">---&gt; </span><span class="mi">88</span> <span class="n">arrays</span> <span class="o">=</span> <span class="p">[</span><span class="n">array</span><span class="p">(</span><span class="n">arr</span><span class="p">)</span> <span class="k">for</span> <span class="n">arr</span> <span class="ow">in</span> <span class="n">arrays</span><span class="p">]</span>
<span class="g g-Whitespace"> </span><span class="mi">89</span> <span class="k">if</span> <span class="ow">not</span> <span class="n">arrays</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">90</span> <span class="k">raise</span> <span class="ne">ValueError</span><span class="p">(</span><span class="s1">&#39;need at least one array to stack&#39;</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/tracer.py:10,</span> in <span class="ni">trace</span><span class="nt">(start_node, fun, x)</span>
<span class="g g-Whitespace"> </span><span class="mi">8</span> <span class="k">with</span> <span class="n">trace_stack</span><span class="o">.</span><span class="n">new_trace</span><span class="p">()</span> <span class="k">as</span> <span class="n">t</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">9</span> <span class="n">start_box</span> <span class="o">=</span> <span class="n">new_box</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">t</span><span class="p">,</span> <span class="n">start_node</span><span class="p">)</span>
<span class="ne">---&gt; </span><span class="mi">10</span> <span class="n">end_box</span> <span class="o">=</span> <span class="n">fun</span><span class="p">(</span><span class="n">start_box</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">11</span> <span class="k">if</span> <span class="n">isbox</span><span class="p">(</span><span class="n">end_box</span><span class="p">)</span> <span class="ow">and</span> <span class="n">end_box</span><span class="o">.</span><span class="n">_trace</span> <span class="o">==</span> <span class="n">start_box</span><span class="o">.</span><span class="n">_trace</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">12</span> <span class="k">return</span> <span class="n">end_box</span><span class="o">.</span><span class="n">_value</span><span class="p">,</span> <span class="n">end_box</span><span class="o">.</span><span class="n">_node</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/core.py:14,</span> in <span class="ni">make_vjp.&lt;locals&gt;.vjp</span><span class="nt">(g)</span>
<span class="ne">---&gt; </span><span class="mi">14</span> <span class="k">def</span> <span class="nf">vjp</span><span class="p">(</span><span class="n">g</span><span class="p">):</span> <span class="k">return</span> <span class="n">backward_pass</span><span class="p">(</span><span class="n">g</span><span class="p">,</span> <span class="n">end_node</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/wrap_util.py:15,</span> in <span class="ni">unary_to_nary.&lt;locals&gt;.nary_operator.&lt;locals&gt;.nary_f.&lt;locals&gt;.unary_f</span><span class="nt">(x)</span>
<span class="g g-Whitespace"> </span><span class="mi">13</span> <span class="k">else</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">14</span> <span class="n">subargs</span> <span class="o">=</span> <span class="n">subvals</span><span class="p">(</span><span class="n">args</span><span class="p">,</span> <span class="nb">zip</span><span class="p">(</span><span class="n">argnum</span><span class="p">,</span> <span class="n">x</span><span class="p">))</span>
<span class="ne">---&gt; </span><span class="mi">15</span> <span class="k">return</span> <span class="n">fun</span><span class="p">(</span><span class="o">*</span><span class="n">subargs</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/core.py:21,</span> in <span class="ni">backward_pass</span><span class="nt">(g, end_node)</span>
<span class="g g-Whitespace"> </span><span class="mi">19</span> <span class="k">for</span> <span class="n">node</span> <span class="ow">in</span> <span class="n">toposort</span><span class="p">(</span><span class="n">end_node</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">20</span> <span class="n">outgrad</span> <span class="o">=</span> <span class="n">outgrads</span><span class="o">.</span><span class="n">pop</span><span class="p">(</span><span class="n">node</span><span class="p">)</span>
<span class="ne">---&gt; </span><span class="mi">21</span> <span class="n">ingrads</span> <span class="o">=</span> <span class="n">node</span><span class="o">.</span><span class="n">vjp</span><span class="p">(</span><span class="n">outgrad</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span>
<span class="g g-Whitespace"> </span><span class="mi">22</span> <span class="k">for</span> <span class="n">parent</span><span class="p">,</span> <span class="n">ingrad</span> <span class="ow">in</span> <span class="nb">zip</span><span class="p">(</span><span class="n">node</span><span class="o">.</span><span class="n">parents</span><span class="p">,</span> <span class="n">ingrads</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">23</span> <span class="n">outgrads</span><span class="p">[</span><span class="n">parent</span><span class="p">]</span> <span class="o">=</span> <span class="n">add_outgrads</span><span class="p">(</span><span class="n">outgrads</span><span class="o">.</span><span class="n">get</span><span class="p">(</span><span class="n">parent</span><span class="p">),</span> <span class="n">ingrad</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/wrap_util.py:15,</span> in <span class="ni">unary_to_nary.&lt;locals&gt;.nary_operator.&lt;locals&gt;.nary_f.&lt;locals&gt;.unary_f</span><span class="nt">(x)</span>
<span class="g g-Whitespace"> </span><span class="mi">13</span> <span class="k">else</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">14</span> <span class="n">subargs</span> <span class="o">=</span> <span class="n">subvals</span><span class="p">(</span><span class="n">args</span><span class="p">,</span> <span class="nb">zip</span><span class="p">(</span><span class="n">argnum</span><span class="p">,</span> <span class="n">x</span><span class="p">))</span>
<span class="ne">---&gt; </span><span class="mi">15</span> <span class="k">return</span> <span class="n">fun</span><span class="p">(</span><span class="o">*</span><span class="n">subargs</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/core.py:67,</span> in <span class="ni">defvjp.&lt;locals&gt;.vjp_argnums.&lt;locals&gt;.&lt;lambda&gt;</span><span class="nt">(g)</span>
<span class="g g-Whitespace"> </span><span class="mi">64</span> <span class="k">raise</span> <span class="ne">NotImplementedError</span><span class="p">(</span>
<span class="g g-Whitespace"> </span><span class="mi">65</span> <span class="s2">&quot;VJP of </span><span class="si">{}</span><span class="s2"> wrt argnum 0 not defined&quot;</span><span class="o">.</span><span class="n">format</span><span class="p">(</span><span class="n">fun</span><span class="o">.</span><span class="vm">__name__</span><span class="p">))</span>
<span class="g g-Whitespace"> </span><span class="mi">66</span> <span class="n">vjp</span> <span class="o">=</span> <span class="n">vjpfun</span><span class="p">(</span><span class="n">ans</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
<span class="ne">---&gt; </span><span class="mi">67</span> <span class="k">return</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="p">(</span><span class="n">vjp</span><span class="p">(</span><span class="n">g</span><span class="p">),)</span>
<span class="g g-Whitespace"> </span><span class="mi">68</span> <span class="k">elif</span> <span class="n">L</span> <span class="o">==</span> <span class="mi">2</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">69</span> <span class="n">argnum_0</span><span class="p">,</span> <span class="n">argnum_1</span> <span class="o">=</span> <span class="n">argnums</span>
<span class="nn">Cell In[9], line 61,</span> in <span class="ni">g_trial</span><span class="nt">(point, P)</span>
<span class="g g-Whitespace"> </span><span class="mi">59</span> <span class="k">def</span> <span class="nf">g_trial</span><span class="p">(</span><span class="n">point</span><span class="p">,</span><span class="n">P</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">60</span> <span class="n">x</span><span class="p">,</span><span class="n">t</span> <span class="o">=</span> <span class="n">point</span>
<span class="ne">---&gt; </span><span class="mi">61</span> <span class="k">return</span> <span class="p">(</span><span class="mi">1</span><span class="o">-</span><span class="n">t</span><span class="p">)</span><span class="o">*</span><span class="n">u</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">+</span> <span class="n">x</span><span class="o">*</span><span class="p">(</span><span class="mi">1</span><span class="o">-</span><span class="n">x</span><span class="p">)</span><span class="o">*</span><span class="n">t</span><span class="o">*</span><span class="n">deep_neural_network</span><span class="p">(</span><span class="n">P</span><span class="p">,</span><span class="n">point</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:423,</span> in <span class="ni">matmul_vjp_1.&lt;locals&gt;.&lt;lambda&gt;</span><span class="nt">(g)</span>
<span class="g g-Whitespace"> </span><span class="mi">421</span> <span class="n">A_ndim</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">ndim</span><span class="p">(</span><span class="n">A</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">422</span> <span class="n">B_meta</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">metadata</span><span class="p">(</span><span class="n">B</span><span class="p">)</span>
<span class="ne">--&gt; </span><span class="mi">423</span> <span class="k">return</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">matmul_adjoint_1</span><span class="p">(</span><span class="n">A</span><span class="p">,</span> <span class="n">g</span><span class="p">,</span> <span class="n">A_ndim</span><span class="p">,</span> <span class="n">B_meta</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:413,</span> in <span class="ni">matmul_adjoint_1</span><span class="nt">(A, G, A_ndim, B_meta)</span>
<span class="g g-Whitespace"> </span><span class="mi">411</span> <span class="k">if</span> <span class="n">B_is_vec</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">412</span> <span class="n">result</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">squeeze</span><span class="p">(</span><span class="n">result</span><span class="p">,</span> <span class="n">anp</span><span class="o">.</span><span class="n">ndim</span><span class="p">(</span><span class="n">G</span><span class="p">)</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)</span>
<span class="ne">--&gt; </span><span class="mi">413</span> <span class="k">return</span> <span class="n">unbroadcast</span><span class="p">(</span><span class="n">result</span><span class="p">,</span> <span class="n">B_meta</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:653,</span> in <span class="ni">unbroadcast</span><span class="nt">(x, target_meta, broadcast_idx)</span>
<span class="g g-Whitespace"> </span><span class="mi">651</span> <span class="k">for</span> <span class="n">axis</span><span class="p">,</span> <span class="n">size</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">target_shape</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">652</span> <span class="k">if</span> <span class="n">size</span> <span class="o">==</span> <span class="mi">1</span><span class="p">:</span>
<span class="ne">--&gt; </span><span class="mi">653</span> <span class="n">x</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="n">axis</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">654</span> <span class="k">if</span> <span class="n">anp</span><span class="o">.</span><span class="n">iscomplexobj</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="ow">and</span> <span class="ow">not</span> <span class="n">target_iscomplex</span><span class="p">:</span>
<span class="g g-Whitespace"> </span><span class="mi">655</span> <span class="n">x</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">real</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_boxes.py:34,</span> in <span class="ni">ArrayBox.__rsub__</span><span class="nt">(self, other)</span>
<span class="ne">---&gt; </span><span class="mi">34</span> <span class="k">def</span> <span class="fm">__rsub__</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">other</span><span class="p">):</span> <span class="k">return</span> <span class="n">anp</span><span class="o">.</span><span class="n">subtract</span><span class="p">(</span><span class="n">other</span><span class="p">,</span> <span class="bp">self</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/tracer.py:45,</span> in <span class="ni">primitive.&lt;locals&gt;.f_wrapped</span><span class="nt">(*args, **kwargs)</span>
<span class="g g-Whitespace"> </span><span class="mi">43</span> <span class="n">argnums</span> <span class="o">=</span> <span class="nb">tuple</span><span class="p">(</span><span class="n">argnum</span> <span class="k">for</span> <span class="n">argnum</span><span class="p">,</span> <span class="n">_</span> <span class="ow">in</span> <span class="n">boxed_args</span><span class="p">)</span>
@@ -2788,12 +2775,33 @@ Using TensorFlow results in a much better execution time. Try it!</p>
<span class="g g-Whitespace"> </span><span class="mi">67</span> <span class="k">return</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="p">(</span><span class="n">vjp</span><span class="p">(</span><span class="n">g</span><span class="p">),)</span>
<span class="g g-Whitespace"> </span><span class="mi">68</span> <span class="k">elif</span> <span class="n">L</span> <span class="o">==</span> <span class="mi">2</span><span class="p">:</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:297,</span> in <span class="ni">grad_np_sum</span><span class="nt">(ans, x, axis, keepdims, dtype)</span>
<span class="g g-Whitespace"> </span><span class="mi">294</span> <span class="k">return</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">anp</span><span class="o">.</span><span class="n">sum</span><span class="p">(</span><span class="n">g</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="n">broadcast_axes</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="kc">True</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">295</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">broadcast_to</span><span class="p">,</span> <span class="n">grad_broadcast_to</span><span class="p">)</span>
<span class="ne">--&gt; </span><span class="mi">297</span> <span class="k">def</span> <span class="nf">grad_np_sum</span><span class="p">(</span><span class="n">ans</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="kc">None</span><span class="p">,</span> <span class="n">keepdims</span><span class="o">=</span><span class="kc">False</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="kc">None</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">298</span> <span class="n">shape</span><span class="p">,</span> <span class="n">dtype</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">shape</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">anp</span><span class="o">.</span><span class="n">result_type</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">299</span> <span class="k">return</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">repeat_to_match_shape</span><span class="p">(</span><span class="n">g</span><span class="p">,</span> <span class="n">shape</span><span class="p">,</span> <span class="n">dtype</span><span class="p">,</span> <span class="n">axis</span><span class="p">,</span> <span class="n">keepdims</span><span class="p">)[</span><span class="mi">0</span><span class="p">]</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:37,</span> in <span class="ni">&lt;lambda&gt;</span><span class="nt">(ans, x, y)</span>
<span class="g g-Whitespace"> </span><span class="mi">32</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">add</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">g</span><span class="p">),</span>
<span class="g g-Whitespace"> </span><span class="mi">33</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">g</span><span class="p">))</span>
<span class="g g-Whitespace"> </span><span class="mi">34</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">multiply</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">y</span> <span class="o">*</span> <span class="n">g</span><span class="p">),</span>
<span class="g g-Whitespace"> </span><span class="mi">35</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">x</span> <span class="o">*</span> <span class="n">g</span><span class="p">))</span>
<span class="g g-Whitespace"> </span><span class="mi">36</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">subtract</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">g</span><span class="p">),</span>
<span class="ne">---&gt; </span><span class="mi">37</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="o">-</span><span class="n">g</span><span class="p">))</span>
<span class="g g-Whitespace"> </span><span class="mi">38</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">divide</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">g</span> <span class="o">/</span> <span class="n">y</span><span class="p">),</span>
<span class="g g-Whitespace"> </span><span class="mi">39</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="o">-</span> <span class="n">g</span> <span class="o">*</span> <span class="n">x</span> <span class="o">/</span> <span class="n">y</span><span class="o">**</span><span class="mi">2</span><span class="p">))</span>
<span class="g g-Whitespace"> </span><span class="mi">40</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">maximum</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">g</span> <span class="o">*</span> <span class="n">balanced_eq</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">ans</span><span class="p">,</span> <span class="n">y</span><span class="p">)),</span>
<span class="g g-Whitespace"> </span><span class="mi">41</span> <span class="k">lambda</span> <span class="n">ans</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">unbroadcast_f</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">g</span> <span class="o">*</span> <span class="n">balanced_eq</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">ans</span><span class="p">,</span> <span class="n">x</span><span class="p">)))</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:659,</span> in <span class="ni">unbroadcast_f</span><span class="nt">(target, f)</span>
<span class="g g-Whitespace"> </span><span class="mi">658</span> <span class="k">def</span> <span class="nf">unbroadcast_f</span><span class="p">(</span><span class="n">target</span><span class="p">,</span> <span class="n">f</span><span class="p">):</span>
<span class="ne">--&gt; </span><span class="mi">659</span> <span class="n">target_meta</span> <span class="o">=</span> <span class="n">anp</span><span class="o">.</span><span class="n">metadata</span><span class="p">(</span><span class="n">target</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">660</span> <span class="k">return</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">unbroadcast</span><span class="p">(</span><span class="n">f</span><span class="p">(</span><span class="n">g</span><span class="p">),</span> <span class="n">target_meta</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/tracer.py:61,</span> in <span class="ni">notrace_primitive.&lt;locals&gt;.f_wrapped</span><span class="nt">(*args, **kwargs)</span>
<span class="g g-Whitespace"> </span><span class="mi">58</span> <span class="nd">@wraps</span><span class="p">(</span><span class="n">f_raw</span><span class="p">)</span>
<span class="g g-Whitespace"> </span><span class="mi">59</span> <span class="k">def</span> <span class="nf">f_wrapped</span><span class="p">(</span><span class="o">*</span><span class="n">args</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">):</span>
<span class="g g-Whitespace"> </span><span class="mi">60</span> <span class="n">argvals</span> <span class="o">=</span> <span class="nb">map</span><span class="p">(</span><span class="n">getval</span><span class="p">,</span> <span class="n">args</span><span class="p">)</span>
<span class="ne">---&gt; </span><span class="mi">61</span> <span class="k">return</span> <span class="n">f_raw</span><span class="p">(</span><span class="o">*</span><span class="n">argvals</span><span class="p">,</span> <span class="o">**</span><span class="n">kwargs</span><span class="p">)</span>
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_wrapper.py:148,</span> in <span class="ni">metadata</span><span class="nt">(A)</span>
<span class="g g-Whitespace"> </span><span class="mi">146</span> <span class="nd">@notrace_primitive</span>
<span class="g g-Whitespace"> </span><span class="mi">147</span> <span class="k">def</span> <span class="nf">metadata</span><span class="p">(</span><span class="n">A</span><span class="p">):</span>
<span class="ne">--&gt; </span><span class="mi">148</span> <span class="k">return</span> <span class="n">_np</span><span class="o">.</span><span class="n">shape</span><span class="p">(</span><span class="n">A</span><span class="p">),</span> <span class="n">_np</span><span class="o">.</span><span class="n">ndim</span><span class="p">(</span><span class="n">A</span><span class="p">),</span> <span class="n">_np</span><span class="o">.</span><span class="n">result_type</span><span class="p">(</span><span class="n">A</span><span class="p">),</span> <span class="n">_np</span><span class="o">.</span><span class="n">iscomplexobj</span><span class="p">(</span><span class="n">A</span><span class="p">)</span>
<span class="ne">KeyboardInterrupt</span>:
</pre></div>