update
This commit is contained in:
@@ -275,7 +275,34 @@ const thebe_selector_output = ".output, .cell_output"
|
||||
</li>
|
||||
<li class="toctree-l1">
|
||||
<a class="reference internal" href="week36.html">
|
||||
Week 36: Linear Rgeression and Statistical interpretations
|
||||
Week 36: Linear Regression and Statistical interpretations
|
||||
</a>
|
||||
</li>
|
||||
<li class="toctree-l1">
|
||||
<a class="reference internal" href="exercisesweek37.html">
|
||||
Exercises week 37
|
||||
</a>
|
||||
</li>
|
||||
<li class="toctree-l1">
|
||||
<a class="reference internal" href="week37.html">
|
||||
Week 37: Statistical interpretations and Resampling Methods
|
||||
</a>
|
||||
</li>
|
||||
<li class="toctree-l1">
|
||||
<a class="reference internal" href="exercisesweek38.html">
|
||||
Exercises week 38
|
||||
</a>
|
||||
</li>
|
||||
</ul>
|
||||
<p aria-level="2" class="caption" role="heading">
|
||||
<span class="caption-text">
|
||||
Projects
|
||||
</span>
|
||||
</p>
|
||||
<ul class="nav bd-sidenav">
|
||||
<li class="toctree-l1">
|
||||
<a class="reference internal" href="project1.html">
|
||||
Project 1 on Machine Learning, deadline October 7 (midnight), 2024
|
||||
</a>
|
||||
</li>
|
||||
</ul>
|
||||
@@ -2577,70 +2604,122 @@ Using TensorFlow results in a much better execution time. Try it!</p>
|
||||
<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">---> </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">Cell In[9], line 79,</span> in <span class="ni">cost_function</span><span class="nt">(P, x, t)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">76</span> <span class="n">point</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">array</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="nn">Cell In[9], line 80,</span> in <span class="ni">cost_function</span><span class="nt">(P, x, t)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">78</span> <span class="n">g_t</span> <span class="o">=</span> <span class="n">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="ne">---> </span><span class="mi">79</span> <span class="n">g_t_jacobian</span> <span class="o">=</span> <span class="n">g_t_jacobian_func</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">80</span> <span class="n">g_t_hessian</span> <span class="o">=</span> <span class="n">g_t_hessian_func</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">79</span> <span class="n">g_t_jacobian</span> <span class="o">=</span> <span class="n">g_t_jacobian_func</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="ne">---> </span><span class="mi">80</span> <span class="n">g_t_hessian</span> <span class="o">=</span> <span class="n">g_t_hessian_func</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">82</span> <span class="n">g_t_dt</span> <span class="o">=</span> <span class="n">g_t_jacobian</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">83</span> <span class="n">g_t_d2x</span> <span class="o">=</span> <span class="n">g_t_hessian</span><span class="p">[</span><span class="mi">0</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/wrap_util.py:20,</span> in <span class="ni">unary_to_nary.<locals>.nary_operator.<locals>.nary_f</span><span class="nt">(*args, **kwargs)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">18</span> <span class="k">else</span><span class="p">:</span>
|
||||
<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">---> </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">---> </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:81,</span> in <span class="ni">hessian</span><span class="nt">(fun, x)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">78</span> <span class="nd">@unary_to_nary</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">79</span> <span class="k">def</span> <span class="nf">hessian</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">80</span> <span class="s2">"Returns a function that computes the exact Hessian."</span>
|
||||
<span class="ne">---> </span><span class="mi">81</span> <span class="k">return</span> <span class="n">jacobian</span><span class="p">(</span><span class="n">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="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'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">---> </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">'need at least one array to stack'</span><span class="p">)</span>
|
||||
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/wrap_util.py:20,</span> in <span class="ni">unary_to_nary.<locals>.nary_operator.<locals>.nary_f</span><span class="nt">(*args, **kwargs)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">18</span> <span class="k">else</span><span class="p">:</span>
|
||||
<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">---> </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/numpy/numpy_wrapper.py:88,</span> in <span class="ni"><listcomp></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'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">---> </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">'need at least one array to stack'</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">"""</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"> """</span>
|
||||
<span class="ne">---> </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/core.py:14,</span> in <span class="ni">make_vjp.<locals>.vjp</span><span class="nt">(g)</span>
|
||||
<span class="ne">---> </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/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">---> </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/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">---> </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/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">---> </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:67,</span> in <span class="ni">defvjp.<locals>.vjp_argnums.<locals>.<lambda></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">"VJP of </span><span class="si">{}</span><span class="s2"> wrt argnum 0 not defined"</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">---> </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">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/wrap_util.py:15,</span> in <span class="ni">unary_to_nary.<locals>.nary_operator.<locals>.nary_f.<locals>.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">---> </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/numpy/numpy_vjps.py:660,</span> in <span class="ni">unbroadcast_f.<locals>.<lambda></span><span class="nt">(g)</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="g g-Whitespace"> </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="ne">--> </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/wrap_util.py:20,</span> in <span class="ni">unary_to_nary.<locals>.nary_operator.<locals>.nary_f</span><span class="nt">(*args, **kwargs)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">18</span> <span class="k">else</span><span class="p">:</span>
|
||||
<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">---> </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/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">--> </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/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">"""</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"> """</span>
|
||||
<span class="ne">---> </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/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">---> </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/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">---> </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/wrap_util.py:15,</span> in <span class="ni">unary_to_nary.<locals>.nary_operator.<locals>.nary_f.<locals>.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">---> </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/wrap_util.py:15,</span> in <span class="ni">unary_to_nary.<locals>.nary_operator.<locals>.nary_f.<locals>.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">---> </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">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">---> </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">Cell In[9], line 48,</span> in <span class="ni">deep_neural_network</span><span class="nt">(deep_params, x)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">45</span> <span class="n">w_output</span> <span class="o">=</span> <span class="n">deep_params</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">]</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">47</span> <span class="c1"># Include bias:</span>
|
||||
<span class="ne">---> </span><span class="mi">48</span> <span class="n">x_prev</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">concatenate</span><span class="p">((</span><span class="n">np</span><span class="o">.</span><span class="n">ones</span><span class="p">((</span><span class="mi">1</span><span class="p">,</span><span class="n">num_points</span><span class="p">)),</span> <span class="n">x_prev</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">50</span> <span class="n">z_output</span> <span class="o">=</span> <span class="n">np</span><span class="o">.</span><span class="n">matmul</span><span class="p">(</span><span class="n">w_output</span><span class="p">,</span> <span class="n">x_prev</span><span class="p">)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">51</span> <span class="n">x_output</span> <span class="o">=</span> <span class="n">z_output</span>
|
||||
|
||||
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_wrapper.py:38,</span> in <span class="ni"><lambda></span><span class="nt">(arr_list, axis)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">35</span> <span class="nd">@primitive</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">36</span> <span class="k">def</span> <span class="nf">concatenate_args</span><span class="p">(</span><span class="n">axis</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">37</span> <span class="k">return</span> <span class="n">_np</span><span class="o">.</span><span class="n">concatenate</span><span class="p">(</span><span class="n">args</span><span class="p">,</span> <span class="n">axis</span><span class="p">)</span><span class="o">.</span><span class="n">view</span><span class="p">(</span><span class="n">ndarray</span><span class="p">)</span>
|
||||
<span class="ne">---> </span><span class="mi">38</span> <span class="n">concatenate</span> <span class="o">=</span> <span class="k">lambda</span> <span class="n">arr_list</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="n">concatenate_args</span><span class="p">(</span><span class="n">axis</span><span class="p">,</span> <span class="o">*</span><span class="n">arr_list</span><span class="p">)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">39</span> <span class="n">vstack</span> <span class="o">=</span> <span class="n">row_stack</span> <span class="o">=</span> <span class="k">lambda</span> <span class="n">tup</span><span class="p">:</span> <span class="n">concatenate</span><span class="p">([</span><span class="n">atleast_2d</span><span class="p">(</span><span class="n">_m</span><span class="p">)</span> <span class="k">for</span> <span class="n">_m</span> <span class="ow">in</span> <span class="n">tup</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">40</span> <span class="k">def</span> <span class="nf">hstack</span><span class="p">(</span><span class="n">tup</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.<locals>.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>
|
||||
@@ -2655,13 +2734,24 @@ Using TensorFlow results in a much better execution time. Try it!</p>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">35</span> <span class="o">.</span><span class="n">format</span><span class="p">(</span><span class="n">fun_name</span><span class="p">,</span> <span class="n">parent_argnums</span><span class="p">))</span>
|
||||
<span class="ne">---> </span><span class="mi">36</span> <span class="bp">self</span><span class="o">.</span><span class="n">vjp</span> <span class="o">=</span> <span class="n">vjpmaker</span><span class="p">(</span><span class="n">parent_argnums</span><span class="p">,</span> <span class="n">value</span><span class="p">,</span> <span class="n">args</span><span class="p">,</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:56,</span> in <span class="ni">defvjp.<locals>.vjp_argnums</span><span class="nt">(argnums, ans, args, kwargs)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">53</span> <span class="n">argnums</span> <span class="o">=</span> <span class="n">kwargs</span><span class="o">.</span><span class="n">get</span><span class="p">(</span><span class="s1">'argnums'</span><span class="p">,</span> <span class="n">count</span><span class="p">())</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">54</span> <span class="n">vjps_dict</span> <span class="o">=</span> <span class="p">{</span><span class="n">argnum</span> <span class="p">:</span> <span class="n">translate_vjp</span><span class="p">(</span><span class="n">vjpmaker</span><span class="p">,</span> <span class="n">fun</span><span class="p">,</span> <span class="n">argnum</span><span class="p">)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">55</span> <span class="k">for</span> <span class="n">argnum</span><span class="p">,</span> <span class="n">vjpmaker</span> <span class="ow">in</span> <span class="nb">zip</span><span class="p">(</span><span class="n">argnums</span><span class="p">,</span> <span class="n">vjpmakers</span><span class="p">)}</span>
|
||||
<span class="ne">---> </span><span class="mi">56</span> <span class="k">def</span> <span class="nf">vjp_argnums</span><span class="p">(</span><span class="n">argnums</span><span class="p">,</span> <span class="n">ans</span><span class="p">,</span> <span class="n">args</span><span class="p">,</span> <span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">57</span> <span class="n">L</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">argnums</span><span class="p">)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">58</span> <span class="c1"># These first two cases are just optimizations</span>
|
||||
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/core.py:48,</span> in <span class="ni">defvjp_argnum.<locals>.vjp_argnums</span><span class="nt">(argnums, *args)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">47</span> <span class="k">def</span> <span class="nf">vjp_argnums</span><span class="p">(</span><span class="n">argnums</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="ne">---> </span><span class="mi">48</span> <span class="n">vjps</span> <span class="o">=</span> <span class="p">[</span><span class="n">vjpmaker</span><span class="p">(</span><span class="n">argnum</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">)</span> <span class="k">for</span> <span class="n">argnum</span> <span class="ow">in</span> <span class="n">argnums</span><span class="p">]</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">49</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="k">for</span> <span class="n">vjp</span> <span class="ow">in</span> <span class="n">vjps</span><span class="p">)</span>
|
||||
|
||||
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/core.py:48,</span> in <span class="ni"><listcomp></span><span class="nt">(.0)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">47</span> <span class="k">def</span> <span class="nf">vjp_argnums</span><span class="p">(</span><span class="n">argnums</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">):</span>
|
||||
<span class="ne">---> </span><span class="mi">48</span> <span class="n">vjps</span> <span class="o">=</span> <span class="p">[</span><span class="n">vjpmaker</span><span class="p">(</span><span class="n">argnum</span><span class="p">,</span> <span class="o">*</span><span class="n">args</span><span class="p">)</span> <span class="k">for</span> <span class="n">argnum</span> <span class="ow">in</span> <span class="n">argnums</span><span class="p">]</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">49</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="k">for</span> <span class="n">vjp</span> <span class="ow">in</span> <span class="n">vjps</span><span class="p">)</span>
|
||||
|
||||
<span class="nn">File ~/miniforge3/envs/myenv/lib/python3.9/site-packages/autograd/numpy/numpy_vjps.py:537,</span> in <span class="ni">grad_concatenate_args</span><span class="nt">(argnum, ans, axis_args, kwargs)</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">532</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">tensordot_adjoint_1</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</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">axes</span><span class="p">,</span> <span class="n">An</span><span class="p">,</span> <span class="n">Bn</span><span class="p">:</span> <span class="k">lambda</span> <span class="n">B</span><span class="p">:</span> <span class="n">match_complex</span><span class="p">(</span><span class="n">A</span><span class="p">,</span> <span class="n">tensordot_adjoint_0</span><span class="p">(</span><span class="n">B</span><span class="p">,</span> <span class="n">G</span><span class="p">,</span> <span class="n">axes</span><span class="p">,</span> <span class="n">An</span><span class="p">,</span> <span class="n">Bn</span><span class="p">)),</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">533</span> <span class="k">lambda</span> <span class="n">ans</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">axes</span><span class="p">,</span> <span class="n">An</span><span class="p">,</span> <span class="n">Bn</span><span class="p">:</span> <span class="k">lambda</span> <span class="n">B</span><span class="p">:</span> <span class="n">match_complex</span><span class="p">(</span><span class="n">G</span><span class="p">,</span> <span class="n">anp</span><span class="o">.</span><span class="n">tensordot</span><span class="p">(</span><span class="n">A</span><span class="p">,</span> <span class="n">B</span><span class="p">,</span> <span class="n">axes</span><span class="p">)))</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">534</span> <span class="n">defvjp</span><span class="p">(</span><span class="n">anp</span><span class="o">.</span><span class="n">outer</span><span class="p">,</span> <span class="k">lambda</span> <span class="n">ans</span><span class="p">,</span> <span class="n">a</span><span class="p">,</span> <span class="n">b</span> <span class="p">:</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">match_complex</span><span class="p">(</span><span class="n">a</span><span class="p">,</span> <span class="n">anp</span><span class="o">.</span><span class="n">dot</span><span class="p">(</span><span class="n">g</span><span class="p">,</span> <span class="n">b</span><span class="o">.</span><span class="n">T</span><span class="p">)),</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">535</span> <span class="k">lambda</span> <span class="n">ans</span><span class="p">,</span> <span class="n">a</span><span class="p">,</span> <span class="n">b</span> <span class="p">:</span> <span class="k">lambda</span> <span class="n">g</span><span class="p">:</span> <span class="n">match_complex</span><span class="p">(</span><span class="n">b</span><span class="p">,</span> <span class="n">anp</span><span class="o">.</span><span class="n">dot</span><span class="p">(</span><span class="n">a</span><span class="o">.</span><span class="n">T</span><span class="p">,</span> <span class="n">g</span><span class="p">)))</span>
|
||||
<span class="ne">--> </span><span class="mi">537</span> <span class="k">def</span> <span class="nf">grad_concatenate_args</span><span class="p">(</span><span class="n">argnum</span><span class="p">,</span> <span class="n">ans</span><span class="p">,</span> <span class="n">axis_args</span><span class="p">,</span> <span class="n">kwargs</span><span class="p">):</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">538</span> <span class="n">axis</span><span class="p">,</span> <span class="n">args</span> <span class="o">=</span> <span class="n">axis_args</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">axis_args</span><span class="p">[</span><span class="mi">1</span><span class="p">:]</span>
|
||||
<span class="g g-Whitespace"> </span><span class="mi">539</span> <span class="n">sizes</span> <span class="o">=</span> <span class="p">[</span><span class="n">anp</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">axis</span><span class="p">]</span> <span class="k">for</span> <span class="n">a</span> <span class="ow">in</span> <span class="n">args</span><span class="p">[:</span><span class="n">argnum</span><span class="p">]]</span>
|
||||
|
||||
<span class="ne">KeyboardInterrupt</span>:
|
||||
</pre></div>
|
||||
|
||||
Reference in New Issue
Block a user