<?xml version="1.0" encoding="utf-8"?><feed xmlns="http://www.w3.org/2005/Atom" ><generator uri="https://jekyllrb.com/" version="3.10.0">Jekyll</generator><link href="https://michaelmoroz.github.io/feed.xml" rel="self" type="application/atom+xml" /><link href="https://michaelmoroz.github.io/" rel="alternate" type="text/html" /><updated>2024-09-12T19:11:47+00:00</updated><id>https://michaelmoroz.github.io/feed.xml</id><title type="html">Mykhailo Moroz</title><subtitle>Computational physics, rendering and other random stuff</subtitle><entry><title type="html">Writing an optimizing tensor compiler from scratch</title><link href="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/" rel="alternate" type="text/html" title="Writing an optimizing tensor compiler from scratch" /><published>2024-09-11T00:00:00+00:00</published><updated>2024-09-11T00:00:00+00:00</updated><id>https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch</id><content type="html" xml:base="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/"><![CDATA[<center>
<table>
  <tr>
    <th>
    <a href="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/#fluid-simulation"><img src="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Demos/fluid_sim.gif?raw=true" height="192px" /></a>
    <a href="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/#fractal-path-tracer"><img src="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Demos/path_tracer.gif?raw=true" height="192px" /></a>
    <a href="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/#n-body-sph-with-a-custom-sphere-rasterizer"><img src="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Demos/n_body.gif?raw=true" height="192px" /></a>
    <a href="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/#texture-embedder-with-small-neural-network"><img src="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Demos/neural_embed.gif?raw=true" height="192px" /></a>
    <a href="https://michaelmoroz.github.io/WritingAnOptimizingTensorCompilerFromScratch/#neural-cellular-automata"><img src="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Demos/nca.gif?raw=true" height="192px" /></a>
    </th>
  </tr>
</table>
</center>
<p>In this blog post I want to talk about the research and development results for a library that I started working on more than a year ago - <a href="https://github.com/MichaelMoroz/TensorFrost">TensorFrost</a>. Under the hood it’s a static optimizing tensor compiler with a focus on being able to do more “shader-like” things while still keeping the ability to do high level linear algebra for ML in Numpy-like syntax with automatic differentiation support. (Click on the example GIF’s for more details!)</p>

<hr />

<p><em>For documentation on basic functionality, read the <a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/README.md">README</a> file in the repo.</em></p>

<ul>
  <li><a href="#so-why-make-a-new-library">So why make a new library?</a></li>
  <li><a href="#architecture">Architecture</a>
    <ul>
      <li><a href="#kernel-fusion">Kernel fusion</a></li>
      <li><a href="#first-prototype">First prototype</a></li>
      <li><a href="#second-prototype">Second prototype</a>
        <ul>
          <li><a href="#optimization-and-generation-of-the-kernels">Optimization and generation of the kernels</a></li>
          <li><a href="#algorithmic-operations">Algorithmic operations</a></li>
          <li><a href="#advanced-kernel-fusion">Advanced kernel fusion</a></li>
          <li><a href="#automatic-differentiation">Automatic differentiation</a></li>
          <li><a href="#ir-under-the-hood">IR under the hood</a></li>
        </ul>
      </li>
    </ul>
  </li>
  <li><a href="#python-frontend">Python frontend</a>
    <ul>
      <li><a href="#main-code">Main code</a></li>
      <li><a href="#host-code">Host code</a>
        <ul>
          <li><a href="#modules">Modules</a></li>
          <li><a href="#optimizer-modules">Optimizer modules</a></li>
        </ul>
      </li>
      <li><a href="#visualization-and-interactivity">Visualization and interactivity</a></li>
    </ul>
  </li>
  <li><a href="#backends">Backends</a>
    <ul>
      <li><a href="#codegen">Codegen</a></li>
      <li><a href="#runtimes">Runtimes</a></li>
    </ul>
  </li>
  <li><a href="#examples-using-tensorfrost">Examples using TensorFrost</a>
    <ul>
      <li><a href="#fluid-simulation">Fluid simulation</a></li>
      <li><a href="#fractal-path-tracer">Fractal path tracer</a></li>
      <li><a href="#texture-embedder-with-small-neural-network">Texture embedder with small neural network</a></li>
      <li><a href="#n-body-sph-with-a-custom-sphere-rasterizer">N-body SPH with a custom sphere rasterizer</a></li>
      <li><a href="#neural-cellular-automata">Neural Cellular Automata</a></li>
    </ul>
  </li>
  <li><a href="#whats-the-current-performance-compared-to-other-tensor-libraries">What’s the current performance compared to other tensor libraries?</a>
    <ul>
      <li><a href="#n-body-simulation">N-body simulation</a></li>
      <li><a href="#mnist-with-a-convolutional-network">MNIST with a convolutional network</a></li>
      <li><a href="#what-about-some-more-advanced-models">What about some more advanced models?</a></li>
    </ul>
  </li>
  <li><a href="#what-is-left-to-do">What is left to do</a></li>
  <li><a href="#conclusion">Conclusion</a></li>
</ul>

<hr />

<p>I started working on this library around 14 months ago, initially I didn’t really plan to do much more than a few matrix operations for an optimization algorithm I wanted to implement in Unity, but there were quite a few things that I wanted to have on top of all of this and it sidetracked me into a writing an entire compiler (hello scope creep 👋).</p>

<p>The thing is, it’s not the first time I tried to make a tensor library, <a href="https://github.com/MichaelMoroz/TensorCL">the first time</a> was a whole 5 years ago and used OpenCL, as I didn’t have an Nvidia GPU at the time. To be honest I’ve been completely unprepared to the magnitude of what it required, and while I did get it to a “somewhat” working state like having basic kernels and somewhat working autodiff using a tape, the lack of good debug tools and actual problems that I wanted to solve pretty much killed it. And for the things that I did want to work on, it was completely unsuited for, as I usually write simulations or graphics, and the overhead of doing a kernel per operation (especially unoptimized kernels) is just too bad to be useful.</p>

<p>Since that time I’ve had a lot of ideas of what I would like a library like that to even look like, and wanted to try working on it again. However I did know, that for this project to survive I would need to make it suitable for projects I usually do, like <a href="https://www.shadertoy.com/user/michael0884">the ones I usually do on Shadertoy</a>. It might seem weird to you as to why I would make a specifically “tensor” library for something that is basically equivalent to writing shaders. But to be honest, shaders are not actually a perfect place for what I do, and a lot of simulation/rendering algorithms can map quite well to high level “tensor-like” operations. While the limitations might force you to come up with creative solutions, for really large or complicated projects it just becomes more of a problem, as it’s very hard to iterate on quickly. This was also one of the main reasons I didn’t really touch ML too much for most of my pet projects (except stuff like <a href="https://www.shadertoy.com/view/DstGDX">Neural Implicit Representations</a>), ML algorithms are usually quite orthogonal to the way you write shaders, usually being split into hundreds of kernel dispatches, while those shader algorithms are effectively just one megakernel most of the time. The only reasonable way to integrate neural networks into those is unrolling the network into a single scalar function, which can be quite unoptimal and limits their size. Not even talking about the fact that training them is completely out of the question. This brings up another problem, shaders don’t have automatic differentiation, which is surprisingly much more useful here than you might think. While its usually used for optimization algorithms like SGD, having the analytic gradient can also be useful for computing normals/curvature of analytic shapes, or forces from potentials in simulations.</p>

<p>So in this library, I hoped to somehow extend the applicability range of a Tensor library to more “shader-like” use cases like rendering and simulations.</p>

<p>And right now I can actually say that at least to some partial degree it did work out. Currently the library is something of a mix of slightly more low-level Numpy-like operations with shader-like control-flow and operations (most of the built-in scalar shader functions are passed through to Python). In terms of where it stands right now, its more low-level than something like JAX or PyTorch, but still not as low-level as Taichi as you still technically operate on something similar to tensors.</p>

<center><img src="/images/high_low.PNG" height="150px" /></center>

<h1 id="so-why-make-a-new-library">So why make a new library?</h1>

<center><a href="https://xkcd.com/927/"><img src="/images/standards.png" height="250px" /></a></center>

<p>It will indeed take a inordinate amount of work to make a library from scratch and get it to a useful state, as I’ve already experienced. Why would I not just use an existing library, as there are seemingly thousands of them? There are a few reasons mostly applicable to my use cases which make both using pure ML libraries or pure shaders annoying.</p>

<p><strong><em>1. Performance scales poorly for non-ML specific operations</em></strong></p>

<p>Of course, nothing stops me from using your usual ML libraries like PyTorch or JAX, but as I mentioned before, they weren’t really designed to be applied to problems that I have. These libraries effectively live in their own realm with their own rules and syntax fine tuned for ML, and hide some of the features the GPU has from the user, they pretty much have 0 crossover with how shaders operate. While technically you can write any algorithm you want in plain tensors, including graphics and simulations, depending on the number or complexity of operations - the performance could get terrible.</p>

<p>Most of the Tensor compiler research that I’ve seen focuses on ML bottlenecks, like efficiently utilizing cache, correctly aligning data for maximum performance of matrix multiplications, convolutions etc. Those usually aren’t a bottleneck when dealing with simulations or rendering, the bottleneck becomes the dynamic nature of the code and its complexity.</p>

<p><strong><em>2. Dynamic control flow is very tedious (if possible) to implement in classic ML libraries</em></strong></p>

<p>The performance is actually not the only issue when writing simulations or graphics, control flow can be very prevalent, but unfortunately it’s quite inconvenient to express in these libraries, if even possible, as native loops in JAX for example, require doing stuff like:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">factorial</span><span class="p">(</span><span class="n">n</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">cond_fun</span><span class="p">(</span><span class="n">state</span><span class="p">):</span>
        <span class="n">i</span><span class="p">,</span> <span class="n">fact</span> <span class="o">=</span> <span class="n">state</span>
        <span class="k">return</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">n</span>

    <span class="k">def</span> <span class="nf">body_fun</span><span class="p">(</span><span class="n">state</span><span class="p">):</span>
        <span class="n">i</span><span class="p">,</span> <span class="n">fact</span> <span class="o">=</span> <span class="n">state</span>
        <span class="k">return</span> <span class="n">i</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">fact</span> <span class="o">*</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">initial_state</span> <span class="o">=</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">final_state</span> <span class="o">=</span> <span class="n">jax</span><span class="p">.</span><span class="n">lax</span><span class="p">.</span><span class="n">while_loop</span><span class="p">(</span><span class="n">cond_fun</span><span class="p">,</span> <span class="n">body_fun</span><span class="p">,</span> <span class="n">initial_state</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">final_state</span><span class="p">[</span><span class="mi">1</span><span class="p">]</span>
</code></pre></div></div>

<p>Which looks completely cursed. Now imagine multiple stacked loops - yikes</p>

<p>Such native loops are required if you want a varying iteration count for some operation, in simulations this could be, for example, summing a force from a varying number of particle neighbors. Even something like Gaussian Splatting requires a variable loop per tile. From a purely classic ML standpoint, this isn’t really a problem, since such cases practically never happen and the computational graph is absolutely static, and whats worse, autodiff gradients of such loops might potentially have atrocious performance (or might just be uncomputable if the loop can be theoretically infinite, as the compiler might not know the specific context where its being used).</p>

<p>You could alternatively write a dynamic mask that will depend on the iteration, and unroll the loop, but this would simply be slower and less readable. Like here, I once tried to make a vectorized Numpy function to computes the mandelbulb SDF:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">mandelbulb_sdf</span><span class="p">(</span><span class="n">pos</span><span class="p">,</span> <span class="n">iter_num</span><span class="o">=</span><span class="n">mandelbulb_iter_num</span><span class="p">,</span> <span class="n">power</span><span class="o">=</span><span class="n">mandelbulb_power</span><span class="p">):</span>
    <span class="n">z</span> <span class="o">=</span> <span class="n">pos</span>
    <span class="n">dr</span> <span class="o">=</span> <span class="mf">1.0</span>
    <span class="n">r</span> <span class="o">=</span> <span class="mf">0.0</span>
    <span class="k">for</span> <span class="n">i</span> <span class="ow">in</span> <span class="nb">range</span><span class="p">(</span><span class="n">iter_num</span><span class="p">):</span>
        <span class="n">r</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">linalg</span><span class="p">.</span><span class="n">norm</span><span class="p">(</span><span class="n">z</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">not_mask</span> <span class="o">=</span> <span class="o">~</span><span class="p">(</span><span class="n">r</span> <span class="o">&gt;</span> <span class="mf">1.5</span><span class="p">)</span>
        
        <span class="n">theta</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">arccos</span><span class="p">(</span><span class="n">z</span><span class="p">[...,</span> <span class="mi">2</span><span class="p">]</span> <span class="o">/</span> <span class="n">r</span><span class="p">)</span>
        <span class="n">phi</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">arctan2</span><span class="p">(</span><span class="n">z</span><span class="p">[...,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">z</span><span class="p">[...,</span> <span class="mi">0</span><span class="p">])</span>
        <span class="n">dr</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">where</span><span class="p">(</span><span class="n">not_mask</span><span class="p">,</span> <span class="n">r</span><span class="o">**</span><span class="p">(</span><span class="n">power</span> <span class="o">-</span> <span class="mf">1.0</span><span class="p">)</span> <span class="o">*</span> <span class="n">power</span> <span class="o">*</span> <span class="n">dr</span> <span class="o">+</span> <span class="mf">1.0</span><span class="p">,</span> <span class="n">dr</span><span class="p">)</span>
        
        <span class="n">zr</span> <span class="o">=</span> <span class="n">r</span><span class="o">**</span><span class="n">power</span>
        <span class="n">theta</span> <span class="o">*=</span> <span class="n">power</span>
        <span class="n">phi</span> <span class="o">*=</span> <span class="n">power</span>
        
        <span class="n">z</span> <span class="o">=</span> <span class="n">np</span><span class="p">.</span><span class="n">where</span><span class="p">(</span><span class="n">not_mask</span><span class="p">[:,</span> <span class="p">:,</span> <span class="p">:,</span> <span class="n">np</span><span class="p">.</span><span class="n">newaxis</span><span class="p">],</span> <span class="n">pos</span> <span class="o">+</span> <span class="n">zr</span><span class="p">[:,</span> <span class="p">:,</span> <span class="p">:,</span> <span class="n">np</span><span class="p">.</span><span class="n">newaxis</span><span class="p">]</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">array</span><span class="p">([</span><span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">theta</span><span class="p">)</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">cos</span><span class="p">(</span><span class="n">phi</span><span class="p">),</span> <span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">theta</span><span class="p">)</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">phi</span><span class="p">),</span> <span class="n">np</span><span class="p">.</span><span class="n">cos</span><span class="p">(</span><span class="n">theta</span><span class="p">)]).</span><span class="n">transpose</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="mi">2</span><span class="p">,</span> <span class="mi">3</span><span class="p">,</span> <span class="mi">0</span><span class="p">),</span> <span class="n">z</span><span class="p">)</span>
    <span class="k">return</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="n">np</span><span class="p">.</span><span class="n">log</span><span class="p">(</span><span class="n">r</span><span class="p">)</span> <span class="o">*</span> <span class="n">r</span> <span class="o">/</span> <span class="n">dr</span>
</code></pre></div></div>

<p><strong><em>3. Extending ML libraries with custom high-performance kernels requires using external tools/APIs like CUDA or Triton</em></strong></p>

<p>Usually, when hitting a performance bottleneck, you have to write custom CUDA kernels, which is doable, but quite inconvenient, as it forces you to use separate environments, CUDA and Python.</p>

<p>There are actually domain specific languages (DSLs) that allow you to write kernels in relatively high-level Python, like <a href="https://www.taichi-lang.org/">Taichi</a> or <a href="https://github.com/NVIDIA/warp">Nvidia’s Warp</a>, they are really nice for simulations or graphics, and Taichi specifically is fine tuned for high performance physics simulations and even has built in automatically optimized sparse grids. And while they are also differentiable, they still are a bit off from what I would consider “perfect”, as you don’t have access to the giant library of ML operations that PyTorch has for example. Of course you can still write ML using them if you really want, there are some<a href="https://github.com/taichi-dev/taichi-nerfs"> Nerf implementations for Taichi</a>, but they require much more code to represent. As a compromise you could also interoperate them with PyTorch for example, but once again, it makes it less convenient to work with. I can also mention <a href="https://github.com/triton-lang/triton">Triton</a> here, but it’s usually used more like a backend for other libraries (like PyTorch) rather than a standalone DSL, at least as far as I have seen.</p>

<p>There is also <a href="https://github.com/shader-slang/slang">Slang</a>, which is quite different from all these from above, as its an improved shader language with added differentiability. It would be nice if it was widely supported. But its even more low level than something like Taichi.</p>

<p><strong><em>4. ML libraries don’t have an easy built-in way to make real-time visualizations with optional interactivity</em></strong></p>

<p>When I want to do some advanced visualizations in Python, the options that are available in ML libraries are often hilariously bad. Usually you just end up making a bunch of matplotlib plots, which are not only slow, if you want to render like hundreds of millions of points or an animation, but also not interactive. (They are fine for papers tho)</p>

<p>In the world of real-time graphics, you can render those points at 60fps, and you could even interact with them in real time. Seeing the things you are working on in real time is in some cases quite useful for quick iteration, and I feel like this is something you miss in your classic ML libs, where the usual interaction you have with your model - is staring at the loss graph for hours.</p>

<p>Though, while I am stating these things, most large ML models are simply not visualizable in real time, and the ones that are, are usually not easy to usefully interpret. Visualizations are usually most applicable to the intersection of ML/Physics/Graphics, like NERFs, diffusion, neural implicit representations, etc. But I still think that even changing hyperparameters in real time and seeing its result on the training loss can also be somewhat interesting, though you do need the model to be rather performant for that.</p>

<p>The lack of a native way to output graphical data from these libraries is even more annoying when you remember that GPU’s are called <strong>Graphics</strong> Processing Units, not Tensor Processing Units. And they have all the required hardware to work with and output graphics.</p>

<p><em>PS. Taichi actually does have a way to this! It has integration with GLFW and ImGUI.</em></p>

<p><strong><em>5. Writing simulations or graphics in a high-level language is could be much easier to iterate on rather than in pure shaders</em></strong></p>

<p>On the other side, in the world of real-time simulations and graphics, I’ve written custom kernels for every specific algorithm something needed: radix sorts, linear solvers, numerical integrators, etc. So when I’m prototyping or having a new idea how to optimize the algorithm globally, it can get annoying to make global changes in the code structure, since they usually require a partial rewrite, creating new kernels and so on, and I don’t really see why this couldn’t be automated from higher-level operations.</p>

<hr />

<p>In the end I was wondering: can I somehow combine the best of both worlds? 
Being able to do both Numpy-like operations while also doing more shader-like things in one place sounds somewhat impossible on paper, but I thought that maybe if you tuned the kernel generation algorithm to specifically be optimal for shader-like operation it might at least work for my use cases? Afterwards I could still support ML use cases well enough even if I just shoehorned the matmul/reduction/convolution kernels separately.</p>

<p>And even if it is impossible, combining everything into a single language would be nice, because the GPU development infrastructure is scattered all over the place, some features are available only in some places, while others are not, they can sometimes not even be interoperable. The development environments are completely separate, there are ML-specific debug tools on Linux, and graphics API specific debug tools on Windows only - and all of these things are using the same hardware - the GPU!</p>

<h1 id="architecture">Architecture</h1>

<h2 id="kernel-fusion">Kernel fusion</h2>

<p>When thinking about the architecture of what a library like this could be, I’ve first thought from the point of view of simple per-element operations and indexing. Nothing really stops you from writing Numpy code as if it was a shader, for example. However that would be extremely slow, why? Numpy usually executes a single operation over the entire tensor at a time, i.e. loads the data, does the single operation, then puts it back into memory. For a modern computer this is <em>extremely</em> inefficient, since the bandwidth and latency of system memory are nowhere near sufficient to keep up with the computational power of the processor. If you want to utilize all available resources - you need to utilize the cache hierarchy to its full extent. If you could cache intermediate results close to the ALU’s of the processor and keep the intermediate results as long as possible there, you could get a massive speedup by reducing latencies by orders of magnitude. This is why modern processors have multiple levels of cache - they try to solve the issue of memory improvements not keeping up with increasing performance. As a matter of fact since 2017 top tier consumer GPU memory bandwidth has only increased from 500Gb/s to 1Tb/s, while performance has spiked from 10 TFlops to 80 TFlops.</p>

<center>
<table>
  <tr>
    <th><img src="/images/cache_hierarchy.png" height="250px" /></th>
  </tr>
  <tr>
    <th><i>Tiers of memory of a typical GPU</i></th>
  </tr>
</table>
</center>

<p>But optimizing this specific aspect when dealing with tensors is actually nothing new, the so called kernel fusion has been around for a while, and is used in tensor compilers like XLA or PyTorch’s <a href="https://dev-discuss.pytorch.org/t/torchinductor-a-pytorch-native-compiler-with-define-by-run-ir-and-symbolic-shapes/747/3">TorchInductor</a>. And the degree to which they can do it is perfectly fine for most ML applications. However, to my knowledge, they do not fuse operations with already existing optimized kernels, even though in some cases this might be beneficial (Well, in TorchInductor there are only 2 operations that don’t - conv and matmul).</p>

<p>If you tried to write algorithms, like the ones I write in Shadertoy, you will eventually start to hit the limits these compilers have. The number of operation nodes you could fuse now rises to the order of thousands, not even mentioning the complex control flow, and ideally they should fit into a single kernel, but its highly likely you will end up with a lot of smaller kernels if you apply fusion naively.</p>

<h2 id="first-prototype">First prototype</h2>

<p>When I initally started prototyping the operation graph compiler in C# in Unity (not exactly your typical place for a compiler prototype, I know), I kept the graph just as a simple Directed Acyclic Graph (DAG), where each node was a single operation with some tensor shape. When I began testing the kernel clustering even on simple physics or rendering, the clustering algorithm quickly started to get out of hand.</p>

<p>Here is an example of a compiled graph with the kernels clusterized.</p>

<center><img src="/images/fluidgraph.png" height="400px" /></center>

<p>This was the operation/kernel cluster graph for this fluid simulation:</p>

<center><iframe width="560" height="315" src="https://www.youtube.com/embed/gRZzMPo1RLg?si=UfZ9HGTNYV51tBtn" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen=""></iframe></center>

<p>It did work, unfortunately there was no good way to procedurally generate shader code in Unity, so I did something rather stupid and implemented a Virtual Machine right in the shader, which made it particularly slow. However, of the good things, there was practically 0 compile time, but unforuntately that was not the goal I had in mind. The way I did reductions to get the average energy here was also rather questionable - I did it purely with atomics. For some cases they are good, but reduction is not one of them. Especially given that there are no natively supported float atomics in HLSL, so you need to emulate them through <code class="language-plaintext highlighter-rouge">InterlockedExchangeCompare</code>, which makes it even slower for this particular use case. One time it even crashed my computer when trying to add a few million floats in parallel to a single element.</p>

<p>The VM basically only knows of 1 dimensional loads and stores, so I also needed to add a compilation stage that converts multidimensional indices into operations that compute the “flattened” index.</p>

<p>At this point the graph had only a few operation types, and I implemented backward-mode autodiff here, which was surprisingly easy. The only not super obvious thing initially was load/store gradients. But those are effectively just atomic_add/load respectively. So in theory, you could implement any ML algorithm even here, but that would be comically slow. Matrix multiplication gradient would be 2 atomic adds per thread in a 3D kernel.</p>

<p>Another thing that turned out to be a huge problem was that the operation graph did not have a “stable” order, it was only topologically sorted (I actually resorted it every time I did something with the graph!), which is good for normal operations, but for inplace operations like stores and atomics this leads to randomly varying results, which is a no-go.</p>

<p>Adding the problem of exponentially growing number of possible kernel fusions and the graphs quicky became an undebuggable node soup, which made this approach not very appealing:</p>

<center><img src="/images/wtf.png" height="400px" /></center>

<p>I also wanted to have at least some sort of control flow, which wasn’t obvious how to add into this specific graph for me, at the time.</p>

<h2 id="second-prototype">Second prototype</h2>

<blockquote>
  <p>Any sufficiently complicated C or Fortran program contains an ad hoc, informally-specified, bug-ridden, slow implementation of half of Common Lisp.</p>
</blockquote>

<p>This time I wrote it in C++ and decided to instead enforce the order of operations to make in-place operations have properly specified order. I constructed the IR like a linked list, order of which is taken from the way the code is parsed. This simplifies the kernel fusion problem to one of finding ranges of fusable operations, and in my case I fuse everything except any node pair that violated rules based on order of memory access, shape of the operations, etc. In some cases this effectively can convert the entire graph into a single kernel, which is exactly what I’m looking for.</p>

<p>Ordering the operations is not the only thing. To implement control flow I needed to have child and parent nodes while still keeping a uniquely specified ordering, I implemented this by using a multi-level linked list. This way operaitons controlled by a conditional statement or by loops become its children.
In multi-level linked lists kernel fusion becomes slightly more tricky, but effectively its just a recursive process of finding the ranges of fusable operations starting from the lowest level of the list, and fusing them level by level.</p>

<center><img src="/images/multilevel.png" height="300px" /></center>

<p>Having children or parents for a tensor is a rather unusual notion, but it is totally fine as long as you make sure that the shape of this node is broadcastable to all of its parents. So you can, for example, have a scalar loop with children of any shape. Or, lets look at a more complex example, a loop of shape [N, 1], with a child loop of shape [] with a child load operation of shape [M]. Totally different shapes, but they are all broadcastable between each other so the shape a generated kernel here will be [N, M] for all these nodes. There are also cases when you have broadcastable shape between parents and children and incompatible shape between neighbors - those are simply incorrect and should lead to a compilation error, as there is no way to generate a kernel here without complex masking, which I don’t plan to support now. Example would be a loop of shape [N] with 2 children of shape [5, N] and [2, N].</p>

<p>The kernels are also represented in the same IR, as children of a “kernel” node. At the code generation stage, everything outside of the kernel nodes is converted into host code (right now C++), and the kernels are converted into device code (OpenMP C++ or GLSL/HLSL). This way the entire program is represented in a single IR, and the compiler can optimize it globally both for CPU and GPU parts.</p>

<p>Since the IR is the same for all compilation stages, you could for example input both high level tensor operations together with explicitly specified kernels, and it can already be done like this:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">A</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">([</span><span class="o">-</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">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
<span class="n">B</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">(</span><span class="n">A</span><span class="p">.</span><span class="n">shape</span><span class="p">,</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="n">C</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">buffer</span><span class="p">(</span><span class="n">A</span><span class="p">.</span><span class="n">shape</span><span class="p">,</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">kernel</span><span class="p">(</span><span class="n">A</span><span class="p">.</span><span class="n">shape</span><span class="p">)</span> <span class="k">as</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">C</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="n">A</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="n">B</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">C</span> <span class="o">=</span> <span class="p">(</span><span class="n">C</span> <span class="o">@</span> <span class="n">A</span><span class="p">.</span><span class="n">T</span><span class="p">)</span> <span class="o">+</span> <span class="n">tf</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">B</span> <span class="o">@</span> <span class="n">A</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">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span>
</code></pre></div></div>

<p>This is already quite nice. (I’ll explain the Python syntax a bit later)</p>

<p><em>(Note: of course, if you tried to compute the gradient here, the compiler would fail, at least at this point in time, as general gradients over control flow are not trivial)</em></p>

<p>Another interesting aspect of this representation is that kernels can be created as children of control flow nodes, meaning you can create a loop of kernels for an iterative algorithm, like for example a bitonic sort!</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">log2N</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">ceil</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">log2</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">float</span><span class="p">(</span><span class="n">element_count</span><span class="p">)))</span>
<span class="n">Nround</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">int</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">exp2</span><span class="p">(</span><span class="n">log2N</span><span class="p">))</span>
<span class="n">sort_id</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">indices</span><span class="p">([</span><span class="n">Nround</span><span class="o">/</span><span class="mi">2</span><span class="p">])[</span><span class="mi">0</span><span class="p">]</span>
<span class="n">steps</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">int</span><span class="p">(</span><span class="n">log2N</span><span class="o">*</span><span class="p">(</span><span class="n">log2N</span> <span class="o">+</span> <span class="mf">1.0</span><span class="p">)</span><span class="o">/</span><span class="mf">2.0</span><span class="p">)</span>

<span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">loop</span><span class="p">(</span><span class="n">steps</span><span class="p">)</span> <span class="k">as</span> <span class="n">step</span><span class="p">:</span>
    <span class="n">j</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">floor</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">float</span><span class="p">(</span><span class="mi">2</span><span class="o">*</span><span class="n">step</span><span class="p">)</span> <span class="o">+</span> <span class="mf">1.0</span><span class="p">)</span> <span class="o">-</span> <span class="mf">0.5</span><span class="p">)</span>
    <span class="n">n</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">round</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">float</span><span class="p">(</span><span class="n">step</span><span class="p">)</span> <span class="o">-</span> <span class="mf">0.5</span><span class="o">*</span><span class="n">j</span><span class="o">*</span><span class="p">(</span><span class="n">j</span><span class="o">+</span><span class="mf">1.0</span><span class="p">))</span>
    <span class="n">B</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">int</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">round</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">exp2</span><span class="p">(</span><span class="n">j</span><span class="o">-</span><span class="n">n</span><span class="p">)))</span>
    <span class="n">mask</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">select</span><span class="p">(</span><span class="n">n</span> <span class="o">&lt;</span> <span class="mf">0.5</span><span class="p">,</span> <span class="mi">2</span><span class="o">*</span><span class="n">B</span> <span class="o">-</span> <span class="mi">1</span><span class="p">,</span> <span class="n">B</span><span class="p">)</span>
    <span class="n">e1</span> <span class="o">=</span> <span class="n">sort_id</span><span class="o">%</span><span class="n">B</span> <span class="o">+</span> <span class="mi">2</span><span class="o">*</span><span class="n">B</span><span class="o">*</span><span class="p">(</span><span class="n">sort_id</span><span class="o">/</span><span class="n">B</span><span class="p">)</span>
    <span class="n">e2</span> <span class="o">=</span> <span class="n">e1</span> <span class="o">^</span> <span class="n">mask</span>

    <span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">if_cond</span><span class="p">((</span><span class="n">e1</span> <span class="o">&lt;</span> <span class="n">element_count</span><span class="p">)</span> <span class="o">&amp;</span> <span class="p">(</span><span class="n">e2</span> <span class="o">&lt;</span> <span class="n">element_count</span><span class="p">)):</span>
        <span class="n">key1</span><span class="p">,</span> <span class="n">key2</span> <span class="o">=</span> <span class="n">keys</span><span class="p">[</span><span class="n">e1</span><span class="p">],</span> <span class="n">keys</span><span class="p">[</span><span class="n">e2</span><span class="p">]</span>

        <span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">if_cond</span><span class="p">(</span><span class="n">key1</span> <span class="o">&lt;</span> <span class="n">key2</span><span class="p">):</span>
            <span class="n">val1</span><span class="p">,</span> <span class="n">val2</span> <span class="o">=</span> <span class="n">values</span><span class="p">[</span><span class="n">e1</span><span class="p">],</span> <span class="n">values</span><span class="p">[</span><span class="n">e2</span><span class="p">]</span>
            <span class="n">keys</span><span class="p">[</span><span class="n">e1</span><span class="p">]</span> <span class="o">=</span> <span class="n">key2</span>
            <span class="n">keys</span><span class="p">[</span><span class="n">e2</span><span class="p">]</span> <span class="o">=</span> <span class="n">key1</span>
            <span class="n">values</span><span class="p">[</span><span class="n">e1</span><span class="p">]</span> <span class="o">=</span> <span class="n">val2</span>
            <span class="n">values</span><span class="p">[</span><span class="n">e2</span><span class="p">]</span> <span class="o">=</span> <span class="n">val1</span>
</code></pre></div></div>

<p>In this case the compiler can spot that you can’t fuse the insides of the loop since they are reading and writing from the same memory, thus creating a kernel region under the loop.</p>

<p>However, while you can do that, in this particular case, I woudn’t recommend relying on the compiler too much, and would put an explicit kernel under the loop, since even changing the loading order from before the stores, to the middle, will split the kernel in 2 and potentially break it right now.</p>

<p>I had <a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Simulation/n-body.ipynb">one specific case</a> when it was an issue when I tried to optimize a software sphere rasterizer, by adding an additional read to check if the atomic min can be skipped I effectively made the compiler think that this part of the code needs to be split into parts and simply just broke the rendering.</p>

<p>This shows that, while powerful, such inference of how the program is structured does not always work, or requires a much more advanced compiler than I have here.
On the other hand, in the case of <a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Simulation/fluid_simulation.ipynb">the 2D fluid simulation example</a>, kernel generation worked quite well, since there is no control flow to confuse the compiler.</p>

<h3 id="optimization-and-generation-of-the-kernels">Optimization and generation of the kernels</h3>

<p>Simply splitting the IR into kernel regions is actually not enough to make sure you dont have a million unneeded loads or stores. One very basic way to optimize the kernels is effectively just copying computations that are cheaper to do than to load from memory. Such things like constants, or very simple arithmetic are examples of this. Doing this optimization not only reduces the number of global memory access a lot usually, but also removes a lot of unneeded kernels that would have just stored a constant or something similar into memory.</p>

<p>I also have the standard “removing of unused computation” here. Since we have the entire IR graph from input to ouput given, this additionally allows to figure out which parts of the computation are influencing the outputs. So I can effectively assume that everything else is unused and can be simply removed.</p>

<p>When generating compute kernels out of such an IR, you can not simply use the N dimensional shape of the kernel, and you need to map the tensor indices to the specific layout of the GPU. In this case, that’s the group indices, and the workgroup thread indices (in more advanced cases, in DX12/Vulkan/CUDA/etc there is also the warp sub-group, but I’ll ignore it for now). To do this mapping we ideally would need to figure out the shape of the workgroup, which must be a compile time constant, from the computations we do. But at the moment I simply estimate the group shape from the last 3 dimensions of the kernel, and clamp them to some predefined constants depending on dimensionality. This is suboptimal, but doing it better would either require having a VM that estimates the range of indices of memory accesses, or an autotuner. The first will take some time to implement, and is in my TODO list, as its also useful for other things, and the second, while easier I’m not currenty considering, as it could increase compile times quite significantly. (And they already reach 30 seconds on Windows for my Variational Monte Carlo solver!)</p>

<h3 id="algorithmic-operations">Algorithmic operations</h3>

<p>The bare IR does not know about complex operations like matrix multiplication, reductions, etc. These are implemented as a special compiler pass, that converts the high level operation nodes into a series of simpler operations. For example reduction sums <code class="language-plaintext highlighter-rouge">tf.sum</code> are converted into a loop of adds, matrix multiplications into a loop of muls and adds, etc.</p>

<p>These simpler operations aren’t yet put into kernels, so the compiler can do additional kernel fusion optimizations on them. As they are written in the same IR and not as separate outside kernels also means that all possible optimizations that that the compiler has - can be applied to them, like inserting simple arithmetic instead of memory loads (useful for operations over procedural data, like random), or adding the activation at the end of the matrix multiplication loop, etc.</p>

<p>(Though to be fair, right now, the compiler optimizes them a bit too aggressively, and can put a lot of needless computation inside a matmul loop for example, this needs to be fixed in the future with better heuristics. Detailed example in <code class="language-plaintext highlighter-rouge">IR under the hood</code> section)</p>

<h3 id="advanced-kernel-fusion">Advanced kernel fusion</h3>

<p>Kernel fusion by splitting into ranges of the multi-level linked list works fine until you get to more complex situations, that do actually happen quite often in ML, for example, reductions of some expressions.</p>

<p>To optimize even these cases you can do something I call tensor load fusion - you replace a loading operation with the recomputed result of the load target with the given load indices replacing the target kernel indices.</p>

<p>This was tricky to implement, but it allows to fuse operations like <code class="language-plaintext highlighter-rouge">tf.sum(A[i,k]*B[k,j], axis=2)</code> into a single 2D kernel, instead of 3D + 2D kernels. This gives a massive speedup in some cases when writing operations in such a naive way. For example this allows you to write out a <code class="language-plaintext highlighter-rouge">conv2d</code> operation in a similar form as a sum over a 6D tensor, and the compiler could fuse it into a single 5D reduction + 4D reduction, while giving comparable performance to native pytorch (for small kernels). Here is an implementation in TensorFrost which is effectively equivalent to PyTorch’es <code class="language-plaintext highlighter-rouge">conv2d</code> without padding:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">conv2d</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">W</span><span class="p">):</span>
    <span class="n">N</span><span class="p">,</span> <span class="n">CIN</span><span class="p">,</span> <span class="n">HI</span><span class="p">,</span> <span class="n">WI</span> <span class="o">=</span> <span class="n">X</span><span class="p">.</span><span class="n">shape</span>
    <span class="n">COUT</span><span class="p">,</span> <span class="n">CIN</span><span class="p">,</span> <span class="n">h</span><span class="p">,</span> <span class="n">w</span> <span class="o">=</span> <span class="n">W</span><span class="p">.</span><span class="n">shape</span>
    <span class="n">bi</span><span class="p">,</span> <span class="n">cout</span><span class="p">,</span> <span class="n">wi</span><span class="p">,</span> <span class="n">hi</span><span class="p">,</span> <span class="n">cin</span><span class="p">,</span> <span class="n">it</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">indices</span><span class="p">([</span><span class="n">N</span><span class="p">,</span> <span class="n">COUT</span><span class="p">,</span> <span class="n">HI</span> <span class="o">-</span> <span class="n">h</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">WI</span> <span class="o">-</span> <span class="n">w</span> <span class="o">+</span> <span class="mi">1</span><span class="p">,</span> <span class="n">CIN</span><span class="p">,</span> <span class="n">h</span> <span class="o">*</span> <span class="n">w</span><span class="p">])</span>
    <span class="n">i</span><span class="p">,</span> <span class="n">j</span> <span class="o">=</span> <span class="n">it</span><span class="o">%</span><span class="n">w</span><span class="p">,</span> <span class="n">it</span><span class="o">/</span><span class="n">w</span>
    <span class="k">return</span>  <span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">X</span><span class="p">[</span><span class="n">bi</span><span class="p">,</span> <span class="n">cin</span><span class="p">,</span> <span class="n">wi</span> <span class="o">+</span> <span class="n">i</span><span class="p">,</span> <span class="n">hi</span> <span class="o">+</span> <span class="n">j</span><span class="p">]</span> <span class="o">*</span> <span class="n">W</span><span class="p">[</span><span class="n">cout</span><span class="p">,</span> <span class="n">cin</span><span class="p">,</span> <span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">]))</span>
</code></pre></div></div>

<p>At the moment only the first <code class="language-plaintext highlighter-rouge">tf.sum(some operations)</code> gets fused into a single kernel. Which is actually completely fine, as we dont want to merge multiple <code class="language-plaintext highlighter-rouge">sum</code> together and we want them to be staged. This is actually why I manually fused the kernel dimensions <code class="language-plaintext highlighter-rouge">w</code> and <code class="language-plaintext highlighter-rouge">h</code> together, as you need to have a balance between the size of the loop and number of reduction stages. This could theoretically be done automatically, but it would require autotune or more advanced heuristics.</p>

<h3 id="automatic-differentiation">Automatic differentiation</h3>

<p>Looking at the example in the section above you might think that autograd would completely fail when differentiating such an expression. Which would be the case if done naively. Whats worse, in the programs I write, loads at addresses are extremely prevalent. If you apply autodiff dirrectly on multidimensional loads you get multidimensional atomic adds. In general, you can’t really do much with them, but in most cases, the indices of the atomic adds are simple, or are constants. 
In those cases you could check what dimension indices of the operation are not used in the atomic add address computation, and conclude that all threads of this atomic add for this dimension add to the same element. This would mean that we can optimize this by transforming this dimension into a sum over it, with a following atomic add to this element in the end. Even more, you could also check if the indices map 1 to 1, meaning that they do not make conflicting writes, and can just be replaced with a simple load, add and store operation. (However, I don’t actually do this at the moment as nonconflicting atomic adds are cheap enough to be ignored for now)</p>

<p>These optimizations improve performance of automatically differentiated operations of this kind by <em>a lot</em>, and makes prototyping differentiable convolution-like algorithms quite easy. (I’m interested in trying to fine tune a differentiable multi-grid Poisson equation solver, with boundary conditions, by making lots of small convolutions like these)</p>

<p>The rest of the automatic gradient algorithm is your run-of-the-mill backwards mode autodiff. The autodiff pass is before the algorithm insertion pass, so the gradients are computed for high-level operations if possible, as they are usually cheaper and more numerically stable. (And also I don’t have a way to compute gradients of control-flow at the moment, so it works as a substitute for that for now)</p>

<p>Right now gradiets are specified like <code class="language-plaintext highlighter-rouge">tf.grad(a,b)</code>, and actually, <code class="language-plaintext highlighter-rouge">a</code> doesn’t need to be a scalar, the default vector jacobian product (VJP) input is always a 1 no matter the dimensionality of a. This is useful, for instance, when computing gradients of a potential, for instance:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">dx</span> <span class="o">=</span> <span class="n">x1</span> <span class="o">-</span> <span class="n">x2</span>
<span class="n">dist</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">dx</span><span class="o">**</span><span class="mi">2</span><span class="p">))</span>
<span class="n">pot</span> <span class="o">=</span> <span class="mf">1.0</span> <span class="o">/</span> <span class="n">dist</span>
<span class="n">force</span> <span class="o">=</span> <span class="o">-</span> <span class="n">tf</span><span class="p">.</span><span class="n">grad</span><span class="p">(</span><span class="n">pot</span><span class="p">,</span> <span class="n">dx</span><span class="p">)</span>
</code></pre></div></div>

<p>Quite nice when doing a particle simulation. In fact I also use this when computing the normals for the SDF in my path tracing example.
I should note that this is only valid behaviour because these computations are independent, for pixels, or for particles, if they were depending on each other, the gradient would be invalid, and in that case you should use a scalar <code class="language-plaintext highlighter-rouge">a</code>.</p>

<p>The compiler when doing the autodiff searches for all unique <code class="language-plaintext highlighter-rouge">a</code>’s and does a full backprop over their dependencies, then all unused gradients are simply removed after this compilation pass.</p>

<p>I still plan to implement forward mode automatic differentiation, in its case its somewhat easier to implement. Since for example I can just do it after all the algorithmic passes, as the form of the computation will be exactly the same as the original, just with slightly different operations.</p>

<p>It does pose the question of what to do when doing a hybrid autodiff, like backward grad of forward grad. In that case the gradients need to be sorted by their order, and done one by one. Unfortunately in this case I would need to implement full jacobian vector product (JVP) (on top of the VJP’s) for all algorithmic operations, not just the base simple non-algorithmic ones, so I’ll probably leave that for the far future.</p>

<h3 id="ir-under-the-hood">IR under the hood</h3>

<p>Let’s look at how the IR looks in the compiler right now. We will look at the bitonic sort example from above to see how control flow is represented in the IR.</p>

<details>
<summary>Parsed/traced input</summary>

<div>

    <div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">int</span> <span class="n">v1_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">4294967295</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">element_count</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">keys</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="n">OutputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_1</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">4294967295</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_2</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">values</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="n">OutputMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_2</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_3</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_4</span> <span class="o">=</span> <span class="n">log2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">log2N</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_4</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_5</span> <span class="o">=</span> <span class="n">exp2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">Nround</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_7</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">Nround</span><span class="p">,</span><span class="n">v1_6</span><span class="p">(</span><span class="mi">2</span><span class="p">)],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">sort_id</span> <span class="o">=</span> <span class="n">dim_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_8</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1065353216</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_9</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N</span><span class="p">,</span><span class="n">v1_8</span><span class="p">(</span><span class="mf">1.0</span><span class="n">f</span><span class="p">)],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_10</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N</span><span class="p">,</span><span class="n">v1_9</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_11</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1073741824</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_12</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_10</span><span class="p">,</span><span class="n">v1_11</span><span class="p">(</span><span class="mf">2.0</span><span class="n">f</span><span class="p">)],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">steps</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_12</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_13</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_14</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">step</span> <span class="o">=</span> <span class="n">loop</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_14</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">steps</span><span class="p">,</span><span class="n">v1_13</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="p">)</span>
<span class="p">{</span>
  <span class="kt">int</span> <span class="n">v2_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_1</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_0</span><span class="p">(</span><span class="mi">2</span><span class="p">),</span><span class="n">step</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_2</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_1</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_3</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1065353216</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_4</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_2</span><span class="p">,</span><span class="n">v2_3</span><span class="p">(</span><span class="mf">1.0</span><span class="n">f</span><span class="p">)],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_5</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_4</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1056964608</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_7</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_5</span><span class="p">,</span><span class="n">v2_6</span><span class="p">(</span><span class="mf">0.5</span><span class="n">f</span><span class="p">)],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">j</span> <span class="o">=</span> <span class="n">floor</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_8</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">step</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_9</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1056964608</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_10</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_9</span><span class="p">(</span><span class="mf">0.5</span><span class="n">f</span><span class="p">),</span><span class="n">j</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_11</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1065353216</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_12</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">j</span><span class="p">,</span><span class="n">v2_11</span><span class="p">(</span><span class="mf">1.0</span><span class="n">f</span><span class="p">)],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_13</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_10</span><span class="p">,</span><span class="n">v2_12</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_14</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_8</span><span class="p">,</span><span class="n">v2_13</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">n</span> <span class="o">=</span> <span class="n">round</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_14</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_15</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">j</span><span class="p">,</span><span class="n">n</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_16</span> <span class="o">=</span> <span class="n">exp2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_15</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_17</span> <span class="o">=</span> <span class="n">round</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_16</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">B</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_17</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_18</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1056964608</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">v2_19</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">n</span><span class="p">,</span><span class="n">v2_18</span><span class="p">(</span><span class="mf">0.5</span><span class="n">f</span><span class="p">)],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_20</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_21</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_20</span><span class="p">(</span><span class="mi">2</span><span class="p">),</span><span class="n">B</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_22</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_23</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_21</span><span class="p">,</span><span class="n">v2_22</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">mask</span> <span class="o">=</span> <span class="n">ternary</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_19</span><span class="p">,</span><span class="n">v2_23</span><span class="p">,</span><span class="n">B</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_24</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">sort_id</span><span class="p">,</span><span class="n">B</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_25</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_26</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_25</span><span class="p">(</span><span class="mi">2</span><span class="p">),</span><span class="n">B</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_27</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">sort_id</span><span class="p">,</span><span class="n">B</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_28</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_26</span><span class="p">,</span><span class="n">v2_27</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">e1</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_24</span><span class="p">,</span><span class="n">v2_28</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">e2</span> <span class="o">=</span> <span class="n">xor</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">mask</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">v2_29</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">element_count</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">v2_30</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">,</span><span class="n">element_count</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">v2_31</span> <span class="o">=</span> <span class="n">and</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_29</span><span class="p">,</span><span class="n">v2_30</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="k">if</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_31</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
  <span class="p">{</span>
    <span class="kt">int</span> <span class="n">key1</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">key2</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">bool</span> <span class="n">v3_0</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">key1</span><span class="p">,</span><span class="n">key2</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
    <span class="k">if</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
    <span class="p">{</span>
      <span class="kt">int</span> <span class="n">val1</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">val2</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
      <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">key2</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
      <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">key1</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
      <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">val2</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
      <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">val1</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="p">)</span>
    <span class="p">}</span>
  <span class="p">}</span>
<span class="p">}</span>
</code></pre></div>    </div>

  </div>
</details>

<p>I made the debug IR representation be somewhat C-like since, at least for me, its easier to read than your usual IR representations. Every line here represents a node in the multilevel linked list. Each <code class="language-plaintext highlighter-rouge">{}</code> scope incapsulates all the child nodes of the previous to <code class="language-plaintext highlighter-rouge">{}</code> node. While the bitonic sort part is basically just a less readible version of the python code above, we now also have some additional nodes in the IR. Specifically <code class="language-plaintext highlighter-rouge">memory</code>, this is the node that represents allocated tensor memory on the device. Here we also see that it has flags signifying that its an output and input of the program. The <code class="language-plaintext highlighter-rouge">RemoveUnusedOperations</code> compilaiton stage removes everything that doesn’t influence those memory nodes.</p>

<p><em>Those of you who know about LLVM would probably question the choices made here, but in my case specifically I was interested in keeping the IR as close to the codegen target as possible, which are C++, CUDA, shading languages, or others. This IR doesn’t really use single assignment form (SSA) or φ nodes, meaning modifications are not versioned. This does pose a problem for autodiff and makes optimization potentially harder, so I do have a compilation pass that can convert at least some in-place operations into versions of the original, in a rather ad-hoc way. I still need to do the same for <code class="language-plaintext highlighter-rouge">stores</code> and <code class="language-plaintext highlighter-rouge">scatters</code> too, since right now autodiff will usually compute the wrong gradients for these operations, as it doesn’t use the correct version for the gradient, or actually, it simply doesn’t have access to it, because it no longer exists in memory. I have a reason why I dont want to version everything - it will potentially result in overly aggressive additional memory allocation (like imagine this sorting algorithm created a copied version of keys/values every iteration), and you would need to optimize for it separately. But this reasoning could be completely wrong, since I haven’t really worked with LLVM and am not sure about how applicable it might be for my use cases.</em></p>

<p>After all the compilation stages, the IR creates kernel nodes, replaces multidimensional indexing with flattened 1D indexing, does some optimizations etc.</p>

<details>
<summary>Final compiled IR</summary>

<div>

    <div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">int</span> <span class="n">element_count</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">keys</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="n">OutputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_0</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">values</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="n">OutputMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="n">InputMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_1</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_2</span> <span class="o">=</span> <span class="n">log2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">17.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">log2N</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">18.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_3</span> <span class="o">=</span> <span class="n">exp2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">34.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">Nround</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">35.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_4</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_5</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">Nround</span><span class="p">,</span><span class="n">v1_4</span><span class="p">(</span><span class="mi">2</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">37.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_6</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_7</span> <span class="o">=</span> <span class="n">log2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_6</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">17.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">log2N_2</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">18.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_8</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1065353216</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_9</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N_2</span><span class="p">,</span><span class="n">v1_8</span><span class="p">(</span><span class="mf">1.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">19.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_10</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N_2</span><span class="p">,</span><span class="n">v1_9</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">38.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_11</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1073741824</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_12</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_10</span><span class="p">,</span><span class="n">v1_11</span><span class="p">(</span><span class="mf">2.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">40.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">steps</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_12</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">41.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_13</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_14</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">step</span> <span class="o">=</span> <span class="n">loop</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_14</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">steps</span><span class="p">,</span><span class="n">v1_13</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">141.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="p">{</span>
  <span class="n">kernel</span><span class="p">(</span><span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
  <span class="p">{</span>
    <span class="kt">float</span> <span class="n">v3_0</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">float</span> <span class="n">v3_1</span> <span class="o">=</span> <span class="n">log2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">17.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">float</span> <span class="n">log2N_3</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">18.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">float</span> <span class="n">v3_2</span> <span class="o">=</span> <span class="n">exp2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">log2N_3</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">34.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">Nround_2</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">35.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_3</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_4</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">Nround_2</span><span class="p">,</span><span class="n">v3_3</span><span class="p">(</span><span class="mi">2</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">37.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_5</span> <span class="o">=</span> <span class="n">block_id</span><span class="p">(</span><span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">256</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_7</span> <span class="o">=</span> <span class="n">block_thread_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_8</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_5</span><span class="p">,</span><span class="n">v3_6</span><span class="p">(</span><span class="mi">256</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">index_0</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_8</span><span class="p">,</span><span class="n">v3_7</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">2.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">bool</span> <span class="n">is_inside_dispatch</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_0</span><span class="p">,</span><span class="n">v3_4</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">40.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
    <span class="k">if</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">is_inside_dispatch</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">140.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
    <span class="p">{</span>
      <span class="kt">int</span> <span class="n">v4_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_1</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_0</span><span class="p">(</span><span class="mi">2</span><span class="p">),</span><span class="n">step</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">142.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_2</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">143.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_3</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1065353216</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_4</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_2</span><span class="p">,</span><span class="n">v4_3</span><span class="p">(</span><span class="mf">1.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">144.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_5</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_4</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">148.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1056964608</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_7</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_5</span><span class="p">,</span><span class="n">v4_6</span><span class="p">(</span><span class="mf">0.5</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">149.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">j</span> <span class="o">=</span> <span class="n">floor</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_7</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">150.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_8</span> <span class="o">=</span> <span class="kt">float</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">step</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">142.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_9</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1056964608</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_10</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_9</span><span class="p">(</span><span class="mf">0.5</span><span class="n">f</span><span class="p">),</span><span class="n">j</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">151.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_11</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1065353216</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_12</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">j</span><span class="p">,</span><span class="n">v4_11</span><span class="p">(</span><span class="mf">1.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">151.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_13</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_10</span><span class="p">,</span><span class="n">v4_12</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">303.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_14</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_8</span><span class="p">,</span><span class="n">v4_13</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">446.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">n</span> <span class="o">=</span> <span class="n">round</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_14</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">447.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_15</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">j</span><span class="p">,</span><span class="n">n</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">598.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_16</span> <span class="o">=</span> <span class="n">exp2</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_15</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">614.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_17</span> <span class="o">=</span> <span class="n">round</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_16</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">615.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">B</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_17</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">616.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_18</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1056964608</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">bool</span> <span class="n">v4_19</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">n</span><span class="p">,</span><span class="n">v4_18</span><span class="p">(</span><span class="mf">0.5</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">448.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_20</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_21</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_20</span><span class="p">(</span><span class="mi">2</span><span class="p">),</span><span class="n">B</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">617.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_22</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_23</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_21</span><span class="p">,</span><span class="n">v4_22</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">618.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">mask</span> <span class="o">=</span> <span class="n">ternary</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_19</span><span class="p">,</span><span class="n">v4_23</span><span class="p">,</span><span class="n">B</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1686.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_24</span> <span class="o">=</span> <span class="n">mod</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_0</span><span class="p">,</span><span class="n">B</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">622.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_25</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_26</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_25</span><span class="p">(</span><span class="mi">2</span><span class="p">),</span><span class="n">B</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">617.000000</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_27</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_0</span><span class="p">,</span><span class="n">B</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">620.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_28</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_26</span><span class="p">,</span><span class="n">v4_27</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1238.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">e1</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_24</span><span class="p">,</span><span class="n">v4_28</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1861.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">e2</span> <span class="o">=</span> <span class="n">xor</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">mask</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3548.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">bool</span> <span class="n">v4_29</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">element_count</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1862.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">bool</span> <span class="n">v4_30</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">,</span><span class="n">element_count</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3549.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">bool</span> <span class="n">v4_31</span> <span class="o">=</span> <span class="n">and</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_29</span><span class="p">,</span><span class="n">v4_30</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5412.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="k">if</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_31</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5512.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
      <span class="p">{</span>
        <span class="kt">int</span> <span class="n">v5_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_1</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">,</span><span class="n">v5_0</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_2</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_3</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">v5_2</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v5_1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1866.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">key1</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v5_3</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1994.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_4</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_5</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">,</span><span class="n">v5_4</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">v5_7</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">,</span><span class="n">v5_6</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v5_5</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3553.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="kt">int</span> <span class="n">key2</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v5_7</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3681.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="kt">bool</span> <span class="n">v5_8</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">key1</span><span class="p">,</span><span class="n">key2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5676.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="k">if</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v5_8</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5776.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="p">{</span>
          <span class="kt">int</span> <span class="n">v6_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_1</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">v6_0</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_2</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_3</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">v6_2</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v6_1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1866.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">val1</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v6_3</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1994.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_4</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_5</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">v6_4</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_7</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">,</span><span class="n">v6_6</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v6_5</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3553.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">val2</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v6_7</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3681.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_8</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_9</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">,</span><span class="n">v6_8</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_10</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_11</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">v6_10</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v6_9</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1866.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">key2</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v6_11</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5675.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_13</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_14</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">element_count</span><span class="p">,</span><span class="n">v6_13</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_15</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_16</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">,</span><span class="n">v6_15</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v6_14</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3553.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">keys</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">key1</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v6_16</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5675.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_18</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_19</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">v6_18</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_20</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_21</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e1</span><span class="p">,</span><span class="n">v6_20</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v6_19</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1866.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">val2</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v6_21</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5675.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_23</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_24</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">v6_23</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_25</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
          <span class="kt">int</span> <span class="n">v6_26</span> <span class="o">=</span> <span class="n">clamp</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">e2</span><span class="p">,</span><span class="n">v6_25</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">v6_24</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">3553.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
          <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">values</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">val1</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v6_26</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">5675.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="p">)</span>
        <span class="p">}</span>
      <span class="p">}</span>
    <span class="p">}</span>
  <span class="p">}</span>
<span class="p">}</span>
</code></pre></div>    </div>

  </div>
</details>

<p>For the most part, the IR here is unchanged, but we have some differences. With the exception of the obviously created kernel under the loop node, we also see that <code class="language-plaintext highlighter-rouge">dim_id</code> nodes, which represented the index of an element at this dimension, are replaced with <code class="language-plaintext highlighter-rouge">block_id</code> and <code class="language-plaintext highlighter-rouge">block_thread_id</code>, on a GPU, for instance, those map to the workgroup index and the internal 3D workgroup thread indices. I didn’t go with a single 1D workgroup thread index, since its somewhat easier to read the generated kernel code, if it maps directly to how compute shaders are written, but this wasn’t really necessary. Additionally the shader compiler might in theory better compile the generated code.</p>

<p>Also notably, by default the compiler clamps all indexing operations to the shape of the <code class="language-plaintext highlighter-rouge">memory</code> input to avoid undefined behaviour. Some internal operations can avoid doing this to not do useless work, but I also plan to give the user access to how the indices could be computed under the hood.</p>

<p>There is also the newly appeared <code class="language-plaintext highlighter-rouge">cost</code> property of the node, which is only used for evaluating which parts to copy or not in the optimization passes, and mostly serves as a heuristic, and doesn’t represent actual cost of executing a node, that would need to be done in a potential autotuner.</p>

<p>Now, lets look at how higher level operation are compiled right now, for example, a really simple matrix multiplication:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">C</span> <span class="o">=</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">A</span><span class="p">)</span> <span class="o">@</span> <span class="n">tf</span><span class="p">.</span><span class="n">cos</span><span class="p">(</span><span class="n">B</span><span class="p">.</span><span class="n">T</span><span class="p">))</span><span class="o">**</span><span class="mf">2.0</span>
</code></pre></div></div>

<p>Which has a pretty simple input IR:</p>

<details>
<summary>Matmul input IR</summary>
<div>

    <div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">int</span> <span class="n">v1_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">4294967295</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_1</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">4294967295</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">K</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">N</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">A</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_2</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">4294967295</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_3</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">B</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">v1_3</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_4</span> <span class="o">=</span> <span class="n">sin</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">A</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_5</span> <span class="o">=</span> <span class="n">transpose</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">B</span><span class="p">],</span> <span class="n">data</span><span class="o">=</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">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_6</span> <span class="o">=</span> <span class="n">cos</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_7</span> <span class="o">=</span> <span class="n">matmul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_4</span><span class="p">,</span><span class="n">v1_6</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_8</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1073741824</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_9</span> <span class="o">=</span> <span class="n">pow</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_7</span><span class="p">,</span><span class="n">v1_8</span><span class="p">(</span><span class="mf">2.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">OutputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
</code></pre></div>    </div>

  </div>
</details>

<p>Close to the beginning of the compilation all nodes that are marked as <code class="language-plaintext highlighter-rouge">Algorithm</code> in the compilers dictionary are replaced with their specific implementations, and after that pass we will be left with an IR like this:</p>

<details>
<summary>After algorithm insertion</summary>
<div>

    <div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">int</span> <span class="n">K</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">N</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">A</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_0</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">B</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">v1_0</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_1</span> <span class="o">=</span> <span class="n">sin</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">A</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_2</span> <span class="o">=</span> <span class="n">dim_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_3</span> <span class="o">=</span> <span class="n">dim_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">transposed</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">B</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v1_3</span><span class="p">,</span><span class="n">v1_2</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v1_4</span> <span class="o">=</span> <span class="n">cos</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">transposed</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_5</span> <span class="o">=</span> <span class="n">dim_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_6</span> <span class="o">=</span> <span class="n">dim_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">matmul_2</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_7</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_8</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_9</span> <span class="o">=</span> <span class="n">loop</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_8</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">K</span><span class="p">,</span><span class="n">v1_7</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="p">)</span>
<span class="p">{</span>
  <span class="kt">float</span> <span class="n">v2_0</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">v1_1</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v1_9</span><span class="p">,</span><span class="n">v1_6</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_1</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">v1_4</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v1_5</span><span class="p">,</span><span class="n">v1_9</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_2</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_0</span><span class="p">,</span><span class="n">v2_1</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">float</span> <span class="n">v2_3</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">matmul_2</span><span class="p">(</span><span class="mf">0.0</span><span class="n">f</span><span class="p">),</span><span class="n">v2_2</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="n">set</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">matmul_2</span><span class="p">(</span><span class="mf">0.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_3</span><span class="p">],</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="p">}</span>
<span class="kt">float</span> <span class="n">v3_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1073741824</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">v3_1</span> <span class="o">=</span> <span class="n">pow</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">matmul_2</span><span class="p">(</span><span class="mf">0.0</span><span class="n">f</span><span class="p">),</span><span class="n">v3_0</span><span class="p">(</span><span class="mf">2.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">OutputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
</code></pre></div>    </div>

  </div>
</details>

<p>As you can see, for the matrix multiplication, it has created a loop that accumulates products like <code class="language-plaintext highlighter-rouge">A[i,k]*B[k,j]</code> over <code class="language-plaintext highlighter-rouge">k</code>, and <code class="language-plaintext highlighter-rouge">transpose</code> is effectively just compiled into a load at transposed indices.
In the future I plan to improve these built-in algorithms to also employ groupshared memory for precaching A and B blocks and sum over them, without this optimization the performance of the matrix multiplication quickly degrades for larger sizes that don’t fit into the cache.
Alternatively, just like in PyTorch, I could simply add natively implemented kernels for these operations, or if I used CUDA, just call cuBLAS, its not even that hard to add. But right now I’m interested in seeing how much I could improve the performance without explicit outside kernels, as it also improves portability, not to mention the fusion. (You don’t have BLAS libraries in graphics API’s unfortunately)</p>

<p>After this compilation pass the kernel generation and fusion happens as well as allocation of memory for the resulting tensor.</p>

<details>
<summary>Final compiled matmul</summary>
<div>

    <div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">int</span> <span class="n">K</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">N</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">A</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">int</span> <span class="n">v1_0</span> <span class="o">=</span> <span class="n">input_shape</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputShapeDim</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="n">InputShapeMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">B</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">InputMemory</span><span class="p">(</span><span class="mi">1</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">K</span><span class="p">,</span><span class="n">v1_0</span><span class="p">],</span> <span class="p">)</span>
<span class="kt">float</span> <span class="n">m0</span> <span class="o">=</span> <span class="n">memory</span><span class="p">(</span><span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="n">OutputMemory</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="n">kernel</span><span class="p">(</span><span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
<span class="p">{</span>
  <span class="kt">int</span> <span class="n">v2_0</span> <span class="o">=</span> <span class="n">block_id</span><span class="p">(</span><span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_1</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">16</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_2</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">16</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_3</span> <span class="o">=</span> <span class="n">block_thread_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_4</span> <span class="o">=</span> <span class="n">block_thread_id</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_5</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">v2_1</span><span class="p">(</span><span class="mi">16</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">1.000000</span><span class="p">,</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_6</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_7</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_5</span><span class="p">,</span><span class="n">v2_6</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">2.000000</span><span class="p">,</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">blocks_shape_0</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_7</span><span class="p">,</span><span class="n">v2_1</span><span class="p">(</span><span class="mi">16</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">4.000000</span><span class="p">,</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_8</span> <span class="o">=</span> <span class="n">div</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_0</span><span class="p">,</span><span class="n">blocks_shape_0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">6.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_9</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_8</span><span class="p">,</span><span class="n">blocks_shape_0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">11.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_10</span> <span class="o">=</span> <span class="n">sub</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_0</span><span class="p">,</span><span class="n">v2_9</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">12.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_11</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_10</span><span class="p">,</span><span class="n">v2_1</span><span class="p">(</span><span class="mi">16</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">13.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">index_0</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_11</span><span class="p">,</span><span class="n">v2_3</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">14.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">v2_12</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_8</span><span class="p">,</span><span class="n">v2_2</span><span class="p">(</span><span class="mi">16</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">7.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">int</span> <span class="n">index_1</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_12</span><span class="p">,</span><span class="n">v2_4</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">8.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">v2_13</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_0</span><span class="p">,</span><span class="n">v1_0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">15.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">v2_14</span> <span class="o">=</span> <span class="n">lt</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_1</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">9.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="kt">bool</span> <span class="n">is_inside_dispatch</span> <span class="o">=</span> <span class="n">and</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v2_13</span><span class="p">,</span><span class="n">v2_14</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">25.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="k">if</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">is_inside_dispatch</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">125.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="p">{</span>
    <span class="kt">float</span> <span class="n">matmul_2</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">flags</span><span class="o">=</span><span class="p">{</span><span class="n">Modified</span><span class="p">,</span> <span class="p">},</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_1</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v3_2</span> <span class="o">=</span> <span class="n">loop</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v3_1</span><span class="p">(</span><span class="mi">0</span><span class="p">),</span><span class="n">K</span><span class="p">,</span><span class="n">v3_0</span><span class="p">(</span><span class="mi">1</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">100.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="p">{</span>
      <span class="kt">int</span> <span class="n">v4_0</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_1</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">9.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_1</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_0</span><span class="p">,</span><span class="n">v3_2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">110.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">A_2</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">A</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v4_1</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">238.000000</span><span class="p">,</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_2</span> <span class="o">=</span> <span class="n">sin</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">A_2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">240.000000</span><span class="p">,</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_3</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_0</span><span class="p">,</span><span class="n">K</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">15.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">int</span> <span class="n">v4_4</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_3</span><span class="p">,</span><span class="n">v3_2</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">116.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">transposed</span> <span class="o">=</span> <span class="n">load</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">B</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v4_4</span><span class="p">],</span> <span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">244.000000</span><span class="p">,</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_5</span> <span class="o">=</span> <span class="n">cos</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">transposed</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">246.000000</span><span class="p">,</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_6</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_2</span><span class="p">,</span><span class="n">v4_5</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">487.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="kt">float</span> <span class="n">v4_7</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">matmul_2</span><span class="p">(</span><span class="mf">0.0</span><span class="n">f</span><span class="p">),</span><span class="n">v4_6</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">488.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
      <span class="n">set</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">matmul_2</span><span class="p">(</span><span class="mf">0.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v4_7</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">489.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
    <span class="p">}</span>
    <span class="kt">float</span> <span class="n">v5_0</span> <span class="o">=</span> <span class="k">const</span><span class="p">(</span><span class="n">data</span><span class="o">=</span><span class="p">[</span><span class="mi">1073741824</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">0.000000</span><span class="p">,</span> <span class="p">)</span>
    <span class="kt">float</span> <span class="n">v5_1</span> <span class="o">=</span> <span class="n">pow</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">matmul_2</span><span class="p">(</span><span class="mf">0.0</span><span class="n">f</span><span class="p">),</span><span class="n">v5_0</span><span class="p">(</span><span class="mf">2.0</span><span class="n">f</span><span class="p">)],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">6.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v5_2</span> <span class="o">=</span> <span class="n">mul</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">index_1</span><span class="p">,</span><span class="n">v1_0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">9.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
    <span class="kt">int</span> <span class="n">v5_3</span> <span class="o">=</span> <span class="n">add</span><span class="p">(</span><span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v5_2</span><span class="p">,</span><span class="n">index_0</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">24.000000</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
    <span class="n">store</span><span class="p">(</span><span class="n">memory</span><span class="o">=</span><span class="p">[</span><span class="n">m0</span><span class="p">],</span> <span class="n">inputs</span><span class="o">=</span><span class="p">[</span><span class="n">v5_1</span><span class="p">],</span> <span class="n">indices</span><span class="o">=</span><span class="p">[</span><span class="n">v5_3</span><span class="p">],</span> <span class="n">cost</span><span class="o">=</span><span class="mf">158.000000</span><span class="p">,</span> <span class="n">indexing_mode</span><span class="o">=</span><span class="n">Unsafe</span><span class="p">,</span> <span class="n">shape</span><span class="o">=</span><span class="p">[</span><span class="n">v1_0</span><span class="p">,</span><span class="n">N</span><span class="p">],</span> <span class="p">)</span>
  <span class="p">}</span>
<span class="p">}</span>

</code></pre></div>    </div>

  </div>
</details>

<p>As you can see, it not only fused the <code class="language-plaintext highlighter-rouge">pow</code> at the end of the matrix multiplication, but also fused the transposition with the sin/cos operations into the summation loop. While this impressively created only a single kernel, this isn’t actually super optimal. Matrix multiplcation is often bottlenecked by arithmetic, not just by memory access. And we effectively instead of doing sin/cos N^2 times, do them N^3 times now! This is something that I still need to fine tune in the fusion heuristics algorithm. In the future though, if you could use groupshared caches - it would be fine enough to do the fusion of these input computations at the matrix “block” load stage, this should be enough to significantly reduce the overhead of doing additional sin/cos, as now we do them only for each “block” of the matrix, not for each product. But I would suspect that for huge matrices, these precomputations will still be a bottleneck, and need to be forcibly unfused. So yes, fewer kernels doesn’t actually mean better sometimes.</p>

<p>In the current compiler version I have disabled matmul load fusion completely until I write a better heuristic, as it usually improves performance.</p>

<h1 id="python-frontend">Python frontend</h1>

<p>Since I’ve already decided to make a tensor library, I wanted for it to have similar syntax to Numpy and be easily usable from Python. The frontend is a Pybind11 wrapper around the C++ library which overloads all operations in Python.</p>

<p>As this library is effecitively a static compiler, the way you use it is split into 2 parts - the compiled function with “virtual” tensors of potentially undefined shape (but defined dimensionality), and the host, where explicit tensor buffers exist.</p>

<h2 id="main-code">Main code</h2>

<p>You basically define a function that looks like this:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">matmul</span><span class="p">():</span>
    <span class="n">A</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">([</span><span class="o">-</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">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
    <span class="n">N</span><span class="p">,</span> <span class="n">M</span> <span class="o">=</span> <span class="n">A</span><span class="p">.</span><span class="n">shape</span>
    <span class="n">B</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">([</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span>  <span class="n">M</span><span class="p">],</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
    <span class="n">K</span> <span class="o">=</span> <span class="n">B</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="n">C</span> <span class="o">=</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">sin</span><span class="p">(</span><span class="n">A</span><span class="p">)</span> <span class="o">@</span> <span class="n">tf</span><span class="p">.</span><span class="n">cos</span><span class="p">(</span><span class="n">B</span><span class="p">.</span><span class="n">T</span><span class="p">))</span><span class="o">**</span><span class="mf">2.0</span>

    <span class="k">return</span> <span class="n">C</span>
</code></pre></div></div>

<p>This will be the core of a <code class="language-plaintext highlighter-rouge">TensorProgram</code>. You probably noticed that it doesn’t have arguments, right now its a rather ad-hoc way to give the ability to restrict shapes of some inputs to other inputs. Technically I could just parse the function python representation and generate <code class="language-plaintext highlighter-rouge">tf.input()</code> automatically, but in either case, you will still need to apply <code class="language-plaintext highlighter-rouge">tf.assert_tensor</code> to enforce their shape.</p>

<p>You are probably already asking why its done is such a weird way, but the main goal was the ability to have undefined tensor shapes at compile time, this way you don’t need to recompile the program every time your input changes. This, as you can see, does add some restrictions, since if you don’t do these shape <code class="language-plaintext highlighter-rouge">assertions</code> the compiler might not be able to figure out that A and B can actually be multiplied together and will throw an error. I could also have gone with the “assume everything is compatible” route and added assertions in the generated IR automatically before applying the operation, but I suspected such behaviour might have very annoying unintended consequences and could also result in fusion of parts of code that shouldn’t have fused, which would either be wrong, or create assertions that might never be valid and always throw errors. Of course, you can still have the shapes as predefined constants everywhere, in that case this particular quirk becomes rather annoying and more of a hindrence. I suspect that in the future I’ll add support for both automatically reading the function arguments and manual specification like here.</p>

<p>So yeah, the input shape of these <code class="language-plaintext highlighter-rouge">tf.input()</code> operations can be <code class="language-plaintext highlighter-rouge">-1</code> for unspecified, and <code class="language-plaintext highlighter-rouge">&gt;0</code> for explicitly specified (or you can use a shape from another input). In some cases having explicit shape might improve performance, as the compiler could do staged reductions and the like.</p>

<p>The way the IR is currently generated is by tracing the Python function, which is significantly easier than parsing Python AST. This does have some interesting side effects. If you do any sort of control flow inside this function, you can only do it with Python values, and on top of that, the result will be unrolled and fixed at compile time.</p>

<p>In fact, all variables, N, M, A, B, etc - are not actual tensors/scalars, but abstractions in the IR, and don’t have a value yet. So doing any kind of print would result in abstract info being spat out, and conversion into Numpy or any other python type would simply be impossible.</p>

<p>As I wanted to have control flow in the IR, I either needed to parse Python AST, or somehow overload existing Python behavour. Initially, all scoped operations, like <code class="language-plaintext highlighter-rouge">if</code> or <code class="language-plaintext highlighter-rouge">loop</code> took a python function as input, which was quite ugly and very unreadable for deep code, in the same way JAX does it, pretty much. But then I discovered that you can actually overload context manager behaviour. I found this trick in <a href="https://github.com/ppenenko/metashade">a library for python based shader generation</a>. This means that I can have custom calls at the beginning and end of a section of code, and I could automatically put it as a child scope.
This is what I did, and I overloaded these for some tensor types, so now you can do contol flow in a nicer way! (Arguably still cursed, since now we have 2 ways to make control flow, with different results, not even mentioning that you need to use the <code class="language-plaintext highlighter-rouge">.set()</code> method or <code class="language-plaintext highlighter-rouge">.val</code> property of a tensor to set its value from these child scopes, as Python doesn’t allow you to overload <code class="language-plaintext highlighter-rouge">=</code> operators, and also scatters and stores are fine too)</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">a</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">const</span><span class="p">(</span><span class="mf">0.0</span><span class="p">)</span>
<span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">if_cond</span><span class="p">(</span><span class="n">A</span> <span class="o">==</span> <span class="n">B</span><span class="p">):</span>
  <span class="n">a</span><span class="p">.</span><span class="n">val</span> <span class="o">+=</span> <span class="mf">1.0</span>

<span class="n">m</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">const</span><span class="p">(</span><span class="o">-</span><span class="mf">1e10</span><span class="p">)</span>
<span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">loop</span><span class="p">(</span><span class="n">begin</span><span class="p">,</span> <span class="n">end</span><span class="p">,</span> <span class="n">step</span><span class="p">)</span> <span class="k">as</span> <span class="n">iteration</span><span class="p">:</span>
  <span class="n">m</span><span class="p">.</span><span class="n">val</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">max</span><span class="p">(</span><span class="n">m</span><span class="p">,</span> <span class="n">data</span><span class="p">[</span><span class="n">iteration</span><span class="p">])</span>

</code></pre></div></div>

<p>I even used them for custom kernel declaration:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">kernel</span><span class="p">([</span><span class="n">M</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="k">as</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">k</span><span class="p">):</span>
  <span class="c1">#stuff
</span></code></pre></div></div>

<p>You can also use <code class="language-plaintext highlighter-rouge">tf.break_loop()</code> to stop the first parent loop. There might be some cases when it doesn’t work, like stopping a CPU loop from within a kernel (what does that even mean?). But usually it works if you don’t do something especially unusual.</p>

<h2 id="host-code">Host code</h2>

<p>Before anything, you should initialize the backend which will be used, like this:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">tf</span><span class="p">.</span><span class="n">initialize</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">cpu</span><span class="p">)</span> <span class="c1">#or tf.opengl
</span></code></pre></div></div>

<p>Additionally you can provide the compiler flag for the C++ compiler, like for example:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">tf</span><span class="p">.</span><span class="n">initialize</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">cpu</span><span class="p">,</span> <span class="s">"-g"</span><span class="p">)</span> <span class="c1">#or /Di on Windows
</span></code></pre></div></div>

<p>After you wrote the main function you can compile it into a <code class="language-plaintext highlighter-rouge">TensorProgram</code> like this</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">matmul_compiled</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">compile</span><span class="p">(</span><span class="n">matmul</span><span class="p">)</span>
</code></pre></div></div>

<p>This traces the function into the IR, compiles the IR into kernel form, converts that into C++/Kernel code, compiles that, and links the compiled library at runtime.
One of the (not) fun things about this on Windows, is that this requires the Microsoft Visual Studio compiler installed, and on top of that its <strong>SLOW</strong> as hell, usually at 20x-50x of the IR compilation time. The compiler requires some totally cursed things to set up the env variables to even work, and the command generator ends up looking like this:</p>

<div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code> <span class="n">ss</span> <span class="o">&lt;&lt;</span> <span class="s">"powershell -command </span><span class="se">\"</span><span class="s">$VisualStudioPath = &amp; </span><span class="se">\\\"</span><span class="s">${Env:ProgramFiles(x86)}</span><span class="se">\\</span><span class="s">Microsoft Visual Studio</span><span class="se">\\</span><span class="s">Installer</span><span class="se">\\</span><span class="s">vswhere.exe</span><span class="se">\\\"</span><span class="s"> -latest -products * -property installationPath; &amp; cmd.exe /C </span><span class="se">\\\"\"\\\"\\\"</span><span class="s">$VisualStudioPath</span><span class="se">\\</span><span class="s">VC</span><span class="se">\\</span><span class="s">Auxiliary</span><span class="se">\\</span><span class="s">Build</span><span class="se">\\</span><span class="s">vcvarsall.bat</span><span class="se">\\\"\\\"</span><span class="s"> x64 &amp;&amp; cl "</span>
       <span class="o">&lt;&lt;</span> <span class="n">kernelCompileOptions</span> <span class="o">&lt;&lt;</span> <span class="s">" /LD "</span> <span class="o">&lt;&lt;</span> <span class="n">tempPath</span>
       <span class="o">&lt;&lt;</span> <span class="n">sourcePath</span> <span class="o">&lt;&lt;</span> <span class="s">" /Fe:"</span> <span class="o">&lt;&lt;</span> <span class="n">dllName</span>
       <span class="o">&lt;&lt;</span> <span class="s">"</span><span class="se">\"\"\\\"\"</span><span class="s">"</span><span class="p">;</span>
</code></pre></div></div>

<p>I have no clue what is going on here, and thankfully I wasn’t the one who wrote this (you should thank @Devaniti for this one). At least it works.</p>

<p>Linux users win with their <code class="language-plaintext highlighter-rouge">gcc</code> in that regard, which is usually also already installed in the system. It is also just 3-4x faster than MVSC.</p>

<p>In the future I want to add Python as an alternative host language (you could also use CUDA too!), as it will speedup the compile times up to 100x, with only slight performance overhead.</p>

<p>These <code class="language-plaintext highlighter-rouge">TensorProgram</code> objects take and output <code class="language-plaintext highlighter-rouge">TensorMemory</code> buffer objects (you can also give Numpy arrays, or <code class="language-plaintext highlighter-rouge">tf.Modules</code> as arguments), which can be created from Numpy arrays like this.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">A</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">([</span><span class="mi">100</span><span class="p">,</span> <span class="mi">100</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">B</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">([</span><span class="mi">100</span><span class="p">,</span> <span class="mi">100</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>
</code></pre></div></div>

<p>Then you can run the program:</p>
<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">C</span> <span class="o">=</span> <span class="n">matmul_compiled</span><span class="p">(</span><span class="n">A</span><span class="p">,</span> <span class="n">B</span><span class="p">)</span>
</code></pre></div></div>

<p>As you can see the inputs are given to the compiled function in the same order as they were executed in the compiled function.</p>

<p>To get the result back into a Numpy array, you can use the <code class="language-plaintext highlighter-rouge">Numpy</code> property:</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">Cnp</span> <span class="o">=</span> <span class="n">C</span><span class="p">.</span><span class="n">numpy</span>
</code></pre></div></div>

<h3 id="modules">Modules</h3>

<p>TensorFrost has a simple module system similar to PyTorch, where you can define a module with parameters (that you can optimize by utilizing the modules from tf.optimizers) and a forward function that computes the output of the module as well as a loss function. Neither of these are actually required, but in some cases the optimizers need a specified <code class="language-plaintext highlighter-rouge">loss</code> function.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">class</span> <span class="nc">SmolNet</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">Module</span><span class="p">):</span>
    <span class="k">def</span> <span class="nf">__init__</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="c1">#specify a custom random scale and offset for the weights when initializing
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">W</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">Parameter</span><span class="p">([</span><span class="mi">16</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">],</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">random_scale</span><span class="o">=</span><span class="mf">0.01</span><span class="p">,</span> <span class="n">random_offset</span><span class="o">=</span><span class="mf">0.0</span><span class="p">)</span>
        <span class="c1">#dont compute gradients for the bias
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">b</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">Parameter</span><span class="p">([</span><span class="o">-</span><span class="mi">1</span><span class="p">],</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">,</span> <span class="n">optimize</span><span class="o">=</span><span class="bp">False</span><span class="p">)</span>
        
    <span class="k">def</span> <span class="nf">assert_parameters</span><span class="p">(</span><span class="bp">self</span><span class="p">):</span>
        <span class="c1">#makes sure that the compiler knows that b has shape compatible with W
</span>        <span class="bp">self</span><span class="p">.</span><span class="n">b</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">assert_tensor</span><span class="p">(</span><span class="bp">self</span><span class="p">.</span><span class="n">b</span><span class="p">,</span> <span class="p">[</span><span class="bp">self</span><span class="p">.</span><span class="n">W</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="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
        
    <span class="k">def</span> <span class="nf">forward</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">):</span>
        <span class="k">return</span> <span class="n">x</span> <span class="o">@</span> <span class="bp">self</span><span class="p">.</span><span class="n">W</span> <span class="o">+</span> <span class="bp">self</span><span class="p">.</span><span class="n">b</span>
    
    <span class="k">def</span> <span class="nf">loss</span><span class="p">(</span><span class="bp">self</span><span class="p">,</span> <span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">):</span>
        <span class="n">y_pred</span> <span class="o">=</span> <span class="bp">self</span><span class="p">.</span><span class="n">forward</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">y</span><span class="p">)</span>
        <span class="k">return</span> <span class="n">tf</span><span class="p">.</span><span class="n">mean</span><span class="p">((</span><span class="n">y</span> <span class="o">-</span> <span class="n">y_pred</span><span class="p">)</span><span class="o">**</span><span class="mi">2</span><span class="p">)</span>
</code></pre></div></div>

<p>When initializing the module you can add 3 types of TensorFrost accessible parameters:</p>
<ul>
  <li><code class="language-plaintext highlighter-rouge">tf.Parameter</code> - a tensor that will be passed to the TensorProgram as an argument</li>
  <li><code class="language-plaintext highlighter-rouge">tf.ParameterArray</code> - a dynamic list of parameters, all of them will be passed to the TensorProgram as arguments</li>
  <li><code class="language-plaintext highlighter-rouge">tf.Module</code> - another module, all of its parameters will be passed to the TensorProgram as arguments</li>
</ul>

<p>The shape argument of the parameter can be a list of integers, where -1 means that the shape is not specified yet, and will be inferred from the input tensor. If you need to compute an operation over several tensors of unspecified shape, you need to assert the shapes in the <code class="language-plaintext highlighter-rouge">assert_parameters</code> function.
<code class="language-plaintext highlighter-rouge">random_scale</code> and <code class="language-plaintext highlighter-rouge">random_offset</code> are used to initialize the weights with random values, and are optional, by default the weights are initialized with Xavier initialization for uniform random values.
<code class="language-plaintext highlighter-rouge">optimize</code> is used to specify if the parameter should be trained or not, by default all parameters are trainable. This argument does not stop you from computing <code class="language-plaintext highlighter-rouge">tf.grad</code> manually, it is just used to specify if the parameter should be updated by the optimizer module.</p>

<p>By itself the module does not do anything, you need to do a second initialization step to either use it inside a TensorProgram, or initialize it as a container for the tensors outside of the program.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="k">def</span> <span class="nf">ComputeForward</span><span class="p">():</span>
    <span class="n">model</span> <span class="o">=</span> <span class="n">SmolNet</span><span class="p">()</span>
    <span class="c1">#creates tf.input tensors from all the parameters of the module
</span>    <span class="n">model</span><span class="p">.</span><span class="n">initialize_input</span><span class="p">()</span>
    <span class="n">X</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">([</span><span class="o">-</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">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">model</span><span class="p">.</span><span class="n">forward</span><span class="p">(</span><span class="n">X</span><span class="p">)</span>

<span class="n">forward</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">compile</span><span class="p">(</span><span class="n">ComputeForward</span><span class="p">)</span>

<span class="n">model_container</span> <span class="o">=</span> <span class="n">SmolNet</span><span class="p">()</span>
<span class="c1">#creates tf.tensor tensors from all the parameters of the module and initializes them
</span><span class="n">model_container</span><span class="p">.</span><span class="n">initialize_parameters</span><span class="p">()</span>
<span class="c1">#you can change them afterwards too
</span><span class="n">model_container</span><span class="p">.</span><span class="n">W</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">([</span><span class="mi">16</span><span class="p">,</span> <span class="mi">100</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">X</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">([</span><span class="mi">100</span><span class="p">,</span> <span class="mi">100</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="c1">#the module is passed as an argument to the compiled function, in the same order as they are created in the function
</span><span class="n">Y</span> <span class="o">=</span> <span class="n">forward</span><span class="p">(</span><span class="n">model_container</span><span class="p">,</span> <span class="n">X</span><span class="p">)</span>
</code></pre></div></div>

<p><code class="language-plaintext highlighter-rouge">model.initialize_input()</code> creates <code class="language-plaintext highlighter-rouge">tf.input()</code> tensors for all the parameters of the module. Afterwards <code class="language-plaintext highlighter-rouge">assert_parameters</code> is automatically called for this and all child modules. This is useful if you want to use the module inside a TensorProgram, as you can just pass the module as an argument to the compiled function, and all the parameters will be automatically created and the shapes will be asserted.
<code class="language-plaintext highlighter-rouge">model.initialize_parameters()</code> creates <code class="language-plaintext highlighter-rouge">tf.tensor()</code> tensors for all the parameters of the module and initializes them with random values. This is useful if you want to use the module outside of a TensorProgram, as you can just pass the module as an argument to the compiled function.</p>

<p>This particular part of the library is still quite early stage, mostly only Python-side, and might change a lot in the future.</p>

<h3 id="optimizer-modules">Optimizer modules</h3>

<p>TensorFrost has a set of built-in optimizer modules that can be used to train the parameters of the module.</p>
<ul>
  <li><code class="language-plaintext highlighter-rouge">tf.optimizers.sgd</code> - Stochastic Gradient Descent, has a <code class="language-plaintext highlighter-rouge">learning_rate</code> and <code class="language-plaintext highlighter-rouge">grad_clip</code> parameters, default values are 0.001 and 0.0 respectively.</li>
  <li><code class="language-plaintext highlighter-rouge">tf.optimizers.adam</code> - Adam optimizer, has a <code class="language-plaintext highlighter-rouge">learning_rate</code>, <code class="language-plaintext highlighter-rouge">beta1</code>, <code class="language-plaintext highlighter-rouge">beta2</code> and <code class="language-plaintext highlighter-rouge">grad_clip</code> parameters, default values are 0.001, 0.9, 0.999 and 0.0 respectively.</li>
  <li><code class="language-plaintext highlighter-rouge">tf.optimizers.rmsprop</code> - RMSProp optimizer, has a <code class="language-plaintext highlighter-rouge">learning_rate</code>, <code class="language-plaintext highlighter-rouge">decay</code> and <code class="language-plaintext highlighter-rouge">grad_clip</code> parameters, default values are 0.001, 0.9 and 0.0 respectively.</li>
</ul>

<p>All optimizer modules are initialized with the module as the first argument, and the training hyperparameters as the rest of the arguments.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">OptimizerStep</span><span class="p">():</span>
    <span class="n">X</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">([</span><span class="o">-</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">tf</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">tf</span><span class="p">.</span><span class="nb">input</span><span class="p">([</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span> <span class="mi">10</span><span class="p">],</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

    <span class="n">model</span> <span class="o">=</span> <span class="n">SmolNet</span><span class="p">()</span>
    <span class="n">opt</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">optimizers</span><span class="p">.</span><span class="n">adam</span><span class="p">(</span><span class="n">model</span><span class="p">,</span> <span class="n">learning_rate</span><span class="o">=</span><span class="mf">0.001</span><span class="p">,</span> <span class="n">beta1</span><span class="o">=</span><span class="mf">0.9</span><span class="p">,</span> <span class="n">beta2</span><span class="o">=</span><span class="mf">0.999</span><span class="p">)</span>
    <span class="n">opt</span><span class="p">.</span><span class="n">initialize_input</span><span class="p">()</span>
    
    <span class="c1">#do a single step of the optimizer (automatically computes gradients and updates the parameters)
</span>    <span class="n">L</span> <span class="o">=</span> <span class="n">opt</span><span class="p">.</span><span class="n">step</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="c1">#or 
</span>    <span class="c1">#L = model.loss(X, Y)
</span>    <span class="c1">#opt.step(L)
</span>
    <span class="n">params</span> <span class="o">=</span> <span class="n">opt</span><span class="p">.</span><span class="n">parameters</span><span class="p">()</span>
    <span class="n">params</span><span class="p">.</span><span class="n">append</span><span class="p">(</span><span class="n">L</span><span class="p">)</span>
    <span class="k">return</span> <span class="n">params</span>

<span class="n">step</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">compile</span><span class="p">(</span><span class="n">OptimizerStep</span><span class="p">)</span>

<span class="n">model_container</span> <span class="o">=</span> <span class="n">SmolNet</span><span class="p">()</span>
<span class="n">opt</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">optimizers</span><span class="p">.</span><span class="n">adam</span><span class="p">(</span><span class="n">model_container</span><span class="p">)</span>
<span class="n">opt</span><span class="p">.</span><span class="n">initialize_parameters</span><span class="p">()</span>

<span class="n">X</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">([</span><span class="mi">100</span><span class="p">,</span> <span class="mi">100</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">tf</span><span class="p">.</span><span class="n">tensor</span><span class="p">(</span><span class="n">np</span><span class="p">.</span><span class="n">zeros</span><span class="p">([</span><span class="mi">100</span><span class="p">,</span> <span class="mi">10</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">out</span> <span class="o">=</span> <span class="n">step</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">opt</span><span class="p">)</span>
<span class="n">opt</span><span class="p">.</span><span class="n">update_parameters</span><span class="p">(</span><span class="n">res</span><span class="p">[:</span><span class="o">-</span><span class="mi">1</span><span class="p">])</span>
<span class="n">loss</span> <span class="o">=</span> <span class="n">res</span><span class="p">[</span><span class="o">-</span><span class="mi">1</span><span class="p">].</span><span class="n">numpy</span><span class="p">[</span><span class="mi">0</span><span class="p">]</span>
</code></pre></div></div>

<p>I’ve also recently added regularizers (reg_type = tf.regularizers.l2 or tf.regularizers.l1) and clipping (tf.clipping.norm or just tf.clipping.clip for a clamp), which can be added like:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">optimizer</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">optimizers</span><span class="p">.</span><span class="n">adam</span><span class="p">(</span><span class="n">model_container</span><span class="p">,</span> <span class="n">beta1</span> <span class="o">=</span> <span class="mf">0.0</span><span class="p">,</span> <span class="n">beta2</span> <span class="o">=</span> <span class="mf">0.999</span><span class="p">,</span> <span class="n">reg_type</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">regularizers</span><span class="p">.</span><span class="n">l2</span><span class="p">,</span> <span class="n">reg</span> <span class="o">=</span> <span class="mf">0.02</span><span class="p">,</span> <span class="n">clip</span> <span class="o">=</span> <span class="mf">0.01</span><span class="p">)</span>
<span class="n">optimizer</span><span class="p">.</span><span class="n">set_clipping_type</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">clipping</span><span class="p">.</span><span class="n">norm</span><span class="p">)</span>
</code></pre></div></div>

<h2 id="visualization-and-interactivity">Visualization and interactivity</h2>

<p>I really wanted a way to output computation results in real time so I decided to add a GLFW + ImGui for a window and simple GUI to the library. (Taichi also did this!) The way it works now is that you can create a window, create the main rendering loop, and then render the tensor as an image. You can do quite a lot of things with this.</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1">#creates a single global window (can only be one at the moment)
</span><span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">show</span><span class="p">(</span><span class="mi">1280</span><span class="p">,</span> <span class="mi">720</span><span class="p">,</span> <span class="s">"a window"</span><span class="p">)</span>

<span class="k">while</span> <span class="ow">not</span> <span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">should_close</span><span class="p">():</span> <span class="c1">#window will close if you press the close button and this will return True
</span>    <span class="n">mx</span><span class="p">,</span> <span class="n">my</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">get_mouse_position</span><span class="p">()</span>
    <span class="n">wx</span><span class="p">,</span> <span class="n">wy</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">get_window_size</span><span class="p">()</span>

    <span class="c1">#simple input example
</span>    <span class="k">if</span> <span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">is_mouse_button_pressed</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">MOUSE_BUTTON_0</span><span class="p">):</span>
        <span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">text</span><span class="p">(</span><span class="s">"Mouse button 0 is pressed"</span><span class="p">)</span>

    <span class="k">if</span> <span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">is_key_pressed</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">KEY_W</span><span class="p">):</span>
        <span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">text</span><span class="p">(</span><span class="s">"W is pressed"</span><span class="p">)</span>

    <span class="c1">#ImGui example
</span>    <span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">begin</span><span class="p">(</span><span class="s">"an imgui window"</span><span class="p">)</span>
    <span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">text</span><span class="p">(</span><span class="s">"some text"</span><span class="p">)</span>
    <span class="n">value</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">slider</span><span class="p">(</span><span class="s">"slider"</span><span class="p">,</span> <span class="n">value</span><span class="p">,</span> <span class="mf">0.0</span><span class="p">,</span> <span class="mf">10.0</span><span class="p">)</span>
    <span class="k">if</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">button</span><span class="p">(</span><span class="s">"a button"</span><span class="p">)):</span>
        <span class="k">print</span><span class="p">(</span><span class="s">"button pressed"</span><span class="p">)</span>
    <span class="n">tf</span><span class="p">.</span><span class="n">imgui</span><span class="p">.</span><span class="n">end</span><span class="p">()</span>

    <span class="c1">#exectute a TensorFrost TensorProgram that outputs a [-1, -1, 3] float32 tensor
</span>    <span class="n">img</span> <span class="o">=</span> <span class="n">render_image</span><span class="p">(...)</span>

    <span class="c1">#you could also just provide a Numpy array as tf.tensor(), this is usually slower tho, as it requires a GPU upload
</span>
    <span class="c1">#display the image (will be stretched to the window size with nearest neighbor interpolation)
</span>    <span class="n">tf</span><span class="p">.</span><span class="n">window</span><span class="p">.</span><span class="n">render_frame</span><span class="p">(</span><span class="n">img</span><span class="p">)</span>
</code></pre></div></div>

<p>In this example you have a window rendering loop, inside of the loop you can query the mouse position and the window size. You can also check if a mouse/keyboard button is pressed. You can create simple ImGUI windows, with text, sliders, checkboxes, buttons and plotlines (in the future I want to integrate <a href="https://github.com/epezent/implot">ImPlot</a> too)</p>

<p>Then to render the frame you pass a tensor to <code class="language-plaintext highlighter-rouge">tf.window.render_frame()</code>.</p>

<h1 id="backends">Backends</h1>

<h2 id="codegen">Codegen</h2>

<p>I have also written my own code generator here, and right now C++, GLSL and HLSL are supported. Adding additional kernel languages is not very difficult. But right now I still need to refactor the host code generation, as its currently hardcoded to only do C++. The library can actually be run in a purely code generation mode if you don’t need to do any work in Python and would rather integrate it into some other project:</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">tf</span><span class="p">.</span><span class="n">initialize</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">codegen</span><span class="p">,</span> <span class="n">kernel_lang</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">hlsl_lang</span><span class="p">)</span> <span class="c1"># or tf.glsl_lang for OpenGL, or tf.cpp_lang for C++
</span></code></pre></div></div>

<p>After you compiled all the tensor programs you need, you can get all the generated code and save it to a file:</p>

<div class="language-python highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1"># Save all the compiled functions
</span><span class="n">cpp_header</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">get_cpp_header</span><span class="p">()</span>
<span class="n">all_main_functions</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">get_all_generated_main_functions</span><span class="p">()</span> <span class="c1">#always in C++
</span><span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="s">'tensorfrost_main.cpp'</span><span class="p">,</span> <span class="s">'w'</span><span class="p">)</span> <span class="k">as</span> <span class="n">f</span><span class="p">:</span>
    <span class="n">f</span><span class="p">.</span><span class="n">write</span><span class="p">(</span><span class="n">cpp_header</span><span class="p">)</span>
    <span class="k">for</span> <span class="n">func</span> <span class="ow">in</span> <span class="n">all_main_functions</span><span class="p">:</span>
        <span class="n">f</span><span class="p">.</span><span class="n">write</span><span class="p">(</span><span class="n">func</span><span class="p">)</span>

<span class="c1"># Save all the compiled kernels
</span><span class="n">all_kernels</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">get_all_generated_kernels</span><span class="p">()</span> <span class="c1">#depends on the kernel_lang
</span><span class="k">for</span> <span class="n">i</span><span class="p">,</span> <span class="n">kernel</span> <span class="ow">in</span> <span class="nb">enumerate</span><span class="p">(</span><span class="n">all_kernels</span><span class="p">):</span>
    <span class="k">with</span> <span class="nb">open</span><span class="p">(</span><span class="s">'generated_kernels/kernel_{}.hlsl'</span><span class="p">.</span><span class="nb">format</span><span class="p">(</span><span class="n">i</span><span class="p">),</span> <span class="s">'w'</span><span class="p">)</span> <span class="k">as</span> <span class="n">f</span><span class="p">:</span>
        <span class="n">f</span><span class="p">.</span><span class="n">write</span><span class="p">(</span><span class="n">kernel</span><span class="p">)</span>
</code></pre></div></div>

<p>This is also not perfect, ideally I would also provide an example implementation of a runtime, but right now it needs to be written by the user. Also in the future I’d want to compile the <code class="language-plaintext highlighter-rouge">TensorPrograms</code> into an archive, that would optionally contain code and compiled binaries, that can be loaded into python immediately. This could be quite useful for debugging code generation in the future.</p>

<p>The generated code right now looks something like this, for the <a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Algorithms/bitonic.ipynb">bitonic sort example</a>.</p>

<details>
<summary>C++ host part</summary>

<div>

    <div class="language-cpp highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">std</span><span class="o">::</span><span class="n">tuple</span><span class="o">&lt;</span><span class="n">TFTensor</span><span class="o">&gt;</span> <span class="n">BitonicSort</span><span class="p">(</span><span class="n">TFContext</span> <span class="n">tf</span><span class="p">,</span> <span class="n">TFTensor</span> <span class="n">input0</span><span class="p">)</span>
<span class="p">{</span>
  <span class="n">tf</span><span class="p">.</span><span class="n">region_begin</span><span class="p">(</span><span class="s">"BitonicSort"</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">N</span> <span class="o">=</span> <span class="n">input0</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">tf</span><span class="p">.</span><span class="n">check_tensor</span><span class="p">(</span><span class="n">input0</span><span class="p">,</span> <span class="s">"input0"</span><span class="p">,</span> <span class="p">{(</span><span class="n">uint</span><span class="p">)</span><span class="n">N</span><span class="p">,</span> <span class="p">(</span><span class="n">uint</span><span class="p">)</span><span class="mi">2</span><span class="p">},</span> <span class="n">TFType</span><span class="o">::</span><span class="n">Int</span><span class="p">);</span>
  <span class="n">TFTensor</span> <span class="n">output0</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">allocate</span><span class="p">(</span><span class="s">"output0"</span><span class="p">,</span> <span class="p">{(</span><span class="n">uint</span><span class="p">)</span><span class="n">N</span><span class="p">,</span> <span class="p">(</span><span class="n">uint</span><span class="p">)</span><span class="mi">2</span><span class="p">},</span> <span class="n">TFType</span><span class="o">::</span><span class="n">Int</span><span class="p">);</span>
  <span class="n">tf</span><span class="p">.</span><span class="n">dispatch</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span> <span class="p">{</span><span class="n">output0</span><span class="p">},</span>  <span class="p">{</span><span class="n">input0</span><span class="p">},</span> <span class="p">{</span><span class="n">asuint</span><span class="p">(</span><span class="n">N</span><span class="p">)},</span> <span class="p">{(</span><span class="n">uint</span><span class="p">)</span><span class="n">N</span><span class="p">,</span> <span class="p">(</span><span class="n">uint</span><span class="p">)</span><span class="mi">2</span><span class="p">},</span> <span class="p">{</span><span class="mi">2</span><span class="p">,</span> <span class="mi">16</span><span class="p">});</span>
  <span class="kt">float</span> <span class="n">log2N</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">log2</span><span class="p">(((</span><span class="kt">float</span><span class="p">)(</span><span class="n">N</span><span class="p">))));</span>
  <span class="kt">int</span> <span class="n">Nround</span> <span class="o">=</span> <span class="p">((</span><span class="kt">int</span><span class="p">)(</span><span class="n">exp2</span><span class="p">(</span><span class="n">log2N</span><span class="p">)));</span>
  <span class="kt">int</span> <span class="n">v4_4</span> <span class="o">=</span> <span class="n">Nround</span> <span class="o">/</span> <span class="mi">2</span><span class="p">;</span>
  <span class="kt">float</span> <span class="n">log2N_2</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">log2</span><span class="p">(((</span><span class="kt">float</span><span class="p">)(</span><span class="n">N</span><span class="p">))));</span>
  <span class="kt">int</span> <span class="n">steps</span> <span class="o">=</span> <span class="p">((</span><span class="kt">int</span><span class="p">)((</span><span class="n">log2N_2</span> <span class="o">*</span> <span class="p">(</span><span class="n">log2N_2</span> <span class="o">+</span> <span class="mf">1.0</span><span class="n">f</span><span class="p">))</span> <span class="o">/</span> <span class="mf">2.0</span><span class="n">f</span><span class="p">));</span>
  <span class="k">for</span> <span class="p">(</span><span class="kt">int</span> <span class="n">step</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">step</span> <span class="o">&lt;</span> <span class="n">steps</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="p">{</span>
    <span class="n">tf</span><span class="p">.</span><span class="n">dispatch</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">{</span><span class="n">output0</span><span class="p">},</span>  <span class="p">{},</span> <span class="p">{</span><span class="n">asuint</span><span class="p">(</span><span class="n">N</span><span class="p">),</span> <span class="n">asuint</span><span class="p">(</span><span class="n">step</span><span class="p">)},</span> <span class="p">{(</span><span class="n">uint</span><span class="p">)</span><span class="n">v4_4</span><span class="p">},</span> <span class="p">{</span><span class="mi">256</span><span class="p">});</span>
  <span class="p">}</span>
  <span class="n">tf</span><span class="p">.</span><span class="n">region_end</span><span class="p">(</span><span class="s">"BitonicSort"</span><span class="p">);</span>
  <span class="k">return</span> <span class="p">{</span><span class="n">output0</span><span class="p">};</span>
<span class="p">}</span>
</code></pre></div>    </div>

  </div>
</details>

<details>
<summary>And generated GLSL shaders</summary>

<div>

    <div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1">//Kernel 1</span>

<span class="k">layout</span> <span class="p">(</span><span class="n">local_size_x</span> <span class="o">=</span> <span class="mi">2</span><span class="p">,</span> <span class="n">local_size_y</span> <span class="o">=</span> <span class="mi">16</span><span class="p">,</span> <span class="n">local_size_z</span> <span class="o">=</span> <span class="mi">1</span><span class="p">)</span> <span class="k">in</span><span class="p">;</span>

<span class="kt">void</span> <span class="nf">main</span><span class="p">()</span> <span class="p">{</span>
  <span class="kt">int</span> <span class="n">block_id</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_WorkGroupID</span><span class="p">.</span><span class="n">x</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">block_thread_id0</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_LocalInvocationID</span><span class="p">.</span><span class="n">x</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">block_thread_id1</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_LocalInvocationID</span><span class="p">.</span><span class="n">y</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">block_thread_id2</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_LocalInvocationID</span><span class="p">.</span><span class="n">z</span><span class="p">);</span>

  <span class="kt">int</span> <span class="n">blocks_shape_0</span> <span class="o">=</span> <span class="p">((</span><span class="mi">2</span> <span class="o">+</span> <span class="mi">2</span><span class="p">)</span> <span class="o">-</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="kt">int</span> <span class="n">vdiv</span> <span class="o">=</span> <span class="n">block_id</span> <span class="o">/</span> <span class="n">blocks_shape_0</span><span class="p">;</span>
  <span class="kt">int</span> <span class="n">index_0</span> <span class="o">=</span> <span class="p">((</span><span class="n">block_id</span> <span class="o">-</span> <span class="p">(</span><span class="n">vdiv</span> <span class="o">*</span> <span class="n">blocks_shape_0</span><span class="p">))</span> <span class="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">block_thread_id0</span><span class="p">;</span>
  <span class="kt">int</span> <span class="n">index_1</span> <span class="o">=</span> <span class="p">(</span><span class="n">vdiv</span> <span class="o">*</span> <span class="mi">16</span><span class="p">)</span> <span class="o">+</span> <span class="n">block_thread_id1</span><span class="p">;</span>
  <span class="kt">bool</span> <span class="n">is_inside_dispatch</span> <span class="o">=</span> <span class="p">(</span><span class="n">index_0</span> <span class="o">&lt;</span> <span class="mi">2</span><span class="p">)</span> <span class="o">&amp;&amp;</span> <span class="p">(</span><span class="n">index_1</span> <span class="o">&lt;</span> <span class="n">var</span><span class="p">.</span><span class="n">N</span><span class="p">);</span>
  <span class="k">if</span> <span class="p">(</span><span class="n">is_inside_dispatch</span><span class="p">)</span>
  <span class="p">{</span>
    <span class="kt">int</span> <span class="n">input0</span> <span class="o">=</span> <span class="n">asint</span><span class="p">(</span><span class="n">input0_mem</span><span class="p">[(</span><span class="n">index_1</span> <span class="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">index_0</span><span class="p">]);</span>
    <span class="kt">int</span> <span class="n">output0</span> <span class="o">=</span> <span class="n">input0</span><span class="p">;</span>
    <span class="n">output0_mem</span><span class="p">[(</span><span class="n">index_1</span> <span class="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">index_0</span><span class="p">]</span> <span class="o">=</span> <span class="n">asuint</span><span class="p">(</span><span class="n">output0</span><span class="p">);</span>
  <span class="p">}</span>
<span class="p">}</span>
</code></pre></div>    </div>

    <div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1">//Kernel 2</span>

<span class="k">layout</span> <span class="p">(</span><span class="n">local_size_x</span> <span class="o">=</span> <span class="mi">256</span><span class="p">,</span> <span class="n">local_size_y</span> <span class="o">=</span> <span class="mi">1</span><span class="p">,</span> <span class="n">local_size_z</span> <span class="o">=</span> <span class="mi">1</span><span class="p">)</span> <span class="k">in</span><span class="p">;</span>

<span class="kt">void</span> <span class="nf">main</span><span class="p">()</span> <span class="p">{</span>
  <span class="kt">int</span> <span class="n">block_id</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_WorkGroupID</span><span class="p">.</span><span class="n">x</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">block_thread_id0</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_LocalInvocationID</span><span class="p">.</span><span class="n">x</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">block_thread_id1</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_LocalInvocationID</span><span class="p">.</span><span class="n">y</span><span class="p">);</span>
  <span class="kt">int</span> <span class="n">block_thread_id2</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">gl_LocalInvocationID</span><span class="p">.</span><span class="n">z</span><span class="p">);</span>

  <span class="kt">float</span> <span class="n">log2N</span> <span class="o">=</span> <span class="n">ceil</span><span class="p">(</span><span class="n">log2</span><span class="p">(</span><span class="kt">float</span><span class="p">(</span><span class="n">var</span><span class="p">.</span><span class="n">N</span><span class="p">)));</span>
  <span class="kt">int</span> <span class="n">Nround</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">exp2</span><span class="p">(</span><span class="n">log2N</span><span class="p">));</span>
  <span class="kt">int</span> <span class="n">index_0</span> <span class="o">=</span> <span class="p">(</span><span class="n">block_id</span> <span class="o">*</span> <span class="mi">256</span><span class="p">)</span> <span class="o">+</span> <span class="n">block_thread_id0</span><span class="p">;</span>
  <span class="kt">bool</span> <span class="n">is_inside_dispatch</span> <span class="o">=</span> <span class="n">index_0</span> <span class="o">&lt;</span> <span class="p">(</span><span class="n">Nround</span> <span class="o">/</span> <span class="mi">2</span><span class="p">);</span>
  <span class="k">if</span> <span class="p">(</span><span class="n">is_inside_dispatch</span><span class="p">)</span>
  <span class="p">{</span>
    <span class="kt">float</span> <span class="n">j</span> <span class="o">=</span> <span class="n">floor</span><span class="p">(</span><span class="n">sqrt</span><span class="p">(</span><span class="kt">float</span><span class="p">(</span><span class="mi">2</span> <span class="o">*</span> <span class="n">var</span><span class="p">.</span><span class="n">step</span><span class="p">)</span> <span class="o">+</span> <span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="n">f</span><span class="p">)</span> <span class="o">-</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="n">f</span><span class="p">);</span>
    <span class="kt">float</span> <span class="n">n</span> <span class="o">=</span> <span class="n">round</span><span class="p">(</span><span class="kt">float</span><span class="p">(</span><span class="n">var</span><span class="p">.</span><span class="n">step</span><span class="p">)</span> <span class="o">-</span> <span class="p">((</span><span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="n">f</span> <span class="o">*</span> <span class="n">j</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">j</span> <span class="o">+</span> <span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="n">f</span><span class="p">)));</span>
    <span class="kt">int</span> <span class="n">B</span> <span class="o">=</span> <span class="kt">int</span><span class="p">(</span><span class="n">round</span><span class="p">(</span><span class="n">exp2</span><span class="p">(</span><span class="n">j</span> <span class="o">-</span> <span class="n">n</span><span class="p">)));</span>
    <span class="kt">int</span> <span class="n">mask</span> <span class="o">=</span> <span class="p">(</span><span class="n">n</span> <span class="o">&lt;</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="n">f</span><span class="p">)</span> <span class="o">?</span> <span class="p">((</span><span class="mi">2</span> <span class="o">*</span> <span class="n">B</span><span class="p">)</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)</span> <span class="o">:</span> <span class="n">B</span><span class="p">;</span>
    <span class="kt">int</span> <span class="n">e1</span> <span class="o">=</span> <span class="p">(</span><span class="n">index_0</span> <span class="o">%</span> <span class="n">B</span><span class="p">)</span> <span class="o">+</span> <span class="p">((</span><span class="mi">2</span> <span class="o">*</span> <span class="n">B</span><span class="p">)</span> <span class="o">*</span> <span class="p">(</span><span class="n">index_0</span> <span class="o">/</span> <span class="n">B</span><span class="p">));</span>
    <span class="kt">int</span> <span class="n">e2</span> <span class="o">=</span> <span class="n">e1</span> <span class="o">^</span> <span class="n">mask</span><span class="p">;</span>
    <span class="k">if</span> <span class="p">((</span><span class="n">e1</span> <span class="o">&lt;</span> <span class="n">var</span><span class="p">.</span><span class="n">N</span><span class="p">)</span> <span class="o">&amp;&amp;</span> <span class="p">(</span><span class="n">e2</span> <span class="o">&lt;</span> <span class="n">var</span><span class="p">.</span><span class="n">N</span><span class="p">))</span>
    <span class="p">{</span>
      <span class="kt">int</span> <span class="n">key1</span> <span class="o">=</span> <span class="n">asint</span><span class="p">(</span><span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e1</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]);</span>
      <span class="kt">int</span> <span class="n">key2</span> <span class="o">=</span> <span class="n">asint</span><span class="p">(</span><span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e2</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]);</span>
      <span class="kt">int</span> <span class="n">val1</span> <span class="o">=</span> <span class="n">asint</span><span class="p">(</span><span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e1</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]);</span>
      <span class="kt">int</span> <span class="n">val2</span> <span class="o">=</span> <span class="n">asint</span><span class="p">(</span><span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e2</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]);</span>
      <span class="k">if</span> <span class="p">(</span><span class="n">key1</span> <span class="o">&gt;</span> <span class="n">key2</span><span class="p">)</span>
      <span class="p">{</span>
        <span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e1</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]</span> <span class="o">=</span> <span class="n">asuint</span><span class="p">(</span><span class="n">key2</span><span class="p">);</span>
        <span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e2</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]</span> <span class="o">=</span> <span class="n">asuint</span><span class="p">(</span><span class="n">key1</span><span class="p">);</span>
        <span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e1</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]</span> <span class="o">=</span> <span class="n">asuint</span><span class="p">(</span><span class="n">val2</span><span class="p">);</span>
        <span class="n">output0_mem</span><span class="p">[(</span><span class="n">clamp</span><span class="p">(</span><span class="n">e2</span><span class="p">,</span> <span class="mi">0</span><span class="p">,</span> <span class="n">var</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="o">*</span> <span class="mi">2</span><span class="p">)</span> <span class="o">+</span> <span class="n">clamp</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="mi">2</span> <span class="o">-</span> <span class="mi">1</span><span class="p">)]</span> <span class="o">=</span> <span class="n">asuint</span><span class="p">(</span><span class="n">val1</span><span class="p">);</span>
      <span class="p">}</span>
    <span class="p">}</span>
  <span class="p">}</span>
<span class="p">}</span>
</code></pre></div>    </div>

  </div>
</details>

<p>I tried to keep the generated code as readable as possible, to some degree it worked, but there are still quite a few ugly parts. I keep all buffers in <code class="language-plaintext highlighter-rouge">uint</code> format, as otherwise you can do <code class="language-plaintext highlighter-rouge">compareexchange</code> atomics on float elements (at least I think you can’t), so because of this there are a lot of <code class="language-plaintext highlighter-rouge">asuint</code> and <code class="language-plaintext highlighter-rouge">asfloat</code>/<code class="language-plaintext highlighter-rouge">asint</code> all over the place in an average generated kernel.</p>

<p>For this specific example its still quite nice, as we dont have a lot of algorithmic operations or automatically generated autodiff slop.</p>

<h2 id="runtimes">Runtimes</h2>

<p>Right now there are only 2 runtime backends - C++/OpenMP and C++/OpenGL. After the compiler generates the C++ code for the shaders and the host, its compiled by the C++ compiler, and also by the OpenGL shader compiler (which is built in the driver and can have horrible bugs, by the way).</p>

<p>I plan on also adding CUDA and Vulkan in the future, for the first one I could just compile everything, host and kernels into a single <code class="language-plaintext highlighter-rouge">.cu</code> file, and its probably relatively straightforward to do (will still need to keep OpenGL for visualization interop), but in the case of Vulkan I would need to write all the boilerplate code for handling basic compute, compiling shaders and memory allocation, that will probably take quite some time.</p>

<h1 id="examples-using-tensorfrost">Examples using TensorFrost</h1>

<h2 id="fluid-simulation">Fluid simulation</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Simulation/fluid_simulation.ipynb">—Link—</a></p>

<center><iframe width="560" height="315" src="https://www.youtube.com/embed/CVF4cZOsMK4?si=SJsZ2R_SIe-yyXgF" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen=""></iframe></center>

<p>Before all the algorithmic Numpy-like stuff, I initially played around with simulations that map nicely to multidimensional arrays, waves and Eulerian (grid based) fluids. For the most interesting example, I implemented a 2D fluid solver with a multigrid pressure solver, RK4 bicubic advection, and vorticity confinement. As this example didn’t require any control flow, but only indexing and basic kernel fusion, it was a nice first test for the compiler.</p>

<p>Also never mind the boundary artifacts or the interpolation issues with the density, its fixable, but I didn’t have enough time to mess around with it. I will probably implement a proper 3D fluid simulation next time anyway.</p>

<h2 id="fractal-path-tracer">Fractal path tracer</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/GUI/interactive_path_tracer.py">—Link—</a></p>

<center><iframe width="560" height="315" src="https://www.youtube.com/embed/ShWO5YSphOY?si=SqgPVOAhQNP_izF6" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen=""></iframe></center>

<p>In this particular case I created my own vec3 class in python here, as without automatic dimension unroll for small constant shape - the compiler will not be able to create a single kernel for it.</p>

<p>Syntactically speaking, this is 1-to-1 the same how I would write a path tracer in a shader, with the exception of there not being any explicit kernel definitions here. I also reused the bicubic interpolation from the fluid example for the camera reprojection, same with the bilinear sampler for the HDRI sky here.</p>

<p>One amazing thing here, is that the normals are computed through backpropagation and not finite differences! It even improved the performance, as FD normals take 4x SDF calculations using the tetrahedron trick, while backwards mode autodiff only takes on the order of 2x. It does only work for unrolled SDF’s like kaleidoscopic iterative fractals (KIF’s), SDF’s that require varying loop iterations or breaks, like Mandelbulb will not be differentiable at the moment.</p>

<p>What about bounding volume hierarchies (BVH) for triangle meshes? Right now you can’t really use them, as the IR does not have a way to declare local arrays, and those will be required for the stack when traversing the BVH tree efficiently (you could actually emulate those with a giant number of local variables and <code class="language-plaintext highlighter-rouge">if</code>’s, but why). I suspect I might add those together with groupshared memory, as they will have similar syntax.</p>

<p>At this point, now that I’ve also implemented modules and optimizers, I could also theoretically implement a <a href="https://research.nvidia.com/publication/2021-06_real-time-neural-radiance-caching-path-tracing">neural radiance cache</a> here, as I can both train the network and trace the rays. But personally I’d probably prefer a sparse hash grid radiance cache, as its a bit more deterministic.</p>

<p><em>PS. Looking at the next example, maybe you could use a hash grid + a small neural network together instead, for the radiance cache? This might improve the quality quite a lot</em></p>

<h2 id="texture-embedder-with-small-neural-network">Texture embedder with small neural network</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Rendering/neural_embed.ipynb">—Link—</a></p>

<center><iframe width="560" height="315" src="https://www.youtube.com/embed/7uzuGftSYKk?si=xt17EQguu4pg_PlV" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen=""></iframe></center>

<p><a href="https://www.shadertoy.com/view/dd33zf">This is very similar to the examples I have done in shadertoy.</a> It’s a small texture, usually 32*32, with an embedding in each pixel. Those embeddings are interpolated at some point in space, concatenated and passed to a small neural network that magically transforms that into a higher resolution image. The interpolation here is also the one from <a href="https://iquilezles.org/articles/texture/">iq’s article “improved texture interpolation”</a>, while that interpolation is “fake” and has 0 gradients at the pixel edges, if you gave the neural net a set of those, but offset by half a pixel, it can actually recreate a proper C2 continuous image without you needing bicubic interpolation! (I didn’t come up with this, all credits to <a href="https://nvlabs.github.io/instant-ngp/assets/mueller2022instant.pdf">Instant-NGP, Appendix A</a>) Surprisingly, my tests show that bicubic is actually somehow worse at representing images here, than this fake smooth interpolaiton. I suspect the neural net has an easier time at “distorting” bilinearly interpolated values, rather than cubically interpolated ones, as fewer pixels are influenced.</p>

<p>I used my Adam optimizer module to optimize the embedding image and the neural net, given a set of random samples from the original image. Of the interesting things, the gradients of the trilinear interpolation are just atomic adds to a tensor of the same shape as the embedding. It is probably as good as you can do it here anyway, given that the interpolation is at random unordered positions.</p>

<h2 id="n-body-sph-with-a-custom-sphere-rasterizer">N-body SPH with a custom sphere rasterizer</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Simulation/n-body.ipynb">—Link—</a></p>

<center><iframe width="560" height="315" src="https://www.youtube.com/embed/AxkabWearoA?si=mz5ZBtK1B1gh4z2L" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen=""></iframe></center>

<p>The code for the simulation here is very simple, first  computes the SPH densities for each particle using a <code class="language-plaintext highlighter-rouge">sum</code> over the gaussian SPH kernels:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">i</span><span class="p">,</span> <span class="n">j</span><span class="p">,</span> <span class="n">k</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">indices</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="mi">3</span><span class="p">])</span>
<span class="n">dx</span> <span class="o">=</span> <span class="n">X</span><span class="p">[</span><span class="n">j</span><span class="p">,</span><span class="n">k</span><span class="p">]</span> <span class="o">-</span> <span class="n">X</span><span class="p">[</span><span class="n">i</span><span class="p">,</span><span class="n">k</span><span class="p">]</span>
<span class="n">dv</span> <span class="o">=</span> <span class="n">V</span><span class="p">[</span><span class="n">j</span><span class="p">,</span><span class="n">k</span><span class="p">]</span> <span class="o">-</span> <span class="n">V</span><span class="p">[</span><span class="n">i</span><span class="p">,</span><span class="n">k</span><span class="p">]</span>

<span class="k">def</span> <span class="nf">sph_kernel</span><span class="p">(</span><span class="n">dist</span><span class="p">,</span> <span class="n">rad</span><span class="p">):</span>
    <span class="k">return</span> <span class="n">tf</span><span class="p">.</span><span class="n">exp</span><span class="p">(</span><span class="o">-</span><span class="p">(</span><span class="n">dist</span> <span class="o">/</span> <span class="n">rad</span><span class="p">)</span><span class="o">**</span><span class="mf">2.0</span><span class="p">)</span>

<span class="c1"># Compute the SPH density
</span><span class="n">dist</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">norm</span><span class="p">(</span><span class="n">dx</span><span class="p">)</span>
<span class="n">rho</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">sph_kernel</span><span class="p">(</span><span class="n">dist</span><span class="p">,</span> <span class="n">sph_rad</span><span class="p">),</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span>
</code></pre></div></div>

<p>And the second part computes the forces, first computes the soft gravity potential, then its negative gradient for the gravitational force. After that it computes the friction (viscosity) and the SPH pressure force, which is the gradient of the SPH kernel times the pressure. There is also the spike force, which is keeping the particles from overlapping to have a nice uniform distribution.</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">def</span> <span class="nf">pressure</span><span class="p">(</span><span class="n">rho</span><span class="p">):</span>
    <span class="k">return</span> <span class="p">(</span><span class="n">rho</span> <span class="o">-</span> <span class="n">rest_density</span><span class="p">)</span>

<span class="c1"># Compute the SPH forces
</span><span class="n">d2</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">dx</span><span class="o">**</span><span class="mf">2.0</span><span class="p">))</span>
<span class="n">dist</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">d2</span> <span class="o">+</span> <span class="mf">1e-4</span><span class="p">)</span> <span class="c1"># soft distance
</span><span class="n">Fg</span> <span class="o">=</span> <span class="o">-</span> <span class="n">tf</span><span class="p">.</span><span class="n">grad</span><span class="p">(</span><span class="n">gravity</span> <span class="o">/</span> <span class="n">dist</span><span class="p">,</span> <span class="n">dx</span><span class="p">)</span>
<span class="n">weight</span> <span class="o">=</span> <span class="n">sph_kernel</span><span class="p">(</span><span class="n">dist</span><span class="p">,</span> <span class="n">sph_rad</span><span class="p">)</span>
<span class="n">weightgrad</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">grad</span><span class="p">(</span><span class="n">weight</span><span class="p">,</span> <span class="n">dx</span><span class="p">)</span>
<span class="n">dvdotdx</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">dot</span><span class="p">(</span><span class="n">dv</span><span class="p">,</span> <span class="n">dx</span><span class="p">))</span> <span class="o">/</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">d2</span><span class="p">)</span> <span class="o">+</span> <span class="mf">1e-5</span><span class="p">)</span>
<span class="n">Fvisc</span> <span class="o">=</span> <span class="o">-</span> <span class="n">viscosity</span> <span class="o">*</span> <span class="n">dvdotdx</span> <span class="o">*</span> <span class="n">weightgrad</span>
<span class="n">Fsph</span> <span class="o">=</span> <span class="n">stiffness</span> <span class="o">*</span> <span class="mf">0.5</span> <span class="o">*</span> <span class="p">(</span><span class="n">pressure</span><span class="p">(</span><span class="n">rho</span><span class="p">[</span><span class="n">i</span><span class="p">])</span> <span class="o">+</span> <span class="n">pressure</span><span class="p">(</span><span class="n">rho</span><span class="p">[</span><span class="n">j</span><span class="p">]))</span> <span class="o">*</span> <span class="n">weightgrad</span>
<span class="n">dist2</span> <span class="o">=</span> <span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">d2</span><span class="p">)</span> <span class="o">+</span> <span class="mf">1e-8</span><span class="p">)</span>
<span class="n">Fspike</span> <span class="o">=</span> <span class="o">-</span> <span class="mf">250.0</span> <span class="o">*</span> <span class="n">sph_kernel</span><span class="p">(</span><span class="n">dist</span><span class="p">,</span> <span class="mf">1.0</span><span class="o">*</span><span class="n">sph_rad</span><span class="p">)</span> <span class="o">*</span> <span class="n">dx</span> <span class="o">/</span> <span class="p">(</span><span class="n">dist2</span><span class="o">*</span><span class="n">dist2</span><span class="p">)</span>
<span class="n">Fij</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">select</span><span class="p">(</span><span class="n">i</span> <span class="o">==</span> <span class="n">j</span><span class="p">,</span> <span class="mf">0.0</span><span class="p">,</span> <span class="n">Fg</span> <span class="o">+</span> <span class="n">Fsph</span> <span class="o">+</span> <span class="n">Fvisc</span> <span class="o">+</span> <span class="n">Fspike</span><span class="p">)</span>
<span class="n">Fi</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">Fij</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">Vnew</span> <span class="o">=</span> <span class="n">V</span> <span class="o">+</span> <span class="n">Fi</span> <span class="o">*</span> <span class="n">dt</span>
<span class="n">Xnew</span> <span class="o">=</span> <span class="n">X</span> <span class="o">+</span> <span class="n">V</span> <span class="o">*</span> <span class="n">dt</span>
</code></pre></div></div>

<p>And thats about it! The compiler actually fuses all these operations and the sum into a single loop, meaning there are only 2 kernels here, 1 for the SPH densities, and another for the forces.</p>

<p>The second part of this example is the rendering, which was also written purely in TensorFrost. Its a custom atomicMin (<code class="language-plaintext highlighter-rouge">tf.scatterMin</code>) rasterizer, which does a loop over all the pixels each particle sphere occupies on the screen, which is done by projecting the particle position onto the screen and doing atomic mins with a packed value high bits of which contain the depth, and low bits contain its index. The last pass after this one is the final shading, it computes the normals of each sphere, and computes its color.</p>

<h2 id="neural-cellular-automata">Neural Cellular Automata</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/ML/NCA/">—Link—</a></p>

<center><iframe width="560" height="315" src="https://www.youtube.com/embed/q9j3lea8Tvs?si=h7njjLrwZCiUpwuN" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture; web-share" referrerpolicy="strict-origin-when-cross-origin" allowfullscreen=""></iframe></center>

<p>I always wanted to recreate the results of the <a href="https://distill.pub/2020/growing-ca/">Growing Neural Cellular Automata</a> article, as the way the model worked was very similar to <a href="https://www.shadertoy.com/view/Wt2BR1">some</a> shadertoys I did!</p>

<p>While implementing it there were some problems with unrolled iterations, the number of kernels got so huge some of the fused ones accessed more buffers than OpenGL supports, which led to a compilation fail. I manually restricted fusion for this to not happen by doing a “hack”. I did that by introducing a <code class="language-plaintext highlighter-rouge">.stop_fusion()</code> method, which is absolutely hilarious given fusion was our initial goal. I guess it would be better to have an automatic way to restrict kernel size in the future.</p>

<p>Ideally I’d want the compiler to be capable to take gradients of loops natively so that it doesn’t generate a thousand kernels, but thats for the future I guess. Right now I keep the iteration count at around 30, while the original used 60-90 as far as I remember. Increasing the fire rate does reduce the required iteration count, as it “technically” increases the average time-step of the simulation, so I did just that.</p>

<p>In my own implementation I did a few changes compared to the original paper, notably I’ve added a laplacian input kernel, as it actually reduces checkerboard artifacts and improves convergense quite drastically. Sobel kernels have a step size of 2*dx, while the laplacian kernel has dx which makes it more accurate. It can also now natively emulate diffusion equations, before it needed to do gradients of gradients to do that.</p>

<p>Also I’ve added quantization and clamping of the fields, which allowed it to be ported to Shadertoy. Surprisingly it also seems to improve training convergence, perhaps due to more normalized channels.</p>

<center><iframe width="640" height="360" frameborder="0" src="https://www.shadertoy.com/embed/XXlyWr?gui=true&amp;t=10&amp;muted=true" allowfullscreen=""></iframe></center>

<p>More examples are in the <a href="https://github.com/MichaelMoroz/TensorFrost/tree/main/examples">examples folder</a>.</p>

<h1 id="whats-the-current-performance-compared-to-other-tensor-libraries">What’s the current performance compared to other tensor libraries?</h1>

<p>I’ll focus on comparing things that are easy to implement in both my library and in PyTorch/JAX. Things like the fluid sim or path tracer, I suspect, have no chance of running good if at all in PyTorch, JAX however might work fine with vmap, but I’m not sure. For these tests, I’ll ignore these use cases, as they will not be easy to port to them anyway, given the quite different syntax (as I didn’t write those vectorized, but more in “shader-like” code)</p>

<h2 id="n-body-simulation">N-body simulation</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Simulation/n-body-benchmark.py">—Link—</a></p>

<p>One of the simplest simulations you can do - is a N-body gravity simulation, which only takes a dozen lines both in TensorFrost, JAX or PyTorch, making this a nice benchmark.</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">dx</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">1</span><span class="p">)</span> <span class="o">-</span> <span class="n">tf</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="n">X</span><span class="p">,</span> <span class="n">axis</span><span class="o">=</span><span class="mi">0</span><span class="p">)</span>

<span class="n">d2</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">unsqueeze</span><span class="p">(</span><span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">dx</span><span class="o">**</span><span class="mf">2.0</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">axis</span><span class="o">=-</span><span class="mi">1</span><span class="p">)</span> <span class="o">+</span> <span class="mf">1e-4</span>
<span class="n">dist</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">d2</span><span class="p">)</span>
<span class="n">Fg</span> <span class="o">=</span> <span class="o">-</span><span class="n">dx</span> <span class="o">*</span> <span class="mf">1.0</span> <span class="o">/</span> <span class="p">(</span><span class="n">d2</span> <span class="o">*</span> <span class="n">dist</span><span class="p">)</span>

<span class="n">Fi</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">sum</span><span class="p">(</span><span class="n">Fg</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">Vnew</span> <span class="o">=</span> <span class="n">V</span> <span class="o">+</span> <span class="n">Fi</span> <span class="o">*</span> <span class="mf">0.001</span>
<span class="n">Xnew</span> <span class="o">=</span> <span class="n">X</span> <span class="o">+</span> <span class="n">Vnew</span> <span class="o">*</span> <span class="mf">0.001</span>
</code></pre></div></div>

<p>This is basically all there is to it, but if you don’t do kernel fusion, you will perform hillariously bad here, as lots of N^2 buffers will be allocated and waste a lot of bandwidth, but just for a demonstration, here is how default eager mode evaluation PyTorch scales compared to TensorFrost:</p>

<center><img src="/images/n-body-bench-torch-default.png" height="400px" /></center>

<p>It is <em>very</em> bad, but this is quite expected, we aren’t fusing any kernels here. Let’s now compare it to compiled PyTorch, and to compiled JAX. (All tests done on Ubuntu/RTX 2060/CUDA)</p>

<center><img src="/images/n-body-bench.png" height="400px" /></center>

<p>The results here are quite surprising, first of all, the compiled version of the function both in Torch and JAX are very close to what TensorFrost achieves here, which I actually didn’t expect, in some cases they even win, like ~1000 particles.  But they still scale slightly worse than the TensorFrost version of the same code curiously. Do they employ some staged reduction here? 
TensorFrost doesn’t actually fuse these operations 100% optimally here, the final generated kernel does 1 thread per each resulting component of the force. Ideally you’d do a single loop for all components, as they easily fit in the registers, and you remove a lot of wasted computations. This is the <code class="language-plaintext highlighter-rouge">explicit loop</code> version of the same calculations. The code looks like this:</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="n">Fi</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="nb">buffer</span><span class="p">([</span><span class="n">N</span><span class="p">,</span> <span class="mi">3</span><span class="p">],</span> <span class="n">tf</span><span class="p">.</span><span class="n">float32</span><span class="p">)</span>

<span class="n">i</span><span class="p">,</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">indices</span><span class="p">([</span><span class="n">N</span><span class="p">])</span>
<span class="n">Fix</span><span class="p">,</span> <span class="n">Fiy</span><span class="p">,</span> <span class="n">Fiz</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">const</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span> <span class="n">tf</span><span class="p">.</span><span class="n">const</span><span class="p">(</span><span class="mf">0.0</span><span class="p">),</span> <span class="n">tf</span><span class="p">.</span><span class="n">const</span><span class="p">(</span><span class="mf">0.0</span><span class="p">)</span>
<span class="n">x0</span><span class="p">,</span> <span class="n">y0</span><span class="p">,</span> <span class="n">z0</span> <span class="o">=</span> <span class="n">X</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">X</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="mi">2</span><span class="p">]</span>
<span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">loop</span><span class="p">(</span><span class="n">N</span><span class="p">)</span> <span class="k">as</span> <span class="n">j</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">z</span> <span class="o">=</span> <span class="n">X</span><span class="p">[</span><span class="n">j</span><span class="p">,</span> <span class="mi">0</span><span class="p">],</span> <span class="n">X</span><span class="p">[</span><span class="n">j</span><span class="p">,</span> <span class="mi">1</span><span class="p">],</span> <span class="n">X</span><span class="p">[</span><span class="n">j</span><span class="p">,</span> <span class="mi">2</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">dz</span> <span class="o">=</span> <span class="n">x</span> <span class="o">-</span> <span class="n">x0</span><span class="p">,</span> <span class="n">y</span> <span class="o">-</span> <span class="n">y0</span><span class="p">,</span> <span class="n">z</span> <span class="o">-</span> <span class="n">z0</span>
    <span class="n">d2</span> <span class="o">=</span> <span class="n">dx</span><span class="o">*</span><span class="n">dx</span> <span class="o">+</span> <span class="n">dy</span><span class="o">*</span><span class="n">dy</span> <span class="o">+</span> <span class="n">dz</span><span class="o">*</span><span class="n">dz</span>
    <span class="n">dist</span> <span class="o">=</span> <span class="n">tf</span><span class="p">.</span><span class="n">sqrt</span><span class="p">(</span><span class="n">d2</span> <span class="o">+</span> <span class="mf">1e-4</span><span class="p">)</span>
    <span class="n">Fg</span> <span class="o">=</span> <span class="o">-</span><span class="n">dx</span> <span class="o">/</span> <span class="p">(</span><span class="n">d2</span> <span class="o">+</span> <span class="mf">1e-4</span><span class="p">)</span> <span class="o">*</span> <span class="mf">1.0</span> <span class="o">/</span> <span class="n">dist</span>
    <span class="n">Fix</span><span class="p">.</span><span class="n">val</span> <span class="o">+=</span> <span class="n">Fg</span> <span class="o">*</span> <span class="n">dx</span>
    <span class="n">Fiy</span><span class="p">.</span><span class="n">val</span> <span class="o">+=</span> <span class="n">Fg</span> <span class="o">*</span> <span class="n">dy</span>
    <span class="n">Fiz</span><span class="p">.</span><span class="n">val</span> <span class="o">+=</span> <span class="n">Fg</span> <span class="o">*</span> <span class="n">dz</span>
    
<span class="n">Fi</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="mi">0</span><span class="p">]</span> <span class="o">=</span> <span class="n">Fix</span>
<span class="n">Fi</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="mi">1</span><span class="p">]</span> <span class="o">=</span> <span class="n">Fiy</span>
<span class="n">Fi</span><span class="p">[</span><span class="n">i</span><span class="p">,</span> <span class="mi">2</span><span class="p">]</span> <span class="o">=</span> <span class="n">Fiz</span>

<span class="n">Vnew</span> <span class="o">=</span> <span class="n">V</span> <span class="o">+</span> <span class="n">Fi</span> <span class="o">*</span> <span class="mf">0.001</span>
<span class="n">Xnew</span> <span class="o">=</span> <span class="n">X</span> <span class="o">+</span> <span class="n">Vnew</span> <span class="o">*</span> <span class="mf">0.001</span>
</code></pre></div></div>

<p>And unsurprisingly it absolutely demolishes every other contestant by a factor 2-3x at the cost of uglier code.</p>

<p>However, this is still nowhere near what a hand-tuned shader version of this would do. As you can actually do block reductions here. Preload a range of particles into groupshared, then do the force computation for this block, and repeat for the next blocks. You can do even better, if you preload something like 4-8 particles into registers, and compute their forces from there, leading to a staged block algorithm. This is how you usually try to optimize reductions that reuse a lot of memory, like <a href="https://siboehm.com/articles/22/CUDA-MMM">matrix multiplications</a>.</p>

<p>Of course, something like Taichi would actually win here against everyone, but its not our comparison target as you can’t write the simulation in vectorized tensor form there, and likely after I implement automatic groupshared cache generation the performance gap might go to 0, or become better, tho I suspect probably not without some really advanced heuristics.</p>

<h2 id="mnist-with-a-convolutional-network">MNIST with a convolutional network</h2>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/ML/MNIST/module.py">–Link to TensorFrost example—</a></p>

<p><a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/ML/MNIST/pytorch.py">–Link to equivalent PyTorch example–</a></p>

<p>After I implemented module support with the optimizers, first thing I made is a convolutional neural net for the Fashion MNIST classification problem.
This is a more classic ML problem and you would probably expect for PyTorch to just straight up win every time. Turns out, actually no, for small network sizes, TensorFrost can actually have a significant win.</p>

<p>Here is a comparison of the number of training iterations per second for the same minibatch size of <code class="language-plaintext highlighter-rouge">128</code> and with the same(ish) ADAM optimizer, for TensorFrost and for PyTorch</p>

<center><img src="/images/MNISTbench.png" height="400px" /></center>

<p>For really tiny networks there is a large win, and the performance drops linearly with more channels/neurons, and at around 8-32-128 becomes slower than eager mode PyTorch. I also have tried compiled PyTorch, but it somehow became slower, did they mistakingly fuse some things incorrectly? I don’t know. I also haven’t tried JAX here, but I suspect its probably somewhat faster than PyTorch. I also wonder if you can compile a training step in PyTorch, thats something that TensorFrost does by defaut, and it’s not a very fair comparison without it.</p>

<p>Honestly speaking, I’m not sure If the win that TensorFrost has at small sizes is due to just having less overhead, or due to better kernels at small size. It might also be that the training step in PyTorch has huge overhead for the minibatch creation. In TensorFrost I have the entire dataset on GPU and use it directly without intermediate steps.</p>

<p>As a bonus I captured the TensorFrost 16-128-512 pass in Nvidia Nsight, with debug regions (ran on RTX 3090)</p>

<center><img src="/images/nsight_mnist.png" height="200px" /></center>

<p><em>PS. It’s pretty annoying to benchmark these, as I usually work on Windows (arguably better for graphics dev), rather than on Linux, and the compiled GPU versions of JAX/PyTorch are only available on Linux. TensorFrost works on both platforms, and it is actually easier to port from Windows to Linux than the other way around</em></p>

<h2 id="what-about-some-more-advanced-models">What about some more advanced models?</h2>

<p>While I could have also tested something like LLM’s or diffusion models, I can pretty much guarantee that for anything that has its bottleneck in matrix multiplication or other linear algebra algorithms TensorFrost will very likely lose by a lot, at least <a href="https://siboehm.com/articles/22/CUDA-MMM">without implementing automation of more advanced optimizations</a> or without just calling external BLAS libraries like cuBLAS, which is doable, but I unfortuanately don’t have enough time to do that, as I focus more on just the compiler itself because its more important for my use cases.</p>

<h1 id="what-is-left-to-do">What is left to do</h1>

<p><strong><em>1. Better handling for small constant shapes</em></strong></p>

<p>One of the things that quite often happens with writing vectorized code for simple particle simulations, like the simple N-body example from above, is that all the vector variables are 3d, and can easily be stored in registers, but the current compiler will generate it as a […, a, b, 3] shaped kernel with a lot of duplicated arithmetic (unless you create your own vec class that operates only on scalars). This is something that could be optimized in the future, by generating the initial kernel with a smaller shape, like […, a, b, 1] and then unrolling the dimensions that are broadcast compared to the kernel’s shape.</p>

<p>Even more advanced optimizations can be performed from that starting point, like automatic caching of data into groupshared memory. Such a caching would require automatic read/write range analysis, for that I will need to implement a VM that executes the operations and computes the range of values they can have with interval arithmetic (perhaps more complex analysis too).</p>

<p><strong><em>2. Improved workgroup utilization</em></strong></p>

<p>At the moment the IR does not have a representation for groupshared memory, which is a big bottleneck for large matrix multiplications, and can be quite useful for optimizing some other algorithms, like <a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Rendering/fft2d.ipynb">FFT</a>/<a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Algorithms/sorting_tests.py">sort</a>/<a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/examples/Rendering/convolution.py">convolutions</a>.</p>

<p>Kernels can be additionally fused at the workgroup level as a post kernel fusion pass, which is quite complicated, and requires additional groupshared memory with group syncs, and a correct estimation of the group size, but for small neural networks it could be a massive speedup.</p>

<p><strong><em>3. Automatic vectorization</em></strong></p>

<p>One nice thing, that I would like to borrow from JAX is <code class="language-plaintext highlighter-rouge">vmap</code>. Writing particle simulations in vectorized form often annoyingly require a lot of <code class="language-plaintext highlighter-rouge">squeeze</code> and <code class="language-plaintext highlighter-rouge">unsqueeze</code> operations, which could be automated if you vectorized a single particle calculation into a target shape. In fact, I could also make the explicit <code class="language-plaintext highlighter-rouge">kernel</code> node usage assume that all its children are scalar (if not - unroll), and it would also behave similarly to <code class="language-plaintext highlighter-rouge">vmap</code> with the exception of it forcibly creating a single kernel no matter what. Implementing <code class="language-plaintext highlighter-rouge">vmap</code> is, I supect, not to difficult, as it only requires padding all the shapes of the child operation with the given shape (with some peculiarities). Syntactically it could look like</p>

<div class="language-py highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="k">with</span> <span class="n">tf</span><span class="p">.</span><span class="n">vmap</span><span class="p">(</span><span class="n">shape</span><span class="p">)</span> <span class="k">as</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="c1">#stuff
</span></code></pre></div></div>
<p>With i,j,k being scalar at trace time, then padded with given shape.</p>

<p><strong><em>4. Better IR?</em></strong></p>

<p>Another thing that is perhaps very high in the priority list right now is to implement repeating computation eliminator, its quite troublesome due to requiring some way to compare entire computation chains (by making a specific computation result node have a unique number), but would remove quite a lot of redundant cruft from the generated kernels.</p>

<p>And on top of that the compilation time scales quadratically with Tensor Program size right now. This is something that I have started to notice with the more complex projects that I tried to do here, like Variational Monte Carlo or Neural Cellular Automata. The number of generated kernels there reaches hundreds, and for NCA sometimes even thousands due to unrolled iteration loop.</p>

<p>While this wasn’t really an issue for most of the simpler stuff, where it was usually bottlenecked by the shader/c++ compiler, this will become a problem for more serious projects in the future.</p>

<p>The way the compiler is written is still very sub-optimal, and more at a research-grade state at the moment. Some operations over every kernel require a complete update of the IR which clearly will make it scale quadratically. This is clearly due to me not having the required experience to properly write a good data structure for a compiler.</p>

<p>So perhaps replacing my own IR with LLVM could make more sense in the long run, though I’m still not sure about the specifics, and how easy it would be to integrate.</p>

<p><strong><em>5. Documentation</em></strong></p>

<p>Right now the only documentation is provided in the <a href="https://github.com/MichaelMoroz/TensorFrost/blob/main/README.md">README.md</a> file in the repository, in the future I should made a separate documentation page for this.</p>

<p><strong><em>6. Easier debugging / profiling</em></strong></p>

<p>While ideally I would have wanted the compiler to generate optimal or close to optimal programs - this is still very much not the case, and some simple things like using/not using reshape might change the performance by an order of magnitude just due to, for example, some reductions now being over one dimension instead of several and the compiler not having a way to optimize this scenario (this was a problem in NCA).</p>

<p>Ideally I’d expose the compiled result in a more readible / easy to access way. Perhaps something like giving the ability to compile the program purely into its IR, so that you can edit it, for example.</p>

<p><strong><em>7. More test cases</em></strong></p>

<p>Compiler tends to fall apart when trying to implement a completely new thing right now - this was kind of expected to be honest, without a large enough set of tests it is almost guaranteed that it will fall apart at some edge cases, and the only way to fix this is just to make a whole lot of example projects which I’m currently working on bit by bit</p>

<p><strong><em>8. Improve the GLFW/ImGui integration</em></strong></p>

<p>I only passed the bare minimum for these to be able to make basic windows and GUI, but ideally all their features should be exposed. This is quite a large task and will take a lot of time I suspect. I also want to integrate ImPlot for much better plots, as for real-time use cases I think those will be more useful than matplotlib.</p>

<p><strong><em>9. Generate python code for the host part of the program</em></strong></p>

<p>As usually the compilation bottleneck is the C++ compiler, I’m considering to just use Python for the kernel dispatch code. This might reduce the performance a bit, but for faster iteration this is certainly going to be useful.</p>

<p><strong><em>10. A whole bunch of basic Numpy functionality is still missing</em></strong></p>

<p>Things like concatenation/stacking/splitting/repeat/etc does not exist yet, and I currently emulate them manually by reindexing, which works but is tedious.</p>

<p>Something like a random number generation module would be quite useful. As they are stateful I want to expose all their internals to the user so I think making them in the form of a <code class="language-plaintext highlighter-rouge">tf.Module</code> makes sense. So something like <code class="language-plaintext highlighter-rouge">rng = tf.random.rng_module(type=..., seed=...)</code> I suppose. Which you use then like <code class="language-plaintext highlighter-rouge">rng.unform(2.0,3.0,shape)</code> or other. You could then pass them around as parts of bigger modules together with their seeds.</p>

<p>Linear algebra algorithms, like QR, SVD, LU, eigendecomposition, inverse, determinants, etc, are not part of the library, which is an issue for more complex data analysis and algorithms. Unfortunately these are <em>very</em> hard to implement from scratch, espectially with any resemblence of performance or numerical precision. So I suppose I have no choice but to use an external library for these in the future.</p>

<p>And there are other ones like FFT, sorting etc. Tho I have them implemented in the examples, I just need to integrate them into the library.</p>

<p><strong><em>11. More tensor formats</em></strong></p>

<p>Right now we only have int32, uint32, float32 and bool, which is not many. Something like custom quantized formats that Taichi provides would be very nice to have. 
Adding the option to pack the last dimension if its constant and small with custom packing would also allow quantized vectors of sorts, the shared exponent format is very useful for simulations.</p>

<p>Perhaps also specifying the behaviour of the tensor on out-of-bounds reads/writes could be nice to add to the format, as the default one is just clamp.</p>

<p>Adding support for HW supported GPU texture formats would be nice for rendering algoirhms, as tensors which are represented as textures will have performance benefits for render-like use cases, not to mention the optional ability to use the HW linear sampler.</p>

<p><em>There are probably a million more things that I forgot, but even without those, this is enough work for years to come.</em></p>

<h1 id="conclusion">Conclusion</h1>

<p>Making a compiler with practically 0 compiler development experience was certainly quite the gamble. I’m 100% sure there are some fundamental architectural decisions that I did wrong, the main of it being perhaps writing the IR from scratch, as I had no clue how LLVM worked even tho I knew that it is used everywhere, and probably for a very good reason. Writing the compiler like that took an astounding amount of time, while initially I’ve expected a few months of work, now I’m at 14 months and some core features I wanted are still missing.</p>

<p>Quite often it was hard to convince myself that this is a useful use of my time, the scope of work is so huge that making progress felt like an eternity, and I could have instead spend that time on making other pet projects. Even now, I can not easily explain why exactly I wanted to make a library/compiler like this, as there are hundreds of little things that come together to make it useful for what I do, and I hope this blog post explains at least some part of my thought process here.</p>

<p>Right now the library is at the point of being somewhat usable. In fact, given the performance tests from above, if you properly improved the performance of the algorithmic primitives I use, like reductions and matmuls, and added more backends like CUDA, this library or a new version of it might become a viable choice for more common ML applications.</p>

<p>If anyone wants to help me with development, PR’s are welcome! There are still like a million things missing, lots of hidden bugs waiting to be found, and by myself it would take a few more years for it to get into a more mature state.</p>

<hr />

<details><summary>Why TensorFrost?</summary><p>
The name was chosen from "tensor" + my surname translated from Ukrainian to English. I know, its not super creative, given how many "TensorSomething" already exist. Also there is the funny problem that LLM's mistakingly assume its TensorFlow. Perhaps I should do `import TensorFrost as fr` instead of `as tf` in my examples.
</p></details>]]></content><author><name></name></author><summary type="html"><![CDATA[In this blog post I want to talk about the research and development results for a library that I started working on more than a year ago - TensorFrost. Under the hood it’s a static optimizing tensor compiler with a focus on being able to do more “shader-like” things while still keeping the ability to do high level linear algebra for ML in Numpy-like syntax with automatic differentiation support. (Click on the example GIF’s for more details!)]]></summary></entry><entry><title type="html">Visualizing General Relativity</title><link href="https://michaelmoroz.github.io/TracingGeodesics/" rel="alternate" type="text/html" title="Visualizing General Relativity" /><published>2022-08-21T00:00:00+00:00</published><updated>2022-08-21T00:00:00+00:00</updated><id>https://michaelmoroz.github.io/TracingGeodesics</id><content type="html" xml:base="https://michaelmoroz.github.io/TracingGeodesics/"><![CDATA[<p>When thinking about renders of things like warp drives and black holes we usually just expect to see a simple approximation or an artist rendition, assuming the math required to pull off something accurate would require someone with at least a PhD in Mathematical Physics. Which I won’t tell that its completely untrue, but in this blog post I’ll try to explain a way to do quite accurate visualizations within a 100 or so lines of code, for basically any kind of space-time for which you can write its metric as code. Although, the detailed mathematical derivation of this approach might be somewhat math heavy.</p>

<p>The main ingredient of any GR render is figuring out how the rays of light move around. Knowing how light moves we can trace rays from the camera into the scene and see where the light came from. So, to render the simplest scene without objects we simply trace a ray for each pixel and assign the color of the pixel to the color of the skybox in the direction in which the ray ends up pointing to.</p>

<ul>
  <li><a href="#what-are-geodesics">What are geodesics?</a></li>
  <li><a href="#mathematical-description-of-shortest-path">Mathematical description of shortest path</a></li>
  <li><a href="#lagrangian-description-of-a-geodesic">Lagrangian description of a geodesic</a></li>
  <li><a href="#hamiltonian-description-of-a-geodesic">Hamiltonian description of a geodesic</a></li>
  <li><a href="#writing-a-hamiltonian-geodesic-tracer-in-glsl">Writing a Hamiltonian geodesic tracer in GLSL</a></li>
  <li><a href="#conclusions">Conclusions</a></li>
  <li><a href="#references">References</a></li>
</ul>

<hr />

<h2 id="what-are-geodesics">What are geodesics?</h2>
<p>How exactly do we trace rays in curves space? Any object inside a curved space follows something called a geodesic.</p>

<p>A geodesic is just a fancy word for, in some sense, a path of shortest length between 2 points inside a space.</p>

<p>I should note that there could be multiple of such paths, which are locally minimal, in the sense that you can’t nudge the path to make it shorter, while globally there might be a shorter path. Also, in Minkowski space-time the definition is a bit more complicated, because of the imaginary distances (when \( ds^2 &lt; 0 \) ).</p>

<p>In our case, instead of paths between 2 points, we are only interested in finding how a ray moves given an initial point and direction, but the definition above will still prove useful when deriving the equations describing a geodesic, which we will use here.</p>

<hr />

<h2 id="mathematical-description-of-shortest-path">Mathematical description of shortest path</h2>
<p><em>I’ll try to quickly go through the derivation. If you wish to skip the math part, jump to <a href="#writing-a-hamiltonian-geodesic-tracer-in-glsl">Writing a Hamiltonian geodesic tracer in GLSL</a>.</em></p>

<p>Mathematically speaking we have some coordinate system, a path, and a way to compute distances between 2 points.</p>

<p>A coordinate system being a set of several numbers labeling each point in the space. A path is a function that takes in the path parameter and outputs a coordinate, in General Relativity the path parameter is usually proper time (like a clock moving with the object, labeling each point), but it can be anything really. The way to compute distances is called a metric (and it’s the main source of scary math here).</p>

<p>In physics, or more generally differential geometry, a metric is defined as an integral(“sum”) of something called the metric tensor. A metric tensor is a bilinear form \( g(a, b) \), it essentially maps pairs of vectors to real numbers, and is a generalization of dot product for curved spaces. So using a metric tensor we can find the length of a vector in space, and also distances \( ds \) between infinitely close points in space.</p>

<p>\begin{equation}
   ds^2 = g(dx, dx)
\end{equation}</p>

<p>In our case, where we describe vectors as a set of numbers, a metric is simply a matrix product of some matrix \( g_{\mu \nu} \) times the vectors. For our infinitesimal distance \( ds \) we get this expression:</p>

<p>\begin{equation}
  ds^2 = \sum_{\mu \nu}^N g_{\mu \nu} dx^\mu dx^\nu
\end{equation}</p>

<p>Usually the sum is just implicitly assumed by <a href="https://en.wikipedia.org/wiki/Einstein_notation">Einstein notation</a> [1].</p>

<p>\begin{equation}
 ds^2 = g_{\mu \nu} dx^\mu dx^\nu 
\end{equation}</p>

<p>Here we can actually see that for some simple choice of \( g_{\mu \nu} \) we can get the distances by Pythagoras’ theorem. Specifically for the case when the metric tensor matrix is a unit matrix.</p>

<p>\begin{equation}
 ds^2 = dx_1^2 + dx_2^2 + dx_3^2 
\end{equation}</p>

<p>For a flat space-time like space, we get something similar but with the exception that the time coordinate component is with a negative sign.</p>

<p>\begin{equation}
 ds^2 = - dx_0^2 + dx_1^2 + dx_2^2 + dx_3^2 
 \label{flat}
\end{equation}</p>

<p>Here I used the \( (- + + +) \) signature, but signs can actually be flipped without changing the geodesics, and in some cases, like for particle physics, it makes more sense to use the opposite \( (+ - - -) \) signature.</p>

<p>Going back to the main question of computing distances, to compute the length between 2 points, A and B, along some path \( x^i(t) \) we simply need to sum the infinitesimal distances together using an integral:</p>

<p>\begin{equation}
 l = \int_A^B \sqrt{g_{\mu \nu} dx^\mu dx^\nu} = \int_A^B \sqrt{g_{\mu \nu} dx^\mu dx^\nu} \frac{dt}{dt} = \int_A^B \sqrt{g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt}} dt
\end{equation}</p>

<p>Where \( \frac{dx^i}{dt} \) is simply how fast the coordinate x changes with respect to the path parameter (“clock”), in some sense can be interpreted as the velocity.</p>

<p>Now our main question is how do we minimize the path length? Here is where we introduce a thing called calculus of variations, which is roughly speaking a way to find how a functional(distance) changes by infinitesimally small variations of its input function(path). Such derivatives have similar properties to normal function derivatives. And in fact, similarly to calculus, to find the extremum of a function (min, max or stationary point), we simply need to equate the variation to 0.</p>

<hr />

<h2 id="lagrangian-description-of-a-geodesic">Lagrangian description of a geodesic</h2>
<p>There is an entire branch of physics related to variational principles, which states that any kind of physical system has a value it likes to minimize (or more generally make unchanging under small variations of path, i.e. “stationary”). That value is called <a href="https://en.wikipedia.org/wiki/Action_(physics)">action</a>, and the function under the integral is called the <a href="https://en.wikipedia.org/wiki/Lagrangian_mechanics">Lagrangian function of the system</a>. The branch of physics studying Lagrangian functions of systems is called Lagrangian mechanics.</p>

<p>In our case the Lagrangian can be written like this:</p>

<p>\begin{equation}
 L = \sqrt{g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt}}
\end{equation}</p>

<p>Turns out we don’t need the square root for the minimum of the functional to be a geodesic, and we can simply remove it from our geodesic Lagrangian:</p>

<p>\begin{equation}
 L = g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} 
\end{equation}</p>

<p>The proof of this you can find <a href="https://physics.stackexchange.com/questions/149082/geodesic-equation-from-variation-is-the-squared-lagrangian-equivalent">here</a> [2]. The only difference such simplification makes is that the parametrization of the path might be different, but the path itself will be the same. Also notably, the equations for this specific case turn out to be the same, meaning the parametrization is also the same.</p>

<p>Additionally, we want this with a 1/2 factor, to simplify the equations down the line.</p>

<p>\begin{equation}
 L = \frac{1}{2} g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} 
\end{equation}</p>

<p>So, our goal right now is to minimize this functional:</p>

<p>\begin{equation}
 S = \int_A^B \frac{1}{2} g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} dt 
\end{equation}</p>

<p>In general the minimum of a functional like this can be found by applying the <a href="https://en.wikipedia.org/wiki/Euler%E2%80%93Lagrange_equation">Euler-Lagrange equations</a> [3]:</p>

<p>\begin{equation}
  \frac{\partial L}{\partial x^i} - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx^i}{dt}} = 0 
  \label{el}
\end{equation}</p>

<hr />

<details>
<summary>Euler-Lagrange equation derivation</summary>

In general, the Action can be written like

\begin{equation*}
  S = \int_A^B  L(t, x(t), \frac{dx(t)}{dt}) dt 
\end{equation*}

Where the Lagrangian is a function of the path parameter, the path itself, and the derivative of the path with respect to the path parameter(also known as the generalized velocity).

To find the minimizing path (or more generally, stationary path) of a functional we need to equate the variation of the action to 0

\begin{equation*}
  \delta S = 0
\end{equation*}

Where the variation of the action is found by adding a small variation \( \delta x \) to the path: \( L(t, x + \delta x, \frac{d(x + \delta x)}{dt}) \) and expanding the 2D Taylor series around the point \( \left( x, \frac{dx(t)}{dt} \right)\)

\begin{equation*}
  \delta S = \int_A^B \left( \frac{\partial L}{\partial x} \delta x + \frac{\partial L}{\partial \frac{dx}{dt}} \frac{d(\delta x)}{dt} \right) dt
\end{equation*}

We dropped the higher order terms since we assume \( \delta x \) to be infinitesimally small.

Then we use integration by parts to get the derivative \( \frac{d}{dt} \) off the path variation

\begin{equation*}
  \delta S = \int_A^B \left( \frac{\partial L}{\partial x} \delta x - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx}{dt}} \delta x \right) dt + \left( \frac{\partial L}{\partial \frac{dx}{dt}} \delta x \right) \biggr \rvert_A^B
\end{equation*}

Since we keep the endpoints of the path stationary the last term is equal to zero:

\begin{equation*}
  \delta S = \int_A^B \left( \frac{\partial L}{\partial x} \delta x - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx}{dt}} \delta x \right) dt =
  \int_A^B \left( \frac{\partial L}{\partial x} - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx}{dt}}  \right) \delta x  dt
\end{equation*}

Equating this to zero we get

\begin{equation*}
 \int_A^B \left( \frac{\partial L}{\partial x} - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx}{dt}}  \right) \delta x dt = 0
\end{equation*}

Which holds true when the path satisfies the expression under the integral

\begin{equation*}
  \frac{\partial L}{\partial x} - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx}{dt}} = 0
\end{equation*}

Which is the Euler-Lagrange equation!


</details>

<p>You can find a more detailed derivation <a href="https://mathworld.wolfram.com/Euler-LagrangeDifferentialEquation.html">here</a> [4]</p>

<hr />

<p>Let’s derive the Euler-Lagrange equations for our geodesic Lagrangian (keep in mind that there is an equation for each coordinate \( x^i \) ):</p>

<p>\begin{equation}
 \frac{\partial L}{\partial \frac{dx^i}{dt}} = 
    \frac{1}{2} \frac{\partial  }{\partial \frac{dx^i}{dt}} g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} = 
    \frac{1}{2} g_{i \nu} \frac{dx^\nu}{dt} + \frac{1}{2} g_{\mu i} \frac{dx^\mu}{dt} = 
    g_{i \nu} \frac{dx^\nu}{dt} 
\label{el0}
\end{equation}</p>

<p>Then we take the derivative with respect to the path parameter:</p>

<p>\begin{equation}
 \frac{d}{dt} \left( g_{i \nu} \frac{dx^\nu}{dt} \right) =   \frac{d g_{i \nu} }{dt}  \frac{dx^\nu}{dt} + g_{i \nu} \frac{d^2x^\nu}{dt^2} = 
 \frac{d g_{i \nu} }{dx^\mu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} +  g_{i \nu} \frac{d^2x^\nu}{dt^2} 
\label{el1}
\end{equation}</p>

<p>And lastly:</p>

<p>\begin{equation}
 \frac{\partial L}{\partial x^i} = \frac{1}{2} \frac{d g_{\mu \nu} }{dx^i} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} 
 \label{el2}
\end{equation}</p>

<p>Substituting \eqref{el1} and \eqref{el2} into Euler-Lagrange equations \eqref{el} leads us to the equation of a geodesic:</p>

<p>\begin{equation}
 \frac{1}{2} \frac{d g_{\mu \nu} }{dx^i} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} - \frac{d g_{i \nu} }{dx^\mu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} - g_{i \nu} \frac{d^2x^\nu}{dt^2} = 0 
\end{equation}</p>

<p>Multiplying by the metric tensor inverse \( - g^{i \nu} \) we get:</p>

<p>\begin{equation}
 \frac{d^2x^i}{dt^2} +  g^{i \nu} \left( \frac{d g_{i \nu} }{dx^\mu} - \frac{1}{2} \frac{d g_{\mu \nu} }{dx^i}  \right) \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} = 0 
\end{equation}</p>

<p>And that’s our final system of equations for a geodesic, we could of course also substitute the Christoffel symbols here, but for our application there is no difference. Of course, we could just use these equations for tracing geodesic rays and call it a day, but unfortunately this would require computing a whole lot of derivatives (in 4d space-time it’s 64 of them to be specific), either manually, or by using numerical differentiation. Thankfully there is a way to avoid this, and in fact simplify the entire algorithm! (At a slight performance cost)</p>

<hr />

<h2 id="hamiltonian-description-of-a-geodesic">Hamiltonian description of a geodesic</h2>
<p>So here comes the star of the show - Hamiltonian mechanics. Hamiltonian equations of motion have a really nice form which allows to easily write a computer program that integrates them by using Euler integration.</p>

<p>\begin{equation}
 \frac{dp^i}{dt} = - \frac{\partial H}{\partial x^i} 
\end{equation}
\begin{equation}
 \frac{dx^i}{dt} =   \frac{\partial H}{\partial p^i} 
\end{equation}</p>

<p>Where \( p \) is the so called generalized momentum, it’s the derivative of the Lagrangian with respect to the coordinate path parameter derivative.</p>

<p>\begin{equation}
 p_i = \frac{\partial L}{\partial \frac{dx^i}{dt} } 
 \label{momentumdef}
\end{equation}</p>

<p>And to get the Hamiltonian itself you need to apply the <a href="https://blog.jessriedel.com/2017/06/28/legendre-transform/">Legendre Transform</a> [6] on the Lagrangian:</p>

<p>\begin{equation}
 H = \sum_{i}^N p^i \frac{dx^i}{dt} - L 
 \label{legandre}
\end{equation}</p>

<hr />

<details>
<summary>Hamilton equations of motion derivation</summary>

Lets start by writing down the Euler-Lagrange equation

\begin{equation*}
  \frac{\partial L}{\partial x} - \frac{d}{dt} \frac{\partial L}{\partial \frac{dx}{dt}} = 0
\end{equation*}

You can see that \( \frac{\partial L}{\partial \frac{dx}{dt} } \) is equal to our definition of generalized momentum \eqref{momentumdef}, so we can substitude it here

\begin{equation*}
  \frac{\partial L}{\partial x} - \frac{dp}{dt} = 0
\end{equation*}

Now lets substitude \(H\) instead of \(L\) by using the definition \eqref{legandre}

\begin{equation*}
  \frac{\partial}{\partial x} \left(  p \frac{dx}{dt} - H \right) - \frac{dp}{dt} = 0
\end{equation*}

The partial derivative of \( p \frac{dx}{dt}  \) with respect to \(x\) is 0, since changing \(x\) doesn't change \( p \) or \( \frac{dx}{dt} \) 

\begin{equation*}
  \frac{\partial}{\partial x} \left(- H \right) - \frac{dp}{dt} = 0
\end{equation*}

After moving things around 

\begin{equation*}
  \frac{dp}{dt} = - \frac{\partial H}{\partial x}
\end{equation*}

Which is our first equation.

Now lets take the partial derivative of \eqref{legandre} with respect to generalized momentum

\begin{equation*}
 \frac{\partial H}{\partial p} =  \frac{\partial}{\partial p} \left( p \frac{dx}{dt}  \right) - \frac{\partial L}{\partial p}  
\end{equation*}

Since \(L\) doesn't depend on \(p\) (as the value of \(L\) doesn't depend on its partial derivative), it means that \( \frac{\partial L}{\partial p} = 0 \), so we get

\begin{equation*}
 \frac{\partial H}{\partial p} =  \frac{\partial}{\partial p} \left( p \frac{dx}{dt} \right)    
\end{equation*}

\begin{equation*}
 \frac{\partial H}{\partial p} =  \frac{dx}{dt} 
\end{equation*}

\begin{equation*}
 \frac{dx}{dt} = \frac{\partial H}{\partial p}
\end{equation*}

Which is our second equation.

</details>

<p>A different derivation of Hamilton’s equations of motion can be found <a href="https://en.wikipedia.org/wiki/Hamiltonian_mechanics#Deriving_Hamilton's_equations">here</a> [5].</p>

<hr />

<p>And for our case the momentum would be the following, which we already computed when writing down the Euler-Lagrange equations \eqref{el0}, and given the definition of generalized momentum \eqref{momentumdef}:</p>

<p>\begin{equation}
 p_i = g_{i j} \frac{dx^j}{dt} 
 \label{momentum}
\end{equation}</p>

<p>To get the “time” derivatives you simply need to multiply both sides by the metric tensor inverse:</p>

<p>\begin{equation}
 \frac{dx^i}{dt} = g^{i j} p_j 
 \label{dxdt}
\end{equation}</p>

<p>And the Hamiltonian itself:</p>

<p>\begin{equation}
 H = \sum_{i}^N \frac{dx^i}{dt} p_i - L =  g_{i j} \frac{dx^i}{dt} \frac{dx^j}{dt} - \frac{1}{2} g_{i j} \frac{dx^i}{dt} \frac{dx^j}{dt} =  \frac{1}{2} g_{i j} \frac{dx^i}{dt} \frac{dx^j}{dt} = L
\end{equation}</p>

<p>Turns out that for this simple choice of a geodesic Lagrangian, the Hamiltonian is equal to the Lagrangian!</p>

<p>Also, we want to know the Hamiltonian as a function of the generalized momentum by substituting \eqref{dxdt} into the Hamiltonian equation:</p>

<p>\begin{equation}
 H = \frac{1}{2} g_{i j} \frac{dx^i}{dt} \frac{dx^j}{dt} = \frac{1}{2} g^{i j} p_i p_j 
 \label{hamiltonian}
\end{equation}</p>

<p>While the equations of motion will simply be:</p>

<p>\begin{equation}
 \frac{dp_i}{dt} = - \frac{\partial H}{\partial x^i} 
 \label{eqmotion1}
\end{equation}</p>

<p>\begin{equation}
 \frac{dx^i}{dt} = g^{i j} p_j 
 \label{eqmotion2}
\end{equation}</p>

<p>This is all we need to write a numerical geodesic integrator!</p>

<hr />

<h2 id="writing-a-hamiltonian-geodesic-tracer-in-glsl">Writing a Hamiltonian geodesic tracer in GLSL</h2>

<p>You might have noticed that in the final Hamilton’s equations of motion I didn’t write out \( \frac{\partial H}{\partial x^i} \), this is actually important! We want to keep the derivative of the Hamiltonian as is, because then instead of computing the 64 derivatives of the metric tensor, we only need 4 to find the Hamiltonian gradient. This is the main simplification of the geodesic tracing algorithm.</p>

<p>Here we will use the GLSL shading language, since it has variables and functions which map quite well to the mathematical operations we will perform here. On top of that we can easily then make a real time GR visualization shader.</p>

<p>First of all, we need a function that evaluates the metric tensor at a 4d point in space and time. Let’s use the <a href="https://en.wikipedia.org/wiki/Alcubierre_drive">Alcubierre warp drive</a> [7] metric as an example, since it is quite simple.</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">mat4</span> <span class="nf">Metric</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">)</span>
<span class="p">{</span>
  <span class="c1">//Alcubierre metric  </span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">R</span> <span class="o">=</span> <span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">sigma</span> <span class="o">=</span> <span class="mi">35</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span> 
  <span class="k">const</span> <span class="kt">float</span> <span class="n">v</span> <span class="o">=</span> <span class="mi">1</span><span class="p">.</span><span class="mi">1</span><span class="p">;</span>

  <span class="kt">float</span> <span class="n">x</span> <span class="o">=</span> <span class="n">v</span><span class="o">*</span><span class="n">x</span><span class="p">.</span><span class="n">x</span><span class="p">;</span>
  <span class="kt">float</span> <span class="n">r</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">sqr</span><span class="p">(</span><span class="n">x</span><span class="p">.</span><span class="n">y</span> <span class="o">-</span> <span class="n">x</span><span class="p">)</span> <span class="o">+</span> <span class="n">x</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">x</span><span class="p">.</span><span class="n">z</span> <span class="o">+</span> <span class="n">x</span><span class="p">.</span><span class="n">w</span><span class="o">*</span><span class="n">x</span><span class="p">.</span><span class="n">w</span><span class="p">);</span>
  <span class="kt">float</span> <span class="n">f</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="p">(</span><span class="n">tanh</span><span class="p">(</span><span class="n">sigma</span><span class="o">*</span><span class="p">(</span><span class="n">r</span> <span class="o">+</span> <span class="n">R</span><span class="p">))</span> <span class="o">-</span> <span class="n">tanh</span><span class="p">(</span><span class="n">sigma</span><span class="o">*</span><span class="p">(</span><span class="n">r</span> <span class="o">-</span> <span class="n">R</span><span class="p">)))</span><span class="o">/</span><span class="n">tanh</span><span class="p">(</span><span class="n">sigma</span><span class="o">*</span><span class="n">R</span><span class="p">);</span>
  <span class="kt">float</span> <span class="n">gtt</span> <span class="o">=</span> <span class="n">v</span><span class="o">*</span><span class="n">v</span><span class="o">*</span><span class="n">f</span><span class="o">*</span><span class="n">f</span> <span class="o">-</span> <span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>
  <span class="kt">float</span> <span class="n">gxt</span> <span class="o">=</span> <span class="o">-</span><span class="n">v</span><span class="o">*</span><span class="n">f</span><span class="p">;</span>
  
  <span class="k">return</span> <span class="kt">mat4</span><span class="p">(</span><span class="n">gtt</span><span class="p">,</span> <span class="n">gxt</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">gxt</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="mi">0</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="mi">1</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="mi">0</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="p">}</span>

</code></pre></div></div>

<p>In our case x is a 4D vector representing position. The first component <code class="language-plaintext highlighter-rouge">x.x</code> or <code class="language-plaintext highlighter-rouge">x[0]</code> being time. As an output we get a 4 by 4 matrix represented by <code class="language-plaintext highlighter-rouge">mat4</code> in GLSL.</p>

<p>Then we need to write down the Hamiltonian \eqref{hamiltonian}. The Hamiltonian is a function that takes 2 things, the position in space-time, and the 4d momentum vector, and outputs a scalar.</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">float</span> <span class="nf">Hamiltonian</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
  <span class="kt">mat4</span> <span class="n">g_inv</span> <span class="o">=</span> <span class="n">inverse</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">));</span>
  <span class="k">return</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="n">dot</span><span class="p">(</span><span class="n">g_inv</span><span class="o">*</span><span class="n">p</span><span class="p">,</span><span class="n">p</span><span class="p">);</span>
<span class="p">}</span>

</code></pre></div></div>

<p>As a bonus here is the Lagrangian</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">float</span> <span class="nf">Lagrangian</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec4</span> <span class="n">dxdt</span><span class="p">)</span>
<span class="p">{</span>
  <span class="k">return</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="n">dot</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">)</span><span class="o">*</span><span class="n">dxdt</span><span class="p">,</span><span class="n">dxdt</span><span class="p">);</span>
<span class="p">}</span>

</code></pre></div></div>

<p>Surprisingly enough that’s it, GLSL already has a matrix inverse function <code class="language-plaintext highlighter-rouge">inverse()</code>, on top of it the Hamiltonian is just the dot product(in GLSL sense) of <code class="language-plaintext highlighter-rouge">g_inv*p</code> and <code class="language-plaintext highlighter-rouge">p</code>, which are the contravariant and covariant momentum vectors respectively. The contravariant momentum actually just being the time derivative of the coordinate <code class="language-plaintext highlighter-rouge">dxdt</code>, i.e. <code class="language-plaintext highlighter-rouge">dot(dxdt,p)</code>.</p>

<p>After this we need to compute the 4D gradient of the Hamiltonian. We can do this by using a forward numerical difference in all 4 spacial directions, using some small value <code class="language-plaintext highlighter-rouge">eps</code>:</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">vec4</span> <span class="nf">HamiltonianGradient</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">eps</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mo">001</span><span class="p">;</span>
  <span class="k">return</span> <span class="p">(</span><span class="kt">vec4</span><span class="p">(</span><span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">eps</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="mi">0</span><span class="p">),</span> <span class="n">p</span><span class="p">),</span>
               <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="n">eps</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">p</span><span class="p">),</span>
               <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</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">eps</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">p</span><span class="p">),</span>
               <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</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="mi">0</span><span class="p">,</span><span class="n">eps</span><span class="p">),</span> <span class="n">p</span><span class="p">))</span> <span class="o">-</span> <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span><span class="p">,</span><span class="n">p</span><span class="p">))</span><span class="o">/</span><span class="n">eps</span><span class="p">;</span>
<span class="p">}</span>

</code></pre></div></div>

<p>Now that we have the Hamiltonian gradient, we can finally write down the equation of motion \eqref{eqmotion1} \eqref{eqmotion2} integration code</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">vec4</span> <span class="nf">IntegrationStep</span><span class="p">(</span><span class="k">inout</span> <span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="k">inout</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">TimeStep</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">1</span><span class="p">;</span>
  <span class="n">p</span> <span class="o">=</span> <span class="n">p</span> <span class="o">-</span> <span class="n">TimeStep</span> <span class="o">*</span> <span class="n">HamiltonianGradient</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">p</span><span class="p">);</span>
  <span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="n">TimeStep</span> <span class="o">*</span> <span class="n">inverse</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">))</span> <span class="o">*</span> <span class="n">p</span><span class="p">;</span>
<span class="p">}</span>

</code></pre></div></div>

<p>You might ask, “wait, that’s it?”, and indeed that is all you need to integrate the geodesic. Of course, it is quite slow since we do a whopping 6 matrix inverse evaluations, which can be optimized down to 1, by replacing most Hamiltonians with Lagrangians which don’t have inverses, since they are equal. Even better is to have the metric inverse already computed analytically, but it’s not possible for every metric, especially for an implicitly defined one.</p>

<p>There is of course the last problem, while initializing the space-time position is easy, how do we initialize the value of the momentum vector <code class="language-plaintext highlighter-rouge">p</code> when starting to trace?</p>

<p>Before tracing the geodesic, you can use the equation \eqref{momentum}</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="n">p</span> <span class="o">=</span> <span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">*</span> <span class="n">dxdt</span><span class="p">;</span>

</code></pre></div></div>
<p>But what is dxdt? It’s nothing more than the 4D direction the ray moves inside space-time. There are 3 categories the directions can fall into:</p>
<ul>
  <li>Time-like, when \( A &lt; 0 \)</li>
  <li>Null, when \( A = 0 \)</li>
  <li>Space-like, when \( A &gt; 0 \)</li>
</ul>

<p>Where \(A\) is
\begin{equation}
  A = g_{\mu \nu} \frac{dx^\mu}{dt} \frac{dx^\nu}{dt} 
\end{equation}</p>

<p>or in GLSL <code class="language-plaintext highlighter-rouge">A = dot(Metric(x) * dxdt, dxdt)</code></p>

<p>When simulating how light travels we just want null directions, which lead to null geodesic solutions. On the other hand, if you want to model an object moving slower than light you need a time-like geodesic. And space-like geodesics for tachyonic stuff, which doesn’t happen in real life, so we ignore it.</p>

<p>So, assuming a flat space metric \eqref{flat} and some 3D direction in space our p for a light ray would be</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">vec4</span> <span class="nf">GetNullMomentum</span><span class="p">(</span><span class="kt">vec3</span> <span class="n">dir</span><span class="p">)</span>
<span class="p">{</span>
  <span class="k">return</span> <span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">*</span> <span class="kt">vec4</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">normalize</span><span class="p">(</span><span class="n">dir</span><span class="p">));</span>
<span class="p">}</span>

</code></pre></div></div>

<p>And the inverse of this operation</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">vec3</span> <span class="nf">GetDirection</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
  <span class="kt">vec4</span> <span class="n">dxdt</span> <span class="o">=</span> <span class="n">inverse</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">))</span> <span class="o">*</span> <span class="n">p</span><span class="p">;</span>
  <span class="k">return</span> <span class="n">normalize</span><span class="p">(</span><span class="n">dxdt</span><span class="p">.</span><span class="n">yzw</span><span class="p">);</span>
<span class="p">}</span>

</code></pre></div></div>

<p>So, in the end the final simple algorithm will look like this:</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">void</span> <span class="nf">TraceGeodesic</span><span class="p">(</span><span class="k">inout</span> <span class="kt">vec3</span> <span class="n">pos</span><span class="p">,</span> <span class="k">inout</span> <span class="kt">vec3</span> <span class="n">dir</span><span class="p">,</span> <span class="k">inout</span> <span class="kt">float</span> <span class="n">time</span><span class="p">)</span>
<span class="p">{</span>
  <span class="kt">vec4</span> <span class="n">x</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">time</span><span class="p">,</span> <span class="n">pos</span><span class="p">);</span>
  <span class="kt">vec4</span> <span class="n">p</span> <span class="o">=</span> <span class="n">GetNullMomentum</span><span class="p">(</span><span class="n">dir</span><span class="p">);</span>

  <span class="k">const</span> <span class="kt">int</span> <span class="n">steps</span> <span class="o">=</span> <span class="mi">256</span><span class="p">;</span>
  <span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">steps</span><span class="p">;</span> <span class="n">i</span><span class="o">++</span><span class="p">)</span>
  <span class="p">{</span>
    <span class="n">IntegrationStep</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">p</span><span class="p">);</span>
    <span class="c1">//you can add a stop condition here when x is below the event horizon for example</span>
  <span class="p">}</span>

  <span class="n">pos</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">yzw</span><span class="p">;</span>
  <span class="n">time</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">x</span><span class="p">;</span>
  <span class="n">dir</span> <span class="o">=</span> <span class="n">GetDirection</span><span class="p">(</span><span class="n">p</span><span class="p">);</span>
<span class="p">}</span>

</code></pre></div></div>

<p>Essentially this is just a 4D ray marching algorithm where the direction of the ray changes every step. In this specific case the size of the step also changes, which can be avoided by normalizing the momentum <code class="language-plaintext highlighter-rouge">p = normalize(p)</code>. This only changes the step length, and doesn’t change the geodesic path, i.e., it works just like a dynamic reparameterization of the path. The time step of the integration can also be varied depending on the metric used. For example, in the case of black holes I change the time step proportionally to the distance to the event horizon, so that the accuracy of the geodesic is roughly proportional to the curvature of space. This is an important optimization to get accurate results, while keeping the computational cost relatively small.</p>

<p>Example shadertoy code:</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">mat4</span> <span class="nf">diag</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">a</span><span class="p">)</span>
<span class="p">{</span>
    <span class="k">return</span> <span class="kt">mat4</span><span class="p">(</span><span class="n">a</span><span class="p">.</span><span class="n">x</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="mi">0</span><span class="p">,</span>
                <span class="mi">0</span><span class="p">,</span><span class="n">a</span><span class="p">.</span><span class="n">y</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="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="n">a</span><span class="p">.</span><span class="n">z</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="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="n">a</span><span class="p">.</span><span class="n">w</span><span class="p">);</span>
<span class="p">}</span>

<span class="kt">mat4</span> <span class="nf">Metric</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">)</span>
<span class="p">{</span>
    <span class="c1">//Kerr-Newman metric in Kerr-Schild coordinates </span>
    <span class="k">const</span> <span class="kt">float</span> <span class="n">a</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">8</span><span class="p">;</span>
    <span class="k">const</span> <span class="kt">float</span> <span class="n">m</span> <span class="o">=</span> <span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>
    <span class="k">const</span> <span class="kt">float</span> <span class="n">Q</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>
    <span class="kt">vec3</span> <span class="n">p</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">yzw</span><span class="p">;</span>
    <span class="kt">float</span> <span class="n">rho</span> <span class="o">=</span> <span class="n">dot</span><span class="p">(</span><span class="n">p</span><span class="p">,</span><span class="n">p</span><span class="p">)</span> <span class="o">-</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="p">;</span>
    <span class="kt">float</span> <span class="n">r2</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="p">(</span><span class="n">rho</span> <span class="o">+</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">rho</span><span class="o">*</span><span class="n">rho</span> <span class="o">+</span> <span class="mi">4</span><span class="p">.</span><span class="mi">0</span><span class="o">*</span><span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="p">));</span>
    <span class="kt">float</span> <span class="n">r</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">r2</span><span class="p">);</span>
    <span class="kt">vec4</span> <span class="n">k</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="n">r</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">x</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">y</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">r2</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="p">),</span> <span class="p">(</span><span class="n">r</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">y</span> <span class="o">-</span> <span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">x</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">r2</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="p">),</span> <span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="o">/</span><span class="n">r</span><span class="p">);</span>
    <span class="kt">float</span> <span class="n">f</span> <span class="o">=</span> <span class="n">r2</span><span class="o">*</span><span class="p">(</span><span class="mi">2</span><span class="p">.</span><span class="mi">0</span><span class="o">*</span><span class="n">m</span><span class="o">*</span><span class="n">r</span> <span class="o">-</span> <span class="n">Q</span><span class="o">*</span><span class="n">Q</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">r2</span><span class="o">*</span><span class="n">r2</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="p">);</span>
    <span class="k">return</span> <span class="n">f</span><span class="o">*</span><span class="kt">mat4</span><span class="p">(</span><span class="n">k</span><span class="p">.</span><span class="n">x</span><span class="o">*</span><span class="n">k</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">y</span><span class="o">*</span><span class="n">k</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">k</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">w</span><span class="o">*</span><span class="n">k</span><span class="p">)</span><span class="o">+</span><span class="n">diag</span><span class="p">(</span><span class="kt">vec4</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">1</span><span class="p">));</span>
<span class="p">}</span>

<span class="kt">float</span> <span class="nf">Hamiltonian</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
    <span class="kt">mat4</span> <span class="n">g_inv</span> <span class="o">=</span> <span class="n">inverse</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">));</span>
    <span class="k">return</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="n">dot</span><span class="p">(</span><span class="n">g_inv</span><span class="o">*</span><span class="n">p</span><span class="p">,</span><span class="n">p</span><span class="p">);</span>
<span class="p">}</span>

<span class="cm">/*
float Lagrangian(vec4 x, vec4 dxdt)
{
    return 0.5*dot(Metric(x)*dxdt,dxdt);
}
*/</span>

<span class="kt">vec4</span> <span class="nf">HamiltonianGradient</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
    <span class="k">const</span> <span class="kt">float</span> <span class="n">eps</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mo">001</span><span class="p">;</span>
    <span class="k">return</span> <span class="p">(</span><span class="kt">vec4</span><span class="p">(</span><span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">eps</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="mi">0</span><span class="p">),</span> <span class="n">p</span><span class="p">),</span>
                 <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="n">eps</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">p</span><span class="p">),</span>
                 <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</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">eps</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span> <span class="n">p</span><span class="p">),</span>
                 <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span> <span class="o">+</span> <span class="kt">vec4</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="mi">0</span><span class="p">,</span><span class="n">eps</span><span class="p">),</span> <span class="n">p</span><span class="p">))</span> <span class="o">-</span> <span class="n">Hamiltonian</span><span class="p">(</span><span class="n">x</span><span class="p">,</span><span class="n">p</span><span class="p">))</span><span class="o">/</span><span class="n">eps</span><span class="p">;</span>
<span class="p">}</span>

<span class="kt">void</span> <span class="nf">IntegrationStep</span><span class="p">(</span><span class="k">inout</span> <span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="k">inout</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
    <span class="k">const</span> <span class="kt">float</span> <span class="n">TimeStep</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">15</span><span class="p">;</span>
    <span class="n">p</span> <span class="o">=</span> <span class="n">p</span> <span class="o">-</span> <span class="n">TimeStep</span> <span class="o">*</span> <span class="n">HamiltonianGradient</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">p</span><span class="p">);</span>
    <span class="n">x</span> <span class="o">=</span> <span class="n">x</span> <span class="o">+</span> <span class="n">TimeStep</span> <span class="o">*</span> <span class="n">inverse</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">))</span> <span class="o">*</span> <span class="n">p</span><span class="p">;</span>
<span class="p">}</span>

<span class="kt">vec4</span> <span class="nf">GetNullMomentum</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec3</span> <span class="n">dir</span><span class="p">)</span>
<span class="p">{</span>
    <span class="k">return</span> <span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">)</span> <span class="o">*</span> <span class="kt">vec4</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">normalize</span><span class="p">(</span><span class="n">dir</span><span class="p">));</span>
<span class="p">}</span>

<span class="kt">vec3</span> <span class="nf">GetDirection</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec4</span> <span class="n">p</span><span class="p">)</span>
<span class="p">{</span>
    <span class="kt">vec4</span> <span class="n">dxdt</span> <span class="o">=</span> <span class="n">inverse</span><span class="p">(</span><span class="n">Metric</span><span class="p">(</span><span class="n">x</span><span class="p">))</span> <span class="o">*</span> <span class="n">p</span><span class="p">;</span>
    <span class="k">return</span> <span class="n">normalize</span><span class="p">(</span><span class="n">dxdt</span><span class="p">.</span><span class="n">yzw</span><span class="p">);</span>
<span class="p">}</span>

<span class="kt">void</span> <span class="nf">TraceGeodesic</span><span class="p">(</span><span class="k">inout</span> <span class="kt">vec3</span> <span class="n">pos</span><span class="p">,</span> <span class="k">inout</span> <span class="kt">vec3</span> <span class="n">dir</span><span class="p">,</span> <span class="k">inout</span> <span class="kt">float</span> <span class="n">time</span><span class="p">)</span>
<span class="p">{</span>
    <span class="kt">vec4</span> <span class="n">x</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">time</span><span class="p">,</span> <span class="n">pos</span><span class="p">);</span>
    <span class="kt">vec4</span> <span class="n">p</span> <span class="o">=</span> <span class="n">GetNullMomentum</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">dir</span><span class="p">);</span>

    <span class="k">const</span> <span class="kt">int</span> <span class="n">steps</span> <span class="o">=</span> <span class="mi">256</span><span class="p">;</span>
    <span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="n">steps</span><span class="p">;</span> <span class="n">i</span><span class="o">++</span><span class="p">)</span>
    <span class="p">{</span>
        <span class="n">IntegrationStep</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">p</span><span class="p">);</span>
        <span class="c1">//you can add a stop condition here when x is below the event horizon for example</span>
    <span class="p">}</span>

    <span class="n">pos</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">yzw</span><span class="p">;</span>
    <span class="n">time</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">x</span><span class="p">;</span>
    <span class="n">dir</span> <span class="o">=</span> <span class="n">GetDirection</span><span class="p">(</span><span class="n">x</span><span class="p">,</span> <span class="n">p</span><span class="p">);</span>
<span class="p">}</span>

<span class="kt">void</span> <span class="nf">mainImage</span><span class="p">(</span><span class="k">out</span> <span class="kt">vec4</span> <span class="n">fragColor</span><span class="p">,</span> <span class="k">in</span> <span class="kt">vec2</span> <span class="n">fragCoord</span><span class="p">)</span> <span class="p">{</span>
    <span class="kt">vec2</span> <span class="n">uv</span> <span class="o">=</span> <span class="mi">2</span><span class="p">.</span><span class="mi">0</span> <span class="o">*</span> <span class="p">(</span><span class="n">fragCoord</span> <span class="o">-</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span> <span class="o">*</span> <span class="n">iResolution</span><span class="p">.</span><span class="n">xy</span><span class="p">)</span> <span class="o">/</span> <span class="n">max</span><span class="p">(</span><span class="n">iResolution</span><span class="p">.</span><span class="n">x</span><span class="p">,</span> <span class="n">iResolution</span><span class="p">.</span><span class="n">y</span><span class="p">);</span>

    <span class="kt">vec3</span> <span class="n">RayPos</span> <span class="o">=</span> <span class="kt">vec3</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="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">,</span>  <span class="mi">32</span><span class="p">.</span><span class="mi">0</span><span class="p">);</span>
    <span class="kt">vec3</span> <span class="n">RayDir</span> <span class="o">=</span> <span class="kt">vec3</span><span class="p">(</span><span class="n">uv</span><span class="p">.</span><span class="n">x</span><span class="p">,</span> <span class="n">uv</span><span class="p">.</span><span class="n">y</span><span class="p">,</span> <span class="o">-</span><span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="p">);</span>
    <span class="kt">float</span>  <span class="n">Time</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>

    <span class="n">RayDir</span> <span class="o">=</span> <span class="n">normalize</span><span class="p">(</span><span class="n">RayDir</span><span class="p">);</span>

    <span class="n">TraceGeodesic</span><span class="p">(</span><span class="n">RayPos</span><span class="p">,</span> <span class="n">RayDir</span><span class="p">,</span> <span class="n">Time</span><span class="p">);</span>

    <span class="n">fragColor</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">texture</span><span class="p">(</span><span class="n">iChannel0</span><span class="p">,</span> <span class="n">RayDir</span><span class="p">).</span><span class="n">rgb</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="p">}</span>
</code></pre></div></div>

<p>You can check out this Shadertoy implementation to see some of the optimizations, like variable timestep, replacing Hamiltonians with Lagrangians, using a symmetric matrix inversion function (a bit faster), reusing some of the computed values (restart if the Shadertoy is black):</p>

<center><iframe width="900" height="500" frameborder="0" src="https://www.shadertoy.com/embed/NtSGWG?gui=true&amp;t=10&amp;paused=false&amp;muted=false" allowfullscreen=""></iframe></center>

<p>The shader above implements both the Alcubierre metric, and the <a href="https://en.wikipedia.org/wiki/Kerr%E2%80%93Newman_metric#Kerr%E2%80%93Schild_coordinates">Kerr–Newman metric in Kerr-Schild coordinates</a> [8] (essentially Cartesian coordinates).</p>

<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code>
<span class="kt">mat4</span> <span class="nf">diag</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">a</span><span class="p">)</span>
<span class="p">{</span>
    <span class="k">return</span> <span class="kt">mat4</span><span class="p">(</span><span class="n">a</span><span class="p">.</span><span class="n">x</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="mi">0</span><span class="p">,</span>
                <span class="mi">0</span><span class="p">,</span><span class="n">a</span><span class="p">.</span><span class="n">y</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="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="n">a</span><span class="p">.</span><span class="n">z</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="mi">0</span><span class="p">,</span><span class="mi">0</span><span class="p">,</span><span class="n">a</span><span class="p">.</span><span class="n">w</span><span class="p">);</span>
<span class="p">}</span>

<span class="kt">mat4</span> <span class="nf">Metric</span><span class="p">(</span><span class="kt">vec4</span> <span class="n">x</span><span class="p">)</span>
<span class="p">{</span>
  <span class="c1">//Kerr-Newman metric in Kerr-Schild coordinates </span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">a</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">8</span><span class="p">;</span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">m</span> <span class="o">=</span> <span class="mi">1</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>
  <span class="k">const</span> <span class="kt">float</span> <span class="n">Q</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span>
  <span class="kt">vec3</span> <span class="n">p</span> <span class="o">=</span> <span class="n">x</span><span class="p">.</span><span class="n">yzw</span><span class="p">;</span>
  <span class="kt">float</span> <span class="n">rho</span> <span class="o">=</span> <span class="n">dot</span><span class="p">(</span><span class="n">p</span><span class="p">,</span><span class="n">p</span><span class="p">)</span> <span class="o">-</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="p">;</span>
  <span class="kt">float</span> <span class="n">r2</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="p">(</span><span class="n">rho</span> <span class="o">+</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">rho</span><span class="o">*</span><span class="n">rho</span> <span class="o">+</span> <span class="mi">4</span><span class="p">.</span><span class="mi">0</span><span class="o">*</span><span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="p">));</span>
  <span class="kt">float</span> <span class="n">r</span> <span class="o">=</span> <span class="n">sqrt</span><span class="p">(</span><span class="n">r2</span><span class="p">);</span>
  <span class="kt">vec4</span> <span class="n">k</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="mi">1</span><span class="p">,</span> <span class="p">(</span><span class="n">r</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">x</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">y</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">r2</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="p">),</span> <span class="p">(</span><span class="n">r</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">y</span> <span class="o">-</span> <span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">x</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">r2</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="p">),</span> <span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="o">/</span><span class="n">r</span><span class="p">);</span>
  <span class="kt">float</span> <span class="n">f</span> <span class="o">=</span> <span class="n">r2</span><span class="o">*</span><span class="p">(</span><span class="mi">2</span><span class="p">.</span><span class="mi">0</span><span class="o">*</span><span class="n">m</span><span class="o">*</span><span class="n">r</span> <span class="o">-</span> <span class="n">Q</span><span class="o">*</span><span class="n">Q</span><span class="p">)</span><span class="o">/</span><span class="p">(</span><span class="n">r2</span><span class="o">*</span><span class="n">r2</span> <span class="o">+</span> <span class="n">a</span><span class="o">*</span><span class="n">a</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">p</span><span class="p">.</span><span class="n">z</span><span class="p">);</span>
  <span class="k">return</span> <span class="n">f</span><span class="o">*</span><span class="kt">mat4</span><span class="p">(</span><span class="n">k</span><span class="p">.</span><span class="n">x</span><span class="o">*</span><span class="n">k</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">y</span><span class="o">*</span><span class="n">k</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">z</span><span class="o">*</span><span class="n">k</span><span class="p">,</span> <span class="n">k</span><span class="p">.</span><span class="n">w</span><span class="o">*</span><span class="n">k</span><span class="p">)</span><span class="o">+</span><span class="n">diag</span><span class="p">(</span><span class="kt">vec4</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">1</span><span class="p">,</span><span class="mi">1</span><span class="p">));</span>    
<span class="p">}</span>

</code></pre></div></div>

<p>This is, of course, not the limit for optimization. The main other optimization is computing analytical inverses of the metric. For a large class of metrics you can use the <a href="https://en.wikipedia.org/wiki/Sherman%E2%80%93Morrison_formula">Sherman–Morrison formula</a> [9]</p>

<h3 id="some-useful-things-to-keep-in-mind">Some useful things to keep in mind</h3>

<p>Since we used numerical finite differences, the results can actually depend quite a lot on the relative values of the float numbers. For example, the accuracy of the numerical derivatives is a lot lower far from the coordinate system center, so you’d need to vary the size of <code class="language-plaintext highlighter-rouge">eps</code> to avoid excessive numerical noise. Also, metrics quite often have numerical singularities which you should avoid, unless you want to get NaN results.</p>

<p>I usually avoid metrics in spherical coordinates due to their polar axis singularity, which has strong visual effects which are extremely hard to avoid even with a tiny varying timestep, although such metrics are usually mathematically simpler and allow for larger timesteps without breaking the look of the Black hole. For spherically symmetric metrics, like non-spinning black holes and wormholes there is a trick to avoid the polar singularity altogether! The thing about spherical symmetry is that the geodesic is always moving inside a 2d plane, which can be mapped to the equatorial plane of the coordinate system, basically reducing the 3d + 1 time problem to 2d + time. (<a href="https://youtu.be/PVO8nvb1o2w">Scott Manley has a video explaining how he rendered wormholes</a>, in there he used this trick + precomputing a lookup table to simplify the computation by a lot)</p>

<p>I’ve also used the dimensionality reduction trick in my wormhole shadertoy:</p>

<center><iframe width="900" height="500" frameborder="0" src="https://www.shadertoy.com/embed/stByz1?gui=true&amp;t=10&amp;paused=true&amp;muted=false" allowfullscreen=""></iframe></center>

<p>There is also a different method that can be used to compute derivatives numerically, and way more accurately. Essentially it’s forward automatic differentiation based on dual numbers, <a href="https://www.shadertoy.com/view/3tGcRt">there are some Shadertoy example which have used this approach</a>.</p>

<p>And finally you could always derive the equations analytically, while this is the most annoying method it is usually the fastest performance-wise. A compromise solution would be derive the equations automatically (like in Wolfram Alpha), this approach is used by <a href="https://github.com/20k/geodesic_raytracing">geodesic_raytracing</a> made by <a href="https://twitter.com/berrow_james">James Berrow</a> (you should follow him on Twitter, he has a lot of cool stuff on this topic).</p>

<p>Figuring out if the ray has fallen inside the event horizon is actually not trivial, and there is no universal method, and while you could just set the color to 0 if the ray is below the event horizon surface, this is incorrect when viewing things from inside the black hole. Tracing the rays should also be done backwards in time, since we trace the rays from the camera, not to the camera, this has a noticeble effect on the resulting render, if not done this also results in completely dark renders inside of black holes, even though light does exist under the event horizon, and can reach from the outside.</p>

<p><a href="https://figshare.com/s/02c8b839dfb53f6e1a59">Redshift in General Relativity</a> can be computed simply from the ratio of the dot products of the velocity of the object at emission/absorption times the momentum of the photon, where the momentum of the photon is parallel to the photons direction of movement. The momentum needs to be parallel transported along the geodesic from emission to absorption, and the good news is - we can just use the geodesic 4-velocity instead, since it also is technically parallel transported along the geodesic and is pointed in the direction of movement. (Just in case, this also means we need to avoid renormalizing the generalized momentum when integrating the equations)</p>

<p>Also notably, the dot product in General Relativity is simply defined from the metric tensor \( g(u,v) \)</p>

<p>\begin{equation}
 u \cdot v = g_{\mu \nu} u^\mu v^\nu 
\end{equation}</p>

<p>Finally, if you only want to compute geodesics for Schwarzschild black holes, you can simply use the equation for a particle with mass 1 in a certain classical force field, details are in 
<a href="https://rantonels.github.io/starless/">this blog post</a>.</p>

<hr />

<h2 id="conclusions">Conclusions</h2>

<center><iframe width="900" height="500" src="https://www.youtube.com/embed/mst0BoDTQdo" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture" allowfullscreen=""></iframe></center>

<p>Using this ray tracing algorithm, you can basically render whatever you want inside any definable space-time. This algorithm was used to render different warped space-times inside of Space Engine, you can check out the blog posts about this:</p>

<ul>
  <li><a href="https://spaceengine.org/news/blog220705/">Kerr black holes</a></li>
  <li><a href="https://spaceengine.org/news/blog220812/">Alcubierre warp fields and wormholes</a></li>
  <li><a href="https://spaceengine.org/news/blog220705/">Volumetric accretion disks around a Kerr black hole</a></li>
</ul>

<p>Fast volumetric ray tracing with geodesics is quite difficult, and we needed to separate the ray marching loop into 2 loops, main loop being the geodesic steps, and the second loop being the volumetric substeps. Since we also use blue noise, it was necessary to keep the steps uniform along the geodesic, otherwise there would be clear artifacts in the volume, which required a few tricks with having a variable number of substeps per geodesic step.</p>

<p>Combining this with SDF’s is somewhat easier, you need to vary the geodesic step to be the min() between the current step size and the SDF. Using this I’ve also tried to make a really simple path tracer in Unity with a Kerr black hole, naturally it was quite slow.</p>

<center><iframe width="900" height="500" src="https://www.youtube.com/embed/_s01oUxTG5I" title="YouTube video player" frameborder="0" allow="accelerometer; autoplay; clipboard-write; encrypted-media; gyroscope; picture-in-picture" allowfullscreen=""></iframe></center>

<p>The project is <a href="https://github.com/The-Order-of-the-Simulation/SpaceTimePathTracer">here</a>, but don’t expect very readable code, this was mostly intended as an experiment.</p>

<p>Note that rendering <strong>moving</strong> objects is waaay harder, and requires either to have a space-time SDF, or some insane acceleration structure for triangles. And on top of that the entire history of the scene’s past needs to be kept in memory, the only simple cases are when the moving objects can be represented as analytical functions you can sample in space and time, like the volumetric accretion disk in Space Engine.</p>

<hr />

<h3 id="references">References</h3>
<ul>
  <li>[1] <a href="https://en.wikipedia.org/wiki/Einstein_notation">Einstein notation</a></li>
  <li>[2] <a href="https://physics.stackexchange.com/questions/149082/geodesic-equation-from-variation-is-the-squared-lagrangian-equivalent">Equivalence of squared Lagrangian to Lagrangian</a></li>
  <li>[3] <a href="https://en.wikipedia.org/wiki/Euler%E2%80%93Lagrange_equation">Euler-Lagrange equations</a></li>
  <li>[4] <a href="https://mathworld.wolfram.com/Euler-LagrangeDifferentialEquation.html">Euler-Lagrange equations derivation</a></li>
  <li>[5] <a href="https://en.wikipedia.org/wiki/Hamiltonian_mechanics#Deriving_Hamilton's_equations">Hamiltonian equations derivation</a></li>
  <li>[6] <a href="https://blog.jessriedel.com/2017/06/28/legendre-transform/">Legendre Transform</a></li>
  <li>[7] <a href="https://en.wikipedia.org/wiki/Alcubierre_drive">Alcubierre metric</a></li>
  <li>[8] <a href="https://en.wikipedia.org/wiki/Kerr%E2%80%93Newman_metric#Kerr%E2%80%93Schild_coordinates">Kerr–Newman metric in Kerr-Schild coordinates</a></li>
  <li>[9] <a href="https://en.wikipedia.org/wiki/Sherman%E2%80%93Morrison_formula">Sherman–Morrison formula</a></li>
  <li>[10] <a href="https://github.com/The-Order-of-the-Simulation/SpaceTimePathTracer">Space-time path tracer</a></li>
</ul>]]></content><author><name></name></author><summary type="html"><![CDATA[When thinking about renders of things like warp drives and black holes we usually just expect to see a simple approximation or an artist rendition, assuming the math required to pull off something accurate would require someone with at least a PhD in Mathematical Physics. Which I won’t tell that its completely untrue, but in this blog post I’ll try to explain a way to do quite accurate visualizations within a 100 or so lines of code, for basically any kind of space-time for which you can write its metric as code. Although, the detailed mathematical derivation of this approach might be somewhat math heavy.]]></summary><media:thumbnail xmlns:media="http://search.yahoo.com/mrss/" url="https://michaelmoroz.github.io/SpaceEngineBH.jpg" /><media:content medium="image" url="https://michaelmoroz.github.io/SpaceEngineBH.jpg" xmlns:media="http://search.yahoo.com/mrss/" /></entry><entry><title type="html">Reintegration tracking</title><link href="https://michaelmoroz.github.io/Reintegration-Tracking/" rel="alternate" type="text/html" title="Reintegration tracking" /><published>2020-08-31T00:00:00+00:00</published><updated>2020-08-31T00:00:00+00:00</updated><id>https://michaelmoroz.github.io/Reintegration-Tracking</id><content type="html" xml:base="https://michaelmoroz.github.io/Reintegration-Tracking/"><![CDATA[<p>In this blog post I’ll explain this advection algorithm and how to use it to make advanced fluid simulations like the ones I made, including <a href="https://www.shadertoy.com/view/WtfyDj">Paint streams</a> and <a href="https://www.shadertoy.com/view/ttBcWm">Everflow</a>.
Before starting, I should give a big thanks to my friend <a href="https://www.shadertoy.com/user/wyatt">Wyatt</a> for giving useful suggestions on building this algorithm.</p>

<center><iframe style="width:640px;height:360px;" frameborder="0" src="https://www.shadertoy.com/embed/WtfyDj?gui=true&amp;t=10&amp;paused=false" allowfullscreen=""></iframe></center>

<p>Looks really smooth, doesn’t it? It even captures droplets with almost close to pixel level precision. And while it does model a fluid with a boundary it does not use particles directly like in SPH, but is actually a semi-Lagrangian grid based algorithm. Oh, and I forgot to say - it’s also super fast.</p>

<p>Let’s start with a brief history, the initial idea for this algorithm came from trying to extend screen space <a href="https://www.facebook.com/groups/shadertoy/?post_id=567902837124080">voronoi particle tracking</a> to avoid particle loss, indeed when trying to store the particle state in screen space as close to the particle location as we can, in the case of overlap one of the particles will be lost, and there is no way to avoid it. But actually we can reformulate the problem, what if we don’t try to avoid particle overlap but tried to conserve the total mass? Like by adding the overlapping particle masses together and weighting their velocities and positions by mass. 
Actually that was already tried by <a href="https://www.shadertoy.com/view/MdtGDX">stb</a>, but it immediately becomes apparent that the total number of particles drops proportionally to the particle/pixel density because of them combining, you might ask if there is a way to separate them back, or if there is a way to achieve approximate particle number conservation? In fact, it is possible, but let’s overview how he does the particle tracking first.</p>

<h3 id="cellular-automaton-particle-tracking">Cellular automaton particle tracking</h3>
<p>This algorithm is similar to a <a href="https://en.wikipedia.org/wiki/Lattice_gas_automaton">lattice gas automaton</a> where the state of each cell is defined by a few discrete levels, is the particle in the cell and in which discrete direction it is moving. In our case instead of using a few discrete states we store the position, velocity and mass of the particle within the cell as floats <em>(or any other number type, depends on the precision you want to achieve, I actually used 2 int16’s per channel since Shadertoy only has a 4 channel output per pixel and I needed to store at least 5 numbers)</em></p>

<center>
<table>
  <tr>
    <th><img src="/images/ParticleCAframe1.JPG" style="width:250px;height:250px;" /></th>
    <th><img src="/images/ParticleCAframe2.JPG" style="width:250px;height:250px;" /></th>
  </tr>
  <tr>
    <th><b>Initial Frame</b></th>
    <th><b>Next frame</b></th>
  </tr>
</table>
</center>

<p>The main idea of the algorithm is to loop over all neighbors of the current cell(including itself), integrate the position of each neighbor particle and add it if it ends up in this cell. Above is a visualization of how it looks like. Each cell has a particle with a mass, the mass is shown by opacity, so 0 mass is invisible and 1 is completely dark. The red cell is our current cell for which we want to find its future state, the future position of the particles is shown by the arrow. We can see that there is more than one particle moving into the red cell. So in the end all the particles that end up in the same cell are summed and their positions and velocities are averaged. In mathematical form it can be written as:</p>

<p>\[ M_{i}^{t+1} = \sum_{j}^\textrm{neighbors} K_{i}(\vec{X}_ {j}^{t} + \Delta t \vec{V}_ {j}^{t}) m_{j}^{t}    —    \textrm{updated mass} \] 
\[ \vec{X}_ {i}^{t+1} = \frac{1}{M_ {i}^{t+1}} \sum_{j}^\textrm{neighbors}  K_{i}(\vec{X}_ {j}^{t} + \Delta t \vec{V}_ {j}^{t}) (\vec{X}_ {j}^{t} + \Delta t \vec{V}_ {j}^{t}) m_{j}^{t}      —     \textrm{updated center of mass} \] 
\[ \vec{V}_ {i}^{t+1} = \frac{1}{M_ {i}^{t+1}} \sum_{j}^\textrm{neighbors}  K_{i}(\vec{X}_ {j}^{t} + \Delta t \vec{V}_ {j}^{t})  \vec{V}_ {j}^{t} m_{j}^{t}    —    \textrm{updated velocity} \] 
Where \(M_{i}^{t}\) is the mass of the particle in cell i on time step t, \(\vec{X}_ {i}^{t}\) is the position of the particle and \(\vec{V}_ {i}^{t}\) is the velocity. The function \(K_{i}(\vec{X})\) is equal to 1 if the point \(\vec{X}\) is inside the cell i, and zero otherwise. \( \Delta t \) is the timestep.</p>

<p>For a square cell the K function is simply
\[ K_{i}(\vec{X}) =  H(\vec{X} - \vec{C}_ {i} + 0.5) H(\vec{C}_ {i} + 0.5 - \vec{X}) \] 
Where \(H\) is the multivariate Heaviside step function. \(\vec{C}_{i}\) is the center of cell i.</p>

<p>It’s quite a simple algorithm, but it has one limitation, to ensure that we counted every possible particle that might end up in this cell we would either need to loop over the entire grid, which is highly expensive, or limit the maximum velocity of the particles to make the search radius finite, and hopefully only 1 pixel wide.</p>

<p>Obviously the second option is much better in our context, and in fact if we want to track particles with velocities so fast that they traverse the grid in 1 frame it’s actually cheaper to do lots of smaller steps with a 1 pixel neighborhood instead of counting all cells in one step, since \( 9 \sqrt{N} &lt; N \) , where N is the number of cells, and the square root is because we need to apply the operation only a linear amount of times, instead of applying it for every cell. And the 9 is the number of neighbors. In 3D it would equivalently be  \( 27 N^{1/3} &lt; N \)  which is even more efficient. I should note that more steps does not mean better quality, since the particles may combine together on the way.
Also we can just do the tracking on CPU, and do it in a forward way, just loop over all particles and add them into the right cells, then we would not care about those things, but we obviously want to use that GPU power to our advantage.</p>

<p>Here is a pseudo-code glsl implementation of this algorithm:</p>
<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1">//this cell position</span>
<span class="kt">ivec2</span> <span class="n">pos</span><span class="p">;</span>
<span class="c1">//values stored in the cell</span>
<span class="kt">vec2</span> <span class="n">velocity</span> <span class="o">=</span> <span class="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">.),</span> <span class="n">position</span> <span class="o">=</span> <span class="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">.);</span>
<span class="kt">float</span> <span class="n">mass</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.;</span>
<span class="c1">//find and average the particles </span>
<span class="c1">//that land in this cell after a time step dt</span>
<span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">x</span> <span class="o">=</span> <span class="o">-</span><span class="n">R</span><span class="p">;</span> <span class="n">x</span> <span class="o">&lt;=</span> <span class="n">R</span><span class="p">;</span> <span class="n">x</span><span class="o">++</span><span class="p">)</span> <span class="c1">//only check the neighbors at radius R</span>
  <span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">y</span> <span class="o">=</span> <span class="o">-</span><span class="n">R</span><span class="p">;</span> <span class="n">y</span> <span class="o">&lt;=</span> <span class="n">R</span><span class="p">;</span> <span class="n">y</span><span class="o">++</span><span class="p">)</span>
  <span class="p">{</span>
      <span class="c1">//get the particle in this neighbor cell from the previous frame</span>
      <span class="n">particle</span> <span class="n">P</span> <span class="o">=</span> <span class="n">getParticle</span><span class="p">(</span><span class="n">pos</span> <span class="o">+</span> <span class="kt">ivec2</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="c1">//integrate the particle position</span>
      <span class="n">P</span><span class="p">.</span><span class="n">X</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">V</span><span class="o">*</span><span class="n">dt</span><span class="p">;</span>
      <span class="c1">//check if the particle is inside of this cell</span>
      <span class="k">if</span><span class="p">(</span><span class="n">inCell</span><span class="p">(</span><span class="n">P</span><span class="p">,</span> <span class="n">pos</span><span class="p">))</span>
      <span class="p">{</span>
        <span class="n">mass</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">M</span><span class="p">;</span> <span class="c1">//add the particle mass to this cell</span>
        <span class="n">position</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">X</span><span class="o">*</span><span class="n">P</span><span class="p">.</span><span class="n">M</span><span class="p">;</span> <span class="c1">//add the particle position weighted by mass</span>
        <span class="n">velocity</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">V</span><span class="o">*</span><span class="n">P</span><span class="p">.</span><span class="n">M</span><span class="p">;</span> <span class="c1">//add the particle velocity weighted by mass(momentum)</span>
      <span class="p">}</span>
  <span class="p">}</span> 

<span class="c1">//normalize</span>
<span class="k">if</span><span class="p">(</span><span class="n">mass</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">)</span> <span class="c1">//if not vacuum</span>
<span class="p">{</span>
  <span class="n">position</span> <span class="o">/=</span> <span class="n">mass</span><span class="p">;</span> <span class="c1">//center of mass</span>
  <span class="n">velocity</span> <span class="o">/=</span> <span class="n">mass</span><span class="p">;</span> <span class="c1">//average velocity</span>
<span class="p">}</span>
</code></pre></div></div>

<h3 id="dividing-particles">Dividing particles</h3>

<p>To combat the problem of a decreasing particle number we can try and divide each particle into M virtual particles with distributed positions that might end up in different cells and thus increase the total number of particles.</p>
<center>
<img src="/images/ParticleCA_div_frame1.JPG" style="width:250px;height:250px;" />
</center>

<p>The radius of the distribution defines how likely is the particle is to multiply. If the all the virtual particles end up inside of a single cell the particle would not divide and stay essentially the same(if the average particle of the virtual particles is equal to the original particle). To make it properly conservative we just need to divide the mass of the particle into a number of equal chunks, and make sure the distribution average is zero.</p>

<p>On the picture above and in the code below we see an example for a 5 virtual particle distribution, the average of the distribution directions is zero and its radius is 0.1 pixels.</p>

<p>Here is the code, the only addition are the diffusion directions and a loop for each of the virtual particles.</p>
<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1">//this cell position</span>
<span class="kt">ivec2</span> <span class="n">pos</span><span class="p">;</span>
<span class="c1">//values stored in the cell</span>
<span class="kt">vec2</span> <span class="n">velocity</span> <span class="o">=</span> <span class="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">.),</span> <span class="n">position</span> <span class="o">=</span> <span class="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">.);</span>
<span class="kt">float</span> <span class="n">mass</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.;</span>

<span class="c1">//diffusion radius</span>
<span class="kt">float</span> <span class="n">difR</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">1</span><span class="p">;</span>
<span class="c1">//diffusion directions</span>
<span class="kt">vec2</span> <span class="n">difDir</span><span class="p">[</span><span class="mi">5</span><span class="p">]</span> <span class="o">=</span> <span class="p">{</span><span class="kt">vec2</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="kt">vec2</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="kt">vec2</span><span class="p">(</span><span class="o">-</span><span class="mi">1</span><span class="p">,</span><span class="mi">0</span><span class="p">),</span><span class="kt">vec2</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="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">,</span><span class="o">-</span><span class="mi">1</span><span class="p">)};</span>

<span class="c1">//find and average the particles </span>
<span class="c1">//that land in this cell after a time step dt</span>
<span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">x</span> <span class="o">=</span> <span class="o">-</span><span class="n">R</span><span class="p">;</span> <span class="n">x</span> <span class="o">&lt;=</span> <span class="n">R</span><span class="p">;</span> <span class="n">x</span><span class="o">++</span><span class="p">)</span> <span class="c1">//only check the neighbors at radius R</span>
  <span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">y</span> <span class="o">=</span> <span class="o">-</span><span class="n">R</span><span class="p">;</span> <span class="n">y</span> <span class="o">&lt;=</span> <span class="n">R</span><span class="p">;</span> <span class="n">y</span><span class="o">++</span><span class="p">)</span>
  <span class="p">{</span>
      <span class="c1">//get the particle in this neighbor cell from the previous frame</span>
      <span class="n">particle</span> <span class="n">P</span> <span class="o">=</span> <span class="n">getParticle</span><span class="p">(</span><span class="n">pos</span> <span class="o">+</span> <span class="kt">ivec2</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="c1">//integrate the particle position</span>
      <span class="n">P</span><span class="p">.</span><span class="n">X</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">V</span><span class="o">*</span><span class="n">dt</span><span class="p">;</span>
      <span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">i</span> <span class="o">=</span> <span class="mi">0</span><span class="p">;</span> <span class="n">i</span> <span class="o">&lt;</span> <span class="mi">5</span><span class="p">;</span> <span class="n">i</span><span class="o">++</span><span class="p">)</span> <span class="c1">//divide particle into 5</span>
      <span class="p">{</span>
        <span class="n">particle</span> <span class="n">difP</span> <span class="o">=</span> <span class="n">P</span><span class="p">;</span>
        <span class="n">difP</span><span class="p">.</span><span class="n">X</span> <span class="o">+=</span> <span class="n">difR</span><span class="o">*</span><span class="n">difDir</span><span class="p">[</span><span class="n">i</span><span class="p">];</span> <span class="c1">//move particle in one of the diffusion directions</span>
        <span class="n">difP</span><span class="p">.</span><span class="n">M</span> <span class="o">/=</span> <span class="mi">5</span><span class="p">.</span><span class="mi">0</span><span class="p">;</span> <span class="c1">//divide mass into 5 particles</span>

        <span class="c1">//check if the divided particle is inside of this cell</span>
        <span class="k">if</span><span class="p">(</span><span class="n">inCell</span><span class="p">(</span><span class="n">difP</span><span class="p">,</span> <span class="n">pos</span><span class="p">))</span>
        <span class="p">{</span>
          <span class="n">mass</span> <span class="o">+=</span> <span class="n">difP</span><span class="p">.</span><span class="n">M</span><span class="p">;</span> <span class="c1">//add the particle mass to this cell</span>
          <span class="n">position</span> <span class="o">+=</span> <span class="n">difP</span><span class="p">.</span><span class="n">X</span><span class="o">*</span><span class="n">difP</span><span class="p">.</span><span class="n">M</span><span class="p">;</span> <span class="c1">//add the particle position weighted by mass</span>
          <span class="n">velocity</span> <span class="o">+=</span> <span class="n">difP</span><span class="p">.</span><span class="n">V</span><span class="o">*</span><span class="n">difP</span><span class="p">.</span><span class="n">M</span><span class="p">;</span> <span class="c1">//add the particle velocity weighted by mass(momentum)</span>
        <span class="p">}</span>
      <span class="p">}</span>
  <span class="p">}</span> 

<span class="c1">//normalize</span>
<span class="k">if</span><span class="p">(</span><span class="n">mass</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">)</span> <span class="c1">//if not vacuum</span>
<span class="p">{</span>
  <span class="n">position</span> <span class="o">/=</span> <span class="n">mass</span><span class="p">;</span> <span class="c1">//center of mass</span>
  <span class="n">velocity</span> <span class="o">/=</span> <span class="n">mass</span><span class="p">;</span> <span class="c1">//average velocity</span>
<span class="p">}</span>
</code></pre></div></div>

<p>You might ask if its possible to achieve a perfect particle number conservation, and not just total mass conservation. Actually to do that we just need to divide the mass into integer chunks that sum into the original mass, the virtual particles need not be the same, so for example one particle with mass 3 can divide into 2 particles with mass 2 and 1. Dividing into more than 2 virtual particles is pretty much the same, even though a bit more complicated.</p>

<h3 id="using-particle-distributions">Using particle distributions</h3>

<p>The previous algorithm did solve the problem sufficiently well, but did have some shortcomings, like the additional loop which makes that algorithm far slower to execute and produces discontinuities in density that made it harder to implement fluid simulations. We can actually do better. Lets look at the limiting case of an infinite number of virtual particles, in such a case we end up with a continuous distribution of mass that we can call \( \rho(\vec{X}) \) (the distribution should also be centered on the particle position).</p>

<p>So to find the amount of mass (and its center) each particle deposits into this cell we need to integrate the distribution within the current cell bounds. (And that is where I got the name - <em>reintegrating tracked</em> particle distributions)
 Our equations are:
\[ M = \int \int_{\Omega} \rho(\vec{X})d\vec{X}    —     \textrm{deposited mass} \] 
\[ \vec{C} = \frac{1}{M} \int \int_{\Omega} \vec{X}\rho(\vec{X})d\vec{X}    —    \textrm{deposited center of mass} \] 
Where \( \Omega \) is the cell region. But which distribution should we use? If we try to use a normal distribution we will end up with the problem of how to compute it, since there is no analytical solution for such an integral, and we would need to numerically integrate the distribution, which is not much better performance-wise than the previous algorithm. The simplest distribution that gives an analytical solution we can use is actually a uniform axis aligned box, for which the mass and the center of mass are trivial to compute analytically. We can also try a uniform circular distribution or a nonuniform circular distribution equal to \( 1 - |\vec{X_0} - \vec{X}|^2 \) for \(  |\vec{X_0} - \vec{X}| \leqslant 1\) and equal to \( 0 \) for \(  |\vec{X_0} - \vec{X}| &gt; 1 \) where \( \vec{X_0} \) is the particle position, but let’s try out the simplest approach.</p>
<center>
<table>
  <tr>
    <th><img src="/images/Reintegration_0_2.JPG" style="width:250px;height:250px;" /></th>
    <th><img src="/images/Reintegration_0_55.JPG" style="width:250px;height:250px;" /></th>
  </tr>
  <tr>
    <th><b>Diffusion radius 0.2</b></th>
    <th><b>Diffusion radius 0.55</b></th>
  </tr>
</table>
</center>
<p>To find the mass and center of mass of the distribution within the bounds of the cell we only need to figure out the overlap box of the cell and the particle distribution. Its relative area will be the relative mass and its center is just the center of mass.
It can be implemented like this:</p>
<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="c1">//this cell position</span>
<span class="kt">ivec2</span> <span class="n">pos</span><span class="p">;</span>
<span class="c1">//values stored in the cell</span>
<span class="kt">vec2</span> <span class="n">velocity</span> <span class="o">=</span> <span class="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">.),</span> <span class="n">position</span> <span class="o">=</span> <span class="kt">vec2</span><span class="p">(</span><span class="mi">0</span><span class="p">.);</span>
<span class="kt">float</span> <span class="n">mass</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.;</span>
<span class="c1">//find and average the particles </span>
<span class="c1">//that land in this cell after a time step dt</span>
<span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">x</span> <span class="o">=</span> <span class="o">-</span><span class="n">R</span><span class="p">;</span> <span class="n">x</span> <span class="o">&lt;=</span> <span class="n">R</span><span class="p">;</span> <span class="n">x</span><span class="o">++</span><span class="p">)</span> <span class="c1">//only check the neighbors at radius R</span>
  <span class="k">for</span><span class="p">(</span><span class="kt">int</span> <span class="n">y</span> <span class="o">=</span> <span class="o">-</span><span class="n">R</span><span class="p">;</span> <span class="n">y</span> <span class="o">&lt;=</span> <span class="n">R</span><span class="p">;</span> <span class="n">y</span><span class="o">++</span><span class="p">)</span>
  <span class="p">{</span>
      <span class="c1">//get the particle in this neighbor cell from the previous frame</span>
      <span class="n">particle</span> <span class="n">P</span> <span class="o">=</span> <span class="n">getParticle</span><span class="p">(</span><span class="n">pos</span> <span class="o">+</span> <span class="kt">ivec2</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="c1">//integrate the particle position</span>
      <span class="n">P</span><span class="p">.</span><span class="n">X</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">V</span><span class="o">*</span><span class="n">dt</span><span class="p">;</span>
      <span class="c1">//find the overlap of the diffused particle distribution with this cell</span>
      
      <span class="kt">vec3</span> <span class="n">ovrlp</span> <span class="o">=</span> <span class="n">overlap</span><span class="p">(</span><span class="n">P</span><span class="p">.</span><span class="n">X</span><span class="p">,</span> <span class="n">pos</span><span class="p">,</span> <span class="n">diffusion_radius</span><span class="p">);</span>
      <span class="kt">float</span> <span class="n">overlapRelativeArea</span> <span class="o">=</span> <span class="n">ovrlp</span><span class="p">.</span><span class="n">z</span><span class="p">;</span>
      <span class="kt">vec2</span> <span class="n">overlapCenterOfMass</span> <span class="o">=</span> <span class="n">ovrlp</span><span class="p">.</span><span class="n">xy</span><span class="p">;</span>
      <span class="kt">float</span> <span class="n">overlapMass</span> <span class="o">=</span> <span class="n">overlapRelativeArea</span><span class="o">*</span><span class="n">P</span><span class="p">.</span><span class="n">M</span><span class="p">;</span>

      <span class="n">mass</span> <span class="o">+=</span> <span class="n">overlapMass</span><span class="p">;</span> <span class="c1">//add the overlap mass to this cell</span>
      <span class="n">position</span> <span class="o">+=</span> <span class="n">overlapCenterOfMass</span><span class="o">*</span><span class="n">overlapMass</span><span class="p">;</span> <span class="c1">//add the overlap center weighted by mass</span>
      <span class="n">velocity</span> <span class="o">+=</span> <span class="n">P</span><span class="p">.</span><span class="n">V</span><span class="o">*</span><span class="n">overlapMass</span><span class="p">;</span> <span class="c1">//add the particle velocity weighted by overlap mass(momentum)</span>
  <span class="p">}</span> 

<span class="c1">//normalize</span>
<span class="k">if</span><span class="p">(</span><span class="n">mass</span> <span class="o">&gt;</span> <span class="mi">0</span><span class="p">.</span><span class="mi">0</span><span class="p">)</span> <span class="c1">//if not vacuum</span>
<span class="p">{</span>
  <span class="n">position</span> <span class="o">/=</span> <span class="n">mass</span><span class="p">;</span> <span class="c1">//center of mass</span>
  <span class="n">velocity</span> <span class="o">/=</span> <span class="n">mass</span><span class="p">;</span> <span class="c1">//average velocity</span>
<span class="p">}</span>
</code></pre></div></div>

<p>The axis alligned box overlap calculation is rather straightforward</p>
<div class="language-glsl highlighter-rouge"><div class="highlight"><pre class="highlight"><code><span class="kt">vec3</span> <span class="nf">overlap</span><span class="p">(</span><span class="kt">vec2</span> <span class="n">x</span><span class="p">,</span> <span class="kt">vec2</span> <span class="n">p</span><span class="p">,</span> <span class="kt">float</span> <span class="n">diffusion_radius</span><span class="p">)</span>
<span class="p">{</span>
    <span class="kt">vec4</span> <span class="n">aabb0</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">p</span> <span class="o">-</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="p">,</span> <span class="n">p</span> <span class="o">+</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="p">);</span> <span class="c1">//cell box</span>
    <span class="kt">vec4</span> <span class="n">aabb1</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">x</span> <span class="o">-</span> <span class="n">diffusion_radius</span><span class="p">,</span> <span class="n">x</span> <span class="o">+</span> <span class="n">diffusion_radius</span><span class="p">);</span> <span class="c1">//particle box</span>
    <span class="kt">vec4</span> <span class="n">aabbX</span> <span class="o">=</span> <span class="kt">vec4</span><span class="p">(</span><span class="n">max</span><span class="p">(</span><span class="n">aabb0</span><span class="p">.</span><span class="n">xy</span><span class="p">,</span> <span class="n">aabb1</span><span class="p">.</span><span class="n">xy</span><span class="p">),</span> <span class="n">min</span><span class="p">(</span><span class="n">aabb0</span><span class="p">.</span><span class="n">zw</span><span class="p">,</span> <span class="n">aabb1</span><span class="p">.</span><span class="n">zw</span><span class="p">));</span> <span class="c1">//overlap box</span>
    <span class="kt">vec2</span> <span class="n">center</span> <span class="o">=</span> <span class="mi">0</span><span class="p">.</span><span class="mi">5</span><span class="o">*</span><span class="p">(</span><span class="n">aabbX</span><span class="p">.</span><span class="n">xy</span> <span class="o">+</span> <span class="n">aabbX</span><span class="p">.</span><span class="n">zw</span><span class="p">);</span> <span class="c1">//center of mass </span>
    <span class="kt">vec2</span> <span class="n">size</span> <span class="o">=</span> <span class="n">max</span><span class="p">(</span><span class="n">aabbX</span><span class="p">.</span><span class="n">zw</span> <span class="o">-</span> <span class="n">aabbX</span><span class="p">.</span><span class="n">xy</span><span class="p">,</span> <span class="mi">0</span><span class="p">.);</span> <span class="c1">//only positive</span>
    <span class="kt">float</span> <span class="n">m</span> <span class="o">=</span> <span class="n">size</span><span class="p">.</span><span class="n">x</span><span class="o">*</span><span class="n">size</span><span class="p">.</span><span class="n">y</span><span class="o">/</span><span class="p">(</span><span class="mi">4</span><span class="p">.</span><span class="mi">0</span><span class="o">*</span><span class="n">diffusion_radius</span><span class="o">*</span><span class="n">diffusion_radius</span><span class="p">);</span> <span class="c1">//relative area</span>
    <span class="c1">//if any of the dimensions are 0 then the mass ratio is 0</span>
    <span class="k">return</span> <span class="kt">vec3</span><span class="p">(</span><span class="n">center</span><span class="p">,</span> <span class="n">m</span><span class="p">);</span>
<span class="p">}</span>
</code></pre></div></div>
<p>As you can see we don’t need to loop over virtual particles anymore since the solution is analytical, so the performance of this approach is the same as in the original cellular automaton particle tracker with the added benefit of giving much smoother results.</p>

<p>And here is a real time visualization of the reintegration process with diffusion radius 0.35</p>
<center><iframe style="width:640px;height:360px;" frameborder="0" src="https://www.shadertoy.com/embed/WlSfWD?gui=true&amp;t=10&amp;paused=false" allowfullscreen=""></iframe></center>

<p>In fact this algorithm has some quite interesting properties, depending on the radius of the distribution the behaviour can change from particle-like to field-like as shown in the simulation below(you need to unpause it). The distribution diameter oscillates between 0.75 and 1.25:</p>
<center><iframe style="width:640px;height:360px;" frameborder="0" src="https://www.shadertoy.com/embed/tl2fWD?gui=true&amp;t=10&amp;paused=true" allowfullscreen=""></iframe></center>
<p>We can see that for a diameter less than 1 pixel the behaviour tends to be particle-like and for a larger one it behaves more like usual advection with numerical diffusion.</p>

<p>Another interesting fact is that if we fix the particle positions to the cell centers and set the distribution radius to 0.5 so that the distribution is exactly as big as the cell we’ll get exactly forward Euler advection! Since we are technically just integrating the cells forward and depositing their contents.</p>

<h3 id="using-the-sph-formulation-instead-of-finite-differences-to-compute-forces">Using the SPH formulation instead of finite differences to compute forces</h3>

<p>Now what are we going to do with this algorithm? We can use the grid and compute finite difference gradients to compute forces. But we are actually losing the sub-cell distribution information - the cell centers of mass.</p>

<p>Since it can model particle systems we can try to adapt particle algorithms here, for example molecular dynamics(it would need exact particle count conservation), or maybe how about using smoothed particle hydrodynamics(SPH)? Actually this algorithm gives pretty much the perfect conditions for SPH, the particles are already uniformly distributed, around 1 particle per cell, and we can quite easily find the particle neighbors, since the grid itself is an acceleration structure! And that is pretty much exactly what was done in <a href="https://www.shadertoy.com/view/WtfyDj">Paint streams</a> or <a href="https://www.shadertoy.com/view/ttBcWm">Everflow</a>. With a large enough distribution radius (0.55-0.6) the natural diffusion is smoothing the particles so that we get away with a relatively small smoothing kernel, about 1.5 pixels wide, and we only need to compute the forces from the closest neighbors which makes it even more efficient.</p>

<p>To implement SPH we just need to integrate the new cell velocity using the reintegrated particle distributions.</p>

<p>\[ \vec{V}_ {i}^{t+1} =  \vec{V}_ {i}^{t} +\Delta t  \frac{\vec{F}_ {i}}{M_ {i}^{t}}  —  \textrm{updated velocity} \] 
Where \(\vec{F}_ {j}\) is the computed SPH force [1] computed for the particle in cell i.</p>

<p>\[ \vec{F}_ {i} =M_ {i}^{t} \sum_{j}^\textrm{neighbors} M_ {j}^{t} \left( \frac{P_{i}}{\rho_{i}^2} + \frac{P_{j}}{\rho_{j}^2} \right)  \nabla_{i} W(\vec{X}_ {j}^{t} - \vec{X}_ {i}^{t}) \]</p>

<p>Where \(P_{i}\) is the pressure in cell i, and \(W(\vec{X})\) is the smoothing kernel. For the density \( \rho_{i} \) we can just use the mass of the cell divided by its volume, assuming the volume is 1 we can just place the mass. It’s ok to do so if the the distribution radius is big enough to smooth out the mass.</p>

<p>Pressure for each cell can just be computed using an equation of state. For a gas its simply just proportional to the cell density times the temperature, but let’s consider only constant temperatures:</p>

<p>\[ P_{i} = k \rho_{i} \]
Where k is some proportionality constant.</p>

<p>For a fluid we can use the <a href="https://en.wikipedia.org/wiki/Cole_equation_of_state">Cole equation of state</a>
\[ P_{i} = k \left( \left(\frac{\rho_{i}}{\rho_{0}} \right)^ \gamma - 1 \right) \]</p>

<p>Where \(\gamma\) is the adiabatic index (\(\gamma = 7.0\) for water), \(\rho_{0}\) is the reference fluid density.</p>

<p>In most of my simulations I just used the following pressure, which worked quite well in this setting.
\[ P_{i} = k \rho_{i} (\rho_{i} - \rho_{0}) \]</p>

<h3 id="storing-more-properties-inside-a-cell">Storing more properties inside a cell</h3>
<p>There is nothing holding us from storing a more advanced description of the insides of the cell, we can store not just a particle - but an entire distribution, and we can also update its size depending on the variations of the centers of mass of other distributions that fell into this cell. We can also store the angular momentum of such a distribution to preserve vorticity in fluids for example, sadly averaging angular momentums of distribution parts is a bit complicated, and <a href="https://www.shadertoy.com/view/WtXcW2">my experiments</a> are not entirely stable and require angular momentum clamping suggesting that the way I added them was not exact (the colors show the curl value in the fluid). 
If I happen to successfully implement such a summation I will write a follow-up article, since perfect angular momentum tracking is really important for nice vortices.</p>

<h3 id="conclusions">Conclusions</h3>
<p>This is a really cool algorithm, and I wanted to share it with other people the moment I had the first successful results, <a href="https://www.shadertoy.com/user/wyatt">Wyatt</a> has already used it for some really cool simulations, including <a href="https://www.shadertoy.com/view/3lffzM">multi-substance interactions</a>. 
Thanks to the mass conserving quality of the advection it can be used to model <a href="https://www.shadertoy.com/view/Wl2yWm">self-gravitating gas</a> too (the angular momentum there is total whack tho, but looks cool)</p>

<p>Other interesting use cases are:</p>
<ul>
  <li><a href="https://www.shadertoy.com/view/ttXcDB">Modelling boilling</a> - very approximately, the equation of state is not that good there.</li>
  <li><a href="https://www.shadertoy.com/view/WtfyW7">Fluid advection</a> with a pressure computed using the Poisson equation.</li>
  <li><a href="https://www.shadertoy.com/view/3llcRj">Rocket Mach diamonds</a> - modelling supersonic gas using a gas equation of state.</li>
  <li><a href="https://www.shadertoy.com/view/WtBcDG">Slime molds</a> - I should probably make a blog post on this too.</li>
  <li><a href="https://www.shadertoy.com/view/Wt2BR1">Life-like cells</a> - ????. Another post needed, yeah.</li>
</ul>

<p>Also <a href="https://github.com/MichaelMoroz/michaelmoroz.github.io/blob/master/files/ReintegrationTracking/ReintegrationTracking.pde">here</a> is a processing implementation of reintegration tracking (no SPH forces)</p>

<h3 id="references">References</h3>
<p>[1] <a href="https://arxiv.org/pdf/1007.1245.pdf">Smoothed Particle Hydrodynamics</a></p>]]></content><author><name></name></author><summary type="html"><![CDATA[In this blog post I’ll explain this advection algorithm and how to use it to make advanced fluid simulations like the ones I made, including Paint streams and Everflow. Before starting, I should give a big thanks to my friend Wyatt for giving useful suggestions on building this algorithm.]]></summary><media:thumbnail xmlns:media="http://search.yahoo.com/mrss/" url="https://michaelmoroz.github.io/ReintegrationTracking.png" /><media:content medium="image" url="https://michaelmoroz.github.io/ReintegrationTracking.png" xmlns:media="http://search.yahoo.com/mrss/" /></entry></feed>