<?xml version="1.0" encoding="utf-8"?><feed xmlns="http://www.w3.org/2005/Atom" xml:lang="en"><generator uri="https://jekyllrb.com/" version="4.4.1">Jekyll</generator><link href="https://ckrapu.github.io/feed.xml" rel="self" type="application/atom+xml"/><link href="https://ckrapu.github.io/" rel="alternate" type="text/html" hreflang="en"/><updated>2026-08-07T00:16:31+00:00</updated><id>https://ckrapu.github.io/feed.xml</id><title type="html">blank</title><subtitle></subtitle><entry><title type="html">A few simple tricks for building chatbots</title><link href="https://ckrapu.github.io/blog/2026/simple-things-ai-engineering/" rel="alternate" type="text/html" title="A few simple tricks for building chatbots"/><published>2026-08-01T00:00:00+00:00</published><updated>2026-08-01T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2026/simple-things-ai-engineering</id><content type="html" xml:base="https://ckrapu.github.io/blog/2026/simple-things-ai-engineering/"><![CDATA[<p>Thanks to LLMs capable of reasoning and using 1M+ tokens in a single prompt, the simple chatbot has grown up. Now it is a proper agent with virtually all the digital capabilities of its human counterpart.</p> <p>That said, hooking up the OpenAI API to a simple back-and-forth conversational system is still the entry point for AI engineering.</p> <p>Even though 99.9% of this system’s complexity is abstracted away behind the inference API, there are a few concepts that I see people take a while to grasp as they learn how to build systems on top of AI.</p> <p>This short note discusses a few of them. I haven’t found good sources or references for these ideas elsewhere, so I hope this note is useful.</p> <h2 id="efficient-citation-via-clipboarding">Efficient citation via clipboarding</h2> <p>LLM-driven search à la Perplexity may seem like a 2024 thing, but it is probably still one of the top three most valuable use cases for AI. In this workstream, being able to reference sources and efficiently show them to the user is critical.</p> <p>Many starter kits for chatbots (e.g., Chainlit and Open WebUI) can easily render Markdown. Markdown allows you to write hyperlinks with a string like <a href="https://www.youtube.com/watch?v=dQw4w9WgXcQ">Please click me!</a>, which is easy to write and simple to remember.</p> <p>LLMs are perfectly capable of writing hyperlink strings like these and require only a few short prompt instructions to do so. However, asking the LLM to write the hyperlink comes with a few limitations:</p> <ul> <li>While the user is waiting for the text to finish streaming, they will see partial hyperlink strings like <code class="language-plaintext highlighter-rouge">[Please click me!](https://www.yo..</code>. Once the string is complete, it will be rendered, causing the partial text to disappear and the hyperlink to appear in its place. This jump is irritating and creates a poor user experience.</li> <li>As system complexity and scope grow, you will likely begin citing sources with very long URLs that take a long time (&gt;100 ms) to write token by token.</li> <li>This approach is token-inefficient because the same result can be produced more cheaply.</li> </ul> <p>The solution is simple. As the reader, you have probably already implemented it or reinvented it on the spot.</p> <p>Regardless, here it is: instruct the LLM to write short citation keys like <code class="language-plaintext highlighter-rouge">@d3</code> or <code class="language-plaintext highlighter-rouge">#a1</code>. A postprocessor sitting on top of the text stream then captures these keys and deterministically replaces them with full hyperlinks. I personally call this the “clipboard” approach in all of my codebases.</p> <p>Some caveats do apply. Namely, you need some mechanism to give the source a nice descriptor (like “Annual Shareholder Report, 2026”). The naive approach from earlier lets you prompt the LLM to create the hyperlinked text on the fly; this approach doesn’t let you do that out of the box.</p> <p>This approach also separates citation styling from prompting logic. That separation is a good idea (IMO), but it can be a pain if you’d rather just write prompts instead of extending application logic. Visually, here’s what it looks like:</p> <p><img src="/images/2026-08-01-simple-things-ai-engineering/streaming-citations.gif" alt="Two citation approaches shown while streaming text, including a citation hover tooltip."/></p> <p>You can also generalize this approach to any string that the LLM might generate inefficiently. If your LLM is powering an agent with tools for accessing APIs or reading files in a structured format, users may ask questions whose answers require the model to regurgitate data nearly verbatim. If the source is a long CSV or JSON file, having the LLM copy its contents token by token is wasteful. Use the clipboard approach to copy and paste the output instead.</p> <p>Your choice of placeholder token is mostly unimportant, so long as it is small and unlikely to show up in the generated text for unrelated reasons. Avoid using standard formats like <code class="language-plaintext highlighter-rouge">[1]</code> or <code class="language-plaintext highlighter-rouge">[2]</code>, since hallucinating LLMs will typically generate false citations in those formats. For richer styling, you could probably use XML or a similar format.</p> <h2 id="constructing-follow-up-questions-or-actions">Constructing follow-up questions or actions</h2> <p>Most domain-specific AI platforms (e.g., Microsoft Copilot, Perplexity, and Glean) now supply recommended or follow-up questions for users to click. This is partly a gimmick from product managers to juice usage numbers, but it can sometimes be genuinely useful.</p> <p>The best follow-up questions (e.g., “Would you like me to also search for restaurants closer to you?”) know both 1) what the user originally asked and 2) how the system responded in the most recent turn. The standard approach I saw in a number of systems from 2025 and earlier was to run a specific prompt to make the LLM generate these questions.</p> <p>Again, this is wasteful and inefficient. Nowadays, using structured generation plus streaming, you can have an LLM answer the user’s question and generate the follow-ups in a single prompt using a schema like the following:</p> <p>(show a schema that uses the OpenAI API to answer and generate follow-ups)</p> <p><img src="/images/2026-08-01-simple-things-ai-engineering/follow-up-api-calls.gif" alt="The naive approach makes separate API calls for the response and follow-ups, while the improved approach returns both over one continuous text stream."/></p> <p>The combined approach can be cost-neutral, as most of the instruction tokens from the follow-up prompt are moved into the main response prompt. In either case, the instructions should remain in the prompt cache 100% of the time, so the prefill cost of additional instructions is low relative to the decode cost. Since both approaches emit the same number of tokens, the overall cost of generation is mostly unchanged.</p> <p>From a user-experience perspective, emitting the follow-up questions from the same stream as the main response probably saves at least 100 milliseconds by avoiding the overhead of a separate API call.</p> <p>The primary weakness of a combined response/follow-up prompt is that it requires the LLM to follow more instructions per prompt, which generally leads to poorer adherence and effectiveness as the prompt grows. Open-weight models released since late 2025 generally do not have this issue in standard workflows. It may still be a problem if you have a very complex prompt with many instructions.</p> <h2 id="using-cheap-llms-to-capture-user-attention">Using cheap LLMs to capture user attention</h2> <p>Every few months, a new high-speed LLM inference provider comes on the market, and the industry reacts by touting the advantages of zero-lag conversations, instant coding agents, and so on. I think the market for that is actually very small. When we have faster LLMs, we don’t simply make existing workflows faster; we generally use the extra latency headroom to make those workflows better. In other words, long-running AI tasks are going to be a fact of life.</p> <p>Consequently, we need to resort to all sorts of tricks to keep the user on our page before they’re emotionally spent from waiting 800 milliseconds for the next dopamine hit.</p> <p>From an AI workflow perspective, this means providing continuous state updates on long-running tasks. Perplexity was an early leader here and had a great design that ensured something interesting was always happening on screen, even if you had to wait for 15+ seconds.</p> <p>My go-to method is to task a very small, cheap LLM, such as Llama 3.2 1B or Gemma 2B, with providing short commentary snippets using prompts like this:</p> <div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Your role is to provide one or two sentences of brief commentary on the state of a running program. Write at a high level, use concise language, and use the present tense.

Original user query:
{query}

Recent actions:
{tool_calls}
</code></pre></div></div> <p>The output should be something like <code class="language-plaintext highlighter-rouge">Searched Salesforce cases from the last two weeks</code>. This can be hugely beneficial for keeping users engaged.</p> <p>Make sure to keep the tool calls ordered from oldest to newest so that you can leverage the prompt cache effectively here. Putting the newest calls first invalidates the prompt cache. If you run this every 3 seconds during a 60-second task, you can easily send this prompt 20+ times, making cache management important <em>even within the scope of a single turn</em>.</p> <p>OpenAI did this quite a bit with ChatGPT, as you can see below:</p> <p><img src="/images/chatgpt-progress.gif" alt="ChatGPT displaying live progress updates while searching for nearby restaurants."/></p> <p>If you have any more ideas to include here, let me know!</p>]]></content><author><name></name></author><category term="llm"/><category term="chatbot"/><summary type="html"><![CDATA[Everyone eventually figures them out on their own, right?]]></summary></entry><entry><title type="html">Poverty Bayes: fitting million-parameter models for pennies with serverless MCMC</title><link href="https://ckrapu.github.io/blog/2026/poverty-bayes-serverless-mcmc/" rel="alternate" type="text/html" title="Poverty Bayes: fitting million-parameter models for pennies with serverless MCMC"/><published>2026-05-27T00:00:00+00:00</published><updated>2026-05-27T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2026/poverty-bayes-serverless-mcmc</id><content type="html" xml:base="https://ckrapu.github.io/blog/2026/poverty-bayes-serverless-mcmc/"><![CDATA[<p>It’s a good time to be an applied probabilist. The deep learning revolution has led to tremendous improvements in the $ / flops department, and we Bayesians can easily hop on this train! During grad school, I used to spend nights and weekends babysitting MCMC runs on my GeForce Titan XP running in my bedroom (by the way, thank you <a href="https://www.nvidia.com/en-us/industries/higher-education-research/academic-grant-program/">NVIDIA Academic Grant Program</a>) while simultaneously trying to keep the waste heat from cooking me as I slept. If you are a newcomer to this field, rejoice in the knowledge that all this suffering is a thing of the past. A slew of companies are rushing to the fore with user-friendly platforms for renting GPUs. For prototyping, I really enjoy working with Modal since I’m cheap and I’m too lazy to keep managing my own fleet.</p> <p>In this post, I’ll show a workflow for using GPU-based inference on Modal for a model which is very large by the standards of Bayesian statisticians \((\vert \theta\vert \gt 10^6)\), by deploying to a datacenter GPU and renting it only for a short time.</p> <h2 id="model--data">Model &amp; data</h2> <p>We’ll use synthetic data for this example.</p> <p>I’ve chosen a hierarchical logistic regression for this post since it has a non-conjugate likelihood, appears commonly in practice, and can easily be assigned more parameters by increasing the number of covariates and/or the number of groups.</p> <p>Let \(i\), \(g\), and \(k\) denote the indices over observations, groups, and covariates. Furthermore, let \(x_i \in \mathbb{R}^{20}\) be the covariate vector for observation \(i\), and let \(g_i \in \{1,\ldots,100000\}\) identify its group. The data-generating process uses population-level slopes \(\beta_k\), group intercept deviations \(\alpha_g\), and group slope deviations \(\gamma_{gk}\). The binary outcome is generated from</p> \[Y_i \sim \operatorname{Bernoulli}(p_i), \qquad \operatorname{logit}(p_i) = \alpha + \alpha_{g_i} + \sum_{k=1}^{K} x_{ik}(\beta_k + \gamma_{g_i k}).\] <p>Essentially, this is a logistic regression with random slopes for 20 covariates and a random intercept for each of 100,000 groups in the data.</p> <p>We’ll use a non-centered parameterization for the group effects. The prior specification is</p> \[\begin{aligned} \alpha &amp;\sim \operatorname{Normal}(0, 1.5), \\ \beta_k &amp;\sim \operatorname{Normal}(0, 1), \\ \sigma_\alpha &amp;\sim \operatorname{HalfNormal}(1), \\ \sigma_{\gamma,k} &amp;\sim \operatorname{HalfNormal}(0.5), \\ z_{\alpha,g} &amp;\sim \operatorname{Normal}(0, 1), \\ z_{\gamma,gk} &amp;\sim \operatorname{Normal}(0, 1), \\ \alpha_g &amp;= \sigma_\alpha z_{\alpha,g}, \\ \gamma_{gk} &amp;= \sigma_{\gamma,k} z_{\gamma,gk}. \end{aligned}\] <p>The code below produces a synthetic dataset; we can control the overall sparsity of the response with the value of <code class="language-plaintext highlighter-rouge">α_true</code>.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>

<span class="n">RANDOM_SEED</span> <span class="o">=</span> <span class="mi">827</span>
<span class="n">rng</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="nf">default_rng</span><span class="p">(</span><span class="n">RANDOM_SEED</span><span class="p">)</span>

<span class="n">N</span> <span class="o">=</span> <span class="mi">1_000_000</span> <span class="c1"># Number of data points
</span><span class="n">G</span> <span class="o">=</span> <span class="mi">100_000</span>   <span class="c1"># Number of groups
</span><span class="n">K</span> <span class="o">=</span> <span class="mi">20</span>      <span class="c1"># Number of covariates / features
</span>
<span class="n">group_idx</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">integers</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">G</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">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>
<span class="n">X</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">normal</span><span class="p">(</span><span class="n">size</span><span class="o">=</span><span class="p">(</span><span class="n">N</span><span class="p">,</span> <span class="n">K</span><span class="p">)).</span><span class="nf">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="n">α_true</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="o">-</span><span class="mf">1.0</span><span class="p">)</span>
<span class="n">β_true</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">normal</span><span class="p">(</span><span class="mf">0.0</span><span class="p">,</span> <span class="mf">0.45</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">K</span><span class="p">).</span><span class="nf">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
<span class="n">σ_α_true</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">0.80</span><span class="p">)</span>
<span class="n">σ_γ_true</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">uniform</span><span class="p">(</span><span class="mf">0.15</span><span class="p">,</span> <span class="mf">0.35</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">K</span><span class="p">).</span><span class="nf">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
<span class="n">α_group_true</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">normal</span><span class="p">(</span><span class="mf">0.0</span><span class="p">,</span> <span class="n">σ_α_true</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="n">G</span><span class="p">).</span><span class="nf">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
<span class="n">γ_group_true</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">normal</span><span class="p">(</span><span class="mf">0.0</span><span class="p">,</span> <span class="n">σ_γ_true</span><span class="p">,</span> <span class="n">size</span><span class="o">=</span><span class="p">(</span><span class="n">G</span><span class="p">,</span> <span class="n">K</span><span class="p">)).</span><span class="nf">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="n">η</span> <span class="o">=</span> <span class="n">α_true</span> <span class="o">+</span> <span class="n">α_group_true</span><span class="p">[</span><span class="n">group_idx</span><span class="p">]</span> <span class="o">+</span> <span class="n">np</span><span class="p">.</span><span class="nf">sum</span><span class="p">(</span><span class="n">X</span> <span class="o">*</span> <span class="p">(</span><span class="n">β_true</span> <span class="o">+</span> <span class="n">γ_group_true</span><span class="p">[</span><span class="n">group_idx</span><span class="p">]),</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
<span class="n">p</span> <span class="o">=</span> <span class="mf">1.0</span> <span class="o">/</span> <span class="p">(</span><span class="mf">1.0</span> <span class="o">+</span> <span class="n">np</span><span class="p">.</span><span class="nf">exp</span><span class="p">(</span><span class="o">-</span><span class="n">η</span><span class="p">))</span>
<span class="n">y</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">binomial</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">p</span><span class="p">).</span><span class="nf">astype</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>

<span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Average of y: </span><span class="si">{</span><span class="n">np</span><span class="p">.</span><span class="nf">mean</span><span class="p">(</span><span class="n">y</span><span class="p">)</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="sh">"</span><span class="p">)</span>
</code></pre></div></div> <div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Average of y: 0.36
</code></pre></div></div> <p>Here, we define our model in PyMC. This is a fairly standard model definition without much nuance or many tricks.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">pymc</span> <span class="k">as</span> <span class="n">pm</span>
<span class="kn">import</span> <span class="n">pytensor</span>
<span class="kn">import</span> <span class="n">pytensor.tensor</span> <span class="k">as</span> <span class="n">pt</span>

<span class="n">pytensor</span><span class="p">.</span><span class="n">config</span><span class="p">.</span><span class="n">floatX</span> <span class="o">=</span> <span class="sh">"</span><span class="s">float32</span><span class="sh">"</span>

<span class="k">def</span> <span class="nf">build_model</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">group_idx</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="nf">asarray</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
    <span class="n">y</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>
    <span class="n">group_idx</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">group_idx</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>

    <span class="n">N</span><span class="p">,</span> <span class="n">K</span> <span class="o">=</span> <span class="n">X</span><span class="p">.</span><span class="n">shape</span>
    <span class="n">G</span> <span class="o">=</span> <span class="nf">int</span><span class="p">(</span><span class="n">group_idx</span><span class="p">.</span><span class="nf">max</span><span class="p">())</span> <span class="o">+</span> <span class="mi">1</span>
    <span class="n">coords</span> <span class="o">=</span> <span class="p">{</span>
        <span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="nf">arange</span><span class="p">(</span><span class="n">N</span><span class="p">),</span>
        <span class="sh">"</span><span class="s">group</span><span class="sh">"</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="nf">arange</span><span class="p">(</span><span class="n">G</span><span class="p">),</span>
        <span class="sh">"</span><span class="s">covariate</span><span class="sh">"</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="nf">arange</span><span class="p">(</span><span class="n">K</span><span class="p">),</span>
    <span class="p">}</span>

    <span class="k">with</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Model</span><span class="p">(</span><span class="n">coords</span><span class="o">=</span><span class="n">coords</span><span class="p">)</span> <span class="k">as</span> <span class="n">model</span><span class="p">:</span>
        <span class="n">X_data</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Data</span><span class="p">(</span><span class="sh">"</span><span class="s">X</span><span class="sh">"</span><span class="p">,</span> <span class="n">X</span><span class="p">,</span> <span class="n">dims</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">covariate</span><span class="sh">"</span><span class="p">))</span>
        <span class="n">group_lookup</span> <span class="o">=</span> <span class="n">pt</span><span class="p">.</span><span class="nf">as_tensor_variable</span><span class="p">(</span><span class="n">group_idx</span><span class="p">,</span> <span class="n">name</span><span class="o">=</span><span class="sh">"</span><span class="s">group_lookup</span><span class="sh">"</span><span class="p">)</span>

        <span class="n">α</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">α</span><span class="sh">"</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">1.5</span><span class="p">))</span>
        <span class="n">β</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">1.0</span><span class="p">),</span> <span class="n">dims</span><span class="o">=</span><span class="sh">"</span><span class="s">covariate</span><span class="sh">"</span><span class="p">)</span>
        <span class="n">σ_α</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">σ_α</span><span class="sh">"</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">1.0</span><span class="p">))</span>
        <span class="n">σ_γ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">σ_γ</span><span class="sh">"</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">0.5</span><span class="p">),</span> <span class="n">dims</span><span class="o">=</span><span class="sh">"</span><span class="s">covariate</span><span class="sh">"</span><span class="p">)</span>

        <span class="n">z_α</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">z_α</span><span class="sh">"</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">1.0</span><span class="p">),</span> <span class="n">dims</span><span class="o">=</span><span class="sh">"</span><span class="s">group</span><span class="sh">"</span><span class="p">)</span>
        <span class="n">z_γ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">z_γ</span><span class="sh">"</span><span class="p">,</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span> <span class="n">np</span><span class="p">.</span><span class="nf">float32</span><span class="p">(</span><span class="mf">1.0</span><span class="p">),</span> <span class="n">dims</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">group</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">covariate</span><span class="sh">"</span><span class="p">))</span>
        <span class="n">α_group</span> <span class="o">=</span> <span class="n">σ_α</span> <span class="o">*</span> <span class="n">z_α</span> 
        <span class="n">γ_group</span> <span class="o">=</span> <span class="n">σ_γ</span> <span class="o">*</span> <span class="n">z_γ</span>

        <span class="n">η</span> <span class="o">=</span> <span class="n">α</span> <span class="o">+</span> <span class="n">α_group</span><span class="p">[</span><span class="n">group_lookup</span><span class="p">]</span> <span class="o">+</span> <span class="n">pm</span><span class="p">.</span><span class="n">math</span><span class="p">.</span><span class="nf">sum</span><span class="p">(</span><span class="n">X_data</span> <span class="o">*</span> <span class="p">(</span><span class="n">β</span> <span class="o">+</span> <span class="n">γ_group</span><span class="p">[</span><span class="n">group_lookup</span><span class="p">]),</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">pm</span><span class="p">.</span><span class="nc">Bernoulli</span><span class="p">(</span><span class="sh">"</span><span class="s">Y</span><span class="sh">"</span><span class="p">,</span> <span class="n">logit_p</span><span class="o">=</span><span class="n">η</span><span class="p">,</span> <span class="n">observed</span><span class="o">=</span><span class="n">y</span><span class="p">,</span> <span class="n">dims</span><span class="o">=</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">)</span>

    <span class="k">return</span> <span class="n">model</span>
</code></pre></div></div> <p>To illustrate why this model could be challenging to fit using only a CPU, we profile the gradient evaluation time. The <code class="language-plaintext highlighter-rouge">dlogp</code> function represents \(\frac{\partial}{\partial \theta}[\log \tilde{p}]\) with \(\tilde{p}\) denoting the unnormalized posterior density; it may be calculated up to 1000 times per sample with the evaluation count increasing with the difficulty of the problem in terms of posterior geometry and/or nonlinearity. Hamiltonian Monte Carlo, the workhorse of most modern Bayesian programming frameworks, requires these gradient evaluations to draw Monte Carlo samples.</p> <p>The code cell below runs the gradient a few times and records the time taken.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">time</span>

<span class="k">def</span> <span class="nf">median_runtime</span><span class="p">(</span><span class="n">fn</span><span class="p">,</span> <span class="n">n</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>
    <span class="nf">fn</span><span class="p">()</span>
    <span class="n">times</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n</span><span class="p">):</span>
        <span class="n">start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="nf">perf_counter</span><span class="p">()</span>
        <span class="nf">fn</span><span class="p">()</span>
        <span class="n">times</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">time</span><span class="p">.</span><span class="nf">perf_counter</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span><span class="p">)</span>
    <span class="k">return</span> <span class="nf">float</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">median</span><span class="p">(</span><span class="n">times</span><span class="p">)),</span> <span class="n">times</span>

<span class="n">local_model</span> <span class="o">=</span> <span class="nf">build_model</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">group_idx</span><span class="p">)</span>
<span class="n">local_point</span> <span class="o">=</span> <span class="n">local_model</span><span class="p">.</span><span class="nf">initial_point</span><span class="p">()</span>

<span class="k">with</span> <span class="n">local_model</span><span class="p">:</span>
    <span class="n">dlogp_fn</span> <span class="o">=</span> <span class="n">local_model</span><span class="p">.</span><span class="nf">compile_dlogp</span><span class="p">()</span>

<span class="n">local_grad_median</span><span class="p">,</span> <span class="n">local_grad_times</span> <span class="o">=</span> <span class="nf">median_runtime</span><span class="p">(</span><span class="k">lambda</span><span class="p">:</span> <span class="nf">dlogp_fn</span><span class="p">(</span><span class="n">local_point</span><span class="p">))</span>
<span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Median grad eval time is </span><span class="si">{</span><span class="n">local_grad_median</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"> seconds</span><span class="sh">"</span><span class="p">)</span>
</code></pre></div></div> <div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Median grad eval time is 0.20 seconds
</code></pre></div></div> <p>If we try to extrapolate to the time required for running a full chain of 1000 samples, we find that assuming ~250 grad evals per sample, we would need 50 seconds per sample and 50,000 seconds (~14 hours) to run a full chain! This simply won’t do.</p> <p>Let’s deploy this remotely and see what we get!</p> <h2 id="running-remotely-on-modal">Running remotely on Modal</h2> <p>For those of you unfamiliar with the magical world of serverless GPU, please take a look at Modal. I will probably never be wealthy enough to outright own a beautiful server rack of datacenter GPUs, so this is the best I can get. Basically, Modal provides APIs for spinning up jobs on GPUs with very little friction and little downtime.</p> <p>To make this example run nicely on GPU, we’ll make a few adjustments. First, we’ll toggle it to use the Jax + NumPyro backend to work out-of-the-box with an NVIDIA GPU.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">modal</span>

<span class="n">modal_image</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">modal</span><span class="p">.</span><span class="n">Image</span><span class="p">.</span><span class="nf">debian_slim</span><span class="p">(</span><span class="n">python_version</span><span class="o">=</span><span class="sh">"</span><span class="s">3.13</span><span class="sh">"</span><span class="p">)</span>
    <span class="p">.</span><span class="nf">uv_pip_install</span><span class="p">(</span>
        <span class="sh">"</span><span class="s">arviz==1.1.0</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">jax[cuda12]==0.7.2</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">numpy==2.3.5</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">numpyro==0.19.0</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">pandas==2.3.3</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">pymc==6.0.1</span><span class="sh">"</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)</span>

<span class="n">app</span> <span class="o">=</span> <span class="n">modal</span><span class="p">.</span><span class="nc">App</span><span class="p">(</span><span class="sh">"</span><span class="s">i-should-be-working-right-now</span><span class="sh">"</span><span class="p">,</span> <span class="n">image</span><span class="o">=</span><span class="n">modal_image</span><span class="p">)</span>
</code></pre></div></div> <p>We’ll set a few environment flags for the float precision and the devices, define a helper to benchmark the execution time, and apply some Jax-isms to prep the compute graph and evaluate it.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">remote_model</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">group_idx</span><span class="p">):</span>
    <span class="kn">import</span> <span class="n">os</span>

    <span class="n">os</span><span class="p">.</span><span class="n">environ</span><span class="p">[</span><span class="sh">"</span><span class="s">JAX_PLATFORMS</span><span class="sh">"</span><span class="p">]</span> <span class="o">=</span> <span class="sh">"</span><span class="s">cuda</span><span class="sh">"</span>

    <span class="kn">import</span> <span class="n">jax</span>
    <span class="kn">import</span> <span class="n">jax.numpy</span> <span class="k">as</span> <span class="n">jnp</span>
    <span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>
    <span class="kn">import</span> <span class="n">pytensor</span>

    <span class="n">pytensor</span><span class="p">.</span><span class="n">config</span><span class="p">.</span><span class="n">floatX</span> <span class="o">=</span> <span class="sh">"</span><span class="s">float32</span><span class="sh">"</span>
    <span class="n">X</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
    <span class="n">y</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">y</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>
    <span class="n">group_idx</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">group_idx</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">int64</span><span class="p">)</span>

    <span class="n">gpu_devices</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="nf">devices</span><span class="p">(</span><span class="sh">"</span><span class="s">gpu</span><span class="sh">"</span><span class="p">)</span>
    <span class="n">gpu_check</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="nf">device_put</span><span class="p">(</span><span class="n">jnp</span><span class="p">.</span><span class="nf">ones</span><span class="p">((</span><span class="mi">512</span><span class="p">,</span> <span class="mi">512</span><span class="p">),</span> <span class="n">dtype</span><span class="o">=</span><span class="n">jnp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span> <span class="n">gpu_devices</span><span class="p">[</span><span class="mi">0</span><span class="p">]).</span><span class="nf">sum</span><span class="p">().</span><span class="nf">block_until_ready</span><span class="p">()</span>
    <span class="n">model</span> <span class="o">=</span> <span class="nf">build_model</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">group_idx</span><span class="p">)</span>

    <span class="n">info</span> <span class="o">=</span> <span class="p">{</span>
        <span class="sh">"</span><span class="s">N</span><span class="sh">"</span><span class="p">:</span> <span class="n">X</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span>
        <span class="sh">"</span><span class="s">G</span><span class="sh">"</span><span class="p">:</span> <span class="nf">int</span><span class="p">(</span><span class="n">group_idx</span><span class="p">.</span><span class="nf">max</span><span class="p">())</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">K</span><span class="sh">"</span><span class="p">:</span> <span class="n">X</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span>
        <span class="sh">"</span><span class="s">jax_backend</span><span class="sh">"</span><span class="p">:</span> <span class="n">jax</span><span class="p">.</span><span class="nf">default_backend</span><span class="p">(),</span>
        <span class="sh">"</span><span class="s">jax_gpu_devices</span><span class="sh">"</span><span class="p">:</span> <span class="p">[</span><span class="nf">str</span><span class="p">(</span><span class="n">device</span><span class="p">)</span> <span class="k">for</span> <span class="n">device</span> <span class="ow">in</span> <span class="n">gpu_devices</span><span class="p">],</span>
        <span class="sh">"</span><span class="s">gpu_check_sum</span><span class="sh">"</span><span class="p">:</span> <span class="nf">float</span><span class="p">(</span><span class="n">gpu_check</span><span class="p">),</span>
        <span class="sh">"</span><span class="s">x_dtype</span><span class="sh">"</span><span class="p">:</span> <span class="nf">str</span><span class="p">(</span><span class="n">X</span><span class="p">.</span><span class="n">dtype</span><span class="p">),</span>
        <span class="sh">"</span><span class="s">value_var_dtypes</span><span class="sh">"</span><span class="p">:</span> <span class="p">{</span><span class="n">var</span><span class="p">.</span><span class="n">name</span><span class="p">:</span> <span class="n">var</span><span class="p">.</span><span class="n">dtype</span> <span class="k">for</span> <span class="n">var</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">value_vars</span><span class="p">},</span>
    <span class="p">}</span>
    <span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">JAX backend: </span><span class="si">{</span><span class="n">info</span><span class="p">[</span><span class="sh">'</span><span class="s">jax_backend</span><span class="sh">'</span><span class="p">]</span><span class="si">}</span><span class="s">; GPU devices: </span><span class="si">{</span><span class="n">info</span><span class="p">[</span><span class="sh">'</span><span class="s">jax_gpu_devices</span><span class="sh">'</span><span class="p">]</span><span class="si">}</span><span class="sh">"</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">model</span><span class="p">,</span> <span class="n">info</span>


<span class="k">def</span> <span class="nf">median_runtime</span><span class="p">(</span><span class="n">fn</span><span class="p">,</span> <span class="n">n</span><span class="o">=</span><span class="mi">5</span><span class="p">,</span> <span class="n">synchronize</span><span class="o">=</span><span class="bp">None</span><span class="p">):</span>
    <span class="kn">import</span> <span class="n">time</span>

    <span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>

    <span class="k">if</span> <span class="n">synchronize</span> <span class="ow">is</span> <span class="bp">None</span><span class="p">:</span>
        <span class="n">synchronize</span> <span class="o">=</span> <span class="k">lambda</span> <span class="n">result</span><span class="p">:</span> <span class="bp">None</span>

    <span class="k">def</span> <span class="nf">timed_run</span><span class="p">():</span>
        <span class="n">start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="nf">perf_counter</span><span class="p">()</span>
        <span class="n">result</span> <span class="o">=</span> <span class="nf">fn</span><span class="p">()</span>
        <span class="nf">synchronize</span><span class="p">(</span><span class="n">result</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">time</span><span class="p">.</span><span class="nf">perf_counter</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span>

    <span class="nf">synchronize</span><span class="p">(</span><span class="nf">fn</span><span class="p">())</span>
    <span class="n">times</span> <span class="o">=</span> <span class="p">[</span><span class="nf">timed_run</span><span class="p">()</span> <span class="k">for</span> <span class="n">_</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n</span><span class="p">)]</span>
    <span class="k">return</span> <span class="nf">float</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">median</span><span class="p">(</span><span class="n">times</span><span class="p">)),</span> <span class="n">times</span>


<span class="k">def</span> <span class="nf">wait_jax</span><span class="p">(</span><span class="n">result</span><span class="p">):</span>
    <span class="kn">import</span> <span class="n">jax</span>

    <span class="n">jax</span><span class="p">.</span><span class="n">tree_util</span><span class="p">.</span><span class="nf">tree_map</span><span class="p">(</span><span class="k">lambda</span> <span class="n">x</span><span class="p">:</span> <span class="n">x</span><span class="p">.</span><span class="nf">block_until_ready</span><span class="p">(),</span> <span class="n">result</span><span class="p">)</span>


<span class="nd">@app.function</span><span class="p">(</span><span class="n">gpu</span><span class="o">=</span><span class="sh">"</span><span class="s">A100</span><span class="sh">"</span><span class="p">,</span> <span class="n">timeout</span><span class="o">=</span><span class="mi">2</span> <span class="o">*</span> <span class="mi">60</span> <span class="o">*</span> <span class="mi">60</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">profile_dlogp</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">group_idx</span><span class="p">,</span> <span class="n">n_evals</span><span class="o">=</span><span class="mi">5</span><span class="p">):</span>

    <span class="kn">import</span> <span class="n">jax</span>
    <span class="kn">import</span> <span class="n">jax.numpy</span> <span class="k">as</span> <span class="n">jnp</span>
    <span class="kn">from</span> <span class="n">pymc.sampling.jax</span> <span class="kn">import</span> <span class="n">get_jaxified_logp</span>

    <span class="n">model</span><span class="p">,</span> <span class="n">info</span> <span class="o">=</span> <span class="nf">remote_model</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">group_idx</span><span class="p">)</span>
    <span class="n">point</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="nf">initial_point</span><span class="p">()</span>

    <span class="k">with</span> <span class="n">model</span><span class="p">:</span>
        <span class="n">pytensor_dlogp</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="nf">compile_dlogp</span><span class="p">()</span>

    <span class="n">pytensor_median</span><span class="p">,</span> <span class="n">pytensor_times</span> <span class="o">=</span> <span class="nf">median_runtime</span><span class="p">(</span><span class="k">lambda</span><span class="p">:</span> <span class="nf">pytensor_dlogp</span><span class="p">(</span><span class="n">point</span><span class="p">),</span> <span class="n">n_evals</span><span class="p">)</span>

    <span class="n">values</span> <span class="o">=</span> <span class="p">[</span><span class="n">jnp</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">point</span><span class="p">[</span><span class="n">var</span><span class="p">.</span><span class="n">name</span><span class="p">])</span> <span class="k">for</span> <span class="n">var</span> <span class="ow">in</span> <span class="n">model</span><span class="p">.</span><span class="n">value_vars</span><span class="p">]</span>
    <span class="n">jax_loss</span> <span class="o">=</span> <span class="nf">get_jaxified_logp</span><span class="p">(</span><span class="n">model</span><span class="p">)</span>
    <span class="n">jax_dlogp</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="nf">jit</span><span class="p">(</span><span class="n">jax</span><span class="p">.</span><span class="nf">value_and_grad</span><span class="p">(</span><span class="n">jax_loss</span><span class="p">))</span>

    <span class="nf">wait_jax</span><span class="p">(</span><span class="nf">jax_dlogp</span><span class="p">(</span><span class="n">values</span><span class="p">))</span>
    <span class="n">jax_median</span><span class="p">,</span> <span class="n">jax_times</span> <span class="o">=</span> <span class="nf">median_runtime</span><span class="p">(</span><span class="k">lambda</span><span class="p">:</span> <span class="nf">jax_dlogp</span><span class="p">(</span><span class="n">values</span><span class="p">),</span> <span class="n">n_evals</span><span class="p">,</span> <span class="n">wait_jax</span><span class="p">)</span>

    <span class="n">profile</span> <span class="o">=</span> <span class="p">{</span>
        <span class="sh">"</span><span class="s">median_pytensor_dlogp_eval_seconds</span><span class="sh">"</span><span class="p">:</span> <span class="n">pytensor_median</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">median_jax_dlogp_eval_seconds</span><span class="sh">"</span><span class="p">:</span> <span class="n">jax_median</span><span class="p">,</span>
    <span class="p">}</span>
    <span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Median dlogp eval: PyTensor=</span><span class="si">{</span><span class="n">pytensor_median</span><span class="si">:</span><span class="p">.</span><span class="mi">3</span><span class="n">f</span><span class="si">}</span><span class="s">s, JAX=</span><span class="si">{</span><span class="n">jax_median</span><span class="si">:</span><span class="p">.</span><span class="mi">3</span><span class="n">f</span><span class="si">}</span><span class="s">s</span><span class="sh">"</span><span class="p">)</span>
    <span class="k">return</span> <span class="p">{</span><span class="o">**</span><span class="n">info</span><span class="p">,</span> <span class="o">**</span><span class="n">profile</span><span class="p">}</span>


<span class="nd">@app.function</span><span class="p">(</span><span class="n">gpu</span><span class="o">=</span><span class="sh">"</span><span class="s">A100</span><span class="sh">"</span><span class="p">,</span> <span class="n">timeout</span><span class="o">=</span><span class="mi">2</span> <span class="o">*</span> <span class="mi">60</span> <span class="o">*</span> <span class="mi">60</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">fit_hierarchical_logistic</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">group_idx</span><span class="p">,</span> <span class="n">draws</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">tune</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">random_seed</span><span class="o">=</span><span class="n">RANDOM_SEED</span><span class="p">):</span>
    <span class="kn">import</span> <span class="n">time</span>
    <span class="kn">import</span> <span class="n">arviz</span> <span class="k">as</span> <span class="n">az</span>
    <span class="kn">import</span> <span class="n">pymc</span> <span class="k">as</span> <span class="n">pm</span>

    <span class="n">model</span><span class="p">,</span> <span class="n">info</span> <span class="o">=</span> <span class="nf">remote_model</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">group_idx</span><span class="p">)</span>

    <span class="k">with</span> <span class="n">model</span><span class="p">:</span>
        <span class="n">start</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="nf">perf_counter</span><span class="p">()</span>
        <span class="n">idata</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nf">sample</span><span class="p">(</span>
            <span class="n">draws</span><span class="p">,</span>
            <span class="n">tune</span><span class="o">=</span><span class="n">tune</span><span class="p">,</span>
            <span class="n">chains</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
            <span class="n">nuts_sampler</span><span class="o">=</span><span class="sh">"</span><span class="s">numpyro</span><span class="sh">"</span><span class="p">,</span>
            <span class="n">random_seed</span><span class="o">=</span><span class="n">random_seed</span><span class="p">,</span>
            <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">α</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ_α</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ_γ</span><span class="sh">"</span><span class="p">],</span>
            <span class="n">progressbar</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
        <span class="p">)</span>
        <span class="n">elapsed_seconds</span> <span class="o">=</span> <span class="n">time</span><span class="p">.</span><span class="nf">perf_counter</span><span class="p">()</span> <span class="o">-</span> <span class="n">start</span>

    <span class="n">β_mean</span> <span class="o">=</span> <span class="n">idata</span><span class="p">.</span><span class="n">posterior</span><span class="p">[</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">].</span><span class="nf">mean</span><span class="p">(</span><span class="n">dim</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">chain</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">draw</span><span class="sh">"</span><span class="p">)).</span><span class="nf">to_numpy</span><span class="p">()</span>

    <span class="k">return</span> <span class="p">{</span>
        <span class="o">**</span><span class="n">info</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">elapsed_seconds</span><span class="sh">"</span><span class="p">:</span> <span class="n">elapsed_seconds</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">summary</span><span class="sh">"</span><span class="p">:</span> <span class="n">az</span><span class="p">.</span><span class="nf">summary</span><span class="p">(</span><span class="n">idata</span><span class="p">,</span> <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">α</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ_α</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ_γ</span><span class="sh">"</span><span class="p">]),</span>
        <span class="sh">"</span><span class="s">β_mean</span><span class="sh">"</span><span class="p">:</span> <span class="n">β_mean</span><span class="p">.</span><span class="nf">tolist</span><span class="p">(),</span>
    <span class="p">}</span>
</code></pre></div></div> <p>With all of this helper logic written, we can finally deploy to the cloud! We will start by running a short job to just profile the logp gradient function.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nf">print</span><span class="p">(</span><span class="sh">"</span><span class="s">Launching Modal GPU dlogp profile</span><span class="sh">"</span><span class="p">)</span>
<span class="k">with</span> <span class="n">app</span><span class="p">.</span><span class="nf">run</span><span class="p">():</span>
    <span class="n">profile_result</span> <span class="o">=</span> <span class="n">profile_dlogp</span><span class="p">.</span><span class="nf">remote</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">group_idx</span><span class="p">)</span>
<span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Finished running; profile results: </span><span class="si">{</span><span class="n">profile_result</span><span class="si">}</span><span class="sh">"</span><span class="p">)</span>
</code></pre></div></div> <div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Launching Modal GPU dlogp profile
Finished running; profile results: {'N': 1000000, 'G': 100000, 'K': 20, 'jax_backend': 'gpu', 'jax_gpu_devices': ['cuda:0'], 'gpu_check_sum': 262144.0, 'x_dtype': 'float32', 'value_var_dtypes': {'α': 'float32', 'β': 'float32', 'σ_α_log__': 'float32', 'σ_γ_log__': 'float32', 'z_α': 'float32', 'z_γ': 'float32'}, 'median_pytensor_dlogp_eval_seconds': 0.2271286229999987, 'median_jax_dlogp_eval_seconds': 0.0013876599999989025}
</code></pre></div></div> <p>Interesting - the gradient takes 0.001 seconds on GPU and 0.22 seconds on CPU. That is around a 200x speedup!</p> <p>Next, we run the Markov chain to completion and retrieve the results.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nf">print</span><span class="p">(</span><span class="sh">"</span><span class="s">Launching Modal GPU run</span><span class="sh">"</span><span class="p">)</span>
<span class="k">with</span> <span class="n">app</span><span class="p">.</span><span class="nf">run</span><span class="p">():</span>
    <span class="n">result</span> <span class="o">=</span> <span class="n">fit_hierarchical_logistic</span><span class="p">.</span><span class="nf">remote</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">group_idx</span><span class="p">)</span>

<span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">MCMC run finished in </span><span class="si">{</span><span class="n">result</span><span class="p">[</span><span class="sh">"</span><span class="s">elapsed_seconds</span><span class="sh">"</span><span class="p">]</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"> seconds</span><span class="sh">"</span><span class="p">)</span>
</code></pre></div></div> <div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>MCMC run finished in 272.16 seconds
</code></pre></div></div> <p>Modal helpfully lists <a href="https://modal.com/pricing">their prices per second</a>; an A100 runs for $0.000583 / second, meaning this run cost me around 15 cents from start to finish.</p> <h2 id="parameter-recovery">Parameter Recovery</h2> <p>We can see all the sampler diagnostics in the posterior summary. We’d need to run more chains to get the \(\hat{R}\) value for this model.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">pandas</span> <span class="k">as</span> <span class="n">pd</span>

<span class="n">posterior_summary</span> <span class="o">=</span> <span class="n">pd</span><span class="p">.</span><span class="nc">DataFrame</span><span class="p">(</span><span class="n">result</span><span class="p">[</span><span class="sh">"</span><span class="s">summary</span><span class="sh">"</span><span class="p">])</span>
<span class="n">posterior_summary</span><span class="p">.</span><span class="n">iloc</span><span class="p">[</span><span class="mi">0</span><span class="p">:</span><span class="mi">5</span><span class="p">]</span>
</code></pre></div></div> <div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>  parameter   mean     sd  eti89_lb  eti89_ub  ess_bulk  ess_tail  r_hat  \
0         α -0.980  0.015    -0.997    -0.952     1.510     5.445    NaN   
1      β[0] -0.275  0.005    -0.282    -0.266     2.042     5.659    NaN   
2      β[1]  0.630  0.010     0.612     0.641     1.627     5.374    NaN   
3      β[2] -0.358  0.007    -0.366    -0.345     2.390     5.593    NaN   
4      β[3]  0.359  0.006     0.348     0.366     1.834     5.607    NaN   

   mcse_mean  mcse_sd  
0      0.012    0.008  
1      0.004    0.003  
2      0.008    0.005  
3      0.005    0.003  
4      0.005    0.003
</code></pre></div></div> <p>A quick diagnostic is to compare the posterior mean of the fixed effects with the values used to simulate the data. With \(10^6\) observations, this model has more than enough data to estimate the fixed-effects.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>

<span class="n">blog_colors</span> <span class="o">=</span> <span class="p">{</span>
    <span class="sh">"</span><span class="s">background</span><span class="sh">"</span><span class="p">:</span> <span class="sh">"</span><span class="s">#1c1c1d</span><span class="sh">"</span><span class="p">,</span>
    <span class="sh">"</span><span class="s">text</span><span class="sh">"</span><span class="p">:</span> <span class="sh">"</span><span class="s">#e8e8e8</span><span class="sh">"</span><span class="p">,</span>
    <span class="sh">"</span><span class="s">accent</span><span class="sh">"</span><span class="p">:</span> <span class="sh">"</span><span class="s">#2698ba</span><span class="sh">"</span><span class="p">,</span>
    <span class="sh">"</span><span class="s">point</span><span class="sh">"</span><span class="p">:</span> <span class="sh">"</span><span class="s">#ffffff</span><span class="sh">"</span><span class="p">,</span>
<span class="p">}</span>
<span class="n">β_posterior_mean</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">result</span><span class="p">[</span><span class="sh">"</span><span class="s">β_mean</span><span class="sh">"</span><span class="p">])</span>
<span class="n">limits</span> <span class="o">=</span> <span class="p">[</span>
    <span class="nf">min</span><span class="p">(</span><span class="n">β_true</span><span class="p">.</span><span class="nf">min</span><span class="p">(),</span> <span class="n">β_posterior_mean</span><span class="p">.</span><span class="nf">min</span><span class="p">())</span> <span class="o">-</span> <span class="mf">0.05</span><span class="p">,</span>
    <span class="nf">max</span><span class="p">(</span><span class="n">β_true</span><span class="p">.</span><span class="nf">max</span><span class="p">(),</span> <span class="n">β_posterior_mean</span><span class="p">.</span><span class="nf">max</span><span class="p">())</span> <span class="o">+</span> <span class="mf">0.05</span><span class="p">,</span>
<span class="p">]</span>

<span class="n">fig</span><span class="p">,</span> <span class="n">ax</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="nf">subplots</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mf">4.8</span><span class="p">,</span> <span class="mf">3.6</span><span class="p">))</span>
<span class="n">fig</span><span class="p">.</span><span class="n">patch</span><span class="p">.</span><span class="nf">set_facecolor</span><span class="p">(</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">background</span><span class="sh">"</span><span class="p">])</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_facecolor</span><span class="p">(</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">background</span><span class="sh">"</span><span class="p">])</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">(</span><span class="n">β_true</span><span class="p">,</span> <span class="n">β_posterior_mean</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">36</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">point</span><span class="sh">"</span><span class="p">],</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.82</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">plot</span><span class="p">(</span><span class="n">limits</span><span class="p">,</span> <span class="n">limits</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">accent</span><span class="sh">"</span><span class="p">],</span> <span class="n">linewidth</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="nf">set_xlim</span><span class="p">(</span><span class="n">limits</span><span class="p">);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">set_ylim</span><span class="p">(</span><span class="n">limits</span><span class="p">);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">grid</span><span class="p">(</span><span class="bp">False</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_xlabel</span><span class="p">(</span><span class="sh">"</span><span class="s">True fixed effect</span><span class="sh">"</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">text</span><span class="sh">"</span><span class="p">]);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">set_ylabel</span><span class="p">(</span><span class="sh">"</span><span class="s">Posterior mean fixed effect</span><span class="sh">"</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">text</span><span class="sh">"</span><span class="p">])</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">tick_params</span><span class="p">(</span><span class="n">colors</span><span class="o">=</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">text</span><span class="sh">"</span><span class="p">])</span>
<span class="k">for</span> <span class="n">spine</span> <span class="ow">in</span> <span class="n">ax</span><span class="p">.</span><span class="n">spines</span><span class="p">.</span><span class="nf">values</span><span class="p">():</span>
    <span class="n">spine</span><span class="p">.</span><span class="nf">set_color</span><span class="p">(</span><span class="n">blog_colors</span><span class="p">[</span><span class="sh">"</span><span class="s">text</span><span class="sh">"</span><span class="p">])</span>
</code></pre></div></div> <p><img src="/images/2026-05-27-poverty-bayes-serverless-mcmc/cell-22-output-0.png" alt="Notebook output" width="620"/></p> <p>Nice! The true value and posterior mean estimates line up perfectly. A job well done for MCMC.</p> <p>I think that Bayes on GPU is tremendously undervalued. If this interests you or you have ideas to chat about, drop me a line.</p>]]></content><author><name></name></author><category term="tutorials"/><category term="statistics"/><category term="python"/><category term="pymc"/><category term="modal"/><category term="gpu"/><summary type="html"><![CDATA[Fitting a million-parameter Bayesian model on Modal with PyMC and NumPyro]]></summary></entry><entry><title type="html">Don’t know where your data is from? Bayesian modeling for unknown coordinates</title><link href="https://ckrapu.github.io/blog/2026/dont-know-where-your-data-is-from/" rel="alternate" type="text/html" title="Don’t know where your data is from? Bayesian modeling for unknown coordinates"/><published>2026-05-24T00:00:00+00:00</published><updated>2026-05-24T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2026/dont-know-where-your-data-is-from</id><content type="html" xml:base="https://ckrapu.github.io/blog/2026/dont-know-where-your-data-is-from/"><![CDATA[<p>An especially strong motivating case for the usage of spatial probability models comes from the mining industry. During exploration for mineral resources, prospectors will take geologic samples by drilling holes and examining the resulting material for presence or concentration of valuable ores. These data typically show strong spatial correlation, but constructing a fully-detailed geophysical model is at times infeasible as we are able to observe very little of the underground conditions, though the advent of remote sensing techniques like ground-penetrating radar and gravimetry has dramatically improved our ability to characterize Earth’s subsurface. To address this challenge, we would like to construct a probability model which uses nearby data to predict a variable of interest at a new location.</p> <p>To illustrate the problem better, we will use a dataset of uranium and vanadium point-referenced concentration measurements from Walker Lake. The data originate from Isaaks and Srivastava’s <a href="https://global.oup.com/academic/product/an-introduction-to-applied-geostatistics-9780195050134"><em>An Introduction to Applied Geostatistics</em></a> and are distributed with the R package <a href="https://r-spatial.github.io/gstat/"><code class="language-plaintext highlighter-rouge">gstat</code></a>.</p> <p>Up to this point, you may have seen lots of examples of how Gaussian process models are use in robotics, spatial statistics, neuroscience, etc. Now, we work through a more exotic example which modifies the Gaussian process model to accommodate the case in which the actual location of our data points is not known precisely, and is only observed with substantial measurement noise. Spatial location error changes the covariance and prediction problem itself, a point emphasized in geostatistical work on location error and GP regression work on noisy spatial inputs by <a href="https://doi.org/10.1214/ss/1081443228">Cressie and Kornak</a> and <a href="https://arxiv.org/abs/1506.08256">Cervone and Pillai</a>.</p> <p>This may seem like an unusual example at first glance, but we include it to give an example of how Bayesian modeling with appropriate priors lets us modify and change nearly any part of the model, given that we have some idea of how to represent our assumptions as part of the model process. Then, we use Monte Carlo methods to turn the inference crank and obtain reliable parameter estimates.</p> <p>Introducing some more notation, let \(\tilde{\mathbf{s}}_i\) denote the recorded coordinate and \(\mathbf{s}_i\) the latent coordinate where the measurement actually occurred. We use \(\mathbf{s}_i = \tilde{\mathbf{s}}_i + \Delta_i\), with \(\Delta_i \sim \operatorname{Normal}(\mathbf{0}, \sigma_s^2 I_2)\), and evaluate the Gaussian process at \(\mathbf{s}_i\) rather than at \(\tilde{\mathbf{s}}_i\). Our choice of coordinate system here is somewhat arbitrary; we could also choose to work with polar coordinates and place priors over the magnitude and the angle of the location error. The scale \(\sigma_s\) is treated as known in this example so that the model represents different assumed levels of coordinate error.</p> \[\begin{aligned} \Delta_i &amp;\sim \mathrm{Normal}(\mathbf{0}, \sigma_s^2\mathbf{I}_2) \\ \mathbf{s}_i &amp;= \tilde{\mathbf{s}}_i + \Delta_i \\ \mu &amp;\sim \mathrm{Normal}(2, 2) \\ \sigma &amp;\sim \mathrm{HalfNormal}(1) \\ \ell &amp;\sim \mathrm{HalfNormal}(100) \\ \sigma_0 &amp;\sim \mathrm{HalfNormal}(0.5) \\ f(\cdot) \mid \sigma,\ell &amp;\sim \mathcal{GP}(0, \sigma^2 c(\cdot,\cdot;\ell)) \\ Y_i \mid f,\mu,\sigma_0,\mathbf{s}_i &amp;\sim \mathrm{Normal}(\mu + f(\mathbf{s}_i), \sigma_0) \end{aligned}\] <p>This model is computationally harder than the fixed-location GP because the covariance matrix changes whenever the latent coordinates change. We use <code class="language-plaintext highlighter-rouge">pm.gp.Marginal</code> so that the latent GP values at the observations are integrated out.</p> <p>We will construct datasets with increasing noise and examine how the model’s parameter estimates change. To begin, we perturb the original coordinates with increasing noise:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="n">pathlib</span> <span class="kn">import</span> <span class="n">Path</span>

<span class="kn">import</span> <span class="n">arviz</span> <span class="k">as</span> <span class="n">az</span>
<span class="kn">import</span> <span class="n">matplotlib</span> <span class="k">as</span> <span class="n">mpl</span>
<span class="kn">import</span> <span class="n">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="n">pandas</span> <span class="k">as</span> <span class="n">pd</span>
<span class="kn">import</span> <span class="n">pymc</span> <span class="k">as</span> <span class="n">pm</span>
<span class="kn">import</span> <span class="n">seaborn</span> <span class="k">as</span> <span class="n">sns</span>
<span class="kn">from</span> <span class="n">matplotlib</span> <span class="kn">import</span> <span class="n">colors</span>
<span class="kn">from</span> <span class="n">matplotlib.patches</span> <span class="kn">import</span> <span class="n">Circle</span>

<span class="kn">from</span> <span class="n">stsp</span> <span class="kn">import</span> <span class="n">use_sepia_gunmetal</span>

<span class="n">RANDOM_SEED</span> <span class="o">=</span> <span class="mi">8927</span>
<span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="nf">seed</span><span class="p">(</span><span class="n">RANDOM_SEED</span><span class="p">)</span>

<span class="n">theme</span><span class="p">,</span> <span class="n">sepia_cmap</span> <span class="o">=</span> <span class="nf">use_sepia_gunmetal</span><span class="p">()</span>
</code></pre></div></div> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">plot_bg_color</span> <span class="o">=</span> <span class="sh">"</span><span class="s">#1c1c1d</span><span class="sh">"</span>
<span class="n">plot_text_color</span> <span class="o">=</span> <span class="sh">"</span><span class="s">#e8e8e8</span><span class="sh">"</span>

<span class="n">plt</span><span class="p">.</span><span class="n">style</span><span class="p">.</span><span class="nf">use</span><span class="p">(</span><span class="sh">"</span><span class="s">dark_background</span><span class="sh">"</span><span class="p">)</span>
<span class="n">mpl</span><span class="p">.</span><span class="n">rcParams</span><span class="p">.</span><span class="nf">update</span><span class="p">(</span>
    <span class="p">{</span>
        <span class="sh">"</span><span class="s">figure.facecolor</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_bg_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">axes.facecolor</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_bg_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">savefig.facecolor</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_bg_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">savefig.edgecolor</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_bg_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">text.color</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_text_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">axes.labelcolor</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_text_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">xtick.color</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_text_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">ytick.color</span><span class="sh">"</span><span class="p">:</span> <span class="n">plot_text_color</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">axes.edgecolor</span><span class="sh">"</span><span class="p">:</span> <span class="sh">"</span><span class="s">#828282</span><span class="sh">"</span><span class="p">,</span>
    <span class="p">}</span>
<span class="p">)</span>

<span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">paper</span><span class="sh">"</span><span class="p">]</span> <span class="o">=</span> <span class="n">plot_bg_color</span>
</code></pre></div></div> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">df</span> <span class="o">=</span> <span class="n">pd</span><span class="p">.</span><span class="nf">read_csv</span><span class="p">(</span><span class="nc">Path</span><span class="p">(</span><span class="sh">"</span><span class="s">../../data/walker/cleaned.csv</span><span class="sh">"</span><span class="p">))</span>

<span class="n">X_walker</span> <span class="o">=</span> <span class="n">df</span><span class="p">[[</span><span class="sh">"</span><span class="s">x_m</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">y_m</span><span class="sh">"</span><span class="p">]].</span><span class="nf">to_numpy</span><span class="p">()</span>
<span class="n">y_walker</span> <span class="o">=</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">u_log10_p1</span><span class="sh">"</span><span class="p">].</span><span class="nf">to_numpy</span><span class="p">()</span>

<span class="n">n_prediction_grid</span> <span class="o">=</span> <span class="mi">70</span>
<span class="n">buffer_fraction</span> <span class="o">=</span> <span class="mf">0.1</span>
<span class="n">x_range</span> <span class="o">=</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">x_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">max</span><span class="p">()</span> <span class="o">-</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">x_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">min</span><span class="p">()</span>
<span class="n">y_range</span> <span class="o">=</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">y_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">max</span><span class="p">()</span> <span class="o">-</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">y_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">min</span><span class="p">()</span>
<span class="n">x_limits</span> <span class="o">=</span> <span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">x_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">min</span><span class="p">()</span> <span class="o">-</span> <span class="n">buffer_fraction</span> <span class="o">*</span> <span class="n">x_range</span><span class="p">,</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">x_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">max</span><span class="p">()</span> <span class="o">+</span> <span class="n">buffer_fraction</span> <span class="o">*</span> <span class="n">x_range</span><span class="p">)</span>
<span class="n">y_limits</span> <span class="o">=</span> <span class="p">(</span><span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">y_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">min</span><span class="p">()</span> <span class="o">-</span> <span class="n">buffer_fraction</span> <span class="o">*</span> <span class="n">y_range</span><span class="p">,</span> <span class="n">df</span><span class="p">[</span><span class="sh">"</span><span class="s">y_m</span><span class="sh">"</span><span class="p">].</span><span class="nf">max</span><span class="p">()</span> <span class="o">+</span> <span class="n">buffer_fraction</span> <span class="o">*</span> <span class="n">y_range</span><span class="p">)</span>
<span class="n">x_new</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">linspace</span><span class="p">(</span><span class="o">*</span><span class="n">x_limits</span><span class="p">,</span> <span class="n">n_prediction_grid</span><span class="p">)</span>
<span class="n">y_new</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">linspace</span><span class="p">(</span><span class="o">*</span><span class="n">y_limits</span><span class="p">,</span> <span class="n">n_prediction_grid</span><span class="p">)</span>
<span class="n">x_new_mesh</span><span class="p">,</span> <span class="n">y_new_mesh</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">meshgrid</span><span class="p">(</span><span class="n">x_new</span><span class="p">,</span> <span class="n">y_new</span><span class="p">)</span>
<span class="n">Xnew</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">column_stack</span><span class="p">([</span><span class="n">x_new_mesh</span><span class="p">.</span><span class="nf">ravel</span><span class="p">(),</span> <span class="n">y_new_mesh</span><span class="p">.</span><span class="nf">ravel</span><span class="p">()])</span>

<span class="n">eps</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="nf">randn</span><span class="p">(</span><span class="o">*</span><span class="n">X_walker</span><span class="p">.</span><span class="n">shape</span><span class="p">)</span>
<span class="n">multipliers</span> <span class="o">=</span> <span class="p">[</span><span class="mf">12.0</span><span class="p">,</span> <span class="mf">25.0</span><span class="p">,</span> <span class="mf">40.0</span><span class="p">]</span>
<span class="n">noisy_xs</span> <span class="o">=</span> <span class="p">[</span><span class="n">X_walker</span> <span class="o">+</span> <span class="n">m</span> <span class="o">*</span> <span class="n">eps</span> <span class="k">for</span> <span class="n">m</span> <span class="ow">in</span> <span class="n">multipliers</span><span class="p">]</span>
</code></pre></div></div> <p>The perturbed coordinates are shown in Figure X. We keep the target variable unchanged in this analysis. Next, we construct the model, using the <code class="language-plaintext highlighter-rouge">pm.Data</code> object as a container to let us swap out coordinates easily.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">location_coords</span> <span class="o">=</span> <span class="p">{</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="nf">arange</span><span class="p">(</span><span class="n">X_walker</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]),</span> <span class="sh">"</span><span class="s">coord</span><span class="sh">"</span><span class="p">:</span> <span class="p">[</span><span class="sh">"</span><span class="s">x</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">y</span><span class="sh">"</span><span class="p">]}</span>
<span class="n">selected_location_idx</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">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="nf">default_rng</span><span class="p">(</span><span class="n">RANDOM_SEED</span><span class="p">).</span><span class="nf">choice</span><span class="p">(</span><span class="n">X_walker</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">size</span><span class="o">=</span><span class="mi">4</span><span class="p">,</span> <span class="n">replace</span><span class="o">=</span><span class="bp">False</span><span class="p">))</span>
<span class="n">location_error_idatas</span> <span class="o">=</span> <span class="p">{}</span>

<span class="k">with</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Model</span><span class="p">(</span><span class="n">coords</span><span class="o">=</span><span class="n">location_coords</span><span class="p">)</span> <span class="k">as</span> <span class="n">location_error_gp_model</span><span class="p">:</span>
    <span class="n">X_noisy</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Data</span><span class="p">(</span><span class="sh">"</span><span class="s">X_noisy</span><span class="sh">"</span><span class="p">,</span> <span class="n">noisy_xs</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">dims</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">coord</span><span class="sh">"</span><span class="p">))</span>
    <span class="n">y</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Data</span><span class="p">(</span><span class="sh">"</span><span class="s">y</span><span class="sh">"</span><span class="p">,</span> <span class="n">y_walker</span><span class="p">,</span> <span class="n">dims</span><span class="o">=</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">)</span>
    <span class="n">σ_s</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Data</span><span class="p">(</span><span class="sh">"</span><span class="s">σ_s</span><span class="sh">"</span><span class="p">,</span> <span class="n">multipliers</span><span class="p">[</span><span class="mi">0</span><span class="p">])</span>

    <span class="n">Δs</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">Δs</span><span class="sh">"</span><span class="p">,</span> <span class="mf">0.0</span><span class="p">,</span> <span class="n">σ_s</span><span class="p">,</span> <span class="n">dims</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">coord</span><span class="sh">"</span><span class="p">))</span>
    <span class="n">X_true</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Deterministic</span><span class="p">(</span><span class="sh">"</span><span class="s">X_true</span><span class="sh">"</span><span class="p">,</span> <span class="n">X_noisy</span> <span class="o">+</span> <span class="n">Δs</span><span class="p">,</span> <span class="n">dims</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">coord</span><span class="sh">"</span><span class="p">))</span>

    <span class="n">μ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">μ</span><span class="sh">"</span><span class="p">,</span> <span class="mf">2.0</span><span class="p">,</span> <span class="mf">2.0</span><span class="p">)</span>
    <span class="n">σ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">σ</span><span class="sh">"</span><span class="p">,</span> <span class="mf">1.0</span><span class="p">)</span>
    <span class="n">ℓ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">ℓ</span><span class="sh">"</span><span class="p">,</span> <span class="mf">100.0</span><span class="p">)</span>
    <span class="n">σ0</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">σ0</span><span class="sh">"</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">)</span>

    <span class="n">cov</span> <span class="o">=</span> <span class="n">σ</span><span class="o">**</span><span class="mi">2</span> <span class="o">*</span> <span class="n">pm</span><span class="p">.</span><span class="n">gp</span><span class="p">.</span><span class="n">cov</span><span class="p">.</span><span class="nc">Matern52</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span> <span class="n">ls</span><span class="o">=</span><span class="n">ℓ</span><span class="p">)</span>
    <span class="n">gp_location</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="n">gp</span><span class="p">.</span><span class="nc">Marginal</span><span class="p">(</span><span class="n">mean_func</span><span class="o">=</span><span class="n">pm</span><span class="p">.</span><span class="n">gp</span><span class="p">.</span><span class="n">mean</span><span class="p">.</span><span class="nc">Constant</span><span class="p">(</span><span class="n">μ</span><span class="p">),</span> <span class="n">cov_func</span><span class="o">=</span><span class="n">cov</span><span class="p">)</span>
    <span class="n">gp_location</span><span class="p">.</span><span class="nf">marginal_likelihood</span><span class="p">(</span><span class="sh">"</span><span class="s">y_obs</span><span class="sh">"</span><span class="p">,</span> <span class="n">X</span><span class="o">=</span><span class="n">X_true</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">sigma</span><span class="o">=</span><span class="n">σ0</span><span class="p">,</span> <span class="n">dims</span><span class="o">=</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">)</span>

    <span class="k">for</span> <span class="n">multiplier</span><span class="p">,</span> <span class="n">X_noisy_value</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">multipliers</span><span class="p">,</span> <span class="n">noisy_xs</span><span class="p">):</span>
        <span class="n">pm</span><span class="p">.</span><span class="nf">set_data</span><span class="p">({</span><span class="sh">"</span><span class="s">X_noisy</span><span class="sh">"</span><span class="p">:</span> <span class="n">X_noisy_value</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ_s</span><span class="sh">"</span><span class="p">:</span> <span class="n">multiplier</span><span class="p">})</span>
        <span class="n">location_error_idatas</span><span class="p">[</span><span class="n">multiplier</span><span class="p">]</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nf">sample</span><span class="p">(</span>
            <span class="n">chains</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span>
            <span class="n">cores</span><span class="o">=</span><span class="mi">2</span><span class="p">,</span>
            <span class="n">target_accept</span><span class="o">=</span><span class="mf">0.95</span><span class="p">,</span>
            <span class="n">mp_ctx</span><span class="o">=</span><span class="sh">"</span><span class="s">spawn</span><span class="sh">"</span><span class="p">,</span>
            <span class="n">random_seed</span><span class="o">=</span><span class="n">RANDOM_SEED</span><span class="p">,</span>
            <span class="n">progressbar</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span>
        <span class="p">)</span>

<span class="n">location_error_parameter_summaries</span> <span class="o">=</span> <span class="p">{</span>
    <span class="n">multiplier</span><span class="p">:</span> <span class="n">az</span><span class="p">.</span><span class="nf">summary</span><span class="p">(</span><span class="n">idata</span><span class="p">,</span> <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">μ</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">ℓ</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ0</span><span class="sh">"</span><span class="p">])</span>
    <span class="k">for</span> <span class="n">multiplier</span><span class="p">,</span> <span class="n">idata</span> <span class="ow">in</span> <span class="n">location_error_idatas</span><span class="p">.</span><span class="nf">items</span><span class="p">()</span>
<span class="p">}</span>
<span class="n">location_error_diagnostics</span> <span class="o">=</span> <span class="n">pd</span><span class="p">.</span><span class="nc">DataFrame</span><span class="p">(</span>
    <span class="p">[</span>
        <span class="p">{</span>
            <span class="sh">"</span><span class="s">Location error SD</span><span class="sh">"</span><span class="p">:</span> <span class="n">multiplier</span><span class="p">,</span>
            <span class="sh">"</span><span class="s">Divergences</span><span class="sh">"</span><span class="p">:</span> <span class="nf">int</span><span class="p">(</span><span class="n">idata</span><span class="p">.</span><span class="n">sample_stats</span><span class="p">[</span><span class="sh">"</span><span class="s">diverging</span><span class="sh">"</span><span class="p">].</span><span class="nf">sum</span><span class="p">()),</span>
            <span class="sh">"</span><span class="s">Max $</span><span class="se">\\</span><span class="s">hat{R}$</span><span class="sh">"</span><span class="p">:</span> <span class="n">location_error_parameter_summaries</span><span class="p">[</span><span class="n">multiplier</span><span class="p">][</span><span class="sh">"</span><span class="s">r_hat</span><span class="sh">"</span><span class="p">].</span><span class="nf">max</span><span class="p">(),</span>
            <span class="sh">"</span><span class="s">Min bulk ESS</span><span class="sh">"</span><span class="p">:</span> <span class="n">location_error_parameter_summaries</span><span class="p">[</span><span class="n">multiplier</span><span class="p">][</span><span class="sh">"</span><span class="s">ess_bulk</span><span class="sh">"</span><span class="p">].</span><span class="nf">min</span><span class="p">(),</span>
            <span class="sh">"</span><span class="s">Mean displacement</span><span class="sh">"</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="nf">sqrt</span><span class="p">((</span><span class="n">idata</span><span class="p">.</span><span class="n">posterior</span><span class="p">[</span><span class="sh">"</span><span class="s">Δs</span><span class="sh">"</span><span class="p">]</span> <span class="o">**</span> <span class="mi">2</span><span class="p">).</span><span class="nf">sum</span><span class="p">(</span><span class="sh">"</span><span class="s">coord</span><span class="sh">"</span><span class="p">)).</span><span class="nf">mean</span><span class="p">().</span><span class="nf">item</span><span class="p">(),</span>
        <span class="p">}</span>
        <span class="k">for</span> <span class="n">multiplier</span><span class="p">,</span> <span class="n">idata</span> <span class="ow">in</span> <span class="n">location_error_idatas</span><span class="p">.</span><span class="nf">items</span><span class="p">()</span>
    <span class="p">]</span>
<span class="p">)</span>
<span class="n">location_error_diagnostics</span>
</code></pre></div></div> <div class="language-text highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Initializing NUTS using jitter+adapt_diag...
Multiprocess sampling (2 chains in 2 jobs)
NUTS: [Δs, μ, σ, ℓ, σ0]
Sampling 2 chains for 1_000 tune and 1_000 draw iterations (2_000 + 2_000 draws total) took 6340 seconds.
There were 24 divergences after tuning. Increase `target_accept` or reparameterize.
Chain 1 reached the maximum tree depth. Increase `max_treedepth`, increase `target_accept` or reparameterize.
We recommend running at least 4 chains for robust computation of convergence diagnostics
The rhat statistic is larger than 1.01 for some parameters. This indicates problems during sampling. See https://arxiv.org/abs/1903.08008 for details
The effective sample size per chain is smaller than 100 for some parameters.  A higher number is needed for reliable rhat and ess computation. See https://arxiv.org/abs/1903.08008 for details
Initializing NUTS using jitter+adapt_diag...
Multiprocess sampling (2 chains in 2 jobs)
NUTS: [Δs, μ, σ, ℓ, σ0]
Sampling 2 chains for 1_000 tune and 1_000 draw iterations (2_000 + 2_000 draws total) took 6376 seconds.
There were 79 divergences after tuning. Increase `target_accept` or reparameterize.
Chain 0 reached the maximum tree depth. Increase `max_treedepth`, increase `target_accept` or reparameterize.
We recommend running at least 4 chains for robust computation of convergence diagnostics
The rhat statistic is larger than 1.01 for some parameters. This indicates problems during sampling. See https://arxiv.org/abs/1903.08008 for details
The effective sample size per chain is smaller than 100 for some parameters.  A higher number is needed for reliable rhat and ess computation. See https://arxiv.org/abs/1903.08008 for details
Initializing NUTS using jitter+adapt_diag...
Multiprocess sampling (2 chains in 2 jobs)
NUTS: [Δs, μ, σ, ℓ, σ0]
Sampling 2 chains for 1_000 tune and 1_000 draw iterations (2_000 + 2_000 draws total) took 7756 seconds.
There were 52 divergences after tuning. Increase `target_accept` or reparameterize.
Chain 0 reached the maximum tree depth. Increase `max_treedepth`, increase `target_accept` or reparameterize.
Chain 1 reached the maximum tree depth. Increase `max_treedepth`, increase `target_accept` or reparameterize.
We recommend running at least 4 chains for robust computation of convergence diagnostics
The rhat statistic is larger than 1.01 for some parameters. This indicates problems during sampling. See https://arxiv.org/abs/1903.08008 for details
The effective sample size per chain is smaller than 100 for some parameters.  A higher number is needed for reliable rhat and ess computation. See https://arxiv.org/abs/1903.08008 for details
</code></pre></div></div> <table> <thead> <tr> <th style="text-align: right"> </th> <th style="text-align: right">Location error SD</th> <th style="text-align: right">Divergences</th> <th style="text-align: right">Max \(\hat{R}\)</th> <th style="text-align: right">Min bulk ESS</th> <th style="text-align: right">Mean displacement</th> </tr> </thead> <tbody> <tr> <td style="text-align: right">0</td> <td style="text-align: right">12.0</td> <td style="text-align: right">24</td> <td style="text-align: right">1.00</td> <td style="text-align: right">268.0</td> <td style="text-align: right">15.280551</td> </tr> <tr> <td style="text-align: right">1</td> <td style="text-align: right">25.0</td> <td style="text-align: right">79</td> <td style="text-align: right">1.03</td> <td style="text-align: right">126.0</td> <td style="text-align: right">31.874539</td> </tr> <tr> <td style="text-align: right">2</td> <td style="text-align: right">40.0</td> <td style="text-align: right">52</td> <td style="text-align: right">1.02</td> <td style="text-align: right">44.0</td> <td style="text-align: right">51.138540</td> </tr> </tbody> </table> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">location_surface_grids</span> <span class="o">=</span> <span class="p">{}</span>
<span class="k">with</span> <span class="n">location_error_gp_model</span><span class="p">:</span>
    <span class="k">for</span> <span class="n">multiplier</span><span class="p">,</span> <span class="n">X_noisy_value</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">multipliers</span><span class="p">,</span> <span class="n">noisy_xs</span><span class="p">):</span>
        <span class="n">pm</span><span class="p">.</span><span class="nf">set_data</span><span class="p">({</span><span class="sh">"</span><span class="s">X_noisy</span><span class="sh">"</span><span class="p">:</span> <span class="n">X_noisy_value</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ_s</span><span class="sh">"</span><span class="p">:</span> <span class="n">multiplier</span><span class="p">})</span>
        <span class="n">posterior_mean_point</span> <span class="o">=</span> <span class="p">{</span>
            <span class="n">name</span><span class="p">:</span> <span class="n">location_error_idatas</span><span class="p">[</span><span class="n">multiplier</span><span class="p">].</span><span class="n">posterior</span><span class="p">[</span><span class="n">name</span><span class="p">].</span><span class="nf">mean</span><span class="p">((</span><span class="sh">"</span><span class="s">chain</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">draw</span><span class="sh">"</span><span class="p">)).</span><span class="n">values</span>
            <span class="k">for</span> <span class="n">name</span> <span class="ow">in</span> <span class="p">[</span><span class="sh">"</span><span class="s">μ</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">ℓ</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">σ0</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">Δs</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">X_true</span><span class="sh">"</span><span class="p">]</span>
        <span class="p">}</span>
        <span class="n">f_mean</span><span class="p">,</span> <span class="n">_</span> <span class="o">=</span> <span class="n">gp_location</span><span class="p">.</span><span class="nf">predict</span><span class="p">(</span><span class="n">Xnew</span><span class="p">,</span> <span class="n">point</span><span class="o">=</span><span class="n">posterior_mean_point</span><span class="p">,</span> <span class="n">diag</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">pred_noise</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="n">location_surface_grids</span><span class="p">[</span><span class="n">multiplier</span><span class="p">]</span> <span class="o">=</span> <span class="n">f_mean</span><span class="p">.</span><span class="nf">reshape</span><span class="p">(</span><span class="n">n_prediction_grid</span><span class="p">,</span> <span class="n">n_prediction_grid</span><span class="p">)</span>

<span class="n">naive_kde_grids</span> <span class="o">=</span> <span class="p">{}</span>
<span class="n">kde_bandwidths</span> <span class="o">=</span> <span class="p">{</span>
    <span class="n">multiplier</span><span class="p">:</span> <span class="n">location_error_idatas</span><span class="p">[</span><span class="n">multiplier</span><span class="p">].</span><span class="n">posterior</span><span class="p">[</span><span class="sh">"</span><span class="s">ℓ</span><span class="sh">"</span><span class="p">].</span><span class="nf">mean</span><span class="p">((</span><span class="sh">"</span><span class="s">chain</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">draw</span><span class="sh">"</span><span class="p">)).</span><span class="nf">item</span><span class="p">()</span>
    <span class="k">for</span> <span class="n">multiplier</span> <span class="ow">in</span> <span class="n">multipliers</span>
<span class="p">}</span>
<span class="k">for</span> <span class="n">multiplier</span><span class="p">,</span> <span class="n">X_noisy_value</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">multipliers</span><span class="p">,</span> <span class="n">noisy_xs</span><span class="p">):</span>
    <span class="n">squared_distance</span> <span class="o">=</span> <span class="p">(</span>
        <span class="p">(</span><span class="n">Xnew</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">,</span> <span class="bp">None</span><span class="p">]</span> <span class="o">-</span> <span class="n">X_noisy_value</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">])</span> <span class="o">**</span> <span class="mi">2</span>
        <span class="o">+</span> <span class="p">(</span><span class="n">Xnew</span><span class="p">[:,</span> <span class="mi">1</span><span class="p">,</span> <span class="bp">None</span><span class="p">]</span> <span class="o">-</span> <span class="n">X_noisy_value</span><span class="p">[:,</span> <span class="mi">1</span><span class="p">])</span> <span class="o">**</span> <span class="mi">2</span>
    <span class="p">)</span>
    <span class="n">weights</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">exp</span><span class="p">(</span><span class="o">-</span><span class="mf">0.5</span> <span class="o">*</span> <span class="n">squared_distance</span> <span class="o">/</span> <span class="n">kde_bandwidths</span><span class="p">[</span><span class="n">multiplier</span><span class="p">]</span><span class="o">**</span><span class="mi">2</span><span class="p">)</span>
    <span class="n">naive_kde_grids</span><span class="p">[</span><span class="n">multiplier</span><span class="p">]</span> <span class="o">=</span> <span class="p">(</span><span class="n">weights</span> <span class="o">@</span> <span class="n">y_walker</span> <span class="o">/</span> <span class="n">weights</span><span class="p">.</span><span class="nf">sum</span><span class="p">(</span><span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)).</span><span class="nf">reshape</span><span class="p">(</span><span class="n">n_prediction_grid</span><span class="p">,</span> <span class="n">n_prediction_grid</span><span class="p">)</span>

<span class="n">location_norm</span> <span class="o">=</span> <span class="n">colors</span><span class="p">.</span><span class="nc">Normalize</span><span class="p">(</span>
    <span class="n">vmin</span><span class="o">=</span><span class="nf">min</span><span class="p">(</span><span class="n">y_walker</span><span class="p">.</span><span class="nf">min</span><span class="p">(),</span> <span class="o">*</span><span class="p">(</span><span class="n">grid</span><span class="p">.</span><span class="nf">min</span><span class="p">()</span> <span class="k">for</span> <span class="n">grid</span> <span class="ow">in</span> <span class="n">location_surface_grids</span><span class="p">.</span><span class="nf">values</span><span class="p">())),</span>
    <span class="n">vmax</span><span class="o">=</span><span class="nf">max</span><span class="p">(</span><span class="n">y_walker</span><span class="p">.</span><span class="nf">max</span><span class="p">(),</span> <span class="o">*</span><span class="p">(</span><span class="n">grid</span><span class="p">.</span><span class="nf">max</span><span class="p">()</span> <span class="k">for</span> <span class="n">grid</span> <span class="ow">in</span> <span class="n">location_surface_grids</span><span class="p">.</span><span class="nf">values</span><span class="p">())),</span>
<span class="p">)</span>
<span class="n">point_colors</span> <span class="o">=</span> <span class="p">[</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">gunmetal</span><span class="sh">"</span><span class="p">],</span> <span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">sepia</span><span class="sh">"</span><span class="p">],</span> <span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">rust</span><span class="sh">"</span><span class="p">],</span> <span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">steel</span><span class="sh">"</span><span class="p">]]</span>
<span class="n">arrow_head_width</span> <span class="o">=</span> <span class="mf">0.018</span> <span class="o">*</span> <span class="nf">max</span><span class="p">(</span><span class="n">x_range</span><span class="p">,</span> <span class="n">y_range</span><span class="p">)</span>

<span class="n">fig</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="nf">figure</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">5</span><span class="p">,</span> <span class="mf">11.2</span> <span class="o">*</span> <span class="p">(</span><span class="mi">5</span> <span class="o">/</span> <span class="mi">7</span><span class="p">)))</span>
<span class="n">gs</span> <span class="o">=</span> <span class="n">fig</span><span class="p">.</span><span class="nf">add_gridspec</span><span class="p">(</span><span class="mi">4</span><span class="p">,</span> <span class="nf">len</span><span class="p">(</span><span class="n">multipliers</span><span class="p">),</span> <span class="n">hspace</span><span class="o">=</span><span class="mf">0.08</span><span class="p">,</span> <span class="n">wspace</span><span class="o">=</span><span class="mf">0.08</span><span class="p">)</span>
<span class="n">axes</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">array</span><span class="p">([[</span><span class="n">fig</span><span class="p">.</span><span class="nf">add_subplot</span><span class="p">(</span><span class="n">gs</span><span class="p">[</span><span class="n">row</span><span class="p">,</span> <span class="n">col</span><span class="p">])</span> <span class="k">for</span> <span class="n">col</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="nf">len</span><span class="p">(</span><span class="n">multipliers</span><span class="p">))]</span> <span class="k">for</span> <span class="n">row</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="mi">4</span><span class="p">)])</span>
<span class="n">panel_idx</span> <span class="o">=</span> <span class="mi">0</span>
<span class="n">surface_meshes</span> <span class="o">=</span> <span class="p">[]</span>

<span class="k">for</span> <span class="n">col</span><span class="p">,</span> <span class="p">(</span><span class="n">multiplier</span><span class="p">,</span> <span class="n">X_noisy_value</span><span class="p">)</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">(</span><span class="nf">zip</span><span class="p">(</span><span class="n">multipliers</span><span class="p">,</span> <span class="n">noisy_xs</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="mi">0</span><span class="p">,</span> <span class="n">col</span><span class="p">]</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">(</span><span class="n">X_walker</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_walker</span><span class="p">[:,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">facecolor</span><span class="o">=</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">paper</span><span class="sh">"</span><span class="p">],</span> <span class="n">edgecolor</span><span class="o">=</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">gunmetal</span><span class="sh">"</span><span class="p">],</span> <span class="n">s</span><span class="o">=</span><span class="mi">14</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">0.35</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="nf">scatter</span><span class="p">(</span><span class="n">X_noisy_value</span><span class="p">[:,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_noisy_value</span><span class="p">[:,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">c</span><span class="o">=</span><span class="n">y_walker</span><span class="p">,</span> <span class="n">cmap</span><span class="o">=</span><span class="n">sepia_cmap</span><span class="p">,</span> <span class="n">norm</span><span class="o">=</span><span class="n">location_norm</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">16</span><span class="p">,</span> <span class="n">edgecolor</span><span class="o">=</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">paper</span><span class="sh">"</span><span class="p">],</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">0.25</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="k">for</span> <span class="n">point_idx</span><span class="p">,</span> <span class="n">point_color</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">selected_location_idx</span><span class="p">,</span> <span class="n">point_colors</span><span class="p">):</span>
        <span class="n">circle</span> <span class="o">=</span> <span class="nc">Circle</span><span class="p">(</span><span class="n">X_noisy_value</span><span class="p">[</span><span class="n">point_idx</span><span class="p">],</span> <span class="n">radius</span><span class="o">=</span><span class="n">multiplier</span><span class="p">,</span> <span class="n">facecolor</span><span class="o">=</span><span class="n">colors</span><span class="p">.</span><span class="nf">to_rgba</span><span class="p">(</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">steel</span><span class="sh">"</span><span class="p">],</span> <span class="mf">0.2</span><span class="p">),</span> <span class="n">edgecolor</span><span class="o">=</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">gunmetal</span><span class="sh">"</span><span class="p">],</span> <span class="n">linewidth</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="nf">add_patch</span><span class="p">(</span><span class="n">circle</span><span class="p">)</span>
        <span class="n">dx</span><span class="p">,</span> <span class="n">dy</span> <span class="o">=</span> <span class="n">X_noisy_value</span><span class="p">[</span><span class="n">point_idx</span><span class="p">]</span> <span class="o">-</span> <span class="n">X_walker</span><span class="p">[</span><span class="n">point_idx</span><span class="p">]</span>
        <span class="n">ax</span><span class="p">.</span><span class="nf">arrow</span><span class="p">(</span><span class="n">X_walker</span><span class="p">[</span><span class="n">point_idx</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_walker</span><span class="p">[</span><span class="n">point_idx</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">dx</span><span class="p">,</span> <span class="n">dy</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">plot_text_color</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">0.8</span><span class="p">,</span> <span class="n">length_includes_head</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">head_width</span><span class="o">=</span><span class="n">arrow_head_width</span><span class="p">,</span> <span class="n">head_length</span><span class="o">=</span><span class="n">arrow_head_width</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">(</span><span class="n">X_walker</span><span class="p">[</span><span class="n">selected_location_idx</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_walker</span><span class="p">[</span><span class="n">selected_location_idx</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">facecolor</span><span class="o">=</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">paper</span><span class="sh">"</span><span class="p">],</span> <span class="n">edgecolor</span><span class="o">=</span><span class="n">plot_text_color</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">44</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">0.8</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">(</span><span class="n">X_noisy_value</span><span class="p">[</span><span class="n">selected_location_idx</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_noisy_value</span><span class="p">[</span><span class="n">selected_location_idx</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">c</span><span class="o">=</span><span class="n">y_walker</span><span class="p">[</span><span class="n">selected_location_idx</span><span class="p">],</span> <span class="n">cmap</span><span class="o">=</span><span class="n">sepia_cmap</span><span class="p">,</span> <span class="n">norm</span><span class="o">=</span><span class="n">location_norm</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">44</span><span class="p">,</span> <span class="n">edgecolor</span><span class="o">=</span><span class="n">plot_text_color</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">0.8</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="n">xaxis</span><span class="p">.</span><span class="nf">set_label_position</span><span class="p">(</span><span class="sh">"</span><span class="s">top</span><span class="sh">"</span><span class="p">);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">set_xlabel</span><span class="p">(</span><span class="sa">rf</span><span class="sh">"</span><span class="s">$\sigma_s = </span><span class="si">{</span><span class="n">multiplier</span><span class="si">:</span><span class="p">.</span><span class="mi">0</span><span class="n">f</span><span class="si">}</span><span class="s">$ m</span><span class="sh">"</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="mi">2</span><span class="p">,</span> <span class="n">col</span><span class="p">]</span>
    <span class="n">mesh</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">pcolormesh</span><span class="p">(</span><span class="n">x_new_mesh</span><span class="p">,</span> <span class="n">y_new_mesh</span><span class="p">,</span> <span class="n">location_surface_grids</span><span class="p">[</span><span class="n">multiplier</span><span class="p">],</span> <span class="n">cmap</span><span class="o">=</span><span class="n">sepia_cmap</span><span class="p">,</span> <span class="n">norm</span><span class="o">=</span><span class="n">location_norm</span><span class="p">,</span> <span class="n">shading</span><span class="o">=</span><span class="sh">"</span><span class="s">auto</span><span class="sh">"</span><span class="p">)</span>
    <span class="n">surface_meshes</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">mesh</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="mi">3</span><span class="p">,</span> <span class="n">col</span><span class="p">]</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">pcolormesh</span><span class="p">(</span><span class="n">x_new_mesh</span><span class="p">,</span> <span class="n">y_new_mesh</span><span class="p">,</span> <span class="n">naive_kde_grids</span><span class="p">[</span><span class="n">multiplier</span><span class="p">],</span> <span class="n">cmap</span><span class="o">=</span><span class="n">sepia_cmap</span><span class="p">,</span> <span class="n">norm</span><span class="o">=</span><span class="n">location_norm</span><span class="p">,</span> <span class="n">shading</span><span class="o">=</span><span class="sh">"</span><span class="s">auto</span><span class="sh">"</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="mi">1</span><span class="p">,</span> <span class="n">col</span><span class="p">]</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">pcolormesh</span><span class="p">(</span><span class="n">x_new_mesh</span><span class="p">,</span> <span class="n">y_new_mesh</span><span class="p">,</span> <span class="n">location_surface_grids</span><span class="p">[</span><span class="n">multiplier</span><span class="p">],</span> <span class="n">cmap</span><span class="o">=</span><span class="n">sepia_cmap</span><span class="p">,</span> <span class="n">norm</span><span class="o">=</span><span class="n">location_norm</span><span class="p">,</span> <span class="n">shading</span><span class="o">=</span><span class="sh">"</span><span class="s">auto</span><span class="sh">"</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.25</span><span class="p">)</span>
    <span class="n">X_true_samples</span> <span class="o">=</span> <span class="n">location_error_idatas</span><span class="p">[</span><span class="n">multiplier</span><span class="p">].</span><span class="n">posterior</span><span class="p">[</span><span class="sh">"</span><span class="s">X_true</span><span class="sh">"</span><span class="p">].</span><span class="nf">stack</span><span class="p">(</span><span class="n">sample</span><span class="o">=</span><span class="p">(</span><span class="sh">"</span><span class="s">chain</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">draw</span><span class="sh">"</span><span class="p">)).</span><span class="nf">transpose</span><span class="p">(</span><span class="sh">"</span><span class="s">obs</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">coord</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">sample</span><span class="sh">"</span><span class="p">)</span>
    <span class="k">for</span> <span class="n">point_idx</span><span class="p">,</span> <span class="n">point_color</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">selected_location_idx</span><span class="p">,</span> <span class="n">point_colors</span><span class="p">):</span>
        <span class="n">x_draws</span> <span class="o">=</span> <span class="n">X_true_samples</span><span class="p">.</span><span class="nf">isel</span><span class="p">(</span><span class="n">obs</span><span class="o">=</span><span class="n">point_idx</span><span class="p">).</span><span class="nf">sel</span><span class="p">(</span><span class="n">coord</span><span class="o">=</span><span class="sh">"</span><span class="s">x</span><span class="sh">"</span><span class="p">).</span><span class="n">values</span>
        <span class="n">y_draws</span> <span class="o">=</span> <span class="n">X_true_samples</span><span class="p">.</span><span class="nf">isel</span><span class="p">(</span><span class="n">obs</span><span class="o">=</span><span class="n">point_idx</span><span class="p">).</span><span class="nf">sel</span><span class="p">(</span><span class="n">coord</span><span class="o">=</span><span class="sh">"</span><span class="s">y</span><span class="sh">"</span><span class="p">).</span><span class="n">values</span>
        <span class="n">sns</span><span class="p">.</span><span class="nf">kdeplot</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">x_draws</span><span class="p">,</span> <span class="n">y</span><span class="o">=</span><span class="n">y_draws</span><span class="p">,</span> <span class="n">levels</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">point_color</span><span class="p">,</span> <span class="n">linewidths</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> <span class="n">fill</span><span class="o">=</span><span class="bp">False</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">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">(</span><span class="n">X_walker</span><span class="p">[</span><span class="n">point_idx</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_walker</span><span class="p">[</span><span class="n">point_idx</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">marker</span><span class="o">=</span><span class="sh">"</span><span class="s">x</span><span class="sh">"</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="n">point_color</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">34</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">1.1</span><span class="p">)</span>
        <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">(</span><span class="n">X_noisy_value</span><span class="p">[</span><span class="n">point_idx</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X_noisy_value</span><span class="p">[</span><span class="n">point_idx</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">marker</span><span class="o">=</span><span class="sh">"</span><span class="s">o</span><span class="sh">"</span><span class="p">,</span> <span class="n">facecolor</span><span class="o">=</span><span class="n">theme</span><span class="p">[</span><span class="sh">"</span><span class="s">paper</span><span class="sh">"</span><span class="p">],</span> <span class="n">edgecolor</span><span class="o">=</span><span class="n">point_color</span><span class="p">,</span> <span class="n">s</span><span class="o">=</span><span class="mi">34</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">1.1</span><span class="p">)</span>

<span class="k">for</span> <span class="n">row</span><span class="p">,</span> <span class="n">row_label</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">([</span><span class="sh">"</span><span class="s">Perturbed coordinates</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">Posterior location density</span><span class="sh">"</span><span class="p">,</span> <span class="sa">r</span><span class="sh">"</span><span class="s">Posterior mean of $f(s)$</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">Naive smoothed KDE</span><span class="sh">"</span><span class="p">]):</span>
    <span class="n">axes</span><span class="p">[</span><span class="n">row</span><span class="p">,</span> <span class="mi">0</span><span class="p">].</span><span class="nf">set_ylabel</span><span class="p">(</span><span class="n">row_label</span><span class="p">)</span>

<span class="k">for</span> <span class="n">ax</span> <span class="ow">in</span> <span class="n">axes</span><span class="p">.</span><span class="nf">ravel</span><span class="p">():</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">set_xlim</span><span class="p">(</span><span class="o">*</span><span class="n">x_limits</span><span class="p">);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">set_ylim</span><span class="p">(</span><span class="o">*</span><span class="n">y_limits</span><span class="p">);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">set_aspect</span><span class="p">(</span><span class="sh">"</span><span class="s">equal</span><span class="sh">"</span><span class="p">);</span> <span class="n">ax</span><span class="p">.</span><span class="nf">grid</span><span class="p">(</span><span class="bp">False</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">tick_params</span><span class="p">(</span><span class="n">axis</span><span class="o">=</span><span class="sh">"</span><span class="s">both</span><span class="sh">"</span><span class="p">,</span> <span class="n">which</span><span class="o">=</span><span class="sh">"</span><span class="s">both</span><span class="sh">"</span><span class="p">,</span> <span class="n">bottom</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">left</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">top</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">right</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">labelbottom</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">labelleft</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">labeltop</span><span class="o">=</span><span class="bp">False</span><span class="p">,</span> <span class="n">labelright</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
    <span class="n">panel_idx</span> <span class="o">+=</span> <span class="mi">1</span>

<span class="n">cbar</span> <span class="o">=</span> <span class="n">fig</span><span class="p">.</span><span class="nf">colorbar</span><span class="p">(</span><span class="n">surface_meshes</span><span class="p">[</span><span class="o">-</span><span class="mi">1</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="nf">ravel</span><span class="p">().</span><span class="nf">tolist</span><span class="p">(),</span> <span class="n">orientation</span><span class="o">=</span><span class="sh">"</span><span class="s">horizontal</span><span class="sh">"</span><span class="p">,</span> <span class="n">fraction</span><span class="o">=</span><span class="mf">0.035</span><span class="p">,</span> <span class="n">pad</span><span class="o">=</span><span class="mf">0.04</span><span class="p">)</span>
<span class="n">cbar</span><span class="p">.</span><span class="nf">set_label</span><span class="p">(</span><span class="sa">r</span><span class="sh">"</span><span class="s">Uranium, $\log_{10}(x + 1)$</span><span class="sh">"</span><span class="p">)</span>

<span class="n">figure_dir</span> <span class="o">=</span> <span class="n">Path</span><span class="p">.</span><span class="nf">home</span><span class="p">()</span> <span class="o">/</span> <span class="sh">"</span><span class="s">ckrapu.github.io/images/2026-05-24-dont-know-where-your-data-is-from</span><span class="sh">"</span>
<span class="n">figure_dir</span><span class="p">.</span><span class="nf">mkdir</span><span class="p">(</span><span class="n">parents</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">exist_ok</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">fig</span><span class="p">.</span><span class="nf">savefig</span><span class="p">(</span><span class="n">figure_dir</span> <span class="o">/</span> <span class="sh">"</span><span class="s">error-in-location-grid.png</span><span class="sh">"</span><span class="p">,</span> <span class="n">dpi</span><span class="o">=</span><span class="mi">220</span><span class="p">,</span> <span class="n">bbox_inches</span><span class="o">=</span><span class="sh">"</span><span class="s">tight</span><span class="sh">"</span><span class="p">,</span> <span class="n">facecolor</span><span class="o">=</span><span class="n">fig</span><span class="p">.</span><span class="nf">get_facecolor</span><span class="p">())</span>
</code></pre></div></div> <div style="text-align: center;"> <img src="/images/2026-05-24-dont-know-where-your-data-is-from/error-in-location-grid.png" alt="Error-in-location model results" width="720"/> </div> <p><em>Figure 6 Data and posterior inference for a GP model with point coordinates which are only approximately observed. The top row shows the true, original coordinates in hollow circles while the noisily observed data points are shaded according to the uranium concentration at that point. The larger circles indicate the spatial scale of the coordinate error. The second row shows contours for the posterior high density region of the true location of four selected points in space; we see that as \(\sigma_s\) increases, these regions gracefully grow larger and larger.</em></p> <p>As we examine Figure 6, we can see that some of the main features of the underlying surface are preserved even as the uncertainty in the coordinates grows. We see a few lighter features in the lower left corner, and a darker region in the top left and lower right. There is variation across these, but it is impressive that we can even work at all with this data given the severity of the perturbations! We include a comparison using a simpler approach, the Nadaraya-Watson Gaussian kernel smoother, with its bandwidth set to match the posterior mean estimate of the length scale under each scenario, and we find that it is unable to provide more than a rough average with little capacity to represent spatial variation.</p>]]></content><author><name></name></author><category term="tutorials"/><category term="statistics"/><category term="python"/><category term="pymc"/><summary type="html"><![CDATA[A PyMC Gaussian process example with uncertain spatial coordinates]]></summary></entry><entry><title type="html">Rolling your own serverless OCR in 40 lines of code</title><link href="https://ckrapu.github.io/blog/2026/ocr-textbooks-modal-deepseek/" rel="alternate" type="text/html" title="Rolling your own serverless OCR in 40 lines of code"/><published>2026-01-14T10:00:00+00:00</published><updated>2026-01-14T10:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2026/ocr-textbooks-modal-deepseek</id><content type="html" xml:base="https://ckrapu.github.io/blog/2026/ocr-textbooks-modal-deepseek/"><![CDATA[<p>A few months ago, I wanted to make my copy of Gelman’s <em>Bayesian Data Analysis</em> searchable for use in a statistics-focused agent.</p> <p>There are some pretty sophisticated OCR tools out there but they tend to have usage limits or get expensive when you’re processing thousands of pages. DeepSeek recently released an <a href="https://arxiv.org/abs/2510.18234">open OCR model</a> that handles mathematical notation well, and I figured I could run it myself if I had access to a GPU. Sadly, my daily driver is a decade-old Titan Xp which no longer supports the latest PyTorch versions and thus can’t run DeepSeek OCR.</p> <p>I ended up using Modal for this.</p> <h3 id="what-is-modal">What is Modal?</h3> <p>Modal is a serverless compute platform that lets you run Python code on cloud infrastructure without managing servers. The killer feature for machine learning work is that you can define a container image, attach a GPU, and pay only for the seconds your code is actually running.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">modal</span>

<span class="n">image</span> <span class="o">=</span> <span class="n">modal</span><span class="p">.</span><span class="n">Image</span><span class="p">.</span><span class="nf">from_registry</span><span class="p">(</span>
    <span class="sh">"</span><span class="s">nvidia/cuda:11.8.0-cudnn8-runtime-ubuntu22.04</span><span class="sh">"</span><span class="p">,</span>
    <span class="n">add_python</span><span class="o">=</span><span class="sh">"</span><span class="s">3.11</span><span class="sh">"</span><span class="p">,</span>
<span class="p">).</span><span class="nf">pip_install</span><span class="p">(</span><span class="sh">"</span><span class="s">torch</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">transformers</span><span class="sh">"</span><span class="p">,</span> <span class="p">...)</span>

<span class="n">app</span> <span class="o">=</span> <span class="n">modal</span><span class="p">.</span><span class="nc">App</span><span class="p">(</span><span class="sh">"</span><span class="s">my-gpu-app</span><span class="sh">"</span><span class="p">)</span>

<span class="nd">@app.function</span><span class="p">(</span><span class="n">image</span><span class="o">=</span><span class="n">image</span><span class="p">,</span> <span class="n">gpu</span><span class="o">=</span><span class="sh">"</span><span class="s">A100</span><span class="sh">"</span><span class="p">)</span>
<span class="k">def</span> <span class="nf">process_something</span><span class="p">():</span>
    <span class="c1"># This runs on an A100 with all your deps installed
</span>    <span class="k">pass</span>
</code></pre></div></div> <p>The decorator pattern is what makes Modal pleasant to use. You write normal Python, sprinkle decorators on the functions that need special hardware, and Modal handles the rest: building the container, provisioning the GPU, routing your requests. For OCR, this is perfect.</p> <h3 id="the-ocr-script">The OCR script</h3> <p>The core idea is simple: deploy a FastAPI server on Modal that accepts images and returns markdown text. Let’s walk through the important pieces.</p> <h3 id="defining-the-container-image">Defining the container image</h3> <p>First, we build a container with all the dependencies. DeepSeek’s OCR model needs PyTorch, transformers, and a few image processing libraries:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="n">pathlib</span> <span class="kn">import</span> <span class="n">Path</span>
<span class="kn">import</span> <span class="n">modal</span>

<span class="n">APP_NAME</span> <span class="o">=</span> <span class="sh">"</span><span class="s">deepseek-ocr-books-api-batch</span><span class="sh">"</span>
<span class="n">ROOT</span> <span class="o">=</span> <span class="nc">Path</span><span class="p">(</span><span class="n">__file__</span><span class="p">).</span><span class="nf">resolve</span><span class="p">().</span><span class="n">parents</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
<span class="n">BOOKS_DIR</span> <span class="o">=</span> <span class="n">ROOT</span> <span class="o">/</span> <span class="sh">"</span><span class="s">references</span><span class="sh">"</span> <span class="o">/</span> <span class="sh">"</span><span class="s">books</span><span class="sh">"</span>
<span class="n">PARSED_DIR</span> <span class="o">=</span> <span class="n">BOOKS_DIR</span> <span class="o">/</span> <span class="sh">"</span><span class="s">parsed</span><span class="sh">"</span>

<span class="n">image</span> <span class="o">=</span> <span class="p">(</span>
    <span class="n">modal</span><span class="p">.</span><span class="n">Image</span><span class="p">.</span><span class="nf">from_registry</span><span class="p">(</span>
        <span class="sh">"</span><span class="s">nvidia/cuda:11.8.0-cudnn8-runtime-ubuntu22.04</span><span class="sh">"</span><span class="p">,</span>
        <span class="n">add_python</span><span class="o">=</span><span class="sh">"</span><span class="s">3.11</span><span class="sh">"</span><span class="p">,</span>
    <span class="p">)</span>
    <span class="p">.</span><span class="nf">apt_install</span><span class="p">(</span><span class="sh">"</span><span class="s">git</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">libgl1</span><span class="sh">"</span><span class="p">,</span> <span class="sh">"</span><span class="s">libglib2.0-0</span><span class="sh">"</span><span class="p">)</span>
    <span class="p">.</span><span class="nf">pip_install</span><span class="p">(</span>
        <span class="sh">"</span><span class="s">torch==2.6.0</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">torchvision==0.21.0</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">transformers==4.46.3</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">PyMuPDF</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">Pillow</span><span class="sh">"</span><span class="p">,</span>
        <span class="sh">"</span><span class="s">numpy</span><span class="sh">"</span><span class="p">,</span>
        <span class="n">extra_index_url</span><span class="o">=</span><span class="sh">"</span><span class="s">https://download.pytorch.org/whl/cu118</span><span class="sh">"</span><span class="p">,</span>
    <span class="p">)</span>
<span class="p">)</span>

<span class="n">app</span> <span class="o">=</span> <span class="n">modal</span><span class="p">.</span><span class="nc">App</span><span class="p">(</span><span class="n">APP_NAME</span><span class="p">)</span>
</code></pre></div></div> <p>The paths at the top let us find PDFs relative to the script location, keeping configuration close to where it’s used.</p> <h3 id="the-fastapi-endpoint">The FastAPI endpoint</h3> <p>Here’s where the magic happens. We wrap a FastAPI server in Modal’s <code class="language-plaintext highlighter-rouge">@modal.asgi_app()</code> decorator, which means Modal will handle spinning up GPU instances and routing HTTP requests to them:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nd">@app.function</span><span class="p">(</span><span class="n">image</span><span class="o">=</span><span class="n">image</span><span class="p">,</span> <span class="n">gpu</span><span class="o">=</span><span class="sh">"</span><span class="s">A100</span><span class="sh">"</span><span class="p">,</span> <span class="n">timeout</span><span class="o">=</span><span class="mi">60</span> <span class="o">*</span> <span class="mi">60</span> <span class="o">*</span><span class="mi">2</span><span class="p">)</span> <span class="c1"># timeout of 2 hours
</span><span class="nd">@modal.asgi_app</span><span class="p">()</span>
<span class="k">def</span> <span class="nf">fastapi_app</span><span class="p">():</span>
    <span class="kn">from</span> <span class="n">fastapi</span> <span class="kn">import</span> <span class="n">FastAPI</span><span class="p">,</span> <span class="n">File</span><span class="p">,</span> <span class="n">UploadFile</span>
    <span class="kn">from</span> <span class="n">PIL</span> <span class="kn">import</span> <span class="n">Image</span>
    <span class="kn">import</span> <span class="n">torch</span>
    <span class="kn">from</span> <span class="n">transformers</span> <span class="kn">import</span> <span class="n">AutoModel</span><span class="p">,</span> <span class="n">AutoTokenizer</span>

    <span class="n">api</span> <span class="o">=</span> <span class="nc">FastAPI</span><span class="p">()</span>
    
    <span class="n">model_name</span> <span class="o">=</span> <span class="sh">"</span><span class="s">deepseek-ai/DeepSeek-OCR</span><span class="sh">"</span>
    <span class="n">tokenizer</span> <span class="o">=</span> <span class="n">AutoTokenizer</span><span class="p">.</span><span class="nf">from_pretrained</span><span class="p">(</span><span class="n">model_name</span><span class="p">,</span> <span class="n">trust_remote_code</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="n">model</span> <span class="o">=</span> <span class="n">AutoModel</span><span class="p">.</span><span class="nf">from_pretrained</span><span class="p">(</span><span class="n">model_name</span><span class="p">,</span> <span class="n">trust_remote_code</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    <span class="n">model</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="nf">cuda</span><span class="p">().</span><span class="nf">to</span><span class="p">(</span><span class="n">torch</span><span class="p">.</span><span class="n">bfloat16</span><span class="p">).</span><span class="nf">eval</span><span class="p">()</span>
</code></pre></div></div> <p>The model loads once when the container starts. Subsequent requests reuse the same loaded model, which is crucial for throughput when processing hundreds of pages.</p> <div class="important-box"> <i class="fas fa-exclamation-triangle"></i> <span>The `trust_remote_code=True` flag is necessary for DeepSeek's model because it includes custom code in the HuggingFace repository</span> </div> <h3 id="handling-batched-inference">Handling batched inference</h3> <p>OCR is embarrassingly parallel since each page is independent. We can process multiple pages in a single forward pass through the model, which is faster than processing them one at a time:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nd">@api.post</span><span class="p">(</span><span class="sh">"</span><span class="s">/ocr_batch</span><span class="sh">"</span><span class="p">)</span>
<span class="k">async</span> <span class="k">def</span> <span class="nf">ocr_batch</span><span class="p">(</span><span class="n">files</span><span class="p">:</span> <span class="nb">list</span><span class="p">[</span><span class="n">UploadFile</span><span class="p">]</span> <span class="o">=</span> <span class="nc">File</span><span class="p">(...))</span> <span class="o">-&gt;</span> <span class="nb">dict</span><span class="p">[</span><span class="nb">str</span><span class="p">,</span> <span class="nb">list</span><span class="p">[</span><span class="nb">str</span><span class="p">]]:</span>
    <span class="n">images</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">for</span> <span class="nb">file</span> <span class="ow">in</span> <span class="n">files</span><span class="p">:</span>
        <span class="n">image_bytes</span> <span class="o">=</span> <span class="k">await</span> <span class="nb">file</span><span class="p">.</span><span class="nf">read</span><span class="p">()</span>
        <span class="n">images</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">Image</span><span class="p">.</span><span class="nf">open</span><span class="p">(</span><span class="n">io</span><span class="p">.</span><span class="nc">BytesIO</span><span class="p">(</span><span class="n">image_bytes</span><span class="p">)).</span><span class="nf">convert</span><span class="p">(</span><span class="sh">"</span><span class="s">RGB</span><span class="sh">"</span><span class="p">))</span>
    <span class="n">batch_items</span> <span class="o">=</span> <span class="p">[</span><span class="nf">prepare_inputs</span><span class="p">(</span><span class="n">image</span><span class="p">)</span> <span class="k">for</span> <span class="n">image</span> <span class="ow">in</span> <span class="n">images</span><span class="p">]</span>
    <span class="n">texts</span> <span class="o">=</span> <span class="nf">run_batch</span><span class="p">(</span><span class="n">batch_items</span><span class="p">)</span>
    <span class="k">return</span> <span class="p">{</span><span class="sh">"</span><span class="s">texts</span><span class="sh">"</span><span class="p">:</span> <span class="n">texts</span><span class="p">}</span>
</code></pre></div></div> <p>The <code class="language-plaintext highlighter-rouge">run_batch</code> function handles the actual model inference. It pads inputs to the same length, runs them through the model in one shot, and decodes the outputs:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">run_batch</span><span class="p">(</span><span class="n">batch_items</span><span class="p">):</span>
    <span class="c1"># Pad sequences to same length
</span>    <span class="n">lengths</span> <span class="o">=</span> <span class="p">[</span><span class="n">item</span><span class="p">[</span><span class="mi">0</span><span class="p">].</span><span class="nf">size</span><span class="p">(</span><span class="mi">0</span><span class="p">)</span> <span class="k">for</span> <span class="n">item</span> <span class="ow">in</span> <span class="n">batch_items</span><span class="p">]</span>
    <span class="n">max_len</span> <span class="o">=</span> <span class="nf">max</span><span class="p">(</span><span class="n">lengths</span><span class="p">)</span>
    
    <span class="n">input_ids</span> <span class="o">=</span> <span class="n">torch</span><span class="p">.</span><span class="nf">full</span><span class="p">((</span><span class="nf">len</span><span class="p">(</span><span class="n">batch_items</span><span class="p">),</span> <span class="n">max_len</span><span class="p">),</span> <span class="n">pad_id</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="nb">long</span><span class="p">)</span>
    <span class="c1"># ... (padding logic)
</span>    
    <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="nf">autocast</span><span class="p">(</span><span class="sh">"</span><span class="s">cuda</span><span class="sh">"</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">torch</span><span class="p">.</span><span class="n">bfloat16</span><span class="p">):</span>
        <span class="k">with</span> <span class="n">torch</span><span class="p">.</span><span class="nf">no_grad</span><span class="p">():</span>
            <span class="n">output_ids</span> <span class="o">=</span> <span class="n">model</span><span class="p">.</span><span class="nf">generate</span><span class="p">(</span>
                <span class="n">input_ids</span><span class="p">.</span><span class="nf">cuda</span><span class="p">(),</span>
                <span class="n">images</span><span class="o">=</span><span class="n">images</span><span class="p">,</span>
                <span class="n">max_new_tokens</span><span class="o">=</span><span class="mi">8192</span><span class="p">,</span>
                <span class="n">temperature</span><span class="o">=</span><span class="mf">0.0</span><span class="p">,</span>
            <span class="p">)</span>
    
    <span class="c1"># Decode outputs
</span>    <span class="n">outputs</span> <span class="o">=</span> <span class="p">[]</span>
    <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">out_ids</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">(</span><span class="n">output_ids</span><span class="p">):</span>
        <span class="n">token_ids</span> <span class="o">=</span> <span class="n">out_ids</span><span class="p">[</span><span class="n">lengths</span><span class="p">[</span><span class="n">i</span><span class="p">]:].</span><span class="nf">tolist</span><span class="p">()</span>
        <span class="n">text</span> <span class="o">=</span> <span class="n">tokenizer</span><span class="p">.</span><span class="nf">decode</span><span class="p">(</span><span class="n">token_ids</span><span class="p">,</span> <span class="n">skip_special_tokens</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        <span class="n">outputs</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">text</span><span class="p">.</span><span class="nf">strip</span><span class="p">())</span>
    <span class="k">return</span> <span class="n">outputs</span>
</code></pre></div></div> <p>Setting <code class="language-plaintext highlighter-rouge">temperature=0.0</code> makes the output deterministic, which helps the model generate results which are more reproducible.</p> <h3 id="the-local-client">The local client</h3> <p>With the server deployed on Modal, we need a client to feed it pages. The <code class="language-plaintext highlighter-rouge">@app.local_entrypoint()</code> decorator marks a function that runs on your local machine but can communicate with the Modal-deployed server:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nd">@app.local_entrypoint</span><span class="p">()</span>
<span class="k">def</span> <span class="nf">main</span><span class="p">(</span><span class="n">api_url</span><span class="p">:</span> <span class="nb">str</span><span class="p">,</span> <span class="n">book</span><span class="p">:</span> <span class="nb">str</span> <span class="o">=</span> <span class="sh">""</span><span class="p">,</span> <span class="n">max_pages</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="bp">None</span><span class="p">,</span> <span class="n">batch_size</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">1</span><span class="p">):</span>
    <span class="kn">import</span> <span class="n">fitz</span>  <span class="c1"># PyMuPDF
</span>    
    <span class="k">if</span> <span class="n">book</span><span class="p">:</span>
        <span class="n">pdf_paths</span> <span class="o">=</span> <span class="p">[</span><span class="n">BOOKS_DIR</span> <span class="o">/</span> <span class="n">book</span><span class="p">]</span>
    <span class="k">else</span><span class="p">:</span>
        <span class="n">pdf_paths</span> <span class="o">=</span> <span class="nf">sorted</span><span class="p">(</span><span class="n">BOOKS_DIR</span><span class="p">.</span><span class="nf">glob</span><span class="p">(</span><span class="sh">"</span><span class="s">*.pdf</span><span class="sh">"</span><span class="p">))</span>
    
    <span class="k">for</span> <span class="n">pdf_path</span> <span class="ow">in</span> <span class="n">pdf_paths</span><span class="p">:</span>
        <span class="k">with</span> <span class="n">fitz</span><span class="p">.</span><span class="nf">open</span><span class="p">(</span><span class="n">pdf_path</span><span class="p">)</span> <span class="k">as</span> <span class="n">doc</span><span class="p">:</span>
            <span class="n">batch_pages</span> <span class="o">=</span> <span class="p">[]</span>
            <span class="k">for</span> <span class="n">page_index</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">doc</span><span class="p">.</span><span class="n">page_count</span><span class="p">):</span>
                <span class="n">page</span> <span class="o">=</span> <span class="n">doc</span><span class="p">[</span><span class="n">page_index</span><span class="p">]</span>
                <span class="n">pix</span> <span class="o">=</span> <span class="n">page</span><span class="p">.</span><span class="nf">get_pixmap</span><span class="p">(</span><span class="n">matrix</span><span class="o">=</span><span class="n">fitz</span><span class="p">.</span><span class="nc">Matrix</span><span class="p">(</span><span class="mi">2</span><span class="p">,</span> <span class="mi">2</span><span class="p">))</span>  <span class="c1"># 2x zoom
</span>                <span class="n">batch_pages</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">pix</span><span class="p">.</span><span class="nf">tobytes</span><span class="p">(</span><span class="sh">"</span><span class="s">png</span><span class="sh">"</span><span class="p">))</span>
                
                <span class="k">if</span> <span class="nf">len</span><span class="p">(</span><span class="n">batch_pages</span><span class="p">)</span> <span class="o">&gt;=</span> <span class="n">batch_size</span><span class="p">:</span>
                    <span class="c1"># Send batch to server
</span>                    <span class="n">response</span> <span class="o">=</span> <span class="n">requests</span><span class="p">.</span><span class="nf">post</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="si">{</span><span class="n">api_url</span><span class="si">}</span><span class="s">/ocr_batch</span><span class="sh">"</span><span class="p">,</span> <span class="n">files</span><span class="o">=</span><span class="n">files</span><span class="p">)</span>
                    <span class="n">texts</span> <span class="o">=</span> <span class="n">response</span><span class="p">.</span><span class="nf">json</span><span class="p">()[</span><span class="sh">"</span><span class="s">texts</span><span class="sh">"</span><span class="p">]</span>
                    <span class="c1"># Save results...
</span></code></pre></div></div> <p>The render-at-2x trick (<code class="language-plaintext highlighter-rouge">fitz.Matrix(2, 2)</code>) is important. Higher resolution input means the OCR model can read smaller text and mathematical subscripts more accurately.</p> <h3 id="cleaning-up-the-output">Cleaning up the output</h3> <p>DeepSeek’s OCR model includes grounding tags, which are coordinates indicating where each piece of text appeared on the page. These can be useful for some applications, but I didn’t need them for searchable text:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">tag_pattern</span> <span class="o">=</span> <span class="n">re</span><span class="p">.</span><span class="nf">compile</span><span class="p">(</span>
    <span class="sa">r</span><span class="sh">"</span><span class="s">&lt;\|ref\|&gt;(.*?)&lt;\|/ref\|&gt;&lt;\|det\|&gt;.*?&lt;\|/det\|&gt;</span><span class="sh">"</span><span class="p">,</span>
    <span class="n">flags</span><span class="o">=</span><span class="n">re</span><span class="p">.</span><span class="n">DOTALL</span><span class="p">,</span>
<span class="p">)</span>

<span class="k">for</span> <span class="n">page_idx</span><span class="p">,</span> <span class="n">text</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">batch_page_indices</span><span class="p">,</span> <span class="n">texts</span><span class="p">):</span>
    <span class="n">text</span> <span class="o">=</span> <span class="n">tag_pattern</span><span class="p">.</span><span class="nf">sub</span><span class="p">(</span><span class="sa">r</span><span class="sh">"</span><span class="s">\1</span><span class="sh">"</span><span class="p">,</span> <span class="n">text</span><span class="p">)</span>  <span class="c1"># Keep text, drop coordinates
</span>    <span class="n">page_path</span> <span class="o">=</span> <span class="n">pages_dir</span> <span class="o">/</span> <span class="sa">f</span><span class="sh">"</span><span class="s">page_</span><span class="si">{</span><span class="n">page_idx</span> <span class="o">+</span> <span class="mi">1</span><span class="si">:</span><span class="mi">04</span><span class="n">d</span><span class="si">}</span><span class="s">.mmd</span><span class="sh">"</span>
    <span class="n">page_path</span><span class="p">.</span><span class="nf">write_text</span><span class="p">(</span><span class="n">text</span><span class="p">,</span> <span class="n">encoding</span><span class="o">=</span><span class="sh">"</span><span class="s">utf-8</span><span class="sh">"</span><span class="p">)</span>
</code></pre></div></div> <p>The <code class="language-plaintext highlighter-rouge">.mmd</code> extension stands for “multimodal markdown,” a convention for markdown that came from an OCR source and might have some artifacts.</p> <h3 id="running-it">Running it</h3> <p>To use this script, you first deploy the server:</p> <div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>modal deploy deepseek_ocr_modal.py
</code></pre></div></div> <p>This gives you a URL like <code class="language-plaintext highlighter-rouge">https://your-workspace--deepseek-ocr-books-api-batch-fastapi-app.modal.run</code>. Then run the client:</p> <div class="language-bash highlighter-rouge"><div class="highlight"><pre class="highlight"><code>modal run deepseek_ocr_modal.py <span class="nt">--api-url</span> <span class="s2">"https://..."</span> <span class="nt">--book</span> <span class="s2">"Gelman - Bayesian Data Analysis.pdf"</span>
</code></pre></div></div> <p>For BDA’s ~600 pages, with a batch size of 4, processing takes about 45 minutes on an A100. The output is a directory of markdown files, one per page, plus a concatenated <code class="language-plaintext highlighter-rouge">document.mmd</code> with page markers. The whole thing cost maybe $2 bucks. I think it’s a great deal. I also now have a setup I can reuse for any PDF, including course notes, papers, and other textbooks.</p> <p>The OCR quality on mathematical content is surprisingly good. Nearly all equations come through intact.</p> <p>The real payoff comes from downstream use. I can now <code class="language-plaintext highlighter-rouge">grep</code> through BDA, paste sections into Claude and ask it to explain the notation, or build a proper search index. All of this from a PDF that was previously just a collection of images.</p> <p>Here’s what the text looks like once it’s parsed and combined into a single file:</p> <div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
## Exponential model

The exponential distribution is commonly used to model 'waiting times' and other continuous, positive, real- valued random variables, often measured on a time scale. The sampling distribution of an outcome \(y\) , given parameter \(\theta\) , is

\[p(y|\theta) = \theta \exp (-y\theta),\mathrm{for} y  0,\]

and \(\theta = 1 / \mathrm{E}(y|\theta)\) is called the 'rate.' Mathematically, the exponential is a special case of the gamma distribution with the parameters \((\alpha ,\beta) = (1,\theta)\) . In this case, however, it is being used as a sampling distribution for an outcome \(y\) , not a prior distribution for a parameter \(\theta\) , as in the Poisson example.

The exponential distribution has a 'memoryless' property that makes it a natural model for survival or lifetime data; the probability that an object survives an additional length of time \(t\) is independent of the time elapsed to this point: \(\operatorname *{Pr}(y  t + s\mid y  s,\theta) = \operatorname *{Pr}(y  t\mid \theta)\) for any \(s,t\) . The conjugate prior distribution for the exponential parameter \(\theta\) , as for the Poisson mean, is \(\operatorname {Gamma}(\theta |\alpha ,\beta)\) with corresponding posterior distribution \(\operatorname {Gamma}(\theta |\alpha +1,\beta +y)\) . The sampling distribution of \(n\) independent exponential observations, \(y = (y_{1},\ldots ,y_{n})\) , with constant rate \(\theta\) is

\[p(y|\theta) = \theta^{n}\exp (-n\bar{y}\theta),\mathrm{for}\bar{y}\geq 0,\]

which when viewed as the likelihood of \(\theta\) , for fixed \(y\) , is proportional to a \(\operatorname {Gamma}(n + 1,n\bar{y})\) density. Thus the \(\operatorname {Gamma}(\alpha ,\beta)\) prior distribution for \(\theta\) can be viewed as \(\alpha - 1\) exponential observations with total waiting time \(\beta\) (see Exercise 2.19).
&lt;--- Page Split ---&gt;
image
image_caption
&lt;centerFigure 2.6 The counties of the United States with the highest \(10\%\) age-standardized death rates for cancer of kidney/ureter for U.S. white males, 1980-1989. Why are most of the shaded counties in the middle of the country? See Section 2.7 for discussion. &lt;/center

### 2.7 Example: informative prior distribution for cancer rates

At the end of Section 2.4, we considered the effect of the prior distribution on inference given a fixed quantity of data. Here, in contrast, we consider a large set of inferences, each based on different data but with a common prior distribution. In addition to illustrating the role of the prior distribution, this example introduces hierarchical modeling, to which we return in Chapter 5.
</code></pre></div></div> <p>Compare that to the original passage from the textbook:</p> <p align="center"> <img src="/images/bda-passage.png" alt="Original passage from Bayesian Data Analysis" width="50%"/> </p> <p>I think it did a pretty good job!</p> <p>If you have a collection of scanned textbooks, this approach might be worth the 5 minutes it takes to set it up.</p>]]></content><author><name></name></author><summary type="html"><![CDATA[Parsing a statistics textbook with Deepseek OCR on Modal]]></summary></entry><entry><title type="html">How I accidentally became a power AI user in big tech</title><link href="https://ckrapu.github.io/blog/2026/how-i-accidentally-became-elite/" rel="alternate" type="text/html" title="How I accidentally became a power AI user in big tech"/><published>2026-01-08T00:00:00+00:00</published><updated>2026-01-08T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2026/how-i-accidentally-became-elite</id><content type="html" xml:base="https://ckrapu.github.io/blog/2026/how-i-accidentally-became-elite/"><![CDATA[<blockquote> <p><strong>TL;DR:</strong> I’m in a non-engineering department and somehow became an extreme outlier for AI-assisted coding usage because I have to wear too many hats and I’m picky about our dependencies.</p> </blockquote> <hr/> <p>Last spring, my company onboarded onto Cursor and shortly after, it was gaining traction in virtually every department. I didn’t see myself as the target audience—my department doesn’t ship product code and, before Cursor, I generally had a low-tech approach: I had autocomplete turned off most of the time, and my AI assistance was mostly limited to ChatGPT buried 50 tabs deep in Chrome.</p> <p>I didn’t know Cursor stats were visible to most managers, and when I found the department-level reports, I was genuinely worried at first. I was one of the more aggressive users of Claude 4 Opus at the time, and I thought I might be reprimanded for incurring too many API fees.</p> <p>The usage dashboard showed that I was in the top 0.XX% of users for activity (as measured by <strong>Agent</strong> requests), but also that I had a very unusual activity pattern of roughly 4x as many <strong>Ask</strong> requests (LLM only explains, like ChatGPT) as compared to <strong>Agent</strong> requests (LLM actually writes code in the project directory structure). This didn’t surprise me much; I enjoy having a back-and-forth dialogue about all the technologies I don’t know yet. I find this almost as much fun as writing the code itself. Most users appeared to be showing a 0.3-0.5x ratio. They were less interested in having the AI explain it and simply wanted the code to be written. I will also note that this is all good-faith activity; I wasn’t churning out slop or using it for non-work purposes.</p> <p>Despite being an early adopter of some AI technology (I started building around GPT-3 in early 2021) I genuinely enjoy typing, writing out functions, scaffolding, and the minutiae of writing code. To me, it’s the same feeling as watching TV or making small talk. Most of it doesn’t require a huge amount of thinking, and it feels pretty nice when you’re in the flow. I never thought of myself as an AI-first, 10x developer type of person.</p> <p>Here’s what drove my usage:</p> <h3 id="1-documentation-gaps-rich-tooling">1. Documentation gaps, rich tooling</h3> <p>We have an internal environment rich with tools, but light on documentation and frequently requiring workarounds. Documentation is lacking at many companies, and it’s super common for a service’s API to drift out of alignment with docs, especially when it’s just an outdated guide someone cobbled together.</p> <p>I spent a <em>lot</em> of effort working around this. AI was a panacea for this problem. Lots of nominally-working but hard-to-use services were now perfectly usable, all of a sudden, because the LLM has the patience to poke around the API, try a dozen examples, read a few different conflicting sources, and figure out how to actually make it work.</p> <h3 id="2-no-devops-engineer-and-thats-fine">2. No devops engineer (and that’s fine)</h3> <p>I don’t have a devops engineer, I will never get one, and that’s okay with me. Most of my projects are deployed via managed cloud services instead of k8s. Early on, I made the judgment call that, even in the best circumstances, our traffic was never going to exceed \(10^4\) active users and cloud costs wouldn’t be an issue, even unoptimized.</p> <p>This means that I can largely control the entire apparatus via Terraform without needing an additional teammate to manage a k8s cluster.</p> <h3 id="3-the-technical-lead-does-analytics">3. The technical lead does analytics</h3> <p>In this case, that was me. Many teams operate by having BA or junior DS-type employees perform analytics and reporting. They’re cheaper, but less likely to automate their work given less coding experience—and I say this as someone with a DS background.</p> <p>Since the reporting needs fell on my shoulders, it was definitely in my interest to vibe code as many SQL queries and dashboarding functions as possible to short-circuit the route from raw data to management decisions.</p> <h3 id="4-not-waiting-for-mcps-when-a-cursorrules-and-some-bash-aliases-would-suffice">4. Not waiting for MCPs when a <code class="language-plaintext highlighter-rouge">.cursorrules</code> and some bash aliases would suffice</h3> <p>Pretty much all companies’ IT departments have been talking about MCP servers for the last year. Some are further along than others. I, being inordinately impatient, just wrote a few bash aliases to make API calls and, for many operational tasks (“run a report against our sales data”, “find all the newest companies we haven’t reached out to yet”, etc.), I could just tell Cursor that it had these command line tools for using Perplexity, Databricks, our company’s Powerpoint template, etc.</p> <h3 id="5-a-great-full-stack-template">5. A great full-stack template</h3> <p>Many of the similar projects I’d seen at my company and elsewhere used software like Streamlit, Gradio, and OpenWebUI. These are fine, but they aren’t ready out-of-the-box for going into production.</p> <p>I invested a <em>lot</em> of time in early 2024 figuring out how we’d deploy an internal chatbot and eventually landed on <a href="https://docs.chainlit.io/get-started/overview">Chainlit</a>. This was a stroke of luck as it was a solid template with batteries included for OAuth, telemetry, human-in-the-loop, and other features that many teams were authoring from scratch at that point. It had quirks (the default schema tanks after 100k conversations) but otherwise was great to work with. It also pushed us to work mostly in FastAPI + LlamaIndex + SQLAlchemy, all of which turned out to be rock-solid in production.</p> <p>Ironically, this OSS project was a much better base than all of the paid consulting/dev shop-produced products that I’d seen.</p> <p>By scaffolding around a quality codebase with sensible defaults, clean API structure, and a minimal but appealing UI, grounding on that foundation gave us much more effective LLM-produced code than if we had vibe coded from scratch.</p> <p>These effects were compounding for me: it produced nice code which I enjoyed working on, which in turn made me want to do more of it instead of getting frustrated by opaque errors.</p>]]></content><author><name></name></author><category term="ai"/><category term="cursor"/><category term="software-development"/><summary type="html"><![CDATA[On becoming 99.9th percentile in token consumption]]></summary></entry><entry><title type="html">Using Antigravity for Statistical Physics in JS</title><link href="https://ckrapu.github.io/blog/2025/antigravity-stat-mech/" rel="alternate" type="text/html" title="Using Antigravity for Statistical Physics in JS"/><published>2025-11-21T00:00:00+00:00</published><updated>2025-11-21T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2025/antigravity-stat-mech</id><content type="html" xml:base="https://ckrapu.github.io/blog/2025/antigravity-stat-mech/"><![CDATA[<p>I like learning about the hidden benchmarks that everyone seems to bring out when a new large language model drops. Mine used to be asking the model about obscure but well-documented people on the internet like family members or acquaintances in the sciences or with IMDB credits. Since ~late 2024, most models are nailing that one so it’s not as interesting. Instead, I’ve moved onto Javascript-based visualizations. of statistical physics</p> <p>Since Gemini 3 and Google’s Antigravity IDE were released recently (and yes, I am aware it is basically Windsurf), I wanted to give it a try with an easy one - the Ising model of ferromagnetism.</p> <p>Here’s what Antigravity with Gemini 3 Pro cooked up in an hour:</p> <style>:root{--bg-color:#1c1c1d;--text-color:#fff;--accent-color:#0fc;--secondary-color:#f0f;--grid-gap:2px;--control-bg:rgba(255,255,255,0.1);--font-family:'Inter',-apple-system,BlinkMacSystemFont,'Segoe UI',Roboto,Oxygen,Ubuntu,Cantarell,'Open Sans','Helvetica Neue',sans-serif}#ising-app.light-theme{--bg-color:#fff;--text-color:#333;--accent-color:#007bff;--secondary-color:#dc3545;--control-bg:rgba(0,0,0,0.05)}#ising-app{background-color:var(--bg-color);color:var(--text-color);font-family:var(--font-family);margin:2rem 0;padding:1.5rem 1rem;border-radius:12px;transition:background-color .3s,color .3s}h1{font-weight:300;letter-spacing:2px;margin-bottom:10px;text-transform:uppercase;font-size:1.5rem}#main-container{display:flex;flex-direction:column;align-items:center;gap:20px;padding-bottom:40px}#canvas-container{position:relative;box-shadow:0 0 50px rgba(0,0,0,0.5);border-radius:12px;overflow:hidden;cursor:none;width:100%;max-width:600px}#simCanvas{display:block;width:100%;height:auto;max-width:100%}#info-overlay{position:absolute;top:0;left:0;width:100%;height:100%;background:rgba(0,0,0,0.85);backdrop-filter:blur(8px);-webkit-backdrop-filter:blur(8px);color:#fff;display:flex;flex-direction:column;justify-content:center;align-items:center;padding:40px;box-sizing:border-box;opacity:0;visibility:hidden;transition:opacity .3s ease,visibility .3s ease;z-index:10;text-align:center;pointer-events:none}#ising-app.light-theme #info-overlay{background:rgba(255,255,255,0.9);color:#333}#info-overlay.visible{opacity:1;visibility:visible;pointer-events:auto}#info-content{max-width:500px}#info-content h2{margin-top:0;font-weight:300;text-transform:uppercase;letter-spacing:2px;margin-bottom:20px;border-bottom:1px solid var(--accent-color);display:inline-block;padding-bottom:5px}#info-content p{font-size:.95rem;line-height:1.6;margin-bottom:15px;opacity:.9}.math-display{font-family:"Times New Roman",Times,serif;font-style:italic;font-size:1.2rem;background:rgba(255,255,255,0.05);padding:15px;border-radius:8px;margin:20px 0;border:1px solid rgba(255,255,255,0.1);display:block}#ising-app.light-theme .math-display{background:rgba(0,0,0,0.03);border:1px solid rgba(0,0,0,0.05)}.math-inline{font-family:"Times New Roman",Times,serif;font-style:italic;padding:0 2px}#controls{display:flex;flex-direction:column;align-items:center;gap:10px;background:var(--control-bg);padding:15px 30px;border-radius:20px;backdrop-filter:blur(10px);border:1px solid rgba(255,255,255,0.1);position:relative}#controls-header{display:flex;align-items:center;gap:8px;margin-bottom:5px}#controls-title{font-size:.9rem;text-transform:uppercase;letter-spacing:1px;font-weight:500;line-height:1}#info-icon{width:16px;height:16px;cursor:help;opacity:.7;transition:opacity .2s;position:relative;display:flex;align-items:center;justify-content:center}#info-icon:hover{opacity:1}#info-icon svg{width:100%;height:100%;fill:currentColor}.math{font-family:'Times New Roman',Times,serif;font-style:italic;background:rgba(255,255,255,0.1);padding:2px 5px;border-radius:4px;display:block;text-align:center;margin:10px 0}#ising-app.light-theme .math{background:rgba(0,0,0,0.05)}#controls-body{display:flex;gap:20px;align-items:center;flex-wrap:wrap;justify-content:center}.control-group{display:flex;flex-direction:column;align-items:center;gap:5px}label{font-size:.7rem;text-transform:uppercase;letter-spacing:1px;opacity:.8}input[type="range"]{-webkit-appearance:none;appearance:none;width:100px;height:4px;background:rgba(128,128,128,0.3);border-radius:2px;outline:0}input[type="range"]::-webkit-slider-thumb{-webkit-appearance:none;width:14px;height:14px;background:var(--text-color);border-radius:50%;cursor:pointer;transition:transform .1s}input[type="range"]::-webkit-slider-thumb:hover{transform:scale(1.2)}.icon-btn{background:0;border:1px solid var(--text-color);color:var(--text-color);width:36px;height:36px;border-radius:50%;cursor:pointer;display:flex;align-items:center;justify-content:center;padding:0;transition:all .2s}.icon-btn:hover{background:var(--text-color);color:var(--bg-color)}.icon-btn svg{width:18px;height:18px;fill:currentColor}#plot-container{width:100%;max-width:600px;height:60px;background:transparent;border-radius:8px;overflow:hidden;position:relative;margin-top:-10px}#plotCanvas{width:100%;height:100%;cursor:default}.value-display{font-family:monospace;font-size:.8rem}@media(max-width:768px){#ising-app{padding:1rem .5rem}#controls{padding:12px 15px}#controls-body{gap:12px}input[type="range"]{width:80px}#info-overlay{padding:20px}#info-content{max-width:100%}}@media(max-width:480px){#controls-body{gap:8px}input[type="range"]{width:70px}.control-group label{font-size:.65rem}}</style> <div id="ising-app"> <div id="main-container"> <div id="canvas-container"> <canvas id="simCanvas" width="600" height="600"></canvas> <div id="info-overlay"> <div id="info-content"> <h2>The Ising Model</h2> <p> A mathematical model of ferromagnetism in statistical mechanics. The grid consists of discrete variables (spins) that can be in one of two states (+1 or -1). </p> <div class="math-display"> H(&sigma;) = -J &sum;<sub>&lt;ij&gt;</sub> &sigma;<sub>i</sub>&sigma;<sub>j</sub> - h &sum;<sub>j</sub> &sigma;<sub>j</sub> </div> <p> <strong>Simulation:</strong> This visualization uses a <em>Random Scan Gibbs Sampler</em>. In each step, a single spin is chosen at random and updated based on the Boltzmann distribution determined by its neighbors and the external field. </p> <p style="font-size: 0.85rem; opacity: 0.7; margin-top: 20px;"> Named after physicist Ernst Ising, who solved the 1D model in his 1924 thesis. </p> </div> </div> </div> <div id="controls"> <div id="controls-header"> <span id="controls-title">Ising Model</span> <div id="info-icon"> <svg viewBox="0 0 24 24"> <path d="M12 2C6.48 2 2 6.48 2 12s4.48 10 10 10 10-4.48 10-10S17.52 2 12 2zm1 15h-2v-6h2v6zm0-8h-2V7h2v2z"/> </svg> </div> </div> <div id="controls-body"> <button id="theme-toggle" class="icon-btn" title="Toggle Theme"> <svg class="sun-icon" viewBox="0 0 24 24" style="display: none;"> <path d="M12 7c-2.76 0-5 2.24-5 5s2.24 5 5 5 5-2.24 5-5-2.24-5-5-5zM2 13h2c.55 0 1-.45 1-1s-.45-1-1-1H2c-.55 0-1 .45-1 1s.45 1 1 1zm18 0h2c.55 0 1-.45 1-1s-.45-1-1-1h-2c-.55 0-1 .45-1 1s.45 1 1 1zM11 2v2c0 .55.45 1 1 1s1-.45 1-1V2c0-.55-.45-1-1-1s-1 .45-1 1zm0 18v2c0 .55.45 1 1 1s1-.45 1-1v-2c0-.55-.45-1-1-1s-1 .45-1 1zM5.99 4.58c-.39-.39-1.03-.39-1.41 0-.39.39-.39 1.03 0 1.41l1.06 1.06c.39.39 1.03.39 1.41 0s.39-1.03 0-1.41L5.99 4.58zm12.37 12.37c-.39-.39-1.03-.39-1.41 0-.39.39-.39 1.03 0 1.41l1.06 1.06c.39.39 1.03.39 1.41 0 .39-.39.39-1.03 0-1.41l-1.06-1.06zm1.06-10.96c.39-.39.39-1.03 0-1.41-.39-.39-1.03-.39-1.41 0l-1.06 1.06c-.39.39-.39 1.03 0 1.41s1.03.39 1.41 0l1.06-1.06zM7.05 18.36c.39-.39.39-1.03 0-1.41-.39-.39-1.03-.39-1.41 0l-1.06 1.06c-.39.39-.39 1.03 0 1.41s1.03.39 1.41 0l1.06-1.06z"/> </svg> <svg class="moon-icon" viewBox="0 0 24 24"> <path d="M9.37 5.51c-.18.64-.27 1.31-.27 1.99 0 4.08 3.32 7.4 7.4 7.4.68 0 1.35-.09 1.99-.27C17.45 18.55 14.4 21 10.75 21c-4.83 0-8.75-3.92-8.75-8.75 0-3.65 2.45-6.7 5.37-7.74z"/> </svg> </button> <div class="control-group"> <label>Temperature <span id="temp-val" class="value-display">2.27</span></label> <input type="range" id="temp-slider" min="0.1" max="5.0" step="0.01" value="2.27"/> </div> <div class="control-group"> <label>External Field <span id="field-val" class="value-display">0.00</span></label> <input type="range" id="field-slider" min="-2.0" max="2.0" step="0.1" value="0.0"/> </div> <div class="control-group"> <label>Speed <span id="speed-val" class="value-display">4</span></label> <input type="range" id="speed-slider" min="1" max="10" step="1" value="4"/> </div> <button id="restart-btn" class="icon-btn" title="Restart Demo"> <svg viewBox="0 0 24 24"> <path d="M17.65 6.35C16.2 4.9 14.21 4 12 4c-4.42 0-7.99 3.58-7.99 8s3.57 8 7.99 8c3.73 0 6.84-2.55 7.73-6h-2.08c-.82 2.33-3.04 4-5.65 4-3.31 0-6-2.69-6-6s2.69-6 6-6c1.66 0 3.14.69 4.22 1.78L13 11h7V4l-2.35 2.35z"/> </svg> </button> </div> </div> <div id="plot-container"> <canvas id="plotCanvas" width="600" height="60"></canvas> </div> </div> </div> <script>function initGrid(){grid=[];for(let e=0;e<GRID_SIZE;e++){let e=[];for(let t=0;t<GRID_SIZE;t++)e.push(Math.random()>.5?1:-1);grid.push(e)}}function resetDemo(){isDemoMode=!0,demoTime=0,initGrid(),magnetizationHistory=new Array(MAX_HISTORY).fill(0)}function draw(){const e=getComputedStyle(app).getPropertyValue("--bg-color").trim(),t=getComputedStyle(app).getPropertyValue("--accent-color").trim(),o=getComputedStyle(app).getPropertyValue("--secondary-color").trim();ctx.fillStyle=e,ctx.fillRect(0,0,canvas.width,canvas.height);for(let e=0;e<GRID_SIZE;e++)for(let n=0;n<GRID_SIZE;n++){const a=n*CELL_SIZE+CELL_SIZE/2,l=e*CELL_SIZE+CELL_SIZE/2,i=CELL_SIZE/2*.8;ctx.beginPath(),ctx.arc(a,l,i,0,2*Math.PI),1===grid[e][n]?(ctx.fillStyle=t,isDarkTheme?(ctx.shadowBlur=10,ctx.shadowColor=t):ctx.shadowBlur=0):(ctx.fillStyle=o,isDarkTheme?(ctx.shadowBlur=10,ctx.shadowColor=o):ctx.shadowBlur=0),ctx.fill(),ctx.shadowBlur=0}isMouseDown&&(ctx.beginPath(),ctx.arc(mouseX,mouseY,BRUSH_RADIUS,0,2*Math.PI),ctx.fillStyle="rgba(128, 128, 128, 0.55)",ctx.fill(),ctx.strokeStyle="rgba(255, 255, 255, 0.5)",ctx.lineWidth=1,ctx.stroke())}function drawPlot(){plotCtx.clearRect(0,0,plotCanvas.width,plotCanvas.height);const e=getComputedStyle(app).getPropertyValue("--accent-color").trim(),t=getComputedStyle(app).getPropertyValue("--secondary-color").trim(),o=getComputedStyle(app).getPropertyValue("--text-color").trim();plotCtx.lineWidth=2,plotCtx.beginPath();for(let e=0;e<MAX_HISTORY;e++){const t=e/(MAX_HISTORY-1)*plotCanvas.width,o=(1-magnetizationHistory[e])/2*plotCanvas.height;0===e?plotCtx.moveTo(t,o):plotCtx.lineTo(t,o)}const n=plotCtx.createLinearGradient(0,plotCanvas.height,0,0);n.addColorStop(0,t),n.addColorStop(1,e),plotCtx.strokeStyle=n,plotCtx.stroke(),plotCtx.fillStyle=o,plotCtx.font="10px sans-serif",plotCtx.textAlign="center",plotCtx.textBaseline="bottom",plotCtx.globalAlpha=.7,plotCtx.fillText("Magnetization",plotCanvas.width/2,plotCanvas.height-2),plotCtx.globalAlpha=1}function update(){const e=Math.floor(Math.pow(speed,3)+10);for(let t=0;t<e;t++){const e=Math.floor(Math.random()*GRID_SIZE),t=Math.floor(Math.random()*GRID_SIZE),o=grid[(e-1+GRID_SIZE)%GRID_SIZE][t],n=grid[(e+1)%GRID_SIZE][t],a=grid[e][(t-1+GRID_SIZE)%GRID_SIZE],l=grid[e][(t+1)%GRID_SIZE],i=1/temperature,s=J*(o+n+a+l)+externalField,r=1/(1+Math.exp(-2*i*s));grid[e][t]=Math.random()<r?1:-1}}function calculateMagnetization(){let e=0;for(let t=0;t<GRID_SIZE;t++)for(let o=0;o<GRID_SIZE;o++)e+=grid[t][o];return e/(GRID_SIZE*GRID_SIZE)}function updateDemo(){isDemoMode&&(demoTime+=.01,temperature=2.27+1*Math.sin(.5*demoTime),externalField=1.5*Math.sin(.8*demoTime),tempSlider.value=temperature,tempVal.textContent=temperature.toFixed(2),fieldSlider.value=externalField,fieldVal.textContent=externalField.toFixed(2))}function loop(){updateDemo(),update(),draw();const e=calculateMagnetization();magnetizationHistory.push(e),magnetizationHistory.shift(),drawPlot(),requestAnimationFrame(loop)}function handleInteraction(e){const t=canvas.getBoundingClientRect(),o=e.clientX-t.left,n=e.clientY-t.top;if(mouseX=o,mouseY=n,isMouseDown){isDemoMode=!1;for(let e=0;e<GRID_SIZE;e++)for(let t=0;t<GRID_SIZE;t++){const a=t*CELL_SIZE+CELL_SIZE/2,l=e*CELL_SIZE+CELL_SIZE/2;Math.sqrt((o-a)**2+(n-l)**2)<BRUSH_RADIUS&&0!==paintState&&(grid[e][t]=paintState)}}}function stopDemo(){isDemoMode=!1}const GRID_SIZE=30,CELL_SIZE=20,J=1,BRUSH_RADIUS=1.5*CELL_SIZE,app=document.getElementById("ising-app");let grid=[],temperature=2.27,externalField=0,speed=4,isMouseDown=!1,isDarkTheme=!0,mouseX=-1e3,mouseY=-1e3,isDemoMode=!0,demoTime=0;const MAX_HISTORY=300;let magnetizationHistory=new Array(MAX_HISTORY).fill(0);const canvas=document.getElementById("simCanvas"),ctx=canvas.getContext("2d"),plotCanvas=document.getElementById("plotCanvas"),plotCtx=plotCanvas.getContext("2d"),tempSlider=document.getElementById("temp-slider"),tempVal=document.getElementById("temp-val"),fieldSlider=document.getElementById("field-slider"),fieldVal=document.getElementById("field-val"),speedSlider=document.getElementById("speed-slider"),speedVal=document.getElementById("speed-val"),themeToggle=document.getElementById("theme-toggle"),restartBtn=document.getElementById("restart-btn"),infoIcon=document.getElementById("info-icon"),infoOverlay=document.getElementById("info-overlay");let paintState=0;canvas.addEventListener("mousedown",e=>{isMouseDown=!0,isDemoMode=!1;const t=canvas.getBoundingClientRect(),o=e.clientX-t.left,n=e.clientY-t.top,a=Math.floor(o/(t.width/GRID_SIZE)),l=Math.floor(n/(t.height/GRID_SIZE));paintState=l>=0&&l<GRID_SIZE&&a>=0&&a<GRID_SIZE?-grid[l][a]:1,handleInteraction(e)}),window.addEventListener("mouseup",()=>{isMouseDown=!1,paintState=0}),canvas.addEventListener("mousemove",e=>{const t=canvas.getBoundingClientRect();mouseX=e.clientX-t.left,mouseY=e.clientY-t.top,handleInteraction(e)}),canvas.addEventListener("touchstart",e=>{isMouseDown=!0,isDemoMode=!1,e.preventDefault();const t=canvas.getBoundingClientRect(),o=e.touches[0],n=o.clientX-t.left,a=o.clientY-t.top,l=Math.floor(n/(t.width/GRID_SIZE)),i=Math.floor(a/(t.height/GRID_SIZE));paintState=i>=0&&i<GRID_SIZE&&l>=0&&l<GRID_SIZE?-grid[i][l]:1,mouseX=n,mouseY=a,handleInteraction(o)},{passive:!1}),canvas.addEventListener("touchmove",e=>{e.preventDefault();const t=canvas.getBoundingClientRect(),o=e.touches[0];mouseX=o.clientX-t.left,mouseY=o.clientY-t.top,handleInteraction(o)},{passive:!1}),window.addEventListener("touchend",()=>{isMouseDown=!1,paintState=0}),tempSlider.addEventListener("input",e=>{stopDemo(),temperature=parseFloat(e.target.value),tempVal.textContent=temperature.toFixed(2)}),fieldSlider.addEventListener("input",e=>{stopDemo(),externalField=parseFloat(e.target.value),fieldVal.textContent=externalField.toFixed(2)}),speedSlider.addEventListener("input",e=>{stopDemo(),speed=parseInt(e.target.value),speedVal.textContent=speed}),themeToggle.addEventListener("click",()=>{isDarkTheme=!isDarkTheme;const e=themeToggle.querySelector(".sun-icon"),t=themeToggle.querySelector(".moon-icon");isDarkTheme?(app.classList.remove("light-theme"),e.style.display="none",t.style.display="block"):(app.classList.add("light-theme"),e.style.display="block",t.style.display="none")}),restartBtn.addEventListener("click",()=>{resetDemo()}),infoIcon.addEventListener("mouseenter",()=>{infoOverlay.classList.add("visible")}),infoIcon.addEventListener("mouseleave",()=>{infoOverlay.classList.remove("visible")}),infoOverlay.addEventListener("mouseenter",()=>{infoOverlay.classList.add("visible")}),infoOverlay.addEventListener("mouseleave",()=>{infoOverlay.classList.remove("visible")}),initGrid(),loop();</script> <p>I’ve gotten decent results out of the last crop of OpenAI and Anthropic models, but the Chrome browser extension for retrieving the DOM really helped too. It’s a great feature, and I expect Cursor to have something similar soon. I think some of the other UI features like showing subtasks and intermediate steps were a little unnecessary. Overall, great work by the former Windsurf team and the other G staff!</p>]]></content><author><name></name></author><summary type="html"><![CDATA[Testing Google's new IDE and top model on a ferromagnetic simulation]]></summary></entry><entry><title type="html">A tutorial on kinematics and optimization using NVIDIA Warp</title><link href="https://ckrapu.github.io/blog/2025/kinematics-and-optimization-warp/" rel="alternate" type="text/html" title="A tutorial on kinematics and optimization using NVIDIA Warp"/><published>2025-08-01T00:00:00+00:00</published><updated>2025-08-01T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2025/kinematics-and-optimization-warp</id><content type="html" xml:base="https://ckrapu.github.io/blog/2025/kinematics-and-optimization-warp/"><![CDATA[<p><a href="https://github.com/NVIDIA/warp">Warp</a> is a promising new open-source software library for performing scientific computing and optimization on GPUs with a Python-native workflow.</p> <p>This notebook provides a simple first example of using Warp to simulate a particle moving along a one-dimensional flow field. We’ll then optimize the direction and magnitude of the flow field to push the particle towards a pair of target locations specified at two distinct timesteps. It doesn’t assume much background in fluid mechanics or machine learning, and shows the basic workflow for manipulating the parameters of a simulation using Warp.</p> <h2 id="simulating-a-particles-movement-in-1d">Simulating a particle’s movement in 1D</h2> <p>We start by importing packages and defining the simulation parameters. For the purposes of this project, Warp can be considered as a drop-in replacement for PyTorch or Jax - it includes both the tensor libraries and automatic differentiation logic.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">from</span> <span class="n">dataclasses</span> <span class="kn">import</span> <span class="n">dataclass</span>
<span class="kn">from</span> <span class="n">IPython.display</span> <span class="kn">import</span> <span class="n">HTML</span>
<span class="kn">from</span> <span class="n">matplotlib.animation</span> <span class="kn">import</span> <span class="n">FuncAnimation</span>
<span class="kn">from</span> <span class="n">tqdm</span> <span class="kn">import</span> <span class="n">tqdm</span>
<span class="kn">from</span> <span class="n">IPython.display</span> <span class="kn">import</span> <span class="n">Image</span>

<span class="kn">import</span> <span class="n">matplotlib</span> <span class="k">as</span> <span class="n">mpl</span>
<span class="kn">import</span> <span class="n">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="n">warp</span> <span class="k">as</span> <span class="n">wp</span>

<span class="n">plt</span><span class="p">.</span><span class="n">style</span><span class="p">.</span><span class="nf">use</span><span class="p">(</span><span class="sh">'</span><span class="s">dark_background</span><span class="sh">'</span><span class="p">)</span>

<span class="n">plot_bg_color</span> <span class="o">=</span> <span class="sh">"</span><span class="s">#1c1c1d</span><span class="sh">"</span>

<span class="n">mpl</span><span class="p">.</span><span class="n">rcParams</span><span class="p">[</span><span class="sh">'</span><span class="s">figure.facecolor</span><span class="sh">'</span><span class="p">]</span> <span class="o">=</span> <span class="n">plot_bg_color</span>
<span class="n">mpl</span><span class="p">.</span><span class="n">rcParams</span><span class="p">[</span><span class="sh">'</span><span class="s">axes.facecolor</span><span class="sh">'</span><span class="p">]</span> <span class="o">=</span> <span class="n">plot_bg_color</span>

<span class="n">np</span><span class="p">.</span><span class="n">random</span><span class="p">.</span><span class="nf">seed</span><span class="p">(</span><span class="mi">8271991</span><span class="p">)</span>

<span class="n">device</span> <span class="o">=</span> <span class="sh">"</span><span class="s">cuda</span><span class="sh">"</span>

<span class="c1"># Instantiate a variable to force Warp to start up
# and show its version + device
</span><span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">([</span><span class="mi">0</span><span class="p">],</span> <span class="n">dtype</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>

</code></pre></div></div> <div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>array(shape=(1,), dtype=float32)
</code></pre></div></div> <p>Here, we specify that the domain will be one meter long and divided into 100 grid cells for numerical simulation. The entire simulation will run for 30 seconds in increments of 100 milliseconds.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nd">@dataclass</span>
<span class="k">class</span> <span class="nc">ParticleExampleConfig</span><span class="p">:</span>
    <span class="n">length_m</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">1.0</span>
    <span class="n">nx</span><span class="p">:</span> <span class="nb">int</span> <span class="o">=</span> <span class="mi">100</span>
    <span class="n">total_time_s</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">30.</span>
    <span class="n">dt</span><span class="p">:</span> <span class="nb">float</span> <span class="o">=</span> <span class="mf">0.1</span>
    <span class="n">start_location_m</span> <span class="o">=</span> <span class="mf">0.1</span>

<span class="n">particle_cfg</span> <span class="o">=</span> <span class="nc">ParticleExampleConfig</span><span class="p">()</span>
</code></pre></div></div> <p>Next, we define the logic for tracking the location of the particle as it is moved around via the velocity field.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">simulate_particle</span><span class="p">(</span><span class="n">particle_cfg</span><span class="p">:</span> <span class="n">ParticleExampleConfig</span><span class="p">,</span> <span class="n">u0</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="n">array</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">wp</span><span class="p">.</span><span class="n">array</span><span class="p">:</span>
      <span class="sh">'''</span><span class="s">
      Runs a simple 1D simulation, taking the flow field (a 1D vector in this case) and produces
      an array with the location of the particle at each timestep.
      </span><span class="sh">'''</span>
      <span class="n">n_steps</span> <span class="o">=</span> <span class="nf">int</span><span class="p">(</span><span class="n">particle_cfg</span><span class="p">.</span><span class="n">total_time_s</span> <span class="o">/</span> <span class="n">particle_cfg</span><span class="p">.</span><span class="n">dt</span><span class="p">)</span>
      <span class="n">positions</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">zeros</span><span class="p">(</span><span class="n">n_steps</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">requires_grad</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>

      <span class="nd">@wp.kernel</span>
      <span class="k">def</span> <span class="nf">integrate_trajectory</span><span class="p">(</span>
          <span class="n">positions</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
          <span class="n">u0</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
          <span class="n">start_pos</span><span class="p">:</span> <span class="nb">float</span><span class="p">,</span>
          <span class="n">dt</span><span class="p">:</span> <span class="nb">float</span><span class="p">,</span>
          <span class="n">dx</span><span class="p">:</span> <span class="nb">float</span><span class="p">,</span>
          <span class="n">nx</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span>
          <span class="n">n_steps</span><span class="p">:</span> <span class="nb">int</span>
      <span class="p">):</span>
          <span class="n">tid</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">tid</span><span class="p">()</span>

          <span class="k">if</span> <span class="n">tid</span> <span class="o">==</span> <span class="mi">0</span><span class="p">:</span>
              <span class="n">positions</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span> <span class="o">=</span> <span class="n">start_pos</span>
              <span class="n">velocity</span> <span class="o">=</span> <span class="nf">float</span><span class="p">(</span><span class="mf">0.0</span><span class="p">)</span>

              <span class="k">for</span> <span class="n">step</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n_steps</span><span class="p">):</span>
                  <span class="n">current_pos</span> <span class="o">=</span> <span class="n">positions</span><span class="p">[</span><span class="n">step</span><span class="p">]</span>

                  <span class="n">cell_idx</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">int32</span><span class="p">(</span><span class="n">current_pos</span> <span class="o">/</span> <span class="n">dx</span><span class="p">)</span>
                  <span class="n">cell_idx</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">clamp</span><span class="p">(</span><span class="n">cell_idx</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">nx</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)</span>

                  <span class="n">flow_velocity</span> <span class="o">=</span> <span class="n">u0</span><span class="p">[</span><span class="n">cell_idx</span><span class="p">]</span>
                  <span class="n">force</span> <span class="o">=</span> <span class="n">flow_velocity</span> <span class="o">-</span> <span class="n">velocity</span>

                  <span class="n">velocity</span> <span class="o">=</span> <span class="n">velocity</span> <span class="o">+</span> <span class="n">force</span> <span class="o">*</span> <span class="n">dt</span>

                  <span class="n">new_pos</span> <span class="o">=</span> <span class="n">current_pos</span> <span class="o">+</span> <span class="n">velocity</span> <span class="o">*</span> <span class="n">dt</span>
                  <span class="n">new_pos</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">clamp</span><span class="p">(</span><span class="n">new_pos</span><span class="p">,</span> <span class="mf">0.0</span><span class="p">,</span> <span class="n">dx</span> <span class="o">*</span> <span class="nf">float</span><span class="p">(</span><span class="n">nx</span><span class="p">))</span>

                  <span class="n">positions</span><span class="p">[</span><span class="n">step</span> <span class="o">+</span> <span class="mi">1</span><span class="p">]</span> <span class="o">=</span> <span class="n">new_pos</span>

      <span class="n">dx</span> <span class="o">=</span> <span class="n">particle_cfg</span><span class="p">.</span><span class="n">length_m</span> <span class="o">/</span> <span class="n">particle_cfg</span><span class="p">.</span><span class="n">nx</span>

      <span class="n">wp</span><span class="p">.</span><span class="nf">launch</span><span class="p">(</span>
          <span class="n">integrate_trajectory</span><span class="p">,</span>
          <span class="n">dim</span><span class="o">=</span><span class="mi">1</span><span class="p">,</span>
          <span class="n">inputs</span><span class="o">=</span><span class="p">[</span>
              <span class="n">positions</span><span class="p">,</span>
              <span class="n">u0</span><span class="p">,</span>
              <span class="n">particle_cfg</span><span class="p">.</span><span class="n">start_location_m</span><span class="p">,</span>
              <span class="n">particle_cfg</span><span class="p">.</span><span class="n">dt</span><span class="p">,</span>
              <span class="n">dx</span><span class="p">,</span>
              <span class="n">particle_cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">,</span>
              <span class="n">n_steps</span>
          <span class="p">]</span>
      <span class="p">)</span>

      <span class="k">return</span> <span class="n">positions</span>


</code></pre></div></div> <p>Let’s go ahead and run it, with a velocity field which is randomly initialized but with a bias pointing towards the center of the domain.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">u0_numpy</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="nf">randn</span><span class="p">(</span><span class="n">particle_cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">)</span> <span class="o">+</span> <span class="n">np</span><span class="p">.</span><span class="nf">linspace</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="n">particle_cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">)</span> <span class="o">*</span> <span class="mf">2.5</span>
<span class="n">u0</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">u0_numpy</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="nb">float</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
<span class="n">positions</span> <span class="o">=</span> <span class="nf">simulate_particle</span><span class="p">(</span><span class="n">particle_cfg</span><span class="p">,</span> <span class="n">u0</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="nf">figure</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">6</span><span class="p">,</span> <span class="mi">2</span><span class="p">))</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">xlabel</span><span class="p">(</span><span class="sh">"</span><span class="s">Time (s)</span><span class="sh">"</span><span class="p">),</span> <span class="n">plt</span><span class="p">.</span><span class="nf">ylabel</span><span class="p">(</span><span class="sh">"</span><span class="s">Position (m)</span><span class="sh">"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">title</span><span class="p">(</span><span class="sh">"</span><span class="s">Particle location over time</span><span class="sh">"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">plot</span><span class="p">(</span><span class="n">positions</span><span class="p">.</span><span class="nf">numpy</span><span class="p">())</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">savefig</span><span class="p">(</span><span class="sh">"</span><span class="s">visuals/location_over_time.png</span><span class="sh">"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">close</span><span class="p">()</span>
<span class="nc">Image</span><span class="p">(</span><span class="sh">"</span><span class="s">visuals/location_over_time.png</span><span class="sh">"</span><span class="p">)</span>
</code></pre></div></div> <div style="text-align: center;"> <img src="/images/2025-08-01-kinematics-and-optimization-warp/positions.png" alt="Locations of particle over time" width="600"/> </div> <p>We can also make an animation to provide a dynamic view of the particle’s behavior subject to the 1D velocity field.</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="n">ax</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="nf">subplots</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">8</span><span class="p">,</span> <span class="mi">2</span><span class="p">))</span>

<span class="n">ax</span><span class="p">.</span><span class="nf">set_xlim</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">particle_cfg</span><span class="p">.</span><span class="n">length_m</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_ylim</span><span class="p">(</span><span class="o">-</span><span class="mf">0.5</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_xlabel</span><span class="p">(</span><span class="sh">"</span><span class="s">Position (m)</span><span class="sh">"</span><span class="p">)</span>

<span class="c1"># Turn off y-axis + spines
</span><span class="n">ax</span><span class="p">.</span><span class="n">axes</span><span class="p">.</span><span class="nf">get_yaxis</span><span class="p">().</span><span class="nf">set_visible</span><span class="p">(</span><span class="bp">False</span><span class="p">)</span>
<span class="k">for</span> <span class="n">spine</span> <span class="ow">in</span> <span class="p">[</span><span class="sh">'</span><span class="s">top</span><span class="sh">'</span><span class="p">,</span> <span class="sh">'</span><span class="s">right</span><span class="sh">'</span><span class="p">,</span> <span class="sh">'</span><span class="s">left</span><span class="sh">'</span><span class="p">]:</span>
    <span class="n">ax</span><span class="p">.</span><span class="n">spines</span><span class="p">[</span><span class="n">spine</span><span class="p">].</span><span class="nf">set_visible</span><span class="p">(</span><span class="bp">False</span><span class="p">)</span>

<span class="c1"># Add tight layout to prevent xlabel cutoff
</span><span class="n">plt</span><span class="p">.</span><span class="nf">tight_layout</span><span class="p">()</span>

<span class="n">particle</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">([],</span> <span class="p">[],</span> <span class="n">s</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">cyan</span><span class="sh">'</span><span class="p">)</span>
<span class="n">time_text</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">text</span><span class="p">(</span><span class="mf">0.02</span><span class="p">,</span> <span class="mf">0.85</span><span class="p">,</span> <span class="sh">''</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">ax</span><span class="p">.</span><span class="n">transAxes</span><span class="p">,</span> <span class="n">fontsize</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>

<span class="c1"># Get position data
</span><span class="n">positions_np</span> <span class="o">=</span> <span class="n">positions</span><span class="p">.</span><span class="nf">numpy</span><span class="p">()</span>
<span class="n">n_frames</span> <span class="o">=</span> <span class="nf">len</span><span class="p">(</span><span class="n">positions_np</span><span class="p">)</span>

<span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
    <span class="n">particle</span><span class="p">.</span><span class="nf">set_offsets</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">empty</span><span class="p">((</span><span class="mi">0</span><span class="p">,</span> <span class="mi">2</span><span class="p">)))</span>
    <span class="n">time_text</span><span class="p">.</span><span class="nf">set_text</span><span class="p">(</span><span class="sh">''</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">particle</span><span class="p">,</span> <span class="n">time_text</span>

<span class="k">def</span> <span class="nf">animate</span><span class="p">(</span><span class="n">frame</span><span class="p">):</span>
    <span class="n">particle</span><span class="p">.</span><span class="nf">set_offsets</span><span class="p">([[</span><span class="n">positions_np</span><span class="p">[</span><span class="n">frame</span><span class="p">],</span> <span class="mi">0</span><span class="p">]])</span>
    <span class="n">time_text</span><span class="p">.</span><span class="nf">set_text</span><span class="p">(</span><span class="sa">f</span><span class="sh">'</span><span class="s">Time: </span><span class="si">{</span><span class="n">frame</span> <span class="o">*</span> <span class="n">particle_cfg</span><span class="p">.</span><span class="n">dt</span><span class="si">:</span><span class="p">.</span><span class="mi">1</span><span class="n">f</span><span class="si">}</span><span class="s">s</span><span class="sh">'</span><span class="p">)</span>
    
    <span class="k">return</span> <span class="n">particle</span><span class="p">,</span> <span class="n">time_text</span>

<span class="n">anim</span> <span class="o">=</span> <span class="nc">FuncAnimation</span><span class="p">(</span><span class="n">fig</span><span class="p">,</span> <span class="n">animate</span><span class="p">,</span> <span class="n">init_func</span><span class="o">=</span><span class="n">init</span><span class="p">,</span> 
                    <span class="n">frames</span><span class="o">=</span><span class="n">n_frames</span><span class="p">,</span> <span class="n">interval</span><span class="o">=</span><span class="mi">40</span><span class="p">,</span> 
                    <span class="n">blit</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">repeat</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">anim</span><span class="p">.</span><span class="nf">save</span><span class="p">(</span><span class="sh">'</span><span class="s">visuals/01-position-animation.gif</span><span class="sh">'</span><span class="p">,</span> <span class="n">writer</span><span class="o">=</span><span class="sh">'</span><span class="s">pillow</span><span class="sh">'</span><span class="p">,</span> <span class="n">fps</span><span class="o">=</span><span class="mi">25</span><span class="p">,</span> <span class="n">dpi</span><span class="o">=</span><span class="mi">80</span><span class="p">)</span>
<span class="nc">HTML</span><span class="p">(</span><span class="n">anim</span><span class="p">.</span><span class="nf">to_jshtml</span><span class="p">())</span>
</code></pre></div></div> <div style="text-align: center;"> <img src="/images/2025-08-01-kinematics-and-optimization-warp/01-position-animation.gif" alt="Locations of particle over time" width="800"/> </div> <p>Great! We’ve got a simple model for a particle moving subject to a velocity field. Let’s move onto the meat of this tutorial - manipulating the model parameters to do what we want.</p> <h2 id="optimizing-the-velocity-field-via-automatic-differentiation">Optimizing the velocity field via automatic differentiation</h2> <p>We’ve passed an initial test of basic functionality for our model; we assume a simple velocity field and let the particle be moved around subject to simple kinematics.</p> <p>Next, we want to show how to optimize an arbitrary model quantity to perform an interesting task. Our new objective is to identify a setting of <code class="language-plaintext highlighter-rouge">u0</code> that pushes the particle to a target location at a given timestep. Our optimization strategy is as follows:</p> <p>For each iteration:</p> <ul> <li>We take the initial velocity field <code class="language-plaintext highlighter-rouge">u0</code> which is just a 1D array for this problem, and run the forward simulation. We use the <code class="language-plaintext highlighter-rouge">tape</code> to record the execution trace</li> <li>Next, we use reverse-mode autodiff in Warp to get <code class="language-plaintext highlighter-rouge">grad</code> which represents \(\nabla_{u_0} L(x_1,...,x_T)\) where \(x_1,...,x_T\) denotes the particle position at each timestep.</li> <li>We apply an update to <code class="language-plaintext highlighter-rouge">u0</code> per the update rule \(u_0^{(i+1)} = u_0^{(i)} - \eta ^{(i)} \cdot \nabla_{u_0^{i}} \left[ L(x_1^{(i)},...,x_T^{(i)}) \right]\) with \(u_0^{(i+1)}\) as the value of the velocity field at the \(i+1\)-th iteration and \(x_1^{(i-1)}\) denoting the value of the position of the particle at timestep 1 in optimization iteration \(i-1\). Our loss function \(L\) is simply the sum of the squared distances between the position locations and the target locations at the provided timesteps.</li> <li>We decrease the update size with \(\eta ^{(i+1)} = \eta ^{(i)} \times 0.99\)</li> </ul> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="sh">'''</span><span class="s">
The configs for this version are mostly the same; we just want
to include a few validation checks.
</span><span class="sh">'''</span>

<span class="nd">@dataclass</span>
<span class="k">class</span> <span class="nc">ParticleOptimizationConfig</span><span class="p">(</span><span class="n">ParticleExampleConfig</span><span class="p">):</span>
    <span class="n">target_times_s</span><span class="p">:</span> <span class="nb">tuple</span><span class="p">[</span><span class="nb">float</span><span class="p">]</span> <span class="o">=</span> <span class="p">(</span><span class="mf">1.</span><span class="p">,</span> <span class="mf">15.</span><span class="p">)</span>
    <span class="n">target_locations_m</span><span class="p">:</span> <span class="nb">tuple</span><span class="p">[</span><span class="nb">float</span><span class="p">]</span> <span class="o">=</span> <span class="p">(</span><span class="mf">0.2</span><span class="p">,</span> <span class="mf">0.75</span><span class="p">)</span>
    <span class="n">step_size</span> <span class="o">=</span> <span class="mf">0.1</span>
    <span class="n">n_iterations</span> <span class="o">=</span> <span class="mi">25</span>
    <span class="n">decay_rate</span> <span class="o">=</span> <span class="mf">0.99</span> <span class="c1"># Decreases the step size by this factor each iteration
</span>
    <span class="k">def</span> <span class="nf">__post_init__</span><span class="p">(</span><span class="n">self</span><span class="p">):</span>
        <span class="k">if</span> <span class="nf">len</span><span class="p">(</span><span class="n">self</span><span class="p">.</span><span class="n">target_times_s</span><span class="p">)</span> <span class="o">!=</span> <span class="nf">len</span><span class="p">(</span><span class="n">self</span><span class="p">.</span><span class="n">target_locations_m</span><span class="p">):</span>
            <span class="k">raise</span> <span class="nc">ValueError</span><span class="p">(</span><span class="sh">"</span><span class="s">target_times_s and target_locations_m must have the same length.</span><span class="sh">"</span><span class="p">)</span>
        <span class="k">if</span> <span class="nf">any</span><span class="p">(</span><span class="n">t</span> <span class="o">&lt;</span> <span class="mi">0</span> <span class="ow">or</span> <span class="n">t</span> <span class="o">&gt;</span> <span class="n">self</span><span class="p">.</span><span class="n">total_time_s</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">self</span><span class="p">.</span><span class="n">target_times_s</span><span class="p">):</span>
            <span class="k">raise</span> <span class="nc">ValueError</span><span class="p">(</span><span class="sh">"</span><span class="s">All target times must be within the total simulation time.</span><span class="sh">"</span><span class="p">)</span>
        <span class="k">if</span> <span class="nf">any</span><span class="p">(</span><span class="n">l</span> <span class="o">&lt;</span> <span class="mi">0</span> <span class="ow">or</span> <span class="n">l</span> <span class="o">&gt;</span> <span class="n">self</span><span class="p">.</span><span class="n">length_m</span> <span class="k">for</span> <span class="n">l</span> <span class="ow">in</span> <span class="n">self</span><span class="p">.</span><span class="n">target_locations_m</span><span class="p">):</span>
            <span class="k">raise</span> <span class="nc">ValueError</span><span class="p">(</span><span class="sh">"</span><span class="s">All target locations must be within the simulation length.</span><span class="sh">"</span><span class="p">)</span>
        
<span class="n">opt_cfg</span> <span class="o">=</span> <span class="nc">ParticleOptimizationConfig</span><span class="p">()</span>
    
</code></pre></div></div> <h3 id="applying-autodiff-and-gradient-descent">Applying autodiff and gradient descent</h3> <p>The above configuration specifies that at 1 and 15 seconds in simulation time we would like the particle to be close to locations of 0.2 m and 0.75 m from the left-hand side of the domain. We’ll use Warp’s implementation of automatic differentiation with gradient descent and a simple decreasing stepsize schedule to find an optimal <code class="language-plaintext highlighter-rouge">u0</code> such that the particle respects the previous conditions.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">u0_history</span> <span class="o">=</span> <span class="p">[]</span>
<span class="n">u0</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</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="nf">randn</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">)</span> <span class="o">*</span> <span class="mf">0.1</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">,</span> <span class="n">requires_grad</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>

<span class="n">target_indices</span> <span class="o">=</span> <span class="p">[</span><span class="nf">int</span><span class="p">(</span><span class="n">t</span> <span class="o">/</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">dt</span><span class="p">)</span> <span class="k">for</span> <span class="n">t</span> <span class="ow">in</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">target_times_s</span><span class="p">]</span>
<span class="n">target_positions</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">.</span><span class="n">target_locations_m</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>

<span class="n">loss_history</span> <span class="o">=</span> <span class="p">[]</span>
<span class="n">grad_norm_history</span> <span class="o">=</span> <span class="p">[]</span>

<span class="n">pbar</span> <span class="o">=</span> <span class="nf">tqdm</span><span class="p">(</span><span class="nf">range</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">.</span><span class="n">n_iterations</span><span class="p">),</span> <span class="n">desc</span><span class="o">=</span><span class="sh">""</span><span class="p">)</span>
<span class="n">step_size</span> <span class="o">=</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">step_size</span>

<span class="k">for</span> <span class="n">iteration</span> <span class="ow">in</span> <span class="n">pbar</span><span class="p">:</span>
    <span class="n">u0_history</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">u0</span><span class="p">.</span><span class="nf">numpy</span><span class="p">().</span><span class="nf">copy</span><span class="p">())</span>
    
    <span class="n">tape</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nc">Tape</span><span class="p">()</span>
    
    <span class="k">with</span> <span class="n">tape</span><span class="p">:</span>
        <span class="n">positions</span> <span class="o">=</span> <span class="nf">simulate_particle</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">,</span> <span class="n">u0</span><span class="p">)</span>
        
        <span class="n">loss</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">zeros</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">,</span> <span class="n">requires_grad</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
        
        <span class="nd">@wp.kernel</span>
        <span class="k">def</span> <span class="nf">compute_l2_loss</span><span class="p">(</span>
            <span class="n">positions</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
            <span class="n">target_positions</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">),</span>
            <span class="n">target_indices</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">int32</span><span class="p">),</span>
            <span class="n">loss</span><span class="p">:</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
        <span class="p">):</span>
            <span class="n">tid</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">tid</span><span class="p">()</span>
            <span class="k">if</span> <span class="n">tid</span> <span class="o">&lt;=</span> <span class="n">target_indices</span><span class="p">.</span><span class="n">shape</span><span class="p">[</span><span class="mi">0</span><span class="p">]:</span>
                <span class="n">idx</span> <span class="o">=</span> <span class="n">target_indices</span><span class="p">[</span><span class="n">tid</span><span class="p">]</span>
                <span class="n">diff</span> <span class="o">=</span> <span class="n">positions</span><span class="p">[</span><span class="n">idx</span><span class="p">]</span> <span class="o">-</span> <span class="n">target_positions</span><span class="p">[</span><span class="n">tid</span><span class="p">]</span>
                <span class="n">wp</span><span class="p">.</span><span class="nf">atomic_add</span><span class="p">(</span><span class="n">loss</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">diff</span> <span class="o">*</span> <span class="n">diff</span><span class="p">)</span>
        
        <span class="n">target_indices_array</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">target_indices</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">int32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
        <span class="n">wp</span><span class="p">.</span><span class="nf">launch</span><span class="p">(</span><span class="n">compute_l2_loss</span><span class="p">,</span> <span class="n">dim</span><span class="o">=</span><span class="nf">len</span><span class="p">(</span><span class="n">target_indices</span><span class="p">),</span> 
                 <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">positions</span><span class="p">,</span> <span class="n">target_positions</span><span class="p">,</span> <span class="n">target_indices_array</span><span class="p">,</span> <span class="n">loss</span><span class="p">])</span>
    
    <span class="n">tape</span><span class="p">.</span><span class="nf">backward</span><span class="p">(</span><span class="n">loss</span><span class="p">)</span>
    
    <span class="n">loss_val</span> <span class="o">=</span> <span class="n">loss</span><span class="p">.</span><span class="nf">numpy</span><span class="p">()[</span><span class="mi">0</span><span class="p">]</span>
    <span class="n">grad</span> <span class="o">=</span> <span class="n">tape</span><span class="p">.</span><span class="n">gradients</span><span class="p">[</span><span class="n">u0</span><span class="p">].</span><span class="nf">numpy</span><span class="p">()</span>
    <span class="n">grad_norm</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="nf">norm</span><span class="p">(</span><span class="n">grad</span><span class="p">)</span>
    
    <span class="n">loss_history</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">loss_val</span><span class="p">)</span>
    <span class="n">grad_norm_history</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">grad_norm</span><span class="p">)</span>
    
    <span class="n">pbar</span><span class="p">.</span><span class="nf">set_description</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Iter </span><span class="si">{</span><span class="n">iteration</span><span class="si">}</span><span class="s">: Loss = </span><span class="si">{</span><span class="n">loss_val</span><span class="si">:</span><span class="p">.</span><span class="mi">6</span><span class="n">f</span><span class="si">}</span><span class="s">, Grad norm = </span><span class="si">{</span><span class="n">grad_norm</span><span class="si">:</span><span class="p">.</span><span class="mi">6</span><span class="n">f</span><span class="si">}</span><span class="sh">"</span><span class="p">)</span>
    
    <span class="n">u0_np</span> <span class="o">=</span> <span class="n">u0</span><span class="p">.</span><span class="nf">numpy</span><span class="p">()</span>
    <span class="n">u0_np</span> <span class="o">-=</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">step_size</span> <span class="o">*</span> <span class="n">grad</span>
    <span class="n">u0</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">u0_np</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">,</span> <span class="n">requires_grad</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
    
    <span class="n">step_size</span> <span class="o">*=</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">decay_rate</span>
</code></pre></div></div> <div class="language-plaintext highlighter-rouge"><div class="highlight"><pre class="highlight"><code>Iter 24: Loss = 0.000754, Grad norm = 0.028800: 100%|██████████| 25/25 [00:00&lt;00:00, 156.56it/s]

Module __main__ 331110c load on device 'cuda:0' took 0.20 ms  (cached)
</code></pre></div></div> <h3 id="animating-the-optimized-trajectory">Animating the optimized trajectory</h3> <p>Our optimization converges to near-zero loss quite quickly. Let’s see what the solution looks like:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">optimized_positions</span> <span class="o">=</span> <span class="nf">simulate_particle</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">,</span> <span class="n">u0</span><span class="p">)</span>

<span class="n">fig</span><span class="p">,</span> <span class="n">ax</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="nf">subplots</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">8</span><span class="p">,</span> <span class="mi">2</span><span class="p">))</span>

<span class="n">ax</span><span class="p">.</span><span class="nf">set_xlim</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">length_m</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_ylim</span><span class="p">(</span><span class="o">-</span><span class="mf">0.5</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_xlabel</span><span class="p">(</span><span class="sh">"</span><span class="s">Position (m)</span><span class="sh">"</span><span class="p">)</span>

<span class="n">ax</span><span class="p">.</span><span class="n">axes</span><span class="p">.</span><span class="nf">get_yaxis</span><span class="p">().</span><span class="nf">set_visible</span><span class="p">(</span><span class="bp">False</span><span class="p">)</span>
<span class="k">for</span> <span class="n">spine</span> <span class="ow">in</span> <span class="p">[</span><span class="sh">'</span><span class="s">top</span><span class="sh">'</span><span class="p">,</span> <span class="sh">'</span><span class="s">right</span><span class="sh">'</span><span class="p">,</span> <span class="sh">'</span><span class="s">left</span><span class="sh">'</span><span class="p">]:</span>
    <span class="n">ax</span><span class="p">.</span><span class="n">spines</span><span class="p">[</span><span class="n">spine</span><span class="p">].</span><span class="nf">set_visible</span><span class="p">(</span><span class="bp">False</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="nf">tight_layout</span><span class="p">()</span>

<span class="n">particle</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">([],</span> <span class="p">[],</span> <span class="n">s</span><span class="o">=</span><span class="mi">100</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">cyan</span><span class="sh">'</span><span class="p">)</span>
<span class="n">time_text</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">text</span><span class="p">(</span><span class="mf">0.02</span><span class="p">,</span> <span class="mf">0.85</span><span class="p">,</span> <span class="sh">''</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">ax</span><span class="p">.</span><span class="n">transAxes</span><span class="p">,</span> <span class="n">fontsize</span><span class="o">=</span><span class="mi">10</span><span class="p">)</span>

<span class="k">for</span> <span class="n">t</span><span class="p">,</span> <span class="n">loc</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">.</span><span class="n">target_times_s</span><span class="p">,</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">target_locations_m</span><span class="p">):</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">axvline</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">loc</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">cyan</span><span class="sh">'</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="sh">'</span><span class="s">--</span><span class="sh">'</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="nf">text</span><span class="p">(</span><span class="n">loc</span><span class="p">,</span> <span class="mf">0.3</span><span class="p">,</span> <span class="sa">f</span><span class="sh">'</span><span class="s">t=</span><span class="si">{</span><span class="n">t</span><span class="si">}</span><span class="s">s</span><span class="sh">'</span><span class="p">,</span> <span class="n">ha</span><span class="o">=</span><span class="sh">'</span><span class="s">center</span><span class="sh">'</span><span class="p">,</span> <span class="n">fontsize</span><span class="o">=</span><span class="mi">8</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">cyan</span><span class="sh">'</span><span class="p">)</span>

<span class="n">positions_np</span> <span class="o">=</span> <span class="n">optimized_positions</span><span class="p">.</span><span class="nf">numpy</span><span class="p">()</span>
<span class="n">n_frames</span> <span class="o">=</span> <span class="nf">len</span><span class="p">(</span><span class="n">positions_np</span><span class="p">)</span>

<span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
    <span class="n">particle</span><span class="p">.</span><span class="nf">set_offsets</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">empty</span><span class="p">((</span><span class="mi">0</span><span class="p">,</span> <span class="mi">2</span><span class="p">)))</span>
    <span class="n">time_text</span><span class="p">.</span><span class="nf">set_text</span><span class="p">(</span><span class="sh">''</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">particle</span><span class="p">,</span> <span class="n">time_text</span>

<span class="k">def</span> <span class="nf">animate</span><span class="p">(</span><span class="n">frame</span><span class="p">):</span>
    <span class="n">particle</span><span class="p">.</span><span class="nf">set_offsets</span><span class="p">([[</span><span class="n">positions_np</span><span class="p">[</span><span class="n">frame</span><span class="p">],</span> <span class="mi">0</span><span class="p">]])</span>
    <span class="n">time_text</span><span class="p">.</span><span class="nf">set_text</span><span class="p">(</span><span class="sa">f</span><span class="sh">'</span><span class="s">Time: </span><span class="si">{</span><span class="n">frame</span> <span class="o">*</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">dt</span><span class="si">:</span><span class="p">.</span><span class="mi">1</span><span class="n">f</span><span class="si">}</span><span class="s">s</span><span class="sh">'</span><span class="p">)</span>
    
    <span class="k">return</span> <span class="n">particle</span><span class="p">,</span> <span class="n">time_text</span>
    
<span class="n">anim</span> <span class="o">=</span> <span class="nc">FuncAnimation</span><span class="p">(</span><span class="n">fig</span><span class="p">,</span> <span class="n">animate</span><span class="p">,</span> <span class="n">init_func</span><span class="o">=</span><span class="n">init</span><span class="p">,</span> 
                    <span class="n">frames</span><span class="o">=</span><span class="n">n_frames</span><span class="p">,</span> <span class="n">interval</span><span class="o">=</span><span class="mi">40</span><span class="p">,</span> 
                    <span class="n">blit</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">repeat</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">anim</span><span class="p">.</span><span class="nf">save</span><span class="p">(</span><span class="sh">'</span><span class="s">visuals/01-optimize-animation.gif</span><span class="sh">'</span><span class="p">,</span> <span class="n">writer</span><span class="o">=</span><span class="sh">'</span><span class="s">pillow</span><span class="sh">'</span><span class="p">,</span> <span class="n">fps</span><span class="o">=</span><span class="mi">25</span><span class="p">,</span> <span class="n">dpi</span><span class="o">=</span><span class="mi">80</span><span class="p">)</span>
<span class="nc">HTML</span><span class="p">(</span><span class="n">anim</span><span class="p">.</span><span class="nf">to_jshtml</span><span class="p">())</span>
</code></pre></div></div> <div style="text-align: center;"> <img src="/images/2025-08-01-kinematics-and-optimization-warp/01-optimize-animation.gif" alt="Results from optimization of simulation parameters" width="800"/> </div> <p>To understand how the solution evolves over iterations, let’s animate the particle movement for several values of <code class="language-plaintext highlighter-rouge">u0</code> obtained at the start, middle, and end of optimization.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">n_rows</span><span class="p">,</span> <span class="n">n_cols</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">6</span>

<span class="n">n_subplots</span> <span class="o">=</span> <span class="n">n_rows</span> <span class="o">*</span> <span class="n">n_cols</span>
<span class="n">step</span> <span class="o">=</span> <span class="nf">max</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">n_iterations</span> <span class="o">//</span> <span class="n">n_subplots</span><span class="p">)</span>
<span class="n">selected_iterations</span> <span class="o">=</span> <span class="nf">list</span><span class="p">(</span><span class="nf">range</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">n_iterations</span><span class="p">,</span> <span class="n">step</span><span class="p">))[:</span><span class="n">n_subplots</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="nf">subplots</span><span class="p">(</span><span class="n">n_rows</span><span class="p">,</span> <span class="n">n_cols</span><span class="p">,</span> <span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">16</span><span class="p">,</span> <span class="n">n_rows</span><span class="o">*</span><span class="mf">1.5</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="nf">flatten</span><span class="p">()</span>

<span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">ax</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">(</span><span class="n">axes</span><span class="p">):</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">set_xlim</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">length_m</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">set_ylim</span><span class="p">(</span><span class="o">-</span><span class="mf">0.5</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">)</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">set_xticks</span><span class="p">([</span><span class="mi">0</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">,</span> <span class="mf">1.0</span><span class="p">])</span>
    <span class="n">ax</span><span class="p">.</span><span class="nf">tick_params</span><span class="p">(</span><span class="n">labelsize</span><span class="o">=</span><span class="mi">8</span><span class="p">)</span>

    <span class="n">ax</span><span class="p">.</span><span class="nf">axis</span><span class="p">(</span><span class="sh">'</span><span class="s">off</span><span class="sh">'</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="nf">tight_layout</span><span class="p">()</span>

<span class="n">particles</span> <span class="o">=</span> <span class="p">[]</span>
<span class="n">titles</span> <span class="o">=</span> <span class="p">[]</span>

<span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">ax</span><span class="p">,</span> <span class="n">iter_idx</span><span class="p">)</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">(</span><span class="nf">zip</span><span class="p">(</span><span class="n">axes</span><span class="p">,</span> <span class="n">selected_iterations</span><span class="p">)):</span>
    <span class="n">particle</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">scatter</span><span class="p">([],</span> <span class="p">[],</span> <span class="n">s</span><span class="o">=</span><span class="mi">50</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">cyan</span><span class="sh">'</span><span class="p">,</span> <span class="n">edgecolor</span><span class="o">=</span><span class="sh">'</span><span class="s">w</span><span class="sh">'</span><span class="p">)</span>
    <span class="n">particles</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">particle</span><span class="p">)</span>
    
    
    <span class="n">title</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">text</span><span class="p">(</span><span class="mf">0.05</span><span class="p">,</span> <span class="mf">0.9</span><span class="p">,</span> <span class="sa">f</span><span class="sh">'</span><span class="s">Iter. </span><span class="si">{</span><span class="n">iter_idx</span><span class="si">}</span><span class="sh">'</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">ax</span><span class="p">.</span><span class="n">transAxes</span><span class="p">,</span> 
                   <span class="n">ha</span><span class="o">=</span><span class="sh">'</span><span class="s">center</span><span class="sh">'</span><span class="p">,</span> <span class="n">fontsize</span><span class="o">=</span><span class="mi">11</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">white</span><span class="sh">'</span><span class="p">)</span>
    <span class="n">titles</span><span class="p">.</span><span class="nf">append</span><span class="p">(</span><span class="n">title</span><span class="p">)</span>
    
    <span class="k">for</span> <span class="n">t</span><span class="p">,</span> <span class="n">loc</span> <span class="ow">in</span> <span class="nf">zip</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">.</span><span class="n">target_times_s</span><span class="p">,</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">target_locations_m</span><span class="p">):</span>
        <span class="n">ax</span><span class="p">.</span><span class="nf">axvline</span><span class="p">(</span><span class="n">x</span><span class="o">=</span><span class="n">loc</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">cyan</span><span class="sh">'</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="sh">'</span><span class="s">--</span><span class="sh">'</span><span class="p">,</span> <span class="n">alpha</span><span class="o">=</span><span class="mf">0.2</span><span class="p">,</span> <span class="n">linewidth</span><span class="o">=</span><span class="mf">3.0</span><span class="p">)</span>

<span class="k">def</span> <span class="nf">init</span><span class="p">():</span>
    <span class="k">for</span> <span class="n">particle</span> <span class="ow">in</span> <span class="n">particles</span><span class="p">:</span>
        <span class="n">particle</span><span class="p">.</span><span class="nf">set_offsets</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">empty</span><span class="p">((</span><span class="mi">0</span><span class="p">,</span> <span class="mi">2</span><span class="p">)))</span>
    <span class="k">return</span> <span class="n">particles</span>

<span class="k">def</span> <span class="nf">animate</span><span class="p">(</span><span class="n">frame</span><span class="p">):</span>
    <span class="n">frame</span> <span class="o">=</span> <span class="n">frame</span> <span class="o">*</span> <span class="mi">2</span>
    
    <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="p">(</span><span class="n">particle</span><span class="p">,</span> <span class="n">iter_idx</span><span class="p">)</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">(</span><span class="nf">zip</span><span class="p">(</span><span class="n">particles</span><span class="p">,</span> <span class="n">selected_iterations</span><span class="p">)):</span>
        <span class="n">u0_iter</span> <span class="o">=</span> <span class="n">wp</span><span class="p">.</span><span class="nf">array</span><span class="p">(</span><span class="n">u0_history</span><span class="p">[</span><span class="n">iter_idx</span><span class="p">],</span> <span class="n">dtype</span><span class="o">=</span><span class="n">wp</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">device</span><span class="o">=</span><span class="n">device</span><span class="p">)</span>
        <span class="n">positions</span> <span class="o">=</span> <span class="nf">simulate_particle</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">,</span> <span class="n">u0_iter</span><span class="p">)</span>
        <span class="n">positions_np</span> <span class="o">=</span> <span class="n">positions</span><span class="p">.</span><span class="nf">numpy</span><span class="p">()</span>
        
        <span class="k">if</span> <span class="n">frame</span> <span class="o">&lt;</span> <span class="nf">len</span><span class="p">(</span><span class="n">positions_np</span><span class="p">):</span>
            <span class="n">particle</span><span class="p">.</span><span class="nf">set_offsets</span><span class="p">([[</span><span class="n">positions_np</span><span class="p">[</span><span class="n">frame</span><span class="p">],</span> <span class="mi">0</span><span class="p">]])</span>
    
    <span class="k">return</span> <span class="n">particles</span>

<span class="n">n_frames</span> <span class="o">=</span> <span class="nf">int</span><span class="p">(</span><span class="n">opt_cfg</span><span class="p">.</span><span class="n">total_time_s</span> <span class="o">/</span> <span class="n">opt_cfg</span><span class="p">.</span><span class="n">dt</span><span class="p">)</span> <span class="o">+</span> <span class="mi">1</span>
<span class="n">anim</span> <span class="o">=</span> <span class="nc">FuncAnimation</span><span class="p">(</span><span class="n">fig</span><span class="p">,</span> <span class="n">animate</span><span class="p">,</span> <span class="n">init_func</span><span class="o">=</span><span class="n">init</span><span class="p">,</span> 
                    <span class="n">frames</span><span class="o">=</span><span class="n">n_frames</span><span class="o">//</span><span class="mi">2</span><span class="p">,</span> <span class="n">interval</span><span class="o">=</span><span class="mi">80</span><span class="p">,</span> 
                    <span class="n">blit</span><span class="o">=</span><span class="bp">True</span><span class="p">,</span> <span class="n">repeat</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">anim</span><span class="p">.</span><span class="nf">save</span><span class="p">(</span><span class="sh">'</span><span class="s">visuals/01-history-animation.gif</span><span class="sh">'</span><span class="p">,</span> <span class="n">writer</span><span class="o">=</span><span class="sh">'</span><span class="s">pillow</span><span class="sh">'</span><span class="p">,</span> <span class="n">fps</span><span class="o">=</span><span class="mi">25</span><span class="p">,</span> <span class="n">dpi</span><span class="o">=</span><span class="mi">80</span><span class="p">)</span>
<span class="nc">HTML</span><span class="p">(</span><span class="n">anim</span><span class="p">.</span><span class="nf">to_jshtml</span><span class="p">())</span>
</code></pre></div></div> <div style="text-align: center;"> <img src="/images/2025-08-01-kinematics-and-optimization-warp/01-history-animation.gif" alt="Results from optimization of simulation parameters" width="1000"/> </div>]]></content><author><name></name></author><category term="tutorials"/><category term="physics"/><category term="simulation"/><category term="python"/><summary type="html"><![CDATA[Showing how to backpropagate through a simple physics simulation]]></summary></entry><entry><title type="html">What I learned from making 200 different LLMs flip coins</title><link href="https://ckrapu.github.io/blog/2025/what-i-learned-from-flipping-coins/" rel="alternate" type="text/html" title="What I learned from making 200 different LLMs flip coins"/><published>2025-08-01T00:00:00+00:00</published><updated>2025-08-01T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2025/what-i-learned-from-flipping-coins</id><content type="html" xml:base="https://ckrapu.github.io/blog/2025/what-i-learned-from-flipping-coins/"><![CDATA[<p><strong>TL;DR: LLMs finetuned for roleplaying give fair coin flips when prompted to do so. The best models for coding and problem solving give very biased coin flips. Factors like architecture and reasoning mode don’t explain much. Code to reproduce this work is provided on <a href="https://colab.research.google.com/drive/1etkT-ADgP6tN9-0KCXQHFqQM9c1yZO_w">Colab</a>.</strong></p> <p>Language models have no reason to flip coins. They are (broadly speaking) <a href="https://people.csail.mit.edu/renda/llm-sampling-paper">not capable of generating random numbers</a> nor should they be able to do so - there’s no step in the data or training pipelines for LLMs that encourages this behavior.</p> <p>But, what if they could? What would it take to produce an LLM capable of drawing a dozen numbers from a simple distribution with statistically reliable results?</p> <p>The point isn’t to see whether \(\mathcal{M}_{transformers}\) is sufficiently expressive to implement a decent random number generator (<a href="https://arxiv.org/abs/2502.10390">surprise, it is</a>) but rather to see if our best computational model for intelligence can be pushed to do it, and if it offers any benefits.</p> <p>Towards this end, I wanted to see what happened if I used a simple prompt for simulating a coin flip with as many models as possible.</p> <p>I acknowledge this is a dumb question. Here is a chance for you to cease reading and go about your day without any ill feelings from me.</p> <h3 id="flipping-coins-across-all-of-openrouter">Flipping coins across all of OpenRouter</h3> <p>OpenRouter supports over 200 large language models’ inference endpoints. Most of them support text and are tuned for instruction-following and chat. There’s a delightful range of models including the tried-and-true <code class="language-plaintext highlighter-rouge">gpt-4-turbo</code>, the roleplaying models, the foreign-language finetunes like <code class="language-plaintext highlighter-rouge">scb10x/llama-3-typhoon-v1.5x-70b-instruct</code>, and even cutting-edge diffusion models from like <code class="language-plaintext highlighter-rouge">mercury/inception</code>.</p> <p>I wrote a Python notebook to run API calls against each of them using a simple prompt borrowed from a <a href="https://arxiv.org/abs/2506.09998#:~:text=We%20introduce%20Verbalized%20Rejection%20Sampling,reduces%20sampling%20bias%20across%20models.">paper on prompting for rejection sampling</a> for flipping coins:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">prompt</span> <span class="o">=</span> <span class="sh">'''</span><span class="s">
  ** Descriptions: ** Your task is to generate a sample consisting of a single draw from the set {0, 1},
  where both 0 and 1 have a probability of 0.5. Please give your output strictly in the following format:

  &lt;format&gt;
  Explanations: [Your step-by-step analyses and results;
  You DO NOT have access to a computer or a random number generator]
  Output: [Your output MUST be an integer in {0, 1}]
  &lt;/format&gt;

  Please ONLY reply according to this format, don</span><span class="sh">'</span><span class="s">t give me any other words.
  Write the output EXACTLY like </span><span class="sh">"</span><span class="s">Output: 0</span><span class="sh">"</span><span class="s"> or </span><span class="sh">"</span><span class="s">Output: 1</span><span class="sh">"</span><span class="s">. Make sure to include the </span><span class="sh">"</span><span class="s">Output: </span><span class="sh">"</span><span class="s"> prefix.
  </span><span class="sh">'''</span>
</code></pre></div></div> <p>Now, there are some problems with this prompt. Ideally, we’d replace <code class="language-plaintext highlighter-rouge">0</code> and <code class="language-plaintext highlighter-rouge">1</code> with randomly generated strings and marginalize over both orderings. This would be more work, so I skipped it.</p> <p>We would also use structured generation to force outputs into <code class="language-plaintext highlighter-rouge">{0, 1}</code>. I skipped the latter because not all endpoints support structured generation.</p> <p>If you were to run a similar analysis, here’s a list of models I would skip:</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">models_to_skip</span> <span class="o">=</span> <span class="p">(</span>
    <span class="sh">'</span><span class="s">eleutherai/llemma_7b</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">google/gemini-2.5-pro-preview-05-06</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">meta-llama/llama-3.1-405b</span><span class="sh">'</span><span class="p">,</span> <span class="c1"># Had some weird errors
</span>    <span class="sh">'</span><span class="s">meta-llama/llama-guard-2-8b</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">meta-llama/llama-guard-3-8b</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">meta-llama/llama-guard-4-12b</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">mistralai/mistral-7b-instruct-v0.2</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">morph/morph-v3</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">morph/morph-v3-fast</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">morph/morph-v3-large</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">openai/gpt-4o-mini-search-preview-2025-03-11</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">openai/gpt-5</span><span class="sh">'</span><span class="p">,</span>  <span class="c1"># Requires a key to access it via OpenRouter
</span>    <span class="sh">'</span><span class="s">openai/o1-mini</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">openai/o1-pro</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">openai/o3-pro</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">openrouter/auto</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">perplexity/sonar-deep-research</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">qwen/qwq-32b-preview</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">switchpoint/router</span><span class="sh">'</span><span class="p">,</span>
    <span class="sh">'</span><span class="s">x-ai/grok-vision-beta</span><span class="sh">'</span><span class="p">,</span>
<span class="p">)</span>
</code></pre></div></div> <p>Some of these models cost too much to run (<code class="language-plaintext highlighter-rouge">sonar-deep-research</code>), others don’t support a standard chat API (<code class="language-plaintext highlighter-rouge">morph-v3</code>) while others are aggressively rate-limited or require additional login info (<code class="language-plaintext highlighter-rouge">google</code> and <code class="language-plaintext highlighter-rouge">openai</code>, respectively).</p> <p>This leaves you with roughly 240 models. A quick random sample of their IDs shows frontier labs, big corpos like <code class="language-plaintext highlighter-rouge">meta</code>, as well as startups like <code class="language-plaintext highlighter-rouge">liquid</code>, <code class="language-plaintext highlighter-rouge">ai21</code> and <code class="language-plaintext highlighter-rouge">kimi</code>. All told, I spent about 80 bucks with a few small test runs. A lot of that was due to me forgetting that <code class="language-plaintext highlighter-rouge">sonar-deep-research</code> was enabled and that is a very expensive API to run.</p> <p>Here are some of the most expensive models to run:</p> <div style="text-align: center;"> <img src="/images/2025-08-15-what-i-learned-from-flipping-coins/openrouter_cost.png" alt="Money spent on models" width="600"/> </div> <p>I was pleasantly surprised that I could get \(n\approx 50\) draws from each model for less than $40. I’d certainly like to do more cross-model experiments using OpenRouter.</p> <h3 id="which-models-are-biased-the-most">Which models are biased the most?</h3> <p>I did a bunch of work on creating predictor features for each model using Perplexity’s API and a prompt like “Categorize model X as dense / mixture of experts and also determine if it uses reasoning…”. I used these as the covariates for a Bayesian binomial regression to see if there was a strong effect discernable from the data for aspects like model architecture, size, the organization that released it, etc.</p> <p>Unfortunately, there weren’t any interesting findings. I tried with different parameterizations and likelihoods. The only effect I repeatedly found was that larger models tended to be less biased.</p> <h4 id="the-actual-results">The actual results</h4> <p>I had a much better experience by just listing the models as ordered by their bias, defined as the absolute value of the average of their draws (confined to <code class="language-plaintext highlighter-rouge">{0,1}</code>) minus 0.5. In other words, it’s the distance of the sample mean from a fair coin flip’s \(p(heads)\). Smaller is better. Here are the 20 least biased models:</p> <table> <thead> <tr> <th>Model Name</th> <th>Bias</th> </tr> </thead> <tbody> <tr> <td>sao10k/l3.1-euryale-70b</td> <td>0.032</td> </tr> <tr> <td>inflection/inflection-3-pi</td> <td>0.032</td> </tr> <tr> <td>openai/gpt-4</td> <td>0.024</td> </tr> <tr> <td>neversleep/noromaid-20b</td> <td>0.023</td> </tr> <tr> <td>mistralai/pixtral-large-2411</td> <td>0.022</td> </tr> <tr> <td>sao10k/l3.3-euryale-70b</td> <td>0.022</td> </tr> <tr> <td>meta-llama/llama-4-maverick</td> <td>0.022</td> </tr> <tr> <td>qwen/qwen2.5-vl-32b-instruct</td> <td>0.022</td> </tr> <tr> <td>scb10x/llama3.1-typhoon2-70b-instruct</td> <td>0.011</td> </tr> <tr> <td>infermatic/mn-inferor-12b</td> <td>0.011</td> </tr> <tr> <td>qwen/qwen-2.5-72b-instruct</td> <td>0.011</td> </tr> <tr> <td>sophosympatheia/midnight-rose-70b</td> <td>0.011</td> </tr> <tr> <td>cohere/command-r-plus-08-2024</td> <td>0.011</td> </tr> <tr> <td>google/gemini-2.5-pro</td> <td>0.0</td> </tr> <tr> <td>cohere/command-a</td> <td>0.0</td> </tr> <tr> <td>nvidia/llama-3.1-nemotron-70b-instruct</td> <td>0.0</td> </tr> <tr> <td>openai/gpt-4-1106-preview</td> <td>0.0</td> </tr> <tr> <td>qwen/qwq-32b</td> <td>0.0</td> </tr> <tr> <td>perplexity/r1-1776</td> <td>0.0</td> </tr> <tr> <td>arliai/qwq-32b-arliai-rpr-v1</td> <td>0.0</td> </tr> </tbody> </table> <p>Something that stood out to me immediately was the multiple roleplaying finetunes, like <code class="language-plaintext highlighter-rouge">infermatic/mn-inferor-12b</code>, <code class="language-plaintext highlighter-rouge">arliai/qwq-32b-arliai-rpr-v1 </code>, and <code class="language-plaintext highlighter-rouge">sao10k/l3.3-euryale-70b</code> among others.</p> <p>None of the RP models show up on the top 20 most-biased:</p> <table> <thead> <tr> <th>Model Name</th> <th>Bias</th> </tr> </thead> <tbody> <tr> <td>inception/mercury-coder</td> <td>0.50</td> </tr> <tr> <td>arcee-ai/spotlight</td> <td>0.50</td> </tr> <tr> <td>cohere/command</td> <td>0.50</td> </tr> <tr> <td>google/gemini-2.0-flash-001</td> <td>0.50</td> </tr> <tr> <td>qwen/qwen-vl-plus</td> <td>0.50</td> </tr> <tr> <td>microsoft/phi-4-multimodal-instruct</td> <td>0.48</td> </tr> <tr> <td>anthropic/claude-3.5-sonnet-20240620</td> <td>0.48</td> </tr> <tr> <td>inception/mercury</td> <td>0.46</td> </tr> <tr> <td>google/gemini-flash-1.5-8b</td> <td>0.46</td> </tr> <tr> <td>openai/gpt-4o-mini-search-preview</td> <td>0.46</td> </tr> <tr> <td>anthropic/claude-3.7-sonnet</td> <td>0.44</td> </tr> <tr> <td>google/gemini-2.5-flash</td> <td>0.44</td> </tr> <tr> <td>amazon/nova-lite-v1</td> <td>0.39</td> </tr> <tr> <td>cohere/command-r</td> <td>0.39</td> </tr> <tr> <td>openai/gpt-5-mini</td> <td>0.37</td> </tr> <tr> <td>deepseek/deepseek-chat-v3-0324</td> <td>0.37</td> </tr> <tr> <td>anthropic/claude-3.5-sonnet</td> <td>0.37</td> </tr> <tr> <td>qwen/qwen3-235b-a22b-thinking-2507</td> <td>0.36</td> </tr> <tr> <td>meta-llama/llama-4-scout</td> <td>0.35</td> </tr> <tr> <td>anthropic/claude-opus-4.1</td> <td>0.35</td> </tr> </tbody> </table> <p>Here, we see many of the heavy-hitters from the last 18 months. The best Claude models are all represented here, as is <code class="language-plaintext highlighter-rouge">deepseek-chat-v3-0324</code> and <code class="language-plaintext highlighter-rouge">qwen3-235b-a22b-thinking-2507</code>. The coding-specialized models fron Mercury also both show up here.</p> <p>This hints at a connection between sampling reliability and creativity. The models which are clearly most tuned for coding/problem solving like Claude and Deepseek are also the worst at giving a balanced sample of coin flips. Conversely, models finetuned for creative and interesting dialogue tend to show less bias.</p>]]></content><author><name></name></author><category term="llm"/><category term="statistics"/><summary type="html"><![CDATA[Probing statistical bias across LLM families and use cases]]></summary></entry><entry><title type="html">Simulating planetary waves</title><link href="https://ckrapu.github.io/blog/2025/simulating-planetary-waves/" rel="alternate" type="text/html" title="Simulating planetary waves"/><published>2025-07-22T00:00:00+00:00</published><updated>2025-07-22T00:00:00+00:00</updated><id>https://ckrapu.github.io/blog/2025/simulating-planetary-waves</id><content type="html" xml:base="https://ckrapu.github.io/blog/2025/simulating-planetary-waves/"><![CDATA[<p>Rossby waves, also known as planetary waves, are large-scale waves that occur in the atmosphere and oceans of planets due to the variation of the Coriolis effect with latitude. On Earth, these waves are a fundamental component of atmospheric and oceanic circulation, playing a crucial role in transporting heat, momentum, and moisture across the globe. They are responsible for many of the large-scale weather patterns we experience, including the formation and movement of high and low-pressure systems, which dictate storm tracks and temperature extremes. Understanding Rossby waves is essential for predicting weather systems and analyzing long-term climate variability.</p> <p>This notebook aims to demonstrate how to perform a series of simplified models for understanding and visualizing these waves. We will start with basic analytical solutions and progress towards more complex simulations to explore their behavior and characteristics.</p> <p>We’ll start by setting up the key parameters for our simulation. These parameters are fundamental to defining the environment and conditions under which the Rossby waves will be modeled. They include spatial dimensions such as the size of the domain in both the x and y directions, which will represent a portion of the Earth’s surface. We also define the grid resolution by specifying the number of grid cells in each direction, <code class="language-plaintext highlighter-rouge">nx</code> and <code class="language-plaintext highlighter-rouge">ny</code>, which directly impacts the level of detail captured in the simulation and the computational cost.</p> <p>The time parameters, such as <code class="language-plaintext highlighter-rouge">max_time_s</code> and the number of <code class="language-plaintext highlighter-rouge">iterations</code>, determine the duration of the simulation and the frequency at which we observe the wave’s evolution. We set the angular velocity, <code class="language-plaintext highlighter-rouge">omega_rad_s</code>, which is the physical property driving the Coriolis effect, a fundamental force behind the generation and behavior of Rossby waves.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="n">tqdm</span>
<span class="kn">import</span> <span class="n">scipy</span>
<span class="kn">import</span> <span class="n">matplotlib.pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">import</span> <span class="n">warnings</span>

<span class="kn">from</span> <span class="n">matplotlib.animation</span> <span class="kn">import</span> <span class="n">FuncAnimation</span>
<span class="kn">from</span> <span class="n">IPython.display</span> <span class="kn">import</span> <span class="n">HTML</span><span class="p">,</span> <span class="n">clear_output</span>
<span class="kn">import</span> <span class="n">matplotlib</span> <span class="k">as</span> <span class="n">mpl</span>
<span class="kn">from</span> <span class="n">dataclasses</span> <span class="kn">import</span> <span class="n">dataclass</span>

<span class="n">plt</span><span class="p">.</span><span class="n">style</span><span class="p">.</span><span class="nf">use</span><span class="p">(</span><span class="sh">'</span><span class="s">dark_background</span><span class="sh">'</span><span class="p">)</span>

<span class="c1"># Numpy complains when we use .real on a complex-valued array, so
# we suppress that warning to keep the notebook short.
</span><span class="n">warnings</span><span class="p">.</span><span class="nf">filterwarnings</span><span class="p">(</span><span class="sh">'</span><span class="s">ignore</span><span class="sh">'</span><span class="p">,</span> <span class="n">category</span><span class="o">=</span><span class="n">np</span><span class="p">.</span><span class="n">exceptions</span><span class="p">.</span><span class="n">ComplexWarning</span><span class="p">)</span>
</code></pre></div></div> <p>Our first version of this will be simplified to use a Cartesian coordinate system roughly matching Earth’s parameters, with a rectangular domain as wide and tall as half the Earth’s circumference, and with an angular velocity roughly matching our own.</p> <p>Note that the role of the <code class="language-plaintext highlighter-rouge">omega_rad_s</code> is only to impart accelerations due to the Coriolis effect. The root cause of planetary waves is the differing strength of the Coriolis affect across latitudes.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nd">@dataclass</span>
<span class="k">class</span> <span class="nc">SimulationConfig</span><span class="p">():</span>
  <span class="sh">'''</span><span class="s">
  Values and parameters which won</span><span class="sh">'</span><span class="s">t be changing during the simulation runs.
  </span><span class="sh">'''</span>
  <span class="n">omega_rad_s</span> <span class="o">=</span> <span class="mf">7.29e-5</span>
  <span class="n">earth_radius_m</span> <span class="o">=</span> <span class="mi">6_371_000</span>
  <span class="n">U</span> <span class="o">=</span> <span class="mi">50</span> <span class="c1"># Average zonal velocity, in meters / s
</span>  <span class="n">x_size_m</span> <span class="o">=</span> <span class="mi">40_075_000</span> <span class="o">/</span> <span class="mi">2</span>
  <span class="n">y_size_m</span> <span class="o">=</span> <span class="mi">40_075_000</span> <span class="o">/</span> <span class="mi">2</span>
  <span class="n">nx</span> <span class="o">=</span> <span class="mi">100</span> <span class="c1"># number of grid cells in horizontal direction
</span>  <span class="n">ny</span> <span class="o">=</span> <span class="mi">100</span> <span class="c1"># number of grid cells vertically
</span>  <span class="n">dx</span> <span class="o">=</span> <span class="n">x_size_m</span> <span class="o">/</span> <span class="n">nx</span>
  <span class="n">dy</span> <span class="o">=</span> <span class="n">y_size_m</span> <span class="o">/</span> <span class="n">ny</span>
  <span class="n">max_time_s</span> <span class="o">=</span> <span class="mi">3600</span><span class="o">*</span><span class="mi">24</span><span class="o">*</span><span class="mi">30</span> <span class="c1"># in terms of seconds, hours, and days
</span>  <span class="n">iterations</span> <span class="o">=</span> <span class="mi">300</span>
  <span class="n">timesteps</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">linspace</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">max_time_s</span><span class="p">,</span> <span class="n">iterations</span><span class="p">)</span>

  <span class="c1"># separate 2d arrays where each element contains the x- or y-location of
</span>  <span class="c1"># that grid cell
</span>  <span class="n">x_arr</span><span class="p">,</span> <span class="n">y_arr</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">meshgrid</span><span class="p">(</span>
      <span class="n">np</span><span class="p">.</span><span class="nf">linspace</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">x_size_m</span><span class="p">,</span> <span class="n">nx</span><span class="p">),</span>
      <span class="n">np</span><span class="p">.</span><span class="nf">linspace</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">y_size_m</span><span class="p">,</span> <span class="n">ny</span><span class="p">)</span>
  <span class="p">)</span>

  <span class="c1"># We also need the latitude values for each grid cell,
</span>  <span class="c1"># in units of radians
</span>  <span class="n">latitude_arr</span>  <span class="o">=</span> <span class="n">y_arr</span> <span class="o">/</span> <span class="n">y_size_m</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">pi</span> <span class="o">/</span> <span class="mi">2</span>
  <span class="n">beta_arr</span> <span class="o">=</span> <span class="mi">2</span> <span class="o">*</span> <span class="n">omega_rad_s</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="nf">cos</span><span class="p">(</span><span class="n">latitude_arr</span><span class="p">)</span> <span class="o">/</span> <span class="n">earth_radius_m</span>

<span class="n">cfg</span> <span class="o">=</span> <span class="nc">SimulationConfig</span><span class="p">()</span>
</code></pre></div></div> <h3 id="run-rossby-wave-simulation">Run Rossby wave simulation</h3> <p>The explanations in this section are based on <a href="https://en.wikipedia.org/wiki/Rossby_wave">the Wikipedia page for planetary waves</a> where you can learn more about their derivation. Note that this section mixes both Cartesian (\(x,y\)) and polar (\(\phi\)) coordinates. To reconcile this, we only use the polar c</p> <p>Our first version of this will be simplified to use a Cartesian coordinate system, which is a two-dimensional flat grid, as opposed to a spherical coordinate system that would more accurately represent the Earth’s curved surface. The rectangular domain used here roughly matches the scale of half the Earth’s circumference in both width and height. The angular velocity parameter is set to a value that approximates Earth’s rotational rate.</p> <p>Note that the role of the <code class="language-plaintext highlighter-rouge">omega_rad_s</code> is primarily to impart accelerations due to the Coriolis effect, which is a fictitious force observed in a rotating frame of reference that deflects moving objects. The root cause of planetary waves, however, is not the Coriolis effect itself, but rather the differing strength of the Coriolis effect across latitudes. This variation is known as the beta effect, and it is this latitudinal gradient that provides the restoring force for Rossby waves, causing them to propagate westward relative to the mean flow. In a Cartesian system, the beta effect is often introduced as a linear variation of the Coriolis parameter with the y-coordinate (representing latitude). The resulting PDE is:</p> \[\frac{\partial\nabla^2 \psi}{\partial t} + U \frac{\partial\nabla^2 \psi}{\partial x} + \beta \frac{\partial \psi}{\partial x}\] <p>where \(U\) denotes the average zonal velocity, also described as the mean westerly flow. Note that \(\beta = \frac{2\omega \cos \phi}{a}\), the Rossby parameter, captures the dependence of the system on latitude, as well as the size of the Earth and the speed of Earth’s rotation.</p> <p>Our first version will use an analytical solution, specified in terms of Cartesian coordinates \(x\), \(y\), wavenumbers \(k\) and \(l\), and angular velocity \(\omega\). The analytical solution provides a direct mathematical expression for the streamfunction \(\psi(x,y,t)\) as a function of spatial position \((x, y)\) and time \(t\).</p> \[\psi(x,y,t) = \psi_0 e^{i(kx + ly - \omega t)}\] <p>Here, \(\psi_0\) represents the amplitude of the wave, and \(k\) and \(l\) are the zonal and meridional wavenumbers, respectively. Wavenumbers describe the spatial frequency of the wave in the x and y directions. A larger wavenumber corresponds to a shorter wavelength and finer spatial structure. The term \(\omega t\) represents the temporal evolution of the wave, with \(\omega\) being the angular frequency of the wave. Note that a realistic solution for atmospheric or oceanic flows would typically involve the superposition of many such terms with different wavenumbers and amplitudes, representing a spectrum of waves. This simulation, for simplicity, uses a combination of a few selected wavenumbers to illustrate the basic wave behavior. Our simulation logic tracks the values of the streamfunction \(\psi\) over time and space. The streamfunction itself is a scalar field, but its gradient \(\nabla \psi\) is physically meaningful as it represents the velocity field of the underlying fluid. Specifically, in this simplified model, the velocity components \((u, v)\) are given by \(u = -\partial\psi / \partial y\) and \(v = \partial\psi / \partial x\). Therefore, we also calculate and store the gradient of the streamfunction in both the x and y directions throughout the simulation.</p> <p>A realistic solution would involve the superposition of terms from multiple wavenumbers. This simulation just uses a few of them. Our simulation logic tracks the values of the streamfunction \(\psi\) as well as its gradient, as the velocity of the underlying fluid field is specified in terms of the gradient \(\nabla \psi\).</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="nd">@dataclass</span>
<span class="k">class</span> <span class="nc">SimulationState</span><span class="p">():</span>
  <span class="sh">'''</span><span class="s">
  Contains the state variables for the simulation.
  </span><span class="sh">'''</span>
  <span class="n">ψ</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="n">ndarray</span>
  <span class="n">u</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="n">ndarray</span>
  <span class="n">v</span><span class="p">:</span> <span class="n">np</span><span class="p">.</span><span class="n">ndarray</span>

<span class="n">state</span> <span class="o">=</span> <span class="nc">SimulationState</span><span class="p">(</span>
  <span class="n">ψ</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">zeros</span><span class="p">((</span><span class="n">cfg</span><span class="p">.</span><span class="n">iterations</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">ny</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">)),</span>
  <span class="n">u</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">zeros</span><span class="p">((</span><span class="n">cfg</span><span class="p">.</span><span class="n">iterations</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">ny</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">)),</span>
  <span class="n">v</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">zeros</span><span class="p">((</span><span class="n">cfg</span><span class="p">.</span><span class="n">iterations</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">ny</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">nx</span><span class="p">))</span>
<span class="p">)</span>

<span class="k">def</span> <span class="nf">rossby_wave</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">ω</span><span class="p">,</span> <span class="n">t</span><span class="p">,</span> <span class="n">ψ_0</span><span class="p">,</span> <span class="n">k</span><span class="p">,</span> <span class="n">l</span><span class="p">):</span>
  <span class="sh">'''</span><span class="s">
  Copied from https://en.wikipedia.org/wiki/Rossby_wave#Free_barotropic_Rossby_waves_under_a_zonal_flow_with_linearized_vorticity_equation
  </span><span class="sh">'''</span>

  <span class="k">return</span> <span class="n">ψ_0</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="nf">exp</span><span class="p">(</span><span class="mf">1j</span> <span class="o">*</span> <span class="p">(</span><span class="n">k</span> <span class="o">*</span> <span class="n">x</span> <span class="o">+</span> <span class="n">l</span> <span class="o">*</span> <span class="n">y</span> <span class="o">-</span> <span class="n">ω</span> <span class="o">*</span> <span class="n">t</span><span class="p">))</span>

<span class="k">def</span> <span class="nf">omega</span><span class="p">(</span><span class="n">k</span><span class="p">,</span> <span class="n">l</span><span class="p">):</span>
  <span class="k">return</span> <span class="n">cfg</span><span class="p">.</span><span class="n">U</span><span class="o">*</span><span class="n">k</span> <span class="o">-</span> <span class="n">cfg</span><span class="p">.</span><span class="n">beta_arr</span> <span class="o">*</span> <span class="n">k</span> <span class="o">/</span> <span class="p">(</span><span class="n">k</span><span class="o">**</span><span class="mi">2</span> <span class="o">+</span> <span class="n">l</span><span class="o">**</span><span class="mi">2</span><span class="p">)</span>

<span class="c1"># Wavenumbers for the Rossby waves
# Define wavenumbers in physical units (rad/m)
</span><span class="n">k_physical</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">pi</span> <span class="o">/</span> <span class="p">(</span><span class="mf">10000e3</span><span class="p">)</span>  <span class="c1"># 10,000 km wavelength
</span><span class="n">l_physical</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">pi</span> <span class="o">/</span> <span class="p">(</span><span class="mf">8000e3</span><span class="p">)</span>   <span class="c1"># 8,000 km wavelength
</span>
<span class="c1"># Convert to grid units
</span><span class="n">ks</span> <span class="o">=</span> <span class="p">[</span><span class="n">k_physical</span> <span class="o">*</span> <span class="n">n</span> <span class="k">for</span> <span class="n">n</span> <span class="ow">in</span> <span class="nf">range</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="c1"># zonal
</span><span class="n">ls</span> <span class="o">=</span> <span class="p">[</span><span class="n">l_physical</span><span class="p">]</span> <span class="c1"># meriodonal
</span>
<span class="c1"># Note that only the gradient of the streamfunction is physically meaningful
# so this parameter is unimportant for the simulation.
</span><span class="n">ψ_0</span> <span class="o">=</span> <span class="mf">1.0</span>


<span class="k">with</span> <span class="n">tqdm</span><span class="p">.</span><span class="nf">tqdm</span><span class="p">(</span><span class="n">total</span><span class="o">=</span><span class="n">cfg</span><span class="p">.</span><span class="n">iterations</span><span class="p">,</span> <span class="n">desc</span><span class="o">=</span><span class="sh">"</span><span class="s">Simulating Rossby waves</span><span class="sh">"</span><span class="p">)</span> <span class="k">as</span> <span class="n">pbar</span><span class="p">:</span>
  <span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">t</span> <span class="ow">in</span> <span class="nf">enumerate</span><span class="p">(</span><span class="n">cfg</span><span class="p">.</span><span class="n">timesteps</span><span class="p">):</span>
    <span class="n">current_ψ</span> <span class="o">=</span> <span class="n">ψ_0</span> <span class="o">*</span> <span class="nf">sum</span><span class="p">(</span>
        <span class="nf">rossby_wave</span><span class="p">(</span><span class="n">cfg</span><span class="p">.</span><span class="n">x_arr</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">y_arr</span><span class="p">,</span> <span class="nf">omega</span><span class="p">(</span><span class="n">k</span><span class="p">,</span> <span class="n">l</span><span class="p">),</span> <span class="n">t</span><span class="p">,</span> <span class="n">ψ_0</span><span class="p">,</span> <span class="n">k</span><span class="p">,</span> <span class="n">l</span><span class="p">)</span> <span class="k">for</span> <span class="n">k</span> <span class="ow">in</span> <span class="n">ks</span> <span class="k">for</span> <span class="n">l</span> <span class="ow">in</span> <span class="n">ls</span>
    <span class="p">)</span>
    <span class="n">state</span><span class="p">.</span><span class="n">ψ</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">=</span> <span class="n">current_ψ</span>
    <span class="n">grady</span><span class="p">,</span> <span class="n">gradx</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">gradient</span><span class="p">(</span><span class="n">ψ</span><span class="p">[</span><span class="n">i</span><span class="p">],</span> <span class="n">cfg</span><span class="p">.</span><span class="n">dy</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">dx</span><span class="p">)</span>
    <span class="n">state</span><span class="p">.</span><span class="n">v</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">=</span> <span class="n">gradx</span>   <span class="c1"># -∂ψ/∂y (eastward velocity)
</span>    <span class="n">state</span><span class="p">.</span><span class="n">u</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">=</span> <span class="o">-</span> <span class="n">grady</span> <span class="c1"># ∂ψ/∂x (northward velocity)
</span>
    <span class="c1"># Calculate RMS average of the streamfunction
</span>    <span class="n">rms_psi</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">sqrt</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">mean</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">abs</span><span class="p">(</span><span class="n">current_ψ</span><span class="p">)</span><span class="o">**</span><span class="mi">2</span><span class="p">))</span>
    <span class="n">pbar</span><span class="p">.</span><span class="nf">set_description</span><span class="p">(</span><span class="sh">"</span><span class="s">RMS(ψ): {:.2f}</span><span class="sh">"</span><span class="p">.</span><span class="nf">format</span><span class="p">(</span><span class="n">rms_psi</span><span class="p">))</span>
    <span class="n">pbar</span><span class="p">.</span><span class="nf">update</span><span class="p">(</span><span class="mi">1</span><span class="p">)</span>
</code></pre></div></div> <h3 id="make-single-plot-of-wave-at-final-timestep">Make single plot of wave at final timestep</h3> <p>Let’s visualize a single timestep from the simulation to understand the spatial structure of the Rossby wave and the associated fluid motion. Here, the color intensity represents the real part of the streamfunction \(\psi\) at the last calculated timestep, providing a spatial map of this scalar field across the domain. Areas of high \(\psi\) are shown in warmer colors, while areas of low \(\psi\) are in cooler colors. The white arrows, displayed using a quiver plot, indicate the direction and magnitude of the gradient of the streamfunction, \(\nabla \psi\), at various points across the grid.</p> <p>As discussed earlier, the gradient of the streamfunction is directly related to the velocity field of the underlying fluid. Specifically, in a 2D incompressible flow represented by a streamfunction, the velocity components \((u, v)\) are given by \(u = -\partial\psi / \partial y\) and \(v = \partial\psi / \partial x\). Therefore, the arrows show the instantaneous velocity vectors \((u, v)\) at each location, revealing the instantaneous flow direction and speed. High values of \(\psi\) often correspond to regions of anticyclonic (clockwise in the Northern Hemisphere) circulation, while low values correspond to cyclonic (counter-clockwise in the Northern Hemisphere) circulation. The arrows show the fluid parcel trajectories at this specific moment in time, revealing the large-scale circulation patterns driven by the wave.</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="n">ax</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="nf">subplots</span><span class="p">()</span>

<span class="c1"># Remember that imshow doesn't know the x-y coords for this data;
# we have to give it the extent.
</span><span class="n">im</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">imshow</span><span class="p">(</span><span class="n">state</span><span class="p">.</span><span class="n">ψ</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">].</span><span class="n">real</span><span class="p">,</span> <span class="n">origin</span><span class="o">=</span><span class="sh">'</span><span class="s">lower</span><span class="sh">'</span><span class="p">,</span> <span class="n">extent</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">x_size_m</span><span class="o">/</span><span class="mi">1000</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">y_size_m</span><span class="o">/</span><span class="mi">1000</span><span class="p">])</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">colorbar</span><span class="p">(</span><span class="n">im</span><span class="p">,</span> <span class="n">label</span> <span class="o">=</span> <span class="sh">"</span><span class="s">Streamfunction $$\psi(x,y)$$</span><span class="sh">"</span><span class="p">)</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">ylabel</span><span class="p">(</span><span class="sh">"</span><span class="s">y (km)</span><span class="sh">"</span><span class="p">),</span> <span class="n">plt</span><span class="p">.</span><span class="nf">xlabel</span><span class="p">(</span><span class="sh">'</span><span class="s">x (km)</span><span class="sh">'</span><span class="p">)</span>

<span class="c1"># Since the grid is quite large, we subset the locations in which we actually
# place arrows using the quiver plot. Picking the right `n` is mostly just trial
# and error.
</span><span class="n">n</span> <span class="o">=</span> <span class="mi">5</span>
<span class="n">q</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">quiver</span><span class="p">(</span><span class="n">cfg</span><span class="p">.</span><span class="n">x_arr</span><span class="p">[::</span><span class="n">n</span><span class="p">,</span> <span class="p">::</span><span class="n">n</span><span class="p">]</span><span class="o">/</span><span class="mi">1000</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">y_arr</span><span class="p">[::</span><span class="n">n</span><span class="p">,</span> <span class="p">::</span><span class="n">n</span><span class="p">]</span><span class="o">/</span><span class="mi">1000</span><span class="p">,</span> <span class="n">state</span><span class="p">.</span><span class="n">u</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="p">::</span><span class="n">n</span><span class="p">,</span> <span class="p">::</span><span class="n">n</span><span class="p">],</span> <span class="n">state</span><span class="p">.</span><span class="n">v</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="p">::</span><span class="n">n</span><span class="p">,</span> <span class="p">::</span><span class="n">n</span><span class="p">],</span>
  <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">w</span><span class="sh">'</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_aspect</span><span class="p">(</span><span class="sh">'</span><span class="s">equal</span><span class="sh">'</span><span class="p">,</span> <span class="n">adjustable</span><span class="o">=</span><span class="sh">'</span><span class="s">box</span><span class="sh">'</span><span class="p">)</span> <span class="c1"># Ensure aspect ratio is equal
</span></code></pre></div></div> <div style="text-align: center;"> <img src="/images/2025-07-22-simulating-planetary-waves/wave-still.png" alt="Rossby wave visualization at final timestep" width="600"/> </div> <h3 id="animating-the-streamfunction-and-velocity-field">Animating the streamfunction and velocity field</h3> <p>To appreciate the dynamical nature of the Rossby wave and observe its characteristic westward propagation, we can create an animation showing its evolution over the course of the simulated period, which spans several days. While the static plot provides a snapshot at a single moment, the animation reveals how the wave structure, represented by the streamfunction contours and the associated velocity field, changes and moves over time. As the animation progresses through the timesteps, we can see the crests and troughs of the streamfunction pattern propagating across the domain.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># Create animation of ψ over time using matplotlib and func animation. Show the animation here
</span>
<span class="c1"># Increase the animation embed limit
</span><span class="n">mpl</span><span class="p">.</span><span class="n">rcParams</span><span class="p">[</span><span class="sh">'</span><span class="s">animation.embed_limit</span><span class="sh">'</span><span class="p">]</span> <span class="o">=</span> <span class="mf">50.0</span>  <span class="c1"># Increase the limit to 50 MB
</span>
<span class="n">fig</span><span class="p">,</span> <span class="n">ax</span> <span class="o">=</span> <span class="n">plt</span><span class="p">.</span><span class="nf">subplots</span><span class="p">(</span><span class="n">figsize</span><span class="o">=</span><span class="p">(</span><span class="mi">8</span><span class="p">,</span> <span class="mi">8</span><span class="p">))</span> <span class="c1"># Create a single figure and axes
</span>
<span class="c1"># Plot for ψ (heatmap)
</span><span class="n">im</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">imshow</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">real</span><span class="p">(</span><span class="n">state</span><span class="p">.</span><span class="n">ψ</span><span class="p">[</span><span class="mi">0</span><span class="p">]),</span> <span class="n">origin</span><span class="o">=</span><span class="sh">'</span><span class="s">lower</span><span class="sh">'</span><span class="p">,</span> <span class="n">extent</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">x_size_m</span><span class="o">/</span><span class="mi">1000</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">y_size_m</span><span class="o">/</span><span class="mi">1000</span><span class="p">])</span>
<span class="n">plt</span><span class="p">.</span><span class="nf">colorbar</span><span class="p">(</span><span class="n">im</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">ax</span><span class="p">.</span><span class="nf">set_title</span><span class="p">(</span><span class="sh">'</span><span class="s">Rossby Wave (ψ) and Flow Direction</span><span class="sh">'</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_xlabel</span><span class="p">(</span><span class="sh">'</span><span class="s">x (km)</span><span class="sh">'</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_ylabel</span><span class="p">(</span><span class="sh">'</span><span class="s">y (km)</span><span class="sh">'</span><span class="p">)</span>


<span class="c1"># Add a text annotation for time
</span><span class="n">time_text</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">text</span><span class="p">(</span><span class="mf">0.02</span><span class="p">,</span> <span class="mf">0.95</span><span class="p">,</span> <span class="sh">''</span><span class="p">,</span> <span class="n">transform</span><span class="o">=</span><span class="n">ax</span><span class="p">.</span><span class="n">transAxes</span><span class="p">,</span> <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">white</span><span class="sh">'</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">bbox</span><span class="o">=</span><span class="nf">dict</span><span class="p">(</span><span class="n">boxstyle</span><span class="o">=</span><span class="sh">'</span><span class="s">round,pad=0.5</span><span class="sh">'</span><span class="p">,</span> <span class="n">fc</span><span class="o">=</span><span class="sh">'</span><span class="s">black</span><span class="sh">'</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="c1"># Plot for gradient (quiver)
# We need to subsample the grid for the quiver plot to avoid overcrowding
</span><span class="n">step</span> <span class="o">=</span> <span class="mi">5</span> <span class="c1"># Adjust this value to change the density of arrows
</span><span class="n">quiver_plot</span> <span class="o">=</span> <span class="n">ax</span><span class="p">.</span><span class="nf">quiver</span><span class="p">(</span><span class="n">cfg</span><span class="p">.</span><span class="n">x_arr</span><span class="p">[::</span><span class="n">step</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">]</span><span class="o">/</span><span class="mi">1000</span><span class="p">,</span> <span class="n">cfg</span><span class="p">.</span><span class="n">y_arr</span><span class="p">[::</span><span class="n">step</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">]</span><span class="o">/</span><span class="mi">1000</span><span class="p">,</span>
                         <span class="n">state</span><span class="p">.</span><span class="n">u</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">],</span> <span class="n">state</span><span class="p">.</span><span class="n">v</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">],</span>
                         <span class="n">color</span><span class="o">=</span><span class="sh">'</span><span class="s">white</span><span class="sh">'</span><span class="p">)</span>
<span class="n">ax</span><span class="p">.</span><span class="nf">set_aspect</span><span class="p">(</span><span class="sh">'</span><span class="s">equal</span><span class="sh">'</span><span class="p">,</span> <span class="n">adjustable</span><span class="o">=</span><span class="sh">'</span><span class="s">box</span><span class="sh">'</span><span class="p">)</span>


<span class="k">def</span> <span class="nf">animate</span><span class="p">(</span><span class="n">i</span><span class="p">):</span>
    <span class="n">im</span><span class="p">.</span><span class="nf">set_data</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">real</span><span class="p">(</span><span class="n">state</span><span class="p">.</span><span class="n">ψ</span><span class="p">[</span><span class="n">i</span><span class="p">]))</span>
    <span class="n">time_in_days</span> <span class="o">=</span> <span class="n">cfg</span><span class="p">.</span><span class="n">timesteps</span><span class="p">[</span><span class="n">i</span><span class="p">]</span> <span class="o">/</span> <span class="p">(</span><span class="mi">3600</span> <span class="o">*</span> <span class="mi">24</span><span class="p">)</span>
    <span class="n">time_text</span><span class="p">.</span><span class="nf">set_text</span><span class="p">(</span><span class="sa">f</span><span class="sh">'</span><span class="s">Time: </span><span class="si">{</span><span class="n">time_in_days</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"> days</span><span class="sh">'</span><span class="p">)</span>

    <span class="n">quiver_plot</span><span class="p">.</span><span class="nf">set_UVC</span><span class="p">(</span><span class="n">state</span><span class="p">.</span><span class="n">u</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">],</span> <span class="n">state</span><span class="p">.</span><span class="n">v</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">,</span> <span class="p">::</span><span class="n">step</span><span class="p">])</span>
    <span class="k">return</span> <span class="p">[</span><span class="n">im</span><span class="p">,</span> <span class="n">time_text</span><span class="p">,</span> <span class="n">quiver_plot</span><span class="p">]</span>

<span class="n">ani</span> <span class="o">=</span> <span class="nc">FuncAnimation</span><span class="p">(</span><span class="n">fig</span><span class="p">,</span> <span class="n">animate</span><span class="p">,</span> <span class="n">frames</span><span class="o">=</span><span class="n">cfg</span><span class="p">.</span><span class="n">iterations</span><span class="p">,</span> <span class="n">interval</span><span class="o">=</span><span class="mi">50</span><span class="p">,</span> <span class="n">blit</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>

<span class="nc">HTML</span><span class="p">(</span><span class="n">ani</span><span class="p">.</span><span class="nf">to_html5_video</span><span class="p">())</span>
</code></pre></div></div> <video controls="" width="85%"> <source src="/images/2025-07-22-simulating-planetary-waves/wave-animation.mp4" type="video/mp4"/> Your browser does not support the video tag. </video>]]></content><author><name></name></author><category term="tutorials"/><category term="physics"/><category term="simulation"/><category term="python"/><summary type="html"><![CDATA[Demonstrating simplified models for understanding and visualizing Rossby waves]]></summary></entry><entry><title type="html">Modeling data with correlated errors across a directed graph</title><link href="https://ckrapu.github.io/blog/2025/modeling-data-with-correlated-errors-across-a-directed-graph/" rel="alternate" type="text/html" title="Modeling data with correlated errors across a directed graph"/><published>2025-04-13T21:57:00+00:00</published><updated>2025-04-13T21:57:00+00:00</updated><id>https://ckrapu.github.io/blog/2025/modeling-data-with-correlated-errors-across-a-directed-graph</id><content type="html" xml:base="https://ckrapu.github.io/blog/2025/modeling-data-with-correlated-errors-across-a-directed-graph/"><![CDATA[<p>Data for objects which can be linked together with a graph are common in fields like epidemiology and sociology. A simple way to model this kind of data is with a linear mixed model with a random effect with some network structure, or an error term which is correlated. The former is appropriate when you have multiple observations for each node in your graph while the latter is applicable when there is usually just one data point for each node.</p> <p>PyMC has functionality for these kinds of models assuming an <strong>undirected</strong> graph. This makes sense if it’s not clear which way the causality goes. For example, if the number of infected people in city \(A\) is correlated with the same quantity in city \(G\) because there is a road from \(A\) to \(G\), then the correlation could be due to people traveling from \(A\) to \(G\), from \(G\) to \(A\), or both processes simultaneously. This is a great opportunity to use a Gaussian Markov random field, synonomous with <a href="https://www.pymc.io/projects/examples/en/latest/spatial/conditional_autoregressive_priors.html">conditional autoregression</a>.</p> <p>Many real-world networks are actually <strong>directed</strong> graphs, however. An example of this type is a river network. Water flows in one direction, and events like floods or pollutant spills which occur downstream will generally not have an effect on the upstream segments of the network.</p> <p>In the case where we have a directed graph, we can use some tricks from linear algebra to account for this network-based correlation to help produce more accurate estimates of regression coefficients and better predictive performance on new data. In this post, we’ll use some basic matrix facts and identities to implement a highly scalable and effective linear mixed model for network data assuming an underlying directed acyclic graph (DAG). To learn more about these terms and concepts, I highly recommend selected chapters from <a href="https://github.com/rmcelreath/stat_rethinking_2024/tree/main?tab=readme-ov-file"><em>Statistical Rethinking</em></a> and <a href="https://www.google.com/books/edition/Probabilistic_Graphical_Models/7dzpHCHzNQ4C?hl=en&amp;gbpv=1&amp;pg=PA1&amp;printsec=frontcover"><em>Probabilistic Graphical Models</em></a>.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kn">import</span> <span class="n">arviz</span> <span class="k">as</span> <span class="n">az</span>
<span class="kn">import</span> <span class="n">numpy</span> <span class="k">as</span> <span class="n">np</span>
<span class="kn">import</span> <span class="n">pymc</span> <span class="k">as</span> <span class="n">pm</span>
<span class="kn">import</span> <span class="n">pytensor.tensor</span> <span class="k">as</span> <span class="n">pt</span>
<span class="kn">import</span> <span class="n">pytensor.sparse</span> <span class="k">as</span> <span class="n">sparse</span>
<span class="kn">import</span> <span class="n">scipy</span> <span class="k">as</span> <span class="n">sp</span>
<span class="kn">import</span> <span class="n">seaborn</span> <span class="k">as</span> <span class="n">sns</span>
<span class="kn">import</span> <span class="n">xarray</span> <span class="k">as</span> <span class="n">xr</span>

<span class="kn">from</span> <span class="n">matplotlib</span> <span class="kn">import</span> <span class="n">pyplot</span> <span class="k">as</span> <span class="n">plt</span>
<span class="kn">from</span> <span class="n">matplotlib.lines</span> <span class="kn">import</span> <span class="n">Line2D</span>
<span class="kn">from</span> <span class="n">numpy.random</span> <span class="kn">import</span> <span class="n">default_rng</span>
<span class="kn">from</span> <span class="n">xarray_einstats</span> <span class="kn">import</span> <span class="n">linalg</span>
<span class="kn">from</span> <span class="n">xarray_einstats.stats</span> <span class="kn">import</span> <span class="n">XrContinuousRV</span>

<span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Running on PyMC v</span><span class="si">{</span><span class="n">pm</span><span class="p">.</span><span class="n">__version__</span><span class="si">}</span><span class="sh">"</span><span class="p">)</span>
<span class="o">%</span><span class="n">config</span> <span class="n">InlineBackend</span><span class="p">.</span><span class="n">figure_format</span> <span class="o">=</span> <span class="sh">'</span><span class="s">retina</span><span class="sh">'</span>
<span class="n">az</span><span class="p">.</span><span class="n">style</span><span class="p">.</span><span class="nf">use</span><span class="p">(</span><span class="sh">"</span><span class="s">arviz-darkgrid</span><span class="sh">"</span><span class="p">)</span>

<span class="n">np</span><span class="p">.</span><span class="nf">set_printoptions</span><span class="p">(</span><span class="n">precision</span><span class="o">=</span><span class="mi">3</span><span class="p">,</span> <span class="n">suppress</span><span class="o">=</span><span class="bp">True</span><span class="p">)</span>
<span class="n">RANDOM_SEED</span> <span class="o">=</span> <span class="mi">8271991</span>
<span class="n">rng</span> <span class="o">=</span> <span class="nf">default_rng</span><span class="p">(</span><span class="n">RANDOM_SEED</span><span class="p">)</span>
</code></pre></div></div> <h2 id="model-structure">Model structure</h2> <p>To illustrate why we should account for any underlying network structure in our dataset, let’s start with an example using synthetic data. We’ll assume the following probabilistic model for \(N\) data points for the rest of the notebook:</p> \[\mathbf{y} \sim \textbf{N}(X \beta, \Sigma) \text{ or equivalently } \mathbf{y} = X \beta + \mathbf{\epsilon}\text{, } \epsilon \sim \textbf{N}(0, \Sigma)\] \[\beta_j \sim \text{Normal}(5) \text{ for all } j \in \{1,..,P\}\] <p>We’ll use \(P = 2\) in this example, though the role of the linear predictor \(X \beta\) is not important. \(\mathbf{y}\) is a vector of observations with correlated errors; the structure of this correlation is encoded in the elements of \(\Sigma\), an \(N\times N\) covariance matrix.</p> <p>Another assumption we will make is that the index \(i \in \{1,...,N\}\) maps individual data points to locations on a directed acyclic graph \(G\). \(G\) can be represented as a binary, asymmetric matrix. Gaussian DAG models (sometimes also described as continuous Bayesian networks) have the nice feature that their underlying correlation structure can be described in terms of an inverse-covariance matrix \(\Lambda\), also referred to as the <em>precision matrix</em>.</p> <p>Assuming that the index \(i\) respects the <a href="https://en.wikipedia.org/wiki/Topological_sorting">topological order</a> of the graph \(G\), the zeros of the precision matrix tell us which elements are conditionally independent. Put differently, if \(\lambda_{ij}\) is zero, then \(\epsilon_i\) and \(\epsilon_j\) are independent given all other elements of \(\mathbf{\epsilon}\)</p> \[\Lambda = \begin{pmatrix} \lambda_1 &amp; \lambda_{21} &amp; 0 &amp; \lambda_{41} &amp; 0 \\ \lambda_{21} &amp; \lambda_2 &amp; \lambda_{32} &amp; 0 &amp; 0 \\ 0 &amp; \lambda_{32} &amp; \lambda_3 &amp; 0 &amp; \lambda_{53} \\ \lambda_{41} &amp; 0 &amp; 0 &amp; \lambda_4 &amp; \lambda_{54} \\ 0 &amp; 0 &amp; \lambda_{53} &amp; \lambda_{54} &amp; \lambda_5 \end{pmatrix}\] <p>To understand this matrix, note that the diagonal elements \(\sigma_1, ...\) control the scale of the random variables while the off-diagonal elements control the correlation between \(y_i\) and \(y_j\). This matrix is equivalent (under Gaussian conditions) to the following graphical diagram:</p> <div style="text-align: center;"> <img src="/images/dag.png" alt="DAG diagram" width="400"/> </div> <p>We can actually simulate data with this correlation structure easily given \(\Lambda\) as shown in the next section. This graph also has an interpretation via the LDL factorization of \(\Lambda\). Noting that \(\Lambda\) is positive semi-definite, it can be written as:</p> <p>\(\Lambda = (I-\mathcal{G})^T\Omega(I-\mathcal{G})\) where \(I\) denotes the identity matrix, \(\mathcal{G}\) is a triangular matrix, and \(\Omega\) is a diagonal matrix containing the scalar precision values on its diagonal. While the elements of \(\mathcal{G}\) can be distinct, having one parameter for every nonzero value of \(\mathcal{G}\) is unnecessary as that implies a unique precision parameter for every edge in the graph \(G\). Instead, we typically work with a parameterization of \(\mathcal{G}\) in terms of a scalar parameter \(\gamma\) and the binary matrix \(G\) with \(\mathcal{G}= \gamma G\). Our interpretation of \(\mathcal{G}\) is quite interesting - it gives a construction for draws of \(\epsilon_j\) via regression on elements occurring earlier in the topological order:</p> \[\epsilon_j = \sum_{i&lt;j} b_{ji} \epsilon_i + \nu_j \quad \text{where} \quad \nu_j \sim \text{N}(0, \omega_j^{-1})\] <p>One may interpret this as saying that the error term for node \(j\) can be written as a linear combination of error terms from nodes that precede it in the topological ordering of the graph, plus an independent noise term \(\nu_j\). The coefficient \(b_{ji}\) represents the strength of the relationship between nodes \(i\) and \(j\), and is nonzero only when there is a directed edge from node \(i\) to node \(j\) in the graph.</p> <h2 id="simulated-data-generation">Simulated data generation</h2> <p>To fully illustrate how to use this model, we’ll begin by creating a dataset reflective of a regression with DAG-structured error terms.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">n</span>       <span class="o">=</span> <span class="mi">29</span> <span class="c1"># Number of samples, also dimension of the covariance matrix
</span><span class="n">n_edges</span> <span class="o">=</span> <span class="n">n</span> <span class="o">*</span> <span class="mi">2</span>

<span class="n">p</span> <span class="o">=</span> <span class="mi">2</span>  <span class="c1"># Number of predictors
</span>
<span class="n">beta_true</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">([</span><span class="o">-</span><span class="mf">0.5</span><span class="p">,</span> <span class="mf">0.5</span><span class="p">])</span>

<span class="c1"># In the directed graph, the coefficients for each element of eps in terms of the
# previous elements with edges in the graph
</span><span class="n">rho</span> <span class="o">=</span> <span class="mf">1.0</span>

<span class="n">x</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">normal</span><span class="p">(</span><span class="n">size</span><span class="o">=</span><span class="p">(</span><span class="n">n</span><span class="p">,</span> <span class="n">p</span><span class="p">))</span>
<span class="n">eps_raw</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">normal</span><span class="p">(</span><span class="n">size</span><span class="o">=</span><span class="p">(</span><span class="n">n</span><span class="p">))</span>

<span class="k">def</span> <span class="nf">sample_sparse_dag</span><span class="p">(</span><span class="n">n</span><span class="p">:</span> <span class="nb">int</span><span class="p">,</span> <span class="n">n_edges</span><span class="p">:</span> <span class="nb">int</span><span class="p">)</span> <span class="o">-&gt;</span> <span class="n">sp</span><span class="p">.</span><span class="n">sparse</span><span class="p">.</span><span class="n">lil_matrix</span><span class="p">:</span>
    <span class="sh">'''</span><span class="s">
    Sample `n_pairs` pairs of edges from a directed graph with `n` nodes.
    Edges are not allowed to point to themselves. Assumes topological ordering
    of the nodes.
    </span><span class="sh">'''</span>
    <span class="n">G</span> <span class="o">=</span> <span class="n">sp</span><span class="p">.</span><span class="n">sparse</span><span class="p">.</span><span class="nf">lil_matrix</span><span class="p">((</span><span class="n">n</span><span class="p">,</span> <span class="n">n</span><span class="p">),</span> <span class="n">dtype</span><span class="o">=</span><span class="nb">int</span><span class="p">)</span>

    <span class="c1"># First, make sure that each node has at least one child
</span>    <span class="c1"># except for the last node
</span>    <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n</span><span class="o">-</span><span class="mi">1</span><span class="p">):</span>
        <span class="n">j</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">integers</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="n">n</span><span class="p">)</span>
        <span class="n">G</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">]</span> <span class="o">=</span> <span class="bp">True</span>

    <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n_edges</span> <span class="o">-</span> <span class="n">n</span><span class="p">):</span>
        <span class="c1"># Sample a pair of nodes
</span>        <span class="c1"># such that i &lt; j
</span>        <span class="n">i</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">integers</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="n">n</span><span class="o">-</span><span class="mi">1</span><span class="p">)</span>
        <span class="n">j</span> <span class="o">=</span> <span class="n">rng</span><span class="p">.</span><span class="nf">integers</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="n">n</span><span class="p">)</span>

        <span class="k">if</span> <span class="n">i</span> <span class="o">==</span> <span class="n">j</span> <span class="ow">or</span> <span class="n">G</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">]:</span>
            <span class="k">continue</span>
        <span class="n">G</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">]</span> <span class="o">=</span> <span class="bp">True</span>

    <span class="k">return</span> <span class="n">G</span><span class="p">.</span><span class="nf">tocsr</span><span class="p">()</span>

<span class="c1"># parent idx is i, child idx is j
</span><span class="n">G</span> <span class="o">=</span> <span class="nf">sample_sparse_dag</span><span class="p">(</span><span class="n">n</span><span class="p">,</span> <span class="n">n_edges</span><span class="p">)</span>
<span class="c1">#G = np.zeros((n, n), dtype=int)
</span><span class="n">G</span> <span class="o">=</span> <span class="n">sp</span><span class="p">.</span><span class="n">sparse</span><span class="p">.</span><span class="nf">csr_matrix</span><span class="p">(</span><span class="n">G</span><span class="p">)</span>

<span class="n">eps</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="nf">zeros</span><span class="p">(</span><span class="n">n</span><span class="p">)</span>

<span class="k">for</span> <span class="n">j</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n</span><span class="p">):</span>
    <span class="n">eps</span><span class="p">[</span><span class="n">j</span><span class="p">]</span> <span class="o">=</span> <span class="n">eps_raw</span><span class="p">[</span><span class="n">j</span><span class="p">]</span>
    <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nf">range</span><span class="p">(</span><span class="n">n</span><span class="p">):</span>
        <span class="k">if</span> <span class="n">G</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">]:</span>
            <span class="n">eps</span><span class="p">[</span><span class="n">j</span><span class="p">]</span> <span class="o">+=</span> <span class="n">rho</span> <span class="o">*</span> <span class="n">eps</span><span class="p">[</span><span class="n">i</span><span class="p">]</span>


<span class="n">y</span> <span class="o">=</span> <span class="n">x</span> <span class="o">@</span> <span class="n">beta_true</span> <span class="o">+</span> <span class="n">eps</span>
</code></pre></div></div> <p>If we visualize the resulting DAG matrix as a heatmap, we can see the sparse, triangular structure as desired.</p> <p>Next, we’ll fit two models. The first model naively assumes that elements of \(\mathbf{y}\) are conditionally independent given \(X\beta\). The second model explicitly accounts for the correlated graph structure. In practice, we will often know whether the elements \(i,j\) share an edge, but we won’t know exactly how strong the correlation between them should be.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">with</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Model</span><span class="p">()</span> <span class="k">as</span> <span class="n">naive_model</span><span class="p">:</span>
    <span class="n">β</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="mi">5</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="n">p</span><span class="p">)</span>
    <span class="n">error_sigma</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">error_sigma</span><span class="sh">"</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span>
    
    <span class="n">_</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">y</span><span class="sh">"</span><span class="p">,</span>
        <span class="n">mu</span><span class="o">=</span><span class="n">pt</span><span class="p">.</span><span class="nf">dot</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">β</span><span class="p">),</span>
        <span class="n">sigma</span><span class="o">=</span><span class="n">error_sigma</span><span class="p">,</span>
        <span class="n">observed</span><span class="o">=</span><span class="n">y</span><span class="p">,</span>
    <span class="p">)</span>

<span class="k">with</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Model</span><span class="p">()</span> <span class="k">as</span> <span class="n">graph_model</span><span class="p">:</span>
    <span class="n">β</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="mi">5</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="n">p</span><span class="p">)</span>

    <span class="n">ω</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">ω</span><span class="sh">"</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span> <span class="c1"># Controls uncorrelated variance
</span>    <span class="n">γ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">γ</span><span class="sh">"</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span> <span class="c1"># Controls correlation between nodes in graph
</span>
    <span class="n">G_pt</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="n">math</span><span class="p">.</span><span class="nf">constant</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="nf">asarray</span><span class="p">(</span><span class="n">G</span><span class="p">.</span><span class="nf">astype</span><span class="p">(</span><span class="nb">float</span><span class="p">).</span><span class="nf">todense</span><span class="p">()))</span>
    <span class="n">Ω</span> <span class="o">=</span> <span class="n">pt</span><span class="p">.</span><span class="nf">diag</span><span class="p">(</span><span class="n">pt</span><span class="p">.</span><span class="nf">ones</span><span class="p">(</span><span class="n">n</span><span class="p">))</span> <span class="o">*</span> <span class="n">ω</span>
    <span class="n">I</span> <span class="o">=</span> <span class="n">pt</span><span class="p">.</span><span class="nf">eye</span><span class="p">(</span><span class="n">n</span><span class="p">)</span>
    
    <span class="n">IminusG</span> <span class="o">=</span> <span class="n">I</span> <span class="o">-</span>  <span class="n">γ</span> <span class="o">*</span> <span class="n">G_pt</span>
    <span class="n">tau</span> <span class="o">=</span> <span class="n">ω</span> <span class="o">*</span> <span class="n">IminusG</span> <span class="o">@</span> <span class="n">IminusG</span><span class="p">.</span><span class="n">T</span> <span class="c1"># Precision matrix
</span>
    <span class="n">_</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">MvNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">y</span><span class="sh">"</span><span class="p">,</span>
        <span class="n">mu</span><span class="o">=</span><span class="n">pt</span><span class="p">.</span><span class="nf">dot</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">β</span><span class="p">),</span>
        <span class="n">tau</span><span class="o">=</span><span class="n">tau</span><span class="p">,</span>
        <span class="n">observed</span><span class="o">=</span><span class="n">y</span><span class="p">,</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">with</span> <span class="n">naive_model</span><span class="p">:</span>
    <span class="n">trace_naive</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nf">sample</span><span class="p">()</span>

<span class="k">with</span> <span class="n">graph_model</span><span class="p">:</span>
    <span class="n">trace_graph</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nf">sample</span><span class="p">()</span>
</code></pre></div></div> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><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="nf">subplots</span><span class="p">(</span><span class="mi">2</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">9</span><span class="p">,</span> <span class="mi">7</span><span class="p">),</span> <span class="n">sharex</span><span class="o">=</span><span class="sh">"</span><span class="s">col</span><span class="sh">"</span><span class="p">)</span>

<span class="n">az</span><span class="p">.</span><span class="nf">plot_posterior</span><span class="p">(</span><span class="n">trace_naive</span><span class="p">,</span> <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">],</span> <span class="n">coords</span><span class="o">=</span><span class="p">{</span><span class="sh">"</span><span class="s">β_dim_0</span><span class="sh">"</span><span class="p">:</span> <span class="mi">0</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="mi">0</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">C0</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</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="nf">axvline</span><span class="p">(</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">k</span><span class="sh">"</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="sh">"</span><span class="s">--</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</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="nf">set_title</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Naive Model - β₀ (True value: </span><span class="si">{</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">0</span><span class="p">]</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="sh">"</span><span class="p">)</span>

<span class="n">az</span><span class="p">.</span><span class="nf">plot_posterior</span><span class="p">(</span><span class="n">trace_naive</span><span class="p">,</span> <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">],</span> <span class="n">coords</span><span class="o">=</span><span class="p">{</span><span class="sh">"</span><span class="s">β_dim_0</span><span class="sh">"</span><span class="p">:</span> <span class="mi">1</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="mi">0</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">C0</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="mi">1</span><span class="p">].</span><span class="nf">axvline</span><span class="p">(</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">k</span><span class="sh">"</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="sh">"</span><span class="s">--</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</span><span class="p">[</span><span class="mi">0</span><span class="p">,</span> <span class="mi">1</span><span class="p">].</span><span class="nf">set_title</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Naive Model - β₁ (True value: </span><span class="si">{</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">1</span><span class="p">]</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="sh">"</span><span class="p">)</span>

<span class="n">az</span><span class="p">.</span><span class="nf">plot_posterior</span><span class="p">(</span><span class="n">trace_graph</span><span class="p">,</span> <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">],</span> <span class="n">coords</span><span class="o">=</span><span class="p">{</span><span class="sh">"</span><span class="s">β_dim_0</span><span class="sh">"</span><span class="p">:</span> <span class="mi">0</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="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">C1</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</span><span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">].</span><span class="nf">axvline</span><span class="p">(</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">k</span><span class="sh">"</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="sh">"</span><span class="s">--</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</span><span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">0</span><span class="p">].</span><span class="nf">set_title</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Graph Model - β₀ (True value: </span><span class="si">{</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">0</span><span class="p">]</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="sh">"</span><span class="p">)</span>

<span class="n">az</span><span class="p">.</span><span class="nf">plot_posterior</span><span class="p">(</span><span class="n">trace_graph</span><span class="p">,</span> <span class="n">var_names</span><span class="o">=</span><span class="p">[</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">],</span> <span class="n">coords</span><span class="o">=</span><span class="p">{</span><span class="sh">"</span><span class="s">β_dim_0</span><span class="sh">"</span><span class="p">:</span> <span class="mi">1</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="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">C1</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</span><span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">].</span><span class="nf">axvline</span><span class="p">(</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">color</span><span class="o">=</span><span class="sh">"</span><span class="s">k</span><span class="sh">"</span><span class="p">,</span> <span class="n">linestyle</span><span class="o">=</span><span class="sh">"</span><span class="s">--</span><span class="sh">"</span><span class="p">,</span> <span class="n">label</span><span class="o">=</span><span class="sh">"</span><span class="s">True value of β₁</span><span class="sh">"</span><span class="p">)</span>
<span class="n">axs</span><span class="p">[</span><span class="mi">1</span><span class="p">,</span> <span class="mi">1</span><span class="p">].</span><span class="nf">set_title</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Graph Model - β₁ (True value: </span><span class="si">{</span><span class="n">beta_true</span><span class="p">[</span><span class="mi">1</span><span class="p">]</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="sh">"</span><span class="p">)</span>

<span class="n">plt</span><span class="p">.</span><span class="nf">tight_layout</span><span class="p">()</span>
</code></pre></div></div> <p>Using the <code class="language-plaintext highlighter-rouge">MvNormal</code> likelihood took longer, but the advantages are clear: after accounting for the correlated errors, our coefficient estimates are much better in terms of HDI width as well as distance between the posterior mean and the true value used to produce the dataset.</p> <p>Unfortunately, using the <code class="language-plaintext highlighter-rouge">MvNormal</code> likelihood in this example incurs a \(\mathcal{O}(N^3)\) computation cost from applying the Cholesky decomposition to \(\Sigma\) in order to compute \(\log p(\mathbf{y} \mid \mathbf{\mu}, \Sigma )\).</p> <p>Having set the basic context for modeling, let’s discuss some concrete computational considerations for conducting inference more efficiently:</p> <ul> <li>First, by assuming which elements of \(\mathbf{\epsilon}\) depend on each other in a DAG-structured manner, we can model the covariance of \(\mathbf{\epsilon}\) in terms of a lower-triangular factorized form of \(\Lambda\) and thus \(\Sigma\).</li> <li>The two main sources of computational cost for modeling with a multivariate normal likelihood are (1) the determinant of the covariance matrix, \(\det(\Sigma)\), and (2) the calculation of the inner product \(\mathbf{\epsilon}^T \Lambda \mathbf{\epsilon}\)</li> </ul> <p>We will use the structural assumptions of our model to make these operations cheaper. Beginning with (1), our assumptions for \(\Lambda\) let us write the following by using some basic matrix identities:</p> \[\begin{align*} \det(\Sigma) &amp;= \det(\Lambda)^{-1}\\ &amp;= \det\left((I-\gamma G)^T\Omega(I- \gamma G)\right)\\ &amp;= \det(\Omega) \left(\det(I-\gamma G)\right)^2\\ &amp;= \det(\Omega) \\ &amp;= \prod_i \omega_{ii} \\ \log \det(\Sigma) &amp;= \sum_i \log \omega_{ii} \end{align*}\] <p>This gives us an \(\mathcal{O}(N)\) formula for \(\det(\Sigma)\) which normally requires \(\mathcal{O}(N^3)\) computation with relaxed assumptions. We are fortunate here that \(\det(I-\gamma G) = 1\) as the determinant of a triangular matrix is given by the product of its diagonals; since \(G\) has no diagonal elements and \(I\) is the identity matrix, this product is equal to 1.</p> <p>For item (2) from above, we can rewrite \(\mathbf{\epsilon}^T \Lambda \mathbf{\epsilon}\) as \(\mathbf{u} = \mathbf{\epsilon}^T (I - \gamma G)^T \Omega^{1/2}\) such that \(\mathbf{\epsilon}^T \Lambda \mathbf{\epsilon} = \mathbf{u}^T\mathbf{u}\). Since \(\Omega\) is diagonal, \(\Omega^{1/2}\) is just the elementwise square root of \(\Omega\). Then, since \(G\) is sparse, the dot product \(\mathbf{\epsilon}^T(I-\gamma G)\) will also be relatively cheap.</p> <p>We can also make a few more simplifications. Until now, we assumed that we would use \(n\) distinct values of \(\omega_{ii}\) for modeling. In practice, this means a separate precision parameter for <em>each</em> observation which is unnecessary. We’ll instead work with a single global precision parameter \(\omega\) so that \(\Omega = \omega I\), letting us simplify a few more operations.</p> <p>We can also simplify this model a bit:</p> \[\begin{align*} \epsilon^T \Lambda \epsilon &amp; = \epsilon^T (I - \gamma G)^T\Omega(I - \gamma G) \epsilon \\ &amp; = \epsilon^T (I - \gamma G)^T\omega I(I - \gamma G) \epsilon \\ &amp; = \omega\left( \epsilon^T (I - \gamma G)^T(I - \gamma G) \epsilon \right) \\ &amp; = \omega\left( (\epsilon - \gamma G \epsilon)^T(\epsilon - \gamma G \epsilon)\right) \\ \end{align*}\] <p>This approach is capable of scaling out to a very large number of dimensions, in theory. Below is a work-in-progress implementation to try to do this. Unfortunately, there are some PyTensor issues that arise when trying to sample with this.</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">I</span> <span class="o">=</span> <span class="n">sp</span><span class="p">.</span><span class="n">sparse</span><span class="p">.</span><span class="nf">eye</span><span class="p">(</span><span class="n">n</span><span class="p">,</span> <span class="nb">format</span><span class="o">=</span><span class="sh">"</span><span class="s">csr</span><span class="sh">"</span><span class="p">,</span> <span class="n">dtype</span><span class="o">=</span><span class="nb">float</span><span class="p">)</span>
<span class="n">n</span> <span class="o">=</span> <span class="nf">len</span><span class="p">(</span><span class="n">y</span><span class="p">)</span>
<span class="k">with</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Model</span><span class="p">()</span> <span class="k">as</span> <span class="n">sparse_graph_model</span><span class="p">:</span>
    <span class="n">β</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">Normal</span><span class="p">(</span><span class="sh">"</span><span class="s">β</span><span class="sh">"</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="mi">5</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="n">p</span><span class="p">)</span>
    <span class="n">μ</span> <span class="o">=</span> <span class="n">pt</span><span class="p">.</span><span class="nf">dot</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">β</span><span class="p">)</span>
    <span class="n">ε</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="n">math</span><span class="p">.</span><span class="nf">constant</span><span class="p">(</span><span class="n">y</span><span class="p">)</span> <span class="o">-</span> <span class="n">μ</span>

    <span class="n">ω</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">'</span><span class="s">ω</span><span class="sh">'</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span> <span class="c1"># controls the diagonal entries of the precision matrix
</span>    <span class="n">γ</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nc">HalfNormal</span><span class="p">(</span><span class="sh">"</span><span class="s">γ</span><span class="sh">"</span><span class="p">,</span> <span class="mi">1</span><span class="p">)</span> <span class="c1"># controls the off-diagonal entries of the precision matrix
</span>
    <span class="n">G_pt</span> <span class="o">=</span> <span class="n">sparse</span><span class="p">.</span><span class="nf">as_sparse_or_tensor_variable</span><span class="p">(</span><span class="n">G</span><span class="p">.</span><span class="nf">astype</span><span class="p">(</span><span class="nb">float</span><span class="p">))</span>
    <span class="n">I_pt</span> <span class="o">=</span> <span class="n">sparse</span><span class="p">.</span><span class="nf">as_sparse_or_tensor_variable</span><span class="p">(</span><span class="n">I</span><span class="p">)</span>

    <span class="c1"># Make ε be shape (n,1)
</span>    <span class="n">ε_col</span> <span class="o">=</span> <span class="n">ε</span><span class="p">[:,</span> <span class="bp">None</span><span class="p">]</span>
    <span class="n">Gε</span> <span class="o">=</span> <span class="n">sparse</span><span class="p">.</span><span class="nf">structured_dot</span><span class="p">(</span><span class="n">G_pt</span><span class="p">,</span> <span class="n">ε_col</span><span class="p">)</span>
    <span class="n">γGε</span> <span class="o">=</span> <span class="n">γ</span> <span class="o">*</span> <span class="n">Gε</span>
    <span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Shape of Gε: </span><span class="si">{</span><span class="n">Gε</span><span class="p">.</span><span class="n">shape</span><span class="p">.</span><span class="nf">eval</span><span class="p">()</span><span class="si">}</span><span class="sh">"</span><span class="p">)</span>
    <span class="nf">print</span><span class="p">(</span><span class="sa">f</span><span class="sh">"</span><span class="s">Shape of γGε: </span><span class="si">{</span><span class="n">γGε</span><span class="p">.</span><span class="n">shape</span><span class="p">.</span><span class="nf">eval</span><span class="p">()</span><span class="si">}</span><span class="sh">"</span><span class="p">)</span>

    <span class="n">resid</span> <span class="o">=</span> <span class="n">ε</span> <span class="o">-</span> <span class="n">γGε</span><span class="p">.</span><span class="nf">squeeze</span><span class="p">()</span>

    <span class="c1"># For quadratic form εᵀ(I-γG)ᵀΩ(I-γG)ε
</span>    <span class="c1"># This is equivalent to ω * (ε - γGε)(ε - γGε)ᵀ
</span>    <span class="n">quadratic_form</span> <span class="o">=</span> <span class="n">ω</span> <span class="o">*</span> <span class="n">pt</span><span class="p">.</span><span class="nf">sum</span><span class="p">(</span><span class="n">resid</span><span class="o">**</span><span class="mi">2</span><span class="p">)</span>
        
    <span class="n">logdet</span> <span class="o">=</span> <span class="n">n</span> <span class="o">*</span> <span class="n">pt</span><span class="p">.</span><span class="nf">log</span><span class="p">(</span><span class="n">ω</span><span class="p">)</span>
    <span class="n">logp</span> <span class="o">=</span>  <span class="o">-</span><span class="mf">0.5</span> <span class="o">*</span> <span class="p">(</span><span class="n">n</span> <span class="o">*</span> <span class="n">pt</span><span class="p">.</span><span class="nf">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">logdet</span> <span class="o">+</span> <span class="n">quadratic_form</span> <span class="p">)</span>
    <span class="n">pm</span><span class="p">.</span><span class="nc">Potential</span><span class="p">(</span><span class="sh">'</span><span class="s">likelihood</span><span class="sh">'</span><span class="p">,</span> <span class="n">logp</span><span class="p">)</span>

    <span class="n">trace_sparse_graph</span> <span class="o">=</span> <span class="n">pm</span><span class="p">.</span><span class="nf">sample</span><span class="p">()</span>
</code></pre></div></div> <p>If we look at the two traces, they are very different. I’m not sure what’s causing the discrepancy. If you do know, please tell me!</p> <div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">az</span><span class="p">.</span><span class="nf">summary</span><span class="p">(</span><span class="n">trace_graph</span><span class="p">)</span>
<span class="n">az</span><span class="p">.</span><span class="nf">summary</span><span class="p">(</span><span class="n">trace_sparse_graph</span><span class="p">)</span>
</code></pre></div></div>]]></content><author><name></name></author><summary type="html"><![CDATA[Using PyMC to model data with DAG-structured error correlations]]></summary></entry></feed>