<?xml version="1.0" encoding="utf-8"?><feed xmlns="http://www.w3.org/2005/Atom" ><generator uri="https://jekyllrb.com/" version="3.10.0">Jekyll</generator><link href="https://mattyred.github.io/feed.xml" rel="self" type="application/atom+xml" /><link href="https://mattyred.github.io/" rel="alternate" type="text/html" /><updated>2026-08-22T16:14:03+00:00</updated><id>https://mattyred.github.io/feed.xml</id><title type="html">Mattia Rosso - Personal page</title><subtitle>personal description</subtitle><author><name>Mattia Rosso</name><email>mattia.rosso@kaust.edu.sa</email></author><entry><title type="html">Double Descent and Bayesian Linear Regression</title><link href="https://mattyred.github.io/posts/" rel="alternate" type="text/html" title="Double Descent and Bayesian Linear Regression" /><published>2026-02-05T00:00:00+00:00</published><updated>2026-02-05T00:00:00+00:00</updated><id>https://mattyred.github.io/future-post</id><content type="html" xml:base="https://mattyred.github.io/posts/"><![CDATA[<p>This notebook explores the double descent phenomenon through the lens of Bayesian linear regression. It demonstrates how model complexity affects generalization by comparing polynomial and Legendre basis expansions.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="nn">warnings</span>
<span class="n">warnings</span><span class="p">.</span><span class="n">filterwarnings</span><span class="p">(</span><span class="s">'ignore'</span><span class="p">)</span>
<span class="kn">import</span> <span class="nn">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="nn">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">import</span> <span class="nn">seaborn</span> <span class="k">as</span> <span class="n">sns</span>
<span class="kn">from</span> <span class="nn">scipy.stats</span> <span class="kn">import</span> <span class="n">multivariate_normal</span>
<span class="kn">from</span> <span class="nn">sklearn</span> <span class="kn">import</span> <span class="n">datasets</span><span class="p">,</span> <span class="n">linear_model</span>
<span class="kn">from</span> <span class="nn">sklearn.metrics</span> <span class="kn">import</span> <span class="n">mean_squared_error</span>
<span class="kn">import</span> <span class="nn">scipy</span>
<span class="kn">import</span> <span class="nn">pandas</span> <span class="k">as</span> <span class="n">pd</span>
<span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">seed</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>

<span class="c1"># Function to generate sinusoidal dataset
</span><span class="k">def</span> <span class="nf">generate_data</span><span class="p">(</span><span class="n">N</span><span class="p">,</span> <span class="n">noise_std</span><span class="o">=</span><span class="mf">0.5</span><span class="p">):</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linspace</span><span class="p">(</span><span class="o">-</span><span class="mi">3</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="n">N</span><span class="p">)</span>
    <span class="n">y_true</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">x</span><span class="p">)</span>
    <span class="n">y</span> <span class="o">=</span> <span class="n">y_true</span> <span class="o">+</span> <span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">normal</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">noise_std</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">N</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">y_true</span>

<span class="c1"># Basis function expansion
</span><span class="k">def</span> <span class="nf">basis_expansion</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="s">'linear'</span><span class="p">,</span> <span class="n">poly_degree</span><span class="o">=</span><span class="mi">3</span><span class="p">):</span>
    <span class="k">if</span> <span class="n">basis</span> <span class="o">==</span> <span class="s">'linear'</span><span class="p">:</span>
        <span class="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">vstack</span><span class="p">((</span><span class="n">np</span><span class="p">.</span><span class="n">ones_like</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">x</span><span class="p">)).</span><span class="n">T</span>
    <span class="k">elif</span> <span class="n">basis</span> <span class="o">==</span> <span class="s">'polynomial'</span><span class="p">:</span>
        <span class="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">vstack</span><span class="p">([</span><span class="n">x</span><span class="o">**</span><span class="n">i</span> <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">poly_degree</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)]).</span><span class="n">T</span>
    <span class="k">elif</span> <span class="n">basis</span> <span class="o">==</span> <span class="s">'legendre'</span><span class="p">:</span>
        <span class="n">degrees</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="n">poly_degree</span> <span class="o">+</span> <span class="mi">1</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">scipy</span><span class="p">.</span><span class="n">special</span><span class="p">.</span><span class="n">eval_legendre</span><span class="p">(</span><span class="n">degrees</span><span class="p">[:,</span> <span class="bp">None</span><span class="p">],</span> <span class="n">x</span><span class="p">).</span><span class="n">T</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="k">raise</span> <span class="nb">ValueError</span><span class="p">(</span><span class="s">"Unknown basis"</span><span class="p">)</span>

<span class="c1"># Compute MLE for linear regression
</span><span class="k">def</span> <span class="nf">compute_mle</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="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">inv</span><span class="p">(</span><span class="n">X</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">X</span><span class="p">)</span> <span class="o">@</span> <span class="n">X</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">y</span>

<span class="c1"># Draw samples from prior
</span><span class="k">def</span> <span class="nf">draw_prior_samples</span><span class="p">(</span><span class="n">prior_mean</span><span class="p">,</span> <span class="n">prior_cov</span><span class="p">,</span> <span class="n">num_samples</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">multivariate_normal</span><span class="p">(</span><span class="n">prior_mean</span><span class="p">,</span> <span class="n">prior_cov</span><span class="p">,</span> <span class="n">num_samples</span><span class="p">)</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">blr</span><span class="p">(</span><span class="n">N</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">gamma</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="s">'polynomial'</span><span class="p">,</span> <span class="n">sigma</span><span class="o">=</span><span class="mf">0.1</span><span class="p">):</span>
    <span class="c1"># 1. Data Generation and Masking
</span>    <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">y_true</span> <span class="o">=</span> <span class="n">generate_data</span><span class="p">(</span><span class="n">N</span><span class="p">,</span> <span class="n">noise_std</span><span class="o">=</span><span class="n">sigma</span><span class="p">)</span>
    <span class="c1">#mask = ~(((x &gt;= 0) &amp; (x &lt;= 3)) | ((x &gt;= -4) &amp; (x &lt;= -3)) | ((x &gt;= -2.5) &amp; (x &lt;= -1)))
</span>    <span class="c1">#x, y, y_true = x[mask], y[mask], y_true[mask]
</span>    
    <span class="n">fig</span><span class="p">,</span> <span class="n">axs</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="n">subplots</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span> <span class="mi">6</span><span class="p">,</span> <span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">12</span><span class="p">,</span> <span class="mi">6</span><span class="p">),</span> <span class="n">sharey</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">sharex</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="n">axs</span> <span class="o">=</span> <span class="n">axs</span><span class="p">.</span><span class="n">flatten</span><span class="p">()</span>
    <span class="n">logML_values</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">mem_values</span> <span class="o">=</span> <span class="p">[]</span>

    <span class="n">degrees</span> <span class="o">=</span> <span class="p">[</span><span class="mi">1</span><span class="p">,</span><span class="mi">3</span><span class="p">,</span><span class="mi">5</span><span class="p">,</span><span class="mi">7</span><span class="p">,</span><span class="mi">9</span><span class="p">,</span><span class="mi">11</span><span class="p">,</span><span class="mi">13</span><span class="p">,</span><span class="mi">15</span><span class="p">,</span><span class="mi">17</span><span class="p">,</span><span class="mi">19</span><span class="p">,</span><span class="mi">21</span><span class="p">,</span><span class="mi">23</span><span class="p">]</span>  <span class="c1"># Polynomial degrees to evaluate
</span>    
    <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">degree</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">degrees</span><span class="p">):</span>
        <span class="n">ax</span> <span class="o">=</span> <span class="n">axs</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>
        <span class="n">X</span> <span class="o">=</span> <span class="n">basis_expansion</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="n">basis</span><span class="p">,</span> <span class="n">poly_degree</span><span class="o">=</span><span class="n">degree</span><span class="p">)</span>
        <span class="n">n</span><span class="p">,</span> <span class="n">d</span> <span class="o">=</span> <span class="n">X</span><span class="p">.</span><span class="n">shape</span> <span class="c1"># n: datapoints, d: parameters
</span>        
        <span class="c1"># 2. Prior and Posterior 
</span>        <span class="c1"># Prior: p(theta) = N(0, (1/gamma)*I)
</span>        <span class="n">S0_inv</span> <span class="o">=</span> <span class="n">gamma</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">eye</span><span class="p">(</span><span class="n">d</span><span class="p">)</span>
        
        <span class="c1"># Posterior covariance Sn: inverse of precision matrix
</span>        <span class="c1"># Precision = Prior_Precision + (sigma^-2 * X.T @ X)
</span>        <span class="n">jitter</span> <span class="o">=</span> <span class="mf">1e-9</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">eye</span><span class="p">(</span><span class="n">d</span><span class="p">)</span>
        <span class="n">Sn_inv</span> <span class="o">=</span> <span class="n">S0_inv</span> <span class="o">+</span> <span class="p">(</span><span class="n">sigma</span><span class="o">**-</span><span class="mi">2</span><span class="p">)</span> <span class="o">*</span> <span class="n">X</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">X</span> <span class="o">+</span> <span class="n">jitter</span>
        <span class="n">Sn</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">inv</span><span class="p">(</span><span class="n">Sn_inv</span><span class="p">)</span>
        <span class="n">mn</span> <span class="o">=</span> <span class="n">Sn</span> <span class="o">@</span> <span class="p">(</span><span class="n">sigma</span><span class="o">**-</span><span class="mi">2</span> <span class="o">*</span> <span class="n">X</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">y</span><span class="p">)</span>
        
        <span class="c1"># 3. Stable Marginal Likelihood
</span>        <span class="n">zn_cov</span> <span class="o">=</span> <span class="n">sigma</span><span class="o">**</span><span class="mi">2</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">eye</span><span class="p">(</span><span class="n">n</span><span class="p">)</span> <span class="o">+</span> <span class="mi">1</span><span class="o">/</span><span class="n">gamma</span> <span class="o">*</span> <span class="n">X</span> <span class="o">@</span> <span class="n">X</span><span class="p">.</span><span class="n">T</span>
        <span class="n">zn_cov</span> <span class="o">+=</span> <span class="mf">1e-6</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">eye</span><span class="p">(</span><span class="n">n</span><span class="p">)</span>
        <span class="n">sign</span><span class="p">,</span> <span class="n">log_det_cov</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">slogdet</span><span class="p">(</span><span class="n">zn_cov</span><span class="p">)</span>
        <span class="n">sol</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">solve</span><span class="p">(</span><span class="n">zn_cov</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
        <span class="n">quad_form</span> <span class="o">=</span> <span class="n">y</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">sol</span>
        <span class="n">log_ml</span> <span class="o">=</span> <span class="o">-</span><span class="n">n</span><span class="o">/</span><span class="mi">2</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="mi">2</span><span class="o">*</span><span class="n">np</span><span class="p">.</span><span class="n">pi</span><span class="p">)</span> <span class="o">-</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="n">log_det_cov</span> <span class="o">-</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="n">quad_form</span> 
        <span class="n">logML_values</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">log_ml</span><span class="p">)</span>

        <span class="c1"># 4. Total Memorization (bits) [cite: 96, 135]
</span>        <span class="c1"># mem = H(prior) - H(posterior)
</span>        <span class="c1"># For Gaussian: 0.5 * log2(|S0| / |Sn|) = 0.5 * log2(|Sn_inv| / |S0_inv|)
</span>        <span class="n">_</span><span class="p">,</span> <span class="n">log_det_S0_inv</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">slogdet</span><span class="p">(</span><span class="n">S0_inv</span><span class="p">)</span>
        <span class="n">_</span><span class="p">,</span> <span class="n">log_det_Sn_inv</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">slogdet</span><span class="p">(</span><span class="n">Sn_inv</span><span class="p">)</span>
        <span class="n">total_mem</span> <span class="o">=</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="p">(</span><span class="n">log_det_Sn_inv</span> <span class="o">-</span> <span class="n">log_det_S0_inv</span><span class="p">)</span> <span class="o">/</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="mi">2</span><span class="p">)</span>
        <span class="n">bits_per_param</span> <span class="o">=</span> <span class="n">total_mem</span> <span class="o">/</span> <span class="n">d</span>
        <span class="n">mem_values</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">bits_per_param</span><span class="p">)</span>

        <span class="c1"># 5. Visualization
</span>        <span class="n">xtest</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linspace</span><span class="p">(</span><span class="o">-</span><span class="mf">3.5</span><span class="p">,</span> <span class="mf">3.5</span><span class="p">,</span> <span class="mi">200</span><span class="p">)</span>
        <span class="n">Xtest</span> <span class="o">=</span> <span class="n">basis_expansion</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="n">basis</span><span class="p">,</span> <span class="n">poly_degree</span><span class="o">=</span><span class="n">degree</span><span class="p">)</span>
        
        <span class="c1"># Draws from the posterior to show generalization vs overfitting [cite: 5, 147]
</span>        <span class="n">posterior_samples</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">multivariate_normal</span><span class="p">(</span><span class="n">mn</span><span class="p">,</span> <span class="n">Sn</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>
        <span class="k">for</span> <span class="n">j</span><span class="p">,</span> <span class="n">w_post</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">posterior_samples</span><span class="p">):</span>
            <span class="n">ax</span><span class="p">.</span><span class="n">plot</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">Xtest</span> <span class="o">@</span> <span class="n">w_post</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'blue'</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.1</span><span class="p">,</span> <span class="n">lw</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>

        <span class="n">ax</span><span class="p">.</span><span class="n">plot</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">Xtest</span> <span class="o">@</span> <span class="n">mn</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'red'</span><span class="p">,</span> <span class="n">lw</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">'Posterior Mean'</span> <span class="k">if</span> <span class="n">i</span><span class="o">==</span><span class="mi">0</span> <span class="k">else</span> <span class="bp">None</span><span class="p">)</span>
        <span class="n">sns</span><span class="p">.</span><span class="n">scatterplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="n">y</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'black'</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">15</span><span class="p">,</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">'Train Data'</span> <span class="k">if</span> <span class="n">i</span><span class="o">==</span><span class="mi">0</span> <span class="k">else</span> <span class="bp">None</span><span class="p">)</span>
        <span class="n">ax</span><span class="p">.</span><span class="n">plot</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">xtest</span><span class="p">),</span> <span class="n">color</span><span class="o">=</span><span class="s">'green'</span><span class="p">,</span> <span class="n">ls</span><span class="o">=</span><span class="s">'--'</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">'Ground Truth'</span> <span class="k">if</span> <span class="n">i</span><span class="o">==</span><span class="mi">0</span> <span class="k">else</span> <span class="bp">None</span><span class="p">)</span>
        
        <span class="n">ax</span><span class="p">.</span><span class="n">set_ylim</span><span class="p">([</span><span class="n">np</span><span class="p">.</span><span class="nb">min</span><span class="p">(</span><span class="n">y</span><span class="p">)</span><span class="o">-</span><span class="mf">1.5</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">y</span><span class="p">)</span><span class="o">+</span><span class="mf">1.5</span><span class="p">])</span>
        <span class="n">ax</span><span class="p">.</span><span class="n">set_title</span><span class="p">(</span><span class="sa">f</span><span class="s">'d=</span><span class="si">{</span><span class="n">d</span><span class="si">}</span><span class="s">, mem=</span><span class="si">{</span><span class="n">bits_per_param</span><span class="si">:</span><span class="p">.</span><span class="mi">2</span><span class="n">f</span><span class="si">}</span><span class="s"> </span><span class="se">\n</span><span class="s"> LML=</span><span class="si">{</span><span class="n">log_ml</span><span class="si">:</span><span class="p">.</span><span class="mi">2</span><span class="n">f</span><span class="si">}</span><span class="s">'</span><span class="p">,</span> <span class="n">fontsize</span><span class="o">=</span><span class="mi">12</span><span class="p">)</span>

    <span class="n">fig</span><span class="p">.</span><span class="n">legend</span><span class="p">(</span><span class="n">loc</span><span class="o">=</span><span class="s">'upper center'</span><span class="p">,</span> <span class="n">ncol</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">bbox_to_anchor</span><span class="o">=</span><span class="p">(</span><span class="mf">0.5</span><span class="p">,</span> <span class="mf">1.05</span><span class="p">))</span>
    <span class="n">plt</span><span class="p">.</span><span class="n">tight_layout</span><span class="p">()</span>
    <span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>
    
    <span class="k">return</span> <span class="n">logML_values</span><span class="p">,</span> <span class="n">mem_values</span><span class="p">,</span> <span class="n">degrees</span>
</code></pre></div></div>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">logML_values</span><span class="p">,</span> <span class="n">mem_values</span><span class="p">,</span> <span class="n">degrees</span> <span class="o">=</span> <span class="n">blr</span><span class="p">(</span><span class="n">N</span><span class="o">=</span><span class="mi">15</span><span class="p">,</span> <span class="n">gamma</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="s">'legendre'</span><span class="p">,</span> <span class="n">sigma</span><span class="o">=</span><span class="mf">0.1</span><span class="p">)</span>
</code></pre></div></div>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_2_0.png" alt="png" /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">degrees</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="n">mem_values</span><span class="p">,</span> <span class="n">marker</span><span class="o">=</span><span class="s">'o'</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">xlabel</span><span class="p">(</span><span class="s">'Polynomial Degree'</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">ylabel</span><span class="p">(</span><span class="s">'Memorization (bits/param)'</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>
</code></pre></div></div>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_3_0.png" alt="png" /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">degrees</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="n">logML_values</span><span class="p">,</span> <span class="n">marker</span><span class="o">=</span><span class="s">'o'</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">xlabel</span><span class="p">(</span><span class="s">'Polynomial Degree'</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">ylabel</span><span class="p">(</span><span class="s">'Log Marginal Likelihood'</span><span class="p">)</span>  
<span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>
</code></pre></div></div>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_4_0.png" alt="png" /></p>

<h2 id="double-descent-in-polynomial-regression">Double descent in polynomial regression</h2>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">compute_y_from_x</span><span class="p">(</span><span class="n">X</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="n">ndarray</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">add</span><span class="p">(</span><span class="mf">2.0</span> <span class="o">*</span> <span class="n">X</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="n">cos</span><span class="p">(</span><span class="n">X</span> <span class="o">*</span> <span class="mi">25</span><span class="p">))[:,</span> <span class="mi">0</span><span class="p">]</span>

<span class="n">low</span><span class="p">,</span> <span class="n">high</span> <span class="o">=</span> <span class="o">-</span><span class="mf">1.0</span><span class="p">,</span> <span class="mf">1.0</span>
<span class="n">num_data</span> <span class="o">=</span> <span class="mi">15</span>
<span class="n">num_features_list</span> <span class="o">=</span> <span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="mi">4</span><span class="p">,</span> <span class="mi">5</span><span class="p">,</span> <span class="mi">6</span><span class="p">,</span> <span class="mi">7</span><span class="p">,</span> <span class="mi">8</span><span class="p">,</span> <span class="mi">9</span><span class="p">,</span> <span class="mi">10</span><span class="p">,</span> <span class="mi">11</span><span class="p">,</span> <span class="mi">12</span><span class="p">,</span> <span class="mi">13</span><span class="p">,</span> <span class="mi">14</span><span class="p">,</span> <span class="mi">15</span><span class="p">,</span> 
                     <span class="mi">16</span><span class="p">,</span> <span class="mi">17</span><span class="p">,</span> <span class="mi">18</span><span class="p">,</span> <span class="mi">19</span><span class="p">,</span> <span class="mi">20</span><span class="p">,</span> <span class="mi">21</span><span class="p">,</span> <span class="mi">22</span><span class="p">,</span> <span class="mi">23</span><span class="p">,</span> <span class="mi">24</span><span class="p">,</span> <span class="mi">25</span><span class="p">,</span> <span class="mi">30</span><span class="p">,</span> <span class="mi">40</span><span class="p">,</span> <span class="mi">50</span><span class="p">,</span> <span class="mi">100</span><span class="p">,</span> <span class="mi">200</span><span class="p">]</span>

<span class="c1"># Fixed training data for consistency across plots
</span><span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">seed</span><span class="p">(</span><span class="mi">42</span><span class="p">)</span> 
<span class="n">X_train</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">uniform</span><span class="p">(</span><span class="n">low</span><span class="o">=</span><span class="n">low</span><span class="p">,</span> <span class="n">high</span><span class="o">=</span><span class="n">high</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="p">(</span><span class="n">num_data</span><span class="p">,</span> <span class="mi">1</span><span class="p">))</span>
<span class="n">y_train</span> <span class="o">=</span> <span class="n">compute_y_from_x</span><span class="p">(</span><span class="n">X_train</span><span class="p">)</span>

<span class="n">X_test</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linspace</span><span class="p">(</span><span class="n">start</span><span class="o">=</span><span class="n">low</span><span class="p">,</span> <span class="n">stop</span><span class="o">=</span><span class="n">high</span><span class="p">,</span> <span class="n">num</span><span class="o">=</span><span class="mi">1000</span><span class="p">).</span><span class="n">reshape</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>
<span class="n">y_test</span> <span class="o">=</span> <span class="n">compute_y_from_x</span><span class="p">(</span><span class="n">X_test</span><span class="p">)</span>

<span class="n">mse_list</span> <span class="o">=</span> <span class="p">[]</span>
<span class="n">preds_list</span> <span class="o">=</span> <span class="p">[]</span> <span class="c1"># Store predictions for subplots
</span>
<span class="c1"># --- Computations ---
</span><span class="k">for</span> <span class="n">num_features</span> <span class="ow">in</span> <span class="n">num_features_list</span><span class="p">:</span>
    <span class="n">feature_degrees</span> <span class="o">=</span> <span class="mi">1</span> <span class="o">+</span> <span class="n">np</span><span class="p">.</span><span class="n">arange</span><span class="p">(</span><span class="n">num_features</span><span class="p">).</span><span class="n">astype</span><span class="p">(</span><span class="nb">int</span><span class="p">)</span>
    
    <span class="c1"># Fit using Legendre Polynomials
</span>    <span class="n">X_train_poly</span> <span class="o">=</span> <span class="n">scipy</span><span class="p">.</span><span class="n">special</span><span class="p">.</span><span class="n">eval_legendre</span><span class="p">(</span><span class="n">feature_degrees</span><span class="p">,</span> <span class="n">X_train</span><span class="p">)</span>
    <span class="n">X_test_poly</span> <span class="o">=</span> <span class="n">scipy</span><span class="p">.</span><span class="n">special</span><span class="p">.</span><span class="n">eval_legendre</span><span class="p">(</span><span class="n">feature_degrees</span><span class="p">,</span> <span class="n">X_test</span><span class="p">)</span>
    <span class="c1">#X_train_poly = np.vander(X_train.flatten(), num_features + 1, increasing=True)[:, 1:]
</span>    <span class="c1">#X_test_poly = np.vander(X_test.flatten(), num_features + 1, increasing=True)[:, 1:]
</span>    
    <span class="c1"># Solve via Moore-Penrose pseudoinverse
</span>    <span class="n">beta_hat</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">pinv</span><span class="p">(</span><span class="n">X_train_poly</span><span class="p">)</span> <span class="o">@</span> <span class="n">y_train</span>
    
    <span class="n">y_train_pred</span> <span class="o">=</span> <span class="n">X_train_poly</span> <span class="o">@</span> <span class="n">beta_hat</span>
    <span class="n">y_test_pred</span> <span class="o">=</span> <span class="n">X_test_poly</span> <span class="o">@</span> <span class="n">beta_hat</span>
    
    <span class="n">mse_list</span><span class="p">.</span><span class="n">append</span><span class="p">({</span>
        <span class="s">"Num. Parameters"</span><span class="p">:</span> <span class="n">num_features</span><span class="p">,</span>
        <span class="s">"Train MSE"</span><span class="p">:</span> <span class="n">mean_squared_error</span><span class="p">(</span><span class="n">y_train</span><span class="p">,</span> <span class="n">y_train_pred</span><span class="p">),</span>
        <span class="s">"Test MSE"</span><span class="p">:</span> <span class="n">mean_squared_error</span><span class="p">(</span><span class="n">y_test</span><span class="p">,</span> <span class="n">y_test_pred</span><span class="p">),</span>
    <span class="p">})</span>
    <span class="n">preds_list</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">y_test_pred</span><span class="p">)</span>

<span class="c1"># --- Visualization: Subplot Grid ---
</span><span class="n">n_plots</span> <span class="o">=</span> <span class="nb">len</span><span class="p">(</span><span class="n">num_features_list</span><span class="p">)</span>
<span class="n">cols</span> <span class="o">=</span> <span class="mi">5</span>
<span class="n">rows</span> <span class="o">=</span> <span class="p">(</span><span class="n">n_plots</span> <span class="o">//</span> <span class="n">cols</span><span class="p">)</span> <span class="o">+</span> <span class="p">(</span><span class="mi">1</span> <span class="k">if</span> <span class="n">n_plots</span> <span class="o">%</span> <span class="n">cols</span> <span class="o">!=</span> <span class="mi">0</span> <span class="k">else</span> <span class="mi">0</span><span class="p">)</span>

<span class="n">fig</span><span class="p">,</span> <span class="n">axes</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="n">subplots</span><span class="p">(</span><span class="n">rows</span><span class="p">,</span> <span class="n">cols</span><span class="p">,</span> <span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">20</span><span class="p">,</span> <span class="mi">4</span> <span class="o">*</span> <span class="n">rows</span><span class="p">),</span> <span class="n">constrained_layout</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">axes</span> <span class="o">=</span> <span class="n">axes</span><span class="p">.</span><span class="n">flatten</span><span class="p">()</span>

<span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">num_features</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">num_features_list</span><span class="p">):</span>
    <span class="n">ax</span> <span class="o">=</span> <span class="n">axes</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>
    <span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">X_test</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">y</span><span class="o">=</span><span class="n">y_test</span><span class="p">,</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'gray'</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.5</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">"True"</span> <span class="k">if</span> <span class="n">i</span><span class="o">==</span><span class="mi">0</span> <span class="k">else</span> <span class="bp">None</span><span class="p">)</span>
    <span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">X_test</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">y</span><span class="o">=</span><span class="n">preds_list</span><span class="p">[</span><span class="n">i</span><span class="p">],</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'red'</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="sa">f</span><span class="s">"Fit (P=</span><span class="si">{</span><span class="n">num_features</span><span class="si">}</span><span class="s">)"</span><span class="p">)</span>
    <span class="n">sns</span><span class="p">.</span><span class="n">scatterplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">X_train</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">y</span><span class="o">=</span><span class="n">y_train</span><span class="p">,</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">20</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'black'</span><span class="p">,</span> <span class="n">edgecolor</span><span class="o">=</span><span class="s">'white'</span><span class="p">)</span>
    
    <span class="n">ax</span><span class="p">.</span><span class="n">set_ylim</span><span class="p">(</span><span class="o">-</span><span class="mi">4</span><span class="p">,</span> <span class="mi">4</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="n">set_title</span><span class="p">(</span><span class="sa">f</span><span class="s">"Features: </span><span class="si">{</span><span class="n">num_features</span><span class="si">}</span><span class="s">"</span><span class="p">)</span>
    <span class="k">if</span> <span class="n">i</span> <span class="o">%</span> <span class="n">cols</span> <span class="o">!=</span> <span class="mi">0</span><span class="p">:</span> <span class="n">ax</span><span class="p">.</span><span class="n">set_ylabel</span><span class="p">(</span><span class="s">""</span><span class="p">)</span>
    <span class="k">if</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="p">(</span><span class="n">n_plots</span> <span class="o">-</span> <span class="n">cols</span><span class="p">):</span> <span class="n">ax</span><span class="p">.</span><span class="n">set_xlabel</span><span class="p">(</span><span class="s">""</span><span class="p">)</span>

<span class="c1"># Hide unused subplots
</span><span class="k">for</span> <span class="n">j</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">i</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="nb">len</span><span class="p">(</span><span class="n">axes</span><span class="p">)):</span>
    <span class="n">axes</span><span class="p">[</span><span class="n">j</span><span class="p">].</span><span class="n">axis</span><span class="p">(</span><span class="s">'off'</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="n">suptitle</span><span class="p">(</span><span class="s">"Polynomial Fits Across Different Model Capacities"</span><span class="p">,</span> <span class="n">fontsize</span><span class="o">=</span><span class="mi">20</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>

<span class="c1"># --- Visualization: MSE Curves ---
</span><span class="n">mse_df</span> <span class="o">=</span> <span class="n">pd</span><span class="p">.</span><span class="n">DataFrame</span><span class="p">(</span><span class="n">mse_list</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">figure</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">10</span><span class="p">,</span> <span class="mi">6</span><span class="p">))</span>
<span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="n">mse_df</span><span class="p">,</span> <span class="n">x</span><span class="o">=</span><span class="s">"Num. Parameters"</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="s">"Test MSE"</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">"Test"</span><span class="p">)</span>
<span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="n">mse_df</span><span class="p">,</span> <span class="n">x</span><span class="o">=</span><span class="s">"Num. Parameters"</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="s">"Train MSE"</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">"Train"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">axvline</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">num_data</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">"black"</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="s">"--"</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">"Interpolation Threshold"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">yscale</span><span class="p">(</span><span class="s">"log"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">ylim</span><span class="p">(</span><span class="n">bottom</span><span class="o">=</span><span class="mf">1e-3</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">title</span><span class="p">(</span><span class="s">"Double Descent Phenomenon"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="n">legend</span><span class="p">()</span>
<span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>
</code></pre></div></div>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_6_0.png" alt="png" /></p>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_6_1.png" alt="png" /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">seed</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span>

<span class="k">def</span> <span class="nf">generate_data</span><span class="p">(</span><span class="n">N</span><span class="p">,</span> <span class="n">noise_std</span><span class="o">=</span><span class="mf">0.5</span><span class="p">):</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">sort</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">uniform</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">N</span><span class="p">))</span>
    <span class="n">y_true</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">x</span> <span class="o">*</span> <span class="mi">5</span><span class="p">)</span>
    <span class="n">y</span> <span class="o">=</span> <span class="n">y_true</span> <span class="o">+</span> <span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">normal</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">noise_std</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">N</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">,</span> <span class="n">y_true</span>

<span class="k">def</span> <span class="nf">basis_expansion</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="s">'legendre'</span><span class="p">,</span> <span class="n">poly_degree</span><span class="o">=</span><span class="mi">3</span><span class="p">):</span>
    <span class="k">if</span> <span class="n">basis</span> <span class="o">==</span> <span class="s">'linear'</span><span class="p">:</span>
        <span class="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">vstack</span><span class="p">((</span><span class="n">np</span><span class="p">.</span><span class="n">ones_like</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">x</span><span class="p">)).</span><span class="n">T</span>
    <span class="k">elif</span> <span class="n">basis</span> <span class="o">==</span> <span class="s">'polynomial'</span><span class="p">:</span>
        <span class="c1"># Using Vandermonde: [x^0, x^1, ..., x^d]
</span>        <span class="k">return</span> <span class="n">np</span><span class="p">.</span><span class="n">vander</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">poly_degree</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">increasing</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="k">elif</span> <span class="n">basis</span> <span class="o">==</span> <span class="s">'legendre'</span><span class="p">:</span>
        <span class="c1"># Correctly evaluate Legendre polynomials for each degree
</span>        <span class="n">X</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">((</span><span class="nb">len</span><span class="p">(</span><span class="n">x</span><span class="p">),</span> <span class="n">poly_degree</span> <span class="o">+</span> <span class="mi">1</span><span class="p">))</span>
        <span class="k">for</span> <span class="n">d</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">poly_degree</span> <span class="o">+</span> <span class="mi">1</span><span class="p">):</span>
            <span class="n">X</span><span class="p">[:,</span> <span class="n">d</span><span class="p">]</span> <span class="o">=</span> <span class="n">scipy</span><span class="p">.</span><span class="n">special</span><span class="p">.</span><span class="n">eval_legendre</span><span class="p">(</span><span class="n">d</span><span class="p">,</span> <span class="n">x</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">X</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="k">raise</span> <span class="nb">ValueError</span><span class="p">(</span><span class="s">"Unknown basis"</span><span class="p">)</span>

<span class="k">def</span> <span class="nf">blr_grid_analysis</span><span class="p">(</span><span class="n">N</span><span class="o">=</span><span class="mi">15</span><span class="p">,</span> <span class="n">sigma</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="s">'legendre'</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">y_true</span> <span class="o">=</span> <span class="n">generate_data</span><span class="p">(</span><span class="n">N</span><span class="p">,</span> <span class="n">noise_std</span><span class="o">=</span><span class="n">sigma</span><span class="p">)</span>
    <span class="n">degrees</span> <span class="o">=</span> <span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="mi">5</span><span class="p">,</span> <span class="mi">10</span><span class="p">,</span> <span class="mi">15</span><span class="p">,</span> <span class="mi">25</span><span class="p">,</span> <span class="mi">50</span><span class="p">,</span> <span class="mi">100</span><span class="p">]</span>
    <span class="n">gammas</span> <span class="o">=</span> <span class="p">[</span><span class="mf">1e-4</span><span class="p">,</span> <span class="mf">1e-3</span><span class="p">,</span> <span class="mf">1e-2</span><span class="p">,</span> <span class="mi">1</span><span class="p">]</span>
    
    <span class="c1"># Grid Plot Setup
</span>    <span class="n">fig</span><span class="p">,</span> <span class="n">axes</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="n">subplots</span><span class="p">(</span><span class="nb">len</span><span class="p">(</span><span class="n">gammas</span><span class="p">),</span> <span class="nb">len</span><span class="p">(</span><span class="n">degrees</span><span class="p">),</span> <span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">24</span><span class="p">,</span> <span class="mi">12</span><span class="p">),</span> <span class="n">sharex</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">sharey</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    
    <span class="n">results</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="n">xtest</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linspace</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">,</span> <span class="mi">400</span><span class="p">)</span>
    <span class="n">ytest_true</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">xtest</span> <span class="o">*</span> <span class="mi">5</span><span class="p">)</span>

    <span class="k">for</span> <span class="n">r</span><span class="p">,</span> <span class="n">gamma</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">gammas</span><span class="p">):</span>
        <span class="k">for</span> <span class="n">c</span><span class="p">,</span> <span class="n">degree</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">degrees</span><span class="p">):</span>
            <span class="n">ax</span> <span class="o">=</span> <span class="n">axes</span><span class="p">[</span><span class="n">r</span><span class="p">,</span> <span class="n">c</span><span class="p">]</span>
            
            <span class="c1"># Expansion
</span>            <span class="n">X</span> <span class="o">=</span> <span class="n">basis_expansion</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="n">basis</span><span class="p">,</span> <span class="n">poly_degree</span><span class="o">=</span><span class="n">degree</span><span class="p">)</span>
            <span class="n">Xtest</span> <span class="o">=</span> <span class="n">basis_expansion</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="n">basis</span><span class="p">,</span> <span class="n">poly_degree</span><span class="o">=</span><span class="n">degree</span><span class="p">)</span>
            <span class="n">n</span><span class="p">,</span> <span class="n">d</span> <span class="o">=</span> <span class="n">X</span><span class="p">.</span><span class="n">shape</span> 
            
            <span class="c1"># Posterior
</span>            <span class="n">S0_inv</span> <span class="o">=</span> <span class="n">gamma</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">eye</span><span class="p">(</span><span class="n">d</span><span class="p">)</span>
            <span class="n">Sn_inv</span> <span class="o">=</span> <span class="n">S0_inv</span> <span class="o">+</span> <span class="p">(</span><span class="n">sigma</span><span class="o">**-</span><span class="mi">2</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">X</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">X</span><span class="p">)</span>
            <span class="n">Sn</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">inv</span><span class="p">(</span><span class="n">Sn_inv</span><span class="p">)</span>
            <span class="n">mn</span> <span class="o">=</span> <span class="n">Sn</span> <span class="o">@</span> <span class="p">(</span><span class="n">sigma</span><span class="o">**-</span><span class="mi">2</span> <span class="o">*</span> <span class="n">X</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">y</span><span class="p">)</span>
            
            <span class="c1"># --- Metrics ---
</span>            <span class="n">y_train_pred</span> <span class="o">=</span> <span class="n">X</span> <span class="o">@</span> <span class="n">mn</span>
            <span class="n">y_test_pred</span> <span class="o">=</span> <span class="n">Xtest</span> <span class="o">@</span> <span class="n">mn</span>
            <span class="n">train_mse</span> <span class="o">=</span> <span class="n">mean_squared_error</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">y_train_pred</span><span class="p">)</span>
            <span class="n">test_mse</span> <span class="o">=</span> <span class="n">mean_squared_error</span><span class="p">(</span><span class="n">ytest_true</span><span class="p">,</span> <span class="n">y_test_pred</span><span class="p">)</span>
            
            <span class="c1"># Marginal Likelihood
</span>            <span class="n">sse</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nb">sum</span><span class="p">((</span><span class="n">y</span> <span class="o">-</span> <span class="n">y_train_pred</span><span class="p">)</span><span class="o">**</span><span class="mi">2</span><span class="p">)</span>
            <span class="n">_</span><span class="p">,</span> <span class="n">log_det_Sn_inv</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">slogdet</span><span class="p">(</span><span class="n">Sn_inv</span><span class="p">)</span>
            <span class="n">log_ml</span> <span class="o">=</span> <span class="p">((</span><span class="n">d</span><span class="o">/</span><span class="mi">2</span><span class="p">)</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="n">gamma</span><span class="p">)</span> <span class="o">-</span> <span class="p">(</span><span class="n">n</span><span class="o">/</span><span class="mi">2</span><span class="p">)</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="mi">2</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">pi</span><span class="p">)</span> <span class="o">+</span> <span class="n">n</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="mi">1</span><span class="o">/</span><span class="n">sigma</span><span class="p">)</span>
                      <span class="o">-</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="n">log_det_Sn_inv</span> <span class="o">-</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="p">((</span><span class="n">sigma</span><span class="o">**-</span><span class="mi">2</span><span class="p">)</span> <span class="o">*</span> <span class="n">sse</span> <span class="o">+</span> <span class="n">gamma</span> <span class="o">*</span> <span class="p">(</span><span class="n">mn</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">mn</span><span class="p">)))</span>

            <span class="c1"># Information Gain (KL Divergence)
</span>            <span class="n">_</span><span class="p">,</span> <span class="n">log_det_Sn</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">slogdet</span><span class="p">(</span><span class="n">Sn</span><span class="p">)</span>
            <span class="n">log_det_S0</span> <span class="o">=</span> <span class="o">-</span><span class="n">d</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="n">gamma</span><span class="p">)</span>
            <span class="n">kl_div</span> <span class="o">=</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="p">(</span><span class="n">log_det_S0</span> <span class="o">-</span> <span class="n">log_det_Sn</span> <span class="o">+</span> <span class="n">gamma</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">trace</span><span class="p">(</span><span class="n">Sn</span><span class="p">)</span> <span class="o">+</span> <span class="n">gamma</span> <span class="o">*</span> <span class="p">(</span><span class="n">mn</span><span class="p">.</span><span class="n">T</span> <span class="o">@</span> <span class="n">mn</span><span class="p">)</span> <span class="o">-</span> <span class="n">d</span><span class="p">)</span>
            <span class="n">bits_per_param</span> <span class="o">=</span> <span class="p">(</span><span class="n">kl_div</span> <span class="o">/</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="mi">2</span><span class="p">))</span> <span class="o">/</span> <span class="n">d</span>

            <span class="n">results</span><span class="p">.</span><span class="n">append</span><span class="p">({</span>
                <span class="s">"Gamma"</span><span class="p">:</span> <span class="n">gamma</span><span class="p">,</span>
                <span class="s">"Degree"</span><span class="p">:</span> <span class="n">degree</span><span class="p">,</span>
                <span class="s">"LML"</span><span class="p">:</span> <span class="n">log_ml</span><span class="p">,</span>
                <span class="s">"Train MSE"</span><span class="p">:</span> <span class="n">train_mse</span><span class="p">,</span>
                <span class="s">"Test MSE"</span><span class="p">:</span> <span class="n">test_mse</span><span class="p">,</span>
                <span class="s">"Bits/Param"</span><span class="p">:</span> <span class="n">bits_per_param</span>
            <span class="p">})</span>

            <span class="c1"># --- Plotting ---
</span>            <span class="n">posterior_samples</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="n">multivariate_normal</span><span class="p">(</span><span class="n">mn</span><span class="p">,</span> <span class="n">Sn</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="mi">8</span><span class="p">)</span>
            <span class="k">for</span> <span class="n">w_post</span> <span class="ow">in</span> <span class="n">posterior_samples</span><span class="p">:</span>
                <span class="n">ax</span><span class="p">.</span><span class="n">plot</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">Xtest</span> <span class="o">@</span> <span class="n">w_post</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'blue'</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.1</span><span class="p">,</span> <span class="n">lw</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>

            <span class="n">ax</span><span class="p">.</span><span class="n">plot</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">y_test_pred</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'red'</span><span class="p">,</span> <span class="n">lw</span><span class="o">=</span><span class="mf">1.5</span><span class="p">)</span>
            <span class="n">ax</span><span class="p">.</span><span class="n">scatter</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">color</span><span class="o">=</span><span class="s">'black'</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">10</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.5</span><span class="p">)</span>
            <span class="n">ax</span><span class="p">.</span><span class="n">plot</span><span class="p">(</span><span class="n">xtest</span><span class="p">,</span> <span class="n">ytest_true</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'green'</span><span class="p">,</span> <span class="n">ls</span><span class="o">=</span><span class="s">'--'</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.6</span><span class="p">)</span>
            <span class="n">ax</span><span class="p">.</span><span class="n">set_ylim</span><span class="p">([</span><span class="o">-</span><span class="mf">2.5</span><span class="p">,</span> <span class="mf">2.5</span><span class="p">])</span>
            
            <span class="k">if</span> <span class="n">r</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="n">ax</span><span class="p">.</span><span class="n">set_title</span><span class="p">(</span><span class="sa">f</span><span class="s">'Deg </span><span class="si">{</span><span class="n">degree</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>
            <span class="k">if</span> <span class="n">c</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span> <span class="n">ax</span><span class="p">.</span><span class="n">set_ylabel</span><span class="p">(</span><span class="sa">f</span><span class="s">'γ = </span><span class="si">{</span><span class="n">gamma</span><span class="si">}</span><span class="s">'</span><span class="p">)</span>

    <span class="n">plt</span><span class="p">.</span><span class="n">tight_layout</span><span class="p">()</span>
    <span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>
    <span class="n">df</span> <span class="o">=</span> <span class="n">pd</span><span class="p">.</span><span class="n">DataFrame</span><span class="p">(</span><span class="n">results</span><span class="p">)</span> 
    <span class="k">return</span> <span class="n">df</span>

<span class="n">N</span> <span class="o">=</span> <span class="mi">15</span>
<span class="n">results_df</span> <span class="o">=</span> <span class="n">blr_grid_analysis</span><span class="p">(</span><span class="n">N</span><span class="o">=</span><span class="n">N</span><span class="p">,</span> <span class="n">sigma</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span> <span class="n">basis</span><span class="o">=</span><span class="s">'legendre'</span><span class="p">)</span>
</code></pre></div></div>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_7_0.png" alt="png" /></p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">fig</span><span class="p">,</span> <span class="p">(</span><span class="n">ax1</span><span class="p">,</span> <span class="n">ax2</span><span class="p">)</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="n">subplots</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">10</span><span class="p">,</span> <span class="mi">3</span><span class="p">))</span>
<span class="n">unique_gammas</span> <span class="o">=</span> <span class="n">results_df</span><span class="p">[</span><span class="s">"Gamma"</span><span class="p">].</span><span class="n">nunique</span><span class="p">()</span>
<span class="n">palette</span> <span class="o">=</span> <span class="n">sns</span><span class="p">.</span><span class="n">color_palette</span><span class="p">(</span><span class="s">"colorblind"</span><span class="p">,</span> <span class="n">n_colors</span><span class="o">=</span><span class="n">unique_gammas</span><span class="p">)</span>
<span class="c1"># 1. Log Marginal Likelihood Plot
</span><span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="n">results_df</span><span class="p">,</span> <span class="n">x</span><span class="o">=</span><span class="s">"Degree"</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="s">"LML"</span><span class="p">,</span> <span class="n">hue</span><span class="o">=</span><span class="s">"Gamma"</span><span class="p">,</span> <span class="n">marker</span><span class="o">=</span><span class="s">"o"</span><span class="p">,</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax1</span><span class="p">,</span> <span class="n">palette</span><span class="o">=</span><span class="n">palette</span><span class="p">,</span> <span class="n">legend</span><span class="o">=</span><span class="bp">None</span><span class="p">)</span>
<span class="n">ax1</span><span class="p">.</span><span class="n">set_title</span><span class="p">(</span><span class="s">"Log Marginal Likelihood"</span><span class="p">)</span>
<span class="n">ax1</span><span class="p">.</span><span class="n">set_xscale</span><span class="p">(</span><span class="s">"log"</span><span class="p">)</span>

<span class="c1"># 2. Test MSE Plot (Double Descent)
</span><span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="n">results_df</span><span class="p">,</span> <span class="n">x</span><span class="o">=</span><span class="s">"Degree"</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="s">"Test MSE"</span><span class="p">,</span> <span class="n">hue</span><span class="o">=</span><span class="s">"Gamma"</span><span class="p">,</span> <span class="n">marker</span><span class="o">=</span><span class="s">"o"</span><span class="p">,</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax2</span><span class="p">,</span> <span class="n">palette</span><span class="o">=</span><span class="n">palette</span><span class="p">)</span>
<span class="n">sns</span><span class="p">.</span><span class="n">lineplot</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="n">results_df</span><span class="p">,</span> <span class="n">x</span><span class="o">=</span><span class="s">"Degree"</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="s">"Train MSE"</span><span class="p">,</span> <span class="n">hue</span><span class="o">=</span><span class="s">"Gamma"</span><span class="p">,</span> <span class="n">marker</span><span class="o">=</span><span class="s">"o"</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="s">'--'</span><span class="p">,</span> <span class="n">ax</span><span class="o">=</span><span class="n">ax2</span><span class="p">,</span> <span class="n">palette</span><span class="o">=</span><span class="n">palette</span><span class="p">,</span> <span class="n">legend</span><span class="o">=</span><span class="bp">None</span><span class="p">)</span>
<span class="n">ax2</span><span class="p">.</span><span class="n">set_title</span><span class="p">(</span><span class="s">"Mean Squared Error"</span><span class="p">)</span>
<span class="n">ax2</span><span class="p">.</span><span class="n">set_yscale</span><span class="p">(</span><span class="s">"log"</span><span class="p">)</span>
<span class="n">ax2</span><span class="p">.</span><span class="n">set_xscale</span><span class="p">(</span><span class="s">"log"</span><span class="p">)</span>
<span class="n">ax2</span><span class="p">.</span><span class="n">set_ylim</span><span class="p">(</span><span class="n">bottom</span><span class="o">=</span><span class="mf">1e-5</span><span class="p">)</span>
<span class="n">ax2</span><span class="p">.</span><span class="n">axvline</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">N</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="s">'black'</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="s">'--'</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="s">'Interpolation Threshold'</span><span class="p">)</span>


<span class="n">handles</span><span class="p">,</span> <span class="n">labels</span> <span class="o">=</span> <span class="n">ax2</span><span class="p">.</span><span class="n">get_legend_handles_labels</span><span class="p">()</span>
<span class="n">ax2</span><span class="p">.</span><span class="n">get_legend</span><span class="p">().</span><span class="n">remove</span><span class="p">()</span> <span class="c1"># Remove the default legend from inside the plot
</span>
<span class="n">fig</span><span class="p">.</span><span class="n">legend</span><span class="p">(</span><span class="n">handles</span><span class="p">,</span> <span class="n">labels</span><span class="p">,</span> <span class="n">loc</span><span class="o">=</span><span class="s">'lower center'</span><span class="p">,</span> <span class="n">title</span><span class="o">=</span><span class="s">'Prior precision'</span><span class="p">,</span> <span class="n">ncol</span><span class="o">=</span><span class="nb">len</span><span class="p">(</span><span class="n">labels</span><span class="p">),</span> 
           <span class="n">bbox_to_anchor</span><span class="o">=</span><span class="p">(</span><span class="mf">0.5</span><span class="p">,</span> <span class="o">-</span><span class="mf">0.15</span><span class="p">),</span> <span class="n">frameon</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="n">tight_layout</span><span class="p">()</span>
<span class="n">plt</span><span class="p">.</span><span class="n">show</span><span class="p">()</span>
</code></pre></div></div>

<p><img src="../images/bayesian_linear_regression_files/bayesian_linear_regression_8_0.png" alt="png" /></p>

<h2 id="references">References</h2>

<p>[1] Schaeffer, Rylan, et al. “Double descent demystified: Identifying, interpreting &amp; ablating the sources of a deep learning puzzle.” arXiv preprint arXiv:2303.14151 (2023).</p>]]></content><author><name>Mattia Rosso</name><email>mattia.rosso@kaust.edu.sa</email></author><category term="Bayesian Deep Learning" /><category term="Double Descent" /><summary type="html"><![CDATA[This notebook explores the double descent phenomenon through the lens of Bayesian linear regression. It demonstrates how model complexity affects generalization by comparing polynomial and Legendre basis expansions.]]></summary></entry></feed>