<script data-pm-proxy="intercept"></script><?xml version="1.0" encoding="UTF-8"?><rss xmlns:dc="http://purl.org/dc/elements/1.1/" xmlns:content="http://purl.org/rss/1.0/modules/content/" xmlns:atom="http://www.w3.org/2005/Atom" version="2.0" xmlns:itunes="http://www.itunes.com/dtds/podcast-1.0.dtd" xmlns:googleplay="http://www.google.com/schemas/play-podcasts/1.0"><channel><title><![CDATA[Just a Byte - AI Compilers, Silicon, and Systems]]></title><description><![CDATA[Technical articles on AI compilers, custom silicon, and AI systems. Exploring JAX and PyTorch targeting custom hardware, with a focus on high-performance LLMs.]]></description><link>https://patricktoulme.substack.com</link><image><url>https://substackcdn.com/image/fetch/$s_!orRT!,w_256,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fcbf6df32-d7d6-4cce-87e9-9b1b510e7929_144x144.png</url><title>Just a Byte - AI Compilers, Silicon, and Systems</title><link>https://patricktoulme.substack.com</link></image><generator>Substack</generator><lastBuildDate>Tue, 01 Sep 2026 15:48:46 GMT</lastBuildDate><atom:link href="/__u/patricktoulme.substack.com/feed" rel="self" type="application/rss+xml"/><copyright><![CDATA[Patrick C. Toulme]]></copyright><language><![CDATA[en]]></language><webMaster><![CDATA[patricktoulme@substack.com]]></webMaster><itunes:owner><itunes:email><![CDATA[patricktoulme@substack.com]]></itunes:email><itunes:name><![CDATA[Patrick C. Toulme]]></itunes:name></itunes:owner><itunes:author><![CDATA[Patrick C. Toulme]]></itunes:author><googleplay:owner><![CDATA[patricktoulme@substack.com]]></googleplay:owner><googleplay:email><![CDATA[patricktoulme@substack.com]]></googleplay:email><googleplay:author><![CDATA[Patrick C. Toulme]]></googleplay:author><itunes:block><![CDATA[Yes]]></itunes:block><item><title><![CDATA[We’ve Built 1% of the AI Compute the World Needs]]></title><description><![CDATA[The world needs a hundred times more compute than exists. The companies that build it are already worth trillions, and they will have to get far bigger to deliver this future.]]></description><link>https://patricktoulme.substack.com/p/weve-built-1-of-the-ai-compute-the</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/weve-built-1-of-the-ai-compute-the</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Tue, 04 Aug 2026 16:47:48 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!hvIA!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p><span>Connect on LinkedIn: </span><a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p><p><span>Follow on X: </span><a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><p>Code reproduction: <a href="https://github.com/patrick-toulme/justabyte/blob/main/ai_compute_1_percent_post/compute_model.py">https://github.com/patrick-toulme/justabyte/blob/main/ai_compute_1_percent_post/compute_model.py</a></p><p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><p>What happens when a software engineer stops using one AI agent and starts commanding a swarm of fifty? When a quarter of adult Americans each have an AI that staffs a team of agents the way a manager staffs a project? In this scenario you don&#8217;t get one agent per person. You get a planet trying to run billions of always-on agents, with nowhere near the silicon to do it.</p><p>This article will argue that the US and world economies are changing rapidly, that AI silicon is the bedrock of that change, and that we have built very little of the compute this future requires.</p><p>Today the world has roughly <strong>19 million H100 equivalent AI chips</strong>. That is every Hopper, every Blackwell, every Google TPU, AWS Trainium, Meta MTIA, AMD, Cerebras, Groq etc. ever shipped and still racked, added up and normalized to one unit. It is the core of the AI economy: about 38 zettaFLOP/s of FP8 compute, most of it built in the last eighteen months.</p><p>Here is the thing people get wrong about the demand side. They picture one agent per human and ask whether that is realistic. No human is going to <em>manage</em> a hundred agents by hand. <strong>You talk to one AI, and it runs the swarm of agents for you.</strong> It plans, spawns subagents, runs them in parallel, and hands back a result. The human stays in a single conversation; the machine fans out into hundreds. That is already how the best agentic coding tools work, and open-source models such as GLM-5.2 and Kimi K3 are now capable of orchestrating these massive swarms of agents.</p><p>So take a grounded case, and a deliberately partial one. Not all 8 billion humans. Not even all of the world&#8217;s roughly 1 billion knowledge workers, because not everyone will adopt this or afford it. Just <strong>half of them, about 500 million people</strong>, each with an AI that runs a modest <strong>50-agent swarm</strong>. How many chips does that need?</p><p><strong>About 830 million chips. Roughly 44 times every AI chip that exists.</strong></p><p>This article shows the math behind that number, and the scale at which the world has to build compute to make this future real.</p><p>That 830 million is not the aggressive case. It is a conservative reading, on a cheap open model, for a deliberately narrow slice of the workforce. I picked one legible group on purpose, as a frame of reference: analysts, engineers, and lawyers using a team of agents at work, and only half of them. One ordinary slice already needs 44&#215; every chip on Earth, and the slices I left out stack fast. The other half of knowledge workers: another 44&#215;. Three billion consumers with even a five agent helper for daily life: another ~26&#215;. The floor climbs past 100&#215;, of which we&#8217;ve built about <strong>1%</strong>. We are not oversupplied on AI compute. We are not even close.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!hvIA!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 424w, /__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 848w, /__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 1272w, /__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!hvIA!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png" width="1080" height="1080" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/f4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:1080,&quot;width&quot;:1080,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;Give half the world's knowledge workers an AI that runs a team of agents: 44&#215; every AI chip on Earth&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="Give half the world's knowledge workers an AI that runs a team of agents: 44&#215; every AI chip on Earth" title="Give half the world's knowledge workers an AI that runs a team of agents: 44&#215; every AI chip on Earth" srcset="/__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 424w, /__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 848w, /__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 1272w, /__u/substackcdn.com/image/fetch/$s_!hvIA!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Ff4a76429-5e54-4089-91d4-bf97c9912cbd_1080x1080.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>I want to make this concrete and reproducible, so I built a small parametric model and sourced every input to <a href="https://www.semianalysis.com/">SemiAnalysis</a> numbers where they exist. The experiment is ~150 lines of Python, linked at the end. This code powers the graphs in this article, and you can change any assumption you disagree with and watch the conclusion move.</p><p>The conclusion does not move much.</p><div><hr></div><h2>The unit, and the question</h2><p>Everything here is in <strong>H100 equivalents (H100e)</strong>, FP8 throughput convention: a Blackwell counts as ~2.3 H100e because it does ~2.3&#215; the FP8 work.</p><p>The thought experiment has one assumption worth stating upfront: <strong>what does &#8220;running AI&#8221; mean per person?</strong> I model <em>always-on agents</em>: not chats you prompt a few times a day, but workers that generate tokens continuously. And every frontier model people actually use now is a <em>reasoning</em> model, one that generates a long hidden chain of thought before it answers, <a href="https://epoch.ai/data-insights/output-length">roughly 8&#215; the tokens of a non-reasoning model</a>. So I count <strong>generated</strong> tokens, reasoning included, at a sustained <strong>40 tokens/sec</strong> per agent, averaged over the full day, with a floor of 15 and an aggressive case of 100.</p><p>The unit of demand is the <strong>swarm</strong>, not the person. Here&#8217;s how that actually works. A user opens one session and states a goal. An <em>orchestrator</em> model reads it, breaks it into pieces, and spawns <strong>subagents</strong> (a planner, a handful of coders, reviewers, test writers, researchers) that run in parallel, each in its own context, and report back. Chain those steps and you have a <em>workflow</em>: the orchestrator fans work out, gathers the results, and decides what to do next. No human is monitoring the 50 agents. The human is managing <strong>one</strong> agent, and that agent is managing the rest. So the thing that scales isn&#8217;t how many people adopt AI. It&#8217;s <strong>how many agents each person&#8217;s AI is running</strong>, and that number is climbing from one toward the hundreds.</p><p>I sweep that swarm size across grounded populations: all software engineers (~28 million, <a href="https://www.slashdata.co/post/global-developer-population-trends-2025-how-many-developers-are-there">per JetBrains/SlashData</a>), a fraction of adult Americans, and half the world&#8217;s ~1 billion knowledge workers. I sweep one more axis underneath everything: <strong>model size</strong>, from 10B to 100T parameters because the entire argument bends on it.</p><div><hr></div><h2>What the world has now</h2><p>The installed compute base is <strong>~15&#8211;20M H100e</strong> (<a href="https://the-decoder.com/global-ai-compute-hits-15-million-h100-equivalents-epoch-ai-finds/">Epoch put global capacity at ~15M in early 2026, and it&#8217;s doubling roughly every seven months</a>). NVIDIA&#8217;s own fleet, its millions of Hopper GPUs plus the fast-ramping Blackwell, comes to roughly <strong>11M H100e</strong>, but it does <em>not</em> own the whole fleet. The custom accelerators are not a rounding error: <a href="https://epoch.ai/data-insights/google-custom-tpus-ai-compute">Google is in fact the single largest owner of AI compute</a>, holding about a quarter of the world&#8217;s total, mostly its own TPUs (Ironwood does 4.6 PFLOPS FP8 with <a href="https://introl.com/blog/google-tpu-architecture-complete-guide-7-generations">~4.3M TPUs shipping in 2026</a>). AWS has <a href="https://techcrunch.com/2026/03/22/an-exclusive-tour-of-amazons-trainium-lab-the-chip-thats-won-over-anthropic-openai-even-apple/">1.4 million Trainium chips deployed across three generations</a>, over a million of them running Anthropic&#8217;s Claude; AMD&#8217;s Instinct is at a <a href="https://www.sec.gov/Archives/edgar/data/0000002488/000000248826000045/pressreleasedatedfebruary2.htm">$10B-a-quarter run rate</a>, and Microsoft&#8217;s Maia, Meta&#8217;s MTIA, and Huawei&#8217;s Ascend add more. Let&#8217;s call the world total <strong>~19M H100e</strong>.</p><p>AI compute production is rapidly growing. NVIDIA went from ~3.5M datacenter GPUs in 2024 to ~6.5&#8211;7M in 2025, up ~55% in units and ~68% in revenue, and <a href="https://arxiv.org/html/2504.16026v1">the fleet&#8217;s </a><em><a href="https://arxiv.org/html/2504.16026v1">compute</a></em><a href="https://arxiv.org/html/2504.16026v1"> is roughly doubling every nine months</a>. Across every vendor, that&#8217;s <strong>~16M H100e added per year</strong> and rising.</p><p>Here is the part that matters: <strong>you cannot simply print these.</strong> The binding constraint on AI chip production is not logic dies; those are abundant, and AI silicon vendors just outbid everyone else for leading node wafers. The constraint is <strong><a href="https://epoch.ai/data-insights/ai-chip-supply-chain-constraints">advanced packaging (CoWoS) and HBM</a></strong>. In 2025 the four largest AI chip designers consumed ~90% of global CoWoS capacity and HBM supply while using only ~12% of advanced logic production. CoWoS has roughly doubled every year (16k &#8594; 40k &#8594; 70k &#8594; 130k wafers/month from 2023 to 2026), and HBM has been sold out through 2026, with SK Hynix at 62% share supplying ~90% of NVIDIA&#8217;s stacks.</p><p>The supply of AI chips is a physical pipeline with significant resource constraints.</p><div><hr></div><h2>What one model costs to serve</h2><p>To serve a model you need two things from your chips: enough <strong>memory</strong> to hold the weights, and enough <strong>memory bandwidth</strong> to stream them per token. For frontier models, bandwidth is the wall.</p><p>Three facts set the per-chip economics, and all three are well grounded in real 2025 deployments:</p><p><strong>1. Frontier models are sparse, and getting sparser.</strong> No one serves a dense 10T model. Frontier models are Mixture of Experts, and the active fraction <em>shrinks as total size grows</em>: <a href="https://arxiv.org/pdf/2412.19437">DeepSeek-V3</a> activates 37B of 671B (5.5%), the open <a href="https://the-decoder.com/zhipu-ais-glm-5-2-closes-in-on-closed-source-leaders-in-coding-marathons/">GLM-5.2</a> activates ~40B of 744B (5.4%), <a href="https://huggingface.co/moonshotai/Kimi-K2-Instruct">Kimi K2</a> activates 32B of 1T (3.2%), and the largest open model yet, <a href="https://huggingface.co/moonshotai/Kimi-K2-Instruct">Kimi K3</a>, activates ~104B of 2.8T (3.7%). The <a href="https://arxiv.org/abs/2501.12370">optimal sparsity curve bends down with scale</a>, and deployed models are more aggressive than optimal because serving economics reward it. Active params grow <em>sub-linearly</em>: I model 10B&#8594;3B active, 100B&#8594;12B, 1T&#8594;40B, 10T&#8594;200B, 100T&#8594;1T.</p><p><strong>2. Decode is memory-bandwidth bound.</strong> To generate one token you stream every active parameter out of HBM once. The roofline is simply <code>tokens/sec &#8776; HBM_bandwidth / (2 &#215; active_param_bytes)</code>. On an H100 (3.35 TB/s) with 37B active params in FP8, that&#8217;s ~45 tokens/sec, single stream. This is not a FLOPs problem. It&#8217;s a &#8220;how fast can you read the weights&#8221; problem, and it&#8217;s why <a href="https://arxiv.org/pdf/2506.04645">bandwidth, not compute, is the decode bottleneck</a> for essentially every model chip pair.</p><p><strong>3. Batching helps, but reasoning caps it.</strong> That 45 tok/s is one user at batch one. Batch many users together and you amortize each weight read across all of them, but only down to the latency your users will tolerate, and reasoning users tolerate very little, as a 25,000-token thinking trace served at low interactivity would take half an hour. So you serve at a usable speed, and <a href="https://newsletter.semianalysis.com/p/inferencemax-open-source-inference">SemiAnalysis&#8217;s InferenceMAX</a> measures exactly what that costs. For DeepSeek-R1 (a reasoning MoE, 37B active) at ~42 tokens/sec/user, <a href="https://inferencex.semianalysis.com/compare/deepseek-r1-b200-vs-h100">an H100 manages just 266 tok/s/GPU and a B200 4,792</a>, about 2,080 per H100-equivalent. Blend a Hopper and Blackwell fleet and you land near <strong>1,300 tok/s/H100e</strong> for a 37B-active model; I scale it inversely with active params.</p><p>(And 1,300 is still generous: it assumes a Blackwell-heavy fleet. An H100-only fleet serves ~5&#215; worse, at 266 tok/s/H100e, and pushing interactivity higher, to 71 or 100 tok/s/user, <a href="https://inferencex.semianalysis.com/compare/deepseek-r1-b200-vs-h100">collapses throughput further still</a>. Every one of those moves makes the numbers below bigger, not smaller.)</p><p>One caveat the other way: I run all of this in FP8. The frontier is sliding to FP4 (<a href="https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/unlocking-high-performance-inference-for-deepseek-with-nvfp4-on-nvidia-blackwell/4497936">DeepSeek quantizes FP8&#8594;NVFP4 with under 1% accuracy loss</a>; Together and Fireworks already serve it, and <a href="https://huggingface.co/moonshotai/Kimi-K3">Kimi K3 now ships its weights natively in MXFP4</a>, quantization-aware trained from the SFT stage on), which roughly <strong>halves</strong> every chip count here.</p><p>Put it together and you get the per-chip serving rate for each model size. Now apply it to a realistic slice of the workforce.</p><div><hr></div><h2>A fraction of people, each running a swarm</h2><p>All humans do not need to use AI agents for the world to be extremely compute constrained. We only need a slice of the workforce running swarms on a cheap, open model, and that is already happening.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!aUsi!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!aUsi!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png" width="1456" height="819" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:819,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;A fraction of professionals, each running a swarm&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="A fraction of professionals, each running a swarm" title="A fraction of professionals, each running a swarm" srcset="/__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!aUsi!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06c03b6f-4299-475a-93ac-f5cfd0dc3c30_1600x900.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>Every line crosses the world&#8217;s entire chip fleet almost immediately. <strong>All ~28 million software engineers</strong>, each driving a 50-agent coding swarm, need ~47 million chips, <strong>2.5&#215; every AI chip on Earth</strong>, and they cross the fleet at about <em>20</em> agents apiece. <strong>A quarter of adult Americans</strong> (~65 million people) on the same 50-agent swarm is <strong>5.7&#215; the fleet</strong>, crossing at roughly <em>9</em> agents each. And <strong>half the world&#8217;s knowledge workers</strong>, ~500 million people whose work happens on a screen, reach the entire global fleet at barely <em>one</em> agent apiece; give each of them a 50-agent swarm and you are at <strong>~830 million chips, 44&#215; everything ever built.</strong></p><p>None of this assumes a giant model, universal adoption, or anyone managing more than one conversation. It assumes a cheap open model, a fraction of the professional world, and the swarm sizes agentic coding tools already ship.</p><p>And notice what this slice leaves out. Half of one profession is a <em>frame of reference</em>, not a forecast. I chose it because it&#8217;s easy to picture and easy to argue with. It counts no one using AI for their health, their kids&#8217; education, their finances, their research, or their small business: the things that will reach billions of people who never call themselves knowledge workers. Each of those is its own swarm story. I left them all out, and the number is still 44&#215; every chip on Earth. Add any one of them back and it only moves one way.</p><div><hr></div><h2>The model is a free download. The silicon isn&#8217;t.</h2><p>The reason a fifty-agent swarm is realistic and not a thought experiment is that the model to run it already shipped, and it&#8217;s open. <a href="https://venturebeat.com/technology/z-ais-open-weights-glm-5-2-beats-gpt-5-5-on-multiple-long-horizon-coding-benchmarks-for-1-6th-the-cost">GLM-5.2</a>, from Zhipu / Z.ai, is a <strong>744 billion parameter model with ~40 billion active</strong>, released in June 2026 under an unrestricted MIT license with a 1-million-token context. It posts <a href="https://the-decoder.com/zhipu-ais-glm-5-2-closes-in-on-closed-source-leaders-in-coding-marathons/">frontier-class agentic coding scores</a>: 62 on SWE-bench Pro and an <a href="https://www.mindstudio.ai/blog/glm-5-2-vs-gpt-5-5-vs-claude-opus-agentic-workflows">MCP-Atlas tool-use score of 77</a>, within a few points of Claude Opus 4.8 on the long-horizon coding marathons. It speaks the Anthropic and OpenAI APIs, so it drops into Claude Code, Cline, or Cursor with a base URL swap, and the hosted API runs about <a href="https://aitoolanalysis.com/glm-coding-plan-review/">$2 per million input tokens and $6 per million output</a>, roughly a sixth of the closed leaders, with $18-a-month coding plans. <a href="https://x.com/PatrickToulme/status/2068134212587184442">I&#8217;ve run it myself</a>: open model, open harness, local serving, frontier coding at essentially zero marginal cost.</p><p>And GLM-5.2 is no longer alone. Weeks after it shipped, <a href="https://huggingface.co/moonshotai/Kimi-K3">Moonshot AI open-sourced Kimi K3</a>: <strong>2.8 trillion total parameters, ~104 billion active</strong>, native vision, the same 1-million-token context, weights free to download under a permissive commercial license. It posts<a href="https://fenxi.fr/en/blog/kimi-k3-moonshot-ai-architecture-benchmarks-explained/"> frontier agentic scores of its own</a>, 42.0 on SWE Marathon against Claude Opus 4.8&#8217;s 40.0, <a href="https://seawork.ai/en/blogs/kimi-k3-for-coding/">88.3 on Terminal-Bench 2.1</a>, with a hosted API at $3 per million input tokens and $15 per million output. Two labs, weeks apart, each handing the world a frontier-class agent orchestrator for nothing but the cost of the silicon to run it. The open frontier isn&#8217;t an event. It&#8217;s a cadence.</p><p>Here&#8217;s the part that matters for compute, and it&#8217;s the opposite of what &#8220;open and free&#8221; sounds like. <strong>Open weights don&#8217;t remove the silicon. They remove the brake.</strong> To actually serve GLM-5.2 you still need the chips: its weights are <a href="https://ofox.ai/blog/glm-5-2-self-host-vllm-hardware-cost-2026/">~744 GB in FP8, a box of roughly eight datacenter GPUs</a> before it generates a single token, and <a href="https://huggingface.co/blog/ResterChed/kimi-k3-model-overview-mxfp4-quantization-open-wei">Kimi K3&#8217;s weigh ~1.4 TB even packed in 4-bit</a>. What open weights remove is the <em>meter</em>. There&#8217;s no per-token bill capping how many agents you run, no vendor rate limit, no one deciding you&#8217;ve had enough tokens for the day. The number of agents you can run stops being a function of your API budget and becomes a function of how many accelerators you can get.</p><p>That is the demand engine for this whole piece. A closed, metered model self-limits: you run the agents you&#8217;re willing to expense. An open model you host yourself has one limit left, and it&#8217;s silicon. Make the best coder a free download and the rational move isn&#8217;t to run one agent and bank the savings; it&#8217;s to run fifty, because the marginal token is nearly free and <em>not</em> running the swarm is leaving value on the table. Open weights are Jevons&#8217; paradox with the last brake cut.</p><p>It also globalizes the demand. A model on Hugging Face runs in any datacenter, any country, any startup&#8217;s basement. When the best model is an open download, no single vendor can even <em>meter</em> the world&#8217;s demand for it, let alone supply it. The only thing standing between an open frontier model and an unbounded swarm is the chip.</p><p>So run every scenario in this piece on GLM-5.2 specifically, with its real 40 billion active params setting the per-chip rate. That&#8217;s a deliberately conservative anchor: Kimi K3 activates ~104 billion params per token to GLM&#8217;s 40 billion, so a world whose swarms run on the newest open flagship instead needs roughly <strong>2.6&#215; every chip count below</strong>. Here is the bill:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!pQLd!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 424w, /__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 848w, /__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 1272w, /__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!pQLd!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png" width="1456" height="907" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:907,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:194407,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/205708085?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 424w, /__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 848w, /__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 1272w, /__u/substackcdn.com/image/fetch/$s_!pQLd!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F7459bb02-b135-417d-99df-4b6c85c557ae_2120x1320.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p> single profession on a fifty-agent swarm already needs more chips than every vendor has ever built. The model is free. The hundred million plus chips to run it are not.</p><p>To put the swarm in everyday units: half the world&#8217;s knowledge workers, 50 agents each, generate on the order of <strong>80,000 trillion tokens a day</strong>. The entire AI industry today, every ChatGPT session and every Gemini call, runs <a href="https://the-decoder.com/google-boasts-1-3-quadrillion-tokens-each-month-but-the-figure-is-mostly-window-dressing/">about 50 trillion tokens/day</a>. One slice of one profession is already <strong>on the order of 1,000&#215; the world&#8217;s current token output</strong>.</p><div><hr></div><h2>What happens when some people want faster inference?</h2><p>Everything above assumes <em>standard</em> reasoning serving: ~40 tokens a second per agent, the InferenceMAX operating point. That&#8217;s fine for a background worker. It is too slow for the thing people actually pay for.</p><p>Watch a developer drive an agentic coding tool and you&#8217;ll see the real demand. They want the answer <strong>now</strong>. Not 40 tokens a second but 100, 200, as fast as the model can think, because a human is sitting there waiting and their time is worth more than the GPU&#8217;s. More importantly, <strong>they will pay for it</strong>. A coding agent that&#8217;s 3&#215; faster is worth a lot more than 3&#215; the price, because for that workload latency <em>is</em> the product.</p><p>But speed and batching are opposites. To make one request faster you pull it out of the batch and hand it more of the chip. On InferenceMAX, pushing DeepSeek-R1 from ~42 to ~100 tokens/sec/user on an H100 <a href="https://inferencex.semianalysis.com/compare/deepseek-r1-b200-vs-h100">cuts throughput from 266 to 23 tok/s/GPU</a>, more than 10&#215; the silicon per token. I&#8217;ll be conservative and price the premium tier at a flat 3&#215;.</p><p>Put a slice of the world on that fast tier and the whole curve lifts:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!409E!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!409E!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png" width="1456" height="819" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/dead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:819,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;The speed premium&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="The speed premium" title="The speed premium" srcset="/__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!409E!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdead7848-4277-4a1e-8871-5acc5c8a41e1_1600x900.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>If that knowledge worker swarm ran on the fast tier, the bill goes from ~830 million chips to <strong>~2.5 billion, 131&#215; the world&#8217;s entire fleet.</strong> Even a modest 20% premium segment adds ~330 million chips on top. And this is the <em>high-margin</em> demand: the customers who pay 2&#8211;3&#215; are precisely the ones who make a chip most worth selling. The latency-sensitive tier doesn&#8217;t just add load. It adds the load people are most willing to fund. Those who are running fast-tier inference for agentic coding will be running a lot of it and paying virtually whatever it costs.</p><div><hr></div><h2>What about AI training? Pretraining is a rounding error. RL isn&#8217;t.</h2><p>You might expect <em>pretraining</em> these large models to dominate the bill. It doesn&#8217;t, and the gap is very illuminating.</p><p>Training cost is <code>6 &#215; active_params &#215; tokens</code>. Run the numbers (FP8, 35% MFU, calibrated to <a href="https://epoch.ai/data-insights/models-over-1e25-flop">GPT-4&#8217;s 2.1e25 FLOP</a> and the <a href="https://newsletter.semianalysis.com/p/100000-h100-clusters-power-network">100k-H100 cluster that retrains GPT-4 in four days</a>): training and continuously refreshing a fleet of <strong>20 frontier model lines</strong> at 10T params costs ~130k H100e as a standing fleet. At 100T, ~880k.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!Bu4S!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!Bu4S!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png" width="1456" height="819" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:819,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;The gap: chips needed vs. chips that exist&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="The gap: chips needed vs. chips that exist" title="The gap: chips needed vs. chips that exist" srcset="/__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!Bu4S!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F346b5066-30cf-41ae-8012-f30c99ac267f_1600x900.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>Against ~830 million chips for serving, training is <strong>under 0.03% of the bill.</strong> The two demand curves above are visually identical because the training contribution vanishes beneath them. This is the inference economy in one fact: <strong>once a model is good, the cost is in running it, forever, for everyone.</strong> The market still prices these chips as if training clusters were the demand. The demand is the billions of agents that haven&#8217;t been spun up yet.</p><p>That covers <em>pretraining</em>. But there&#8217;s a much more compute-hungry paradigm of training: <strong>RL post-training</strong>.</p><p>The scale up of reinforcement learning is the steepest curve in AI right now. Reinforcement learning was <a href="https://epoch.ai/gradient-updates/why-gpt5-used-less-training-compute-than-gpt45-but-gpt6-probably-wont">~1% of pretraining compute before 2024</a>; then <a href="https://epoch.ai/gradient-updates/the-promise-of-reasoning-models">under 10% for OpenAI&#8217;s o1</a>, <a href="https://epoch.ai/gradient-updates/what-went-into-training-deepseek-r1">~20% for DeepSeek-R1</a>, and for <a href="https://www.lesswrong.com/posts/xpj6KhDM9bJybdnEe/how-well-does-rl-scale">xAI&#8217;s Grok 4, roughly </a><em><a href="https://www.lesswrong.com/posts/xpj6KhDM9bJybdnEe/how-well-does-rl-scale">parity</a></em><a href="https://www.lesswrong.com/posts/xpj6KhDM9bJybdnEe/how-well-does-rl-scale">, with RL run at pretraining scale</a> on a 200,000-GPU cluster. Epoch estimates RL compute is growing <a href="https://epoch.ai/gradient-updates/how-far-can-reasoning-models-scale">~10&#215; every three to five months</a>, against ~4&#215; a year for pretraining. The two are converging now.</p><p>And here&#8217;s why it belongs in <em>this</em> article and not a training footnote: <strong>RL is inference-shaped.</strong> Its cost isn&#8217;t the gradient step; it&#8217;s <em>rollout generation</em>. To run RL you sample many long answers per prompt and grade them (<a href="https://arxiv.org/abs/2501.12948">DeepSeek-R1 drew 16 completions per question, up to 32,768 tokens each</a>), and that generation, autoregressive memory-bandwidth-bound decode, is <a href="https://blog.guanghan.ai/post/260208_rl_infra/">70&#8211;90% of the wall-clock</a>. As SemiAnalysis puts it, inference <a href="https://newsletter.semianalysis.com/p/scaling-reinforcement-learning-environments-reward-hacking-agents-scaling-data">is &#8220;no longer just the end of the pipeline, but an integral part of training.&#8221;</a> RL doesn&#8217;t dent the &#8220;serving dominates&#8221; claim. It <em>is</em> serving, pointed inward towards improving the model.</p><p>Pretraining plus RL across 20 frontier lines is ~260k chips, a rounding error against the ~830 million for serving. But at <em>today&#8217;s</em> scale it&#8217;s the whole story. Most of the current fleet isn&#8217;t serving the public at all: <a href="https://epoch.ai/data-insights/openai-compute-spend">only ~30% of OpenAI&#8217;s 2024 compute went to inference</a>, the rest to training, RL, and experiments, and <a href="https://epoch.ai/gradient-updates/three-issues-undermining-compute-based-ai-policies">~40% of the world&#8217;s AI compute is training of one kind or another</a>. That&#8217;s exactly why only ~60% of the fleet is free to serve, and why we&#8217;re so far from one agent per human. The compute to <em>build</em> the models is still eating the compute that would <em>run</em> them, and the fastest-growing part of that, RL, is itself inference.</p><div><hr></div><h2>You can&#8217;t ignore the gap in compute</h2><p>The realistic scenarios need anywhere from a few times to a few hundred times the world&#8217;s AI silicon. How fast can we build it?</p><p>At ~16M H100e of net new production per year, closing this gap of ~830 million chips is <strong>~50 years of the entire planet&#8217;s AI chip output</strong>, devoted to this one slice of the workforce, ignoring replacement of dying hardware. Production is growing ~1.6&#215; a year in units, faster in H100e, so the real timeline is shorter. But the throat is CoWoS and HBM, and those double roughly <em>once a year</em>, not once a quarter. You cannot software your way out of a packaging constraint. You cannot fine-tune your way past an HBM shortage. The gap between &#8220;19 million chips&#8221; and &#8220;what a fraction of the workforce running swarms needs&#8221; is not a market sentiment. It&#8217;s a decade plus of the hardest manufacturing on Earth running flat out.</p><p>This is the whole argument, and I&#8217;ve kept it in units of silicon on purpose. I&#8217;m not going to model market caps or P/E multiples. The installed base is the <em>floor</em> of the AI economy&#8217;s capacity, every credible demand scenario sits one to two orders of magnitude above that floor, and the floor is bottlenecked on two of the most capital-intensive, slowest-to-scale processes in existence.</p><div><hr></div><h2>And then you have to power them</h2><p>All in (GPU, networking, cooling, grid losses), a modern AI chip pulls on the order of a kilowatt from the wall; a <a href="https://newsletter.semianalysis.com/p/gb200-hardware-architecture-and-component">GB200 NVL72 rack is 120 kW for 72 GPUs</a>, and since perf per watt keeps improving I&#8217;ll call it ~1 kW per H100 equivalent and be generous. Serving half the world&#8217;s knowledge workers a 50-agent swarm each, ~830 million chips, is then <strong>about 830 gigawatts</strong> of continuous draw.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!_CiG!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 424w, /__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 848w, /__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 1272w, /__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!_CiG!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png" width="1080" height="1080" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:1080,&quot;width&quot;:1080,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;830 GW: about 15&#215; every data center on Earth&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="830 GW: about 15&#215; every data center on Earth" title="830 GW: about 15&#215; every data center on Earth" srcset="/__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 424w, /__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 848w, /__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 1272w, /__u/substackcdn.com/image/fetch/$s_!_CiG!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F8b537d8e-8cdd-43c2-adea-9178308c2824_1080x1080.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>Every data center on Earth combined draws <a href="https://www.iea.org/reports/energy-and-ai/energy-demand-from-ai">~55 GW today</a> (&#8776;485 TWh/yr, ~1.5% of global electricity). This one realistic case is <strong>roughly 15 times the entire planet&#8217;s current data center load</strong>, nearly double all the electricity the United States consumes, for a single slice of one profession. Push the swarm to 500 agents and it&#8217;s ~8,300 GW, many times US generating capacity.</p><p>And power moves slower than silicon. New generation and <a href="https://www.cudocompute.com/blog/ai-data-center-capacity-planning">grid interconnects take 24 to 72 months</a>; fewer than 5% of existing data centers can even feed a 50 kW rack, never mind 120. CoWoS doubles in a year; a substation does not. The demand curve isn&#8217;t sitting above one slow physical wall; it&#8217;s sitting above two in series, packaging <em>and</em> power, and it clears neither for a decade. The chips are useless without the watts, and the watts are harder to get than the chips.</p><div><hr></div><h2>Hardware vendors aren&#8217;t competing. There is no pie.</h2><p>Here&#8217;s the pervasive story I want to kill. Every quarter someone writes that NVIDIA, Google&#8217;s TPU, AWS Trainium, AMD, and Cerebras are locked in a fight for the AI accelerator market: a fixed pie, each new entrant taking a slice from the rest, winners and losers.</p><p><strong>This is wrong. There is effectively infinite demand for AI silicon. One vendor cannot supply the entire market.</strong></p><p>Putting the numbers on it:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!5et9!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 424w, /__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 848w, /__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 1272w, /__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!5et9!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png" width="1456" height="874" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/b34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:874,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;Every vendor combined vs. one model per human&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="Every vendor combined vs. one model per human" title="Every vendor combined vs. one model per human" srcset="/__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 424w, /__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 848w, /__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 1272w, /__u/substackcdn.com/image/fetch/$s_!5et9!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb34e55b4-358f-481a-b527-9fc077a2031d_1600x960.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>Add up <em>every</em> AI chip from <em>every</em> vendor: ~11M H100 equivalents of NVIDIA, ~4.8M of Google TPU, ~1.2M of AWS Trainium, AMD&#8217;s ~1M, and the rest (Maia, MTIA, Huawei). You get about <strong>19 million chips</strong>. That total, the entire industry&#8217;s combined output, doesn&#8217;t even cover <strong>all the world&#8217;s software engineers running a 50-agent coding swarm</strong>, which needs ~47 million chips, <strong>2.5&#215; every chip every vendor has ever made</strong>. A quarter of adult Americans on the same swarm needs ~109 million, <strong>~6&#215;</strong>; half the world&#8217;s knowledge workers, ~830 million chips, <strong>44&#215;</strong>.</p><p>When demand sits anywhere from a couple times to many hundreds of times above the <em>combined</em> supply of every manufacturer, those manufacturers are not competing for customers. There aren&#8217;t enough chips in existence to fight over one. NVIDIA at ~58% share isn&#8217;t winning a zero-sum game against Trainium; they are both selling everything they can build into a market that would swallow dozens of times the volume. Cerebras and Groq aren&#8217;t stealing NVIDIA&#8217;s inference business. They&#8217;re additional capacity aimed squarely at the low-latency premium tier from two sections ago, which is <em>additive</em> demand NVIDIA can&#8217;t fully serve anyway.</p><p>The real competition isn&#8217;t vendor versus vendor for customers. It&#8217;s all vendors versus physics for <strong>CoWoS and HBM</strong>, the shared, sold-out chokepoint upstream of all of them. They&#8217;re not fighting over the pie. Which is also why portability is a myth: when every vendor is supply limited, each one optimizes its own stack to the metal and sells out regardless. <a href="/__u/patricktoulme.substack.com/p/portability-is-a-myth-why-the-best">Nobody is losing deals over a missing CUDA shim</a>; they&#8217;re losing them over packaging substrate.</p><div><hr></div><h2>How to break this</h2><p>I held the <a href="https://github.com/patrick-toulme/justabyte/blob/main/ai_compute_1_percent_post/compute_model.py">whole model</a> to one rule: every uncertain assumption breaks <em>toward</em> supply. So the real numbers are likelier to be worse, not better. The knobs, and which way they push:</p><ul><li><p><strong>Throughput.</strong> I used ~1,300 tok/s/H100e: InferenceMAX at 42 tok/s/user, blended across a Blackwell-heavy fleet. An older H100-heavy fleet serves ~5&#215; worse (266/H100e), and the faster interactivity coding agents demand cuts it further. Realistic serving pushes the numbers up, not down.</p></li><li><p><strong>Sparsity.</strong> If active fractions <em>don&#8217;t</em> keep shrinking (if a 10T model needs 500B active, not 200B), serving costs roughly 2.5&#215; more. My schedule is optimistic for supply here too: Kimi K3 already runs above it, at ~104B active where my curve implies ~80B for a 2.8T model.</p></li><li><p><strong>Population intensity.</strong> Every number above uses the mid rate of 40 tok/s/agent. The floor is 15 (chips &#247; ~2.7); the aggressive case is 100 (chips &#215; 2.5). Even the floor blows past the fleet at modest swarm sizes.</p></li><li><p><strong>Model size is the real variable.</strong> This is the one that helps supply. If a few billion parameter model turned out to be all anyone needed, today&#8217;s fleet could nearly cover one agent per human. The entire bull case for &#8220;we have enough chips&#8221; is &#8220;models get <em>smaller</em> and we stop wanting more of them.&#8221;</p></li><li><p><strong>What I don&#8217;t model.</strong> Fleet depreciation and failures (GPUs last ~3&#8211;5 years, so a slice of annual production just replaces dead silicon), node-level memory and networking limits, and downtime. Each of these shrinks the <em>effective</em> fleet, so leaving them out is again supply-favorable.</p></li></ul><p>What I&#8217;m <em>not</em> claiming: that a fifty-agent swarm lands in every knowledge worker&#8217;s hands next year. Adoption is slow, latency is hard, and most of the workforce isn&#8217;t running agents yet. This is a demand <em>curve</em> made of physics, population, and the economics of cheap agents, not a forecast of next year&#8217;s bookings.</p><div><hr></div><h2>&#8220;But models will get more efficient&#8221;</h2><p>This is the last objection standing, and the one people are most certain of. Models will get smaller. A 10 billion parameter model will eventually match today&#8217;s 10 trillion parameter frontier through distillation, better data, and architectures we haven&#8217;t found yet. Serving collapses 60&#215;, and the shortage solves itself.</p><p>It doesn&#8217;t. It inverts. This is <a href="https://en.wikipedia.org/wiki/Jevons_paradox">Jevons&#8217; paradox</a>, the most durable pattern in the history of resource use: make something cheaper to consume and total consumption goes <strong>up</strong>, not down. Cheaper steam engines burned more coal, not less. Cheaper compute didn&#8217;t empty the datacenters; it built hyperscale. Make an agent 60&#215; cheaper and nobody runs one agent and banks the savings. They run sixty. Then six hundred. Then a thousand, because at that price the only expensive thing left is restraint. We are watching it happen in real time. GLM-5.2 already put a frontier coder at a sixth the price as an open download, and the response wasn&#8217;t fewer agents. It was swarms. And the labs are now selling efficiency itself as the headline: Moonshot pitches Kimi K3 as <a href="https://x.com/Kimi_Moonshot/status/2081760186235289764">2.5&#215; the intelligence per unit of compute</a> of its predecessor. That gain won&#8217;t be banked either. It will be spent, on more agents.</p><p>So watch what a breakthrough actually does. Say a 100B model matches today&#8217;s 10T quality, and people respond the obvious way, not with one agent each but with a swarm:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!a6i1!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!a6i1!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png" width="1456" height="819" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/e3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:819,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;Jevons: a cheaper model just means everyone runs a swarm&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="Jevons: a cheaper model just means everyone runs a swarm" title="Jevons: a cheaper model just means everyone runs a swarm" srcset="/__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 424w, /__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 848w, /__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 1272w, /__u/substackcdn.com/image/fetch/$s_!a6i1!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fe3d3c3d8-254b-4fa5-87e8-37db212ccfd3_1600x900.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>A 100B model is 16&#215; cheaper per agent than the 10T. So nobody who pockets that saving runs <em>one</em> agent and banks it; they grow the swarm. Take the same half-billion knowledge workers and give each a hundred agents on the cheaper model and you need <strong>~500 million chips</strong>; push to a thousand, the number cheap agents actually invite, and it&#8217;s <strong>~5 billion chips, 263&#215; the entire world fleet</strong>, six times the realistic headline we started from. Even the most efficient case, a 10B model at frontier quality with a thousand agents each, is ~1.25 billion chips, 66&#215; the fleet.</p><p>The efficiency breakthrough didn&#8217;t shrink the problem. It detonated it. And it compounds with everything else: a cheaper model is also a cheaper <em>fast</em> model, so more of that swarm runs on the premium tier instead of the batched one. Every gain in efficiency gets spent on more agents, running faster, for more people, until it slams back into the one wall that doesn&#8217;t move: how fast the world can package silicon and stack memory.</p><p>This is the deep reason the chips are the floor and not the ceiling. You cannot make models efficient enough to escape it, because <strong>efficiency is the thing that creates the demand.</strong> The cheaper inference gets, the more of it the world buys.</p><div><hr></div><h2>Does this mean the stocks are cheap?</h2><p>Regardless of what you believe about a future in which humans are running agents 24/7, the <em>direction</em> the world is taking is undeniable. <strong>Models keep getting more useful</strong>. Agents multiply rather than consolidate; every coding tool already spawns dozens at once. Token volume is growing severalfold a year; <a href="https://techcrunch.com/2026/02/27/chatgpt-reaches-900m-weekly-active-users">ChatGPT went from 400M to 900M weekly users in twelve months</a>. You don&#8217;t need my exact scenario. You need <em>any</em> point on that curve a few times past today, and every one of them sits far above the installed base, behind two walls that take years to move.</p><p>That&#8217;s the valuation argument, and it&#8217;s why I kept the whole piece in chips. When demand is structurally multiples above supply, and supply is gated by fabs and power on multi year clocks, the gap clears exactly one way: the companies that make the packaging, the HBM, the accelerators, and the power have to get <em>much bigger</em>, in absolute terms, for a long time. &#8220;Undervalued&#8221; here isn&#8217;t a P/E call or an entry price. It means the market is implicitly pricing this fleet as if it were near its mature size, when the arithmetic says it&#8217;s a <strong>small multiple of the floor</strong>.</p><div><hr></div><h2>Conclusion: There is no bubble in AI.</h2><p>The reflexive take on AI compute is that it&#8217;s a bubble: too many GPUs, too much capex, a glut waiting to clear. Run the demand side in units of silicon and the opposite falls out. Today&#8217;s entire global fleet can&#8217;t cover even a fraction of the workforce running modest agent swarms on a cheap open model. That is the most grounded scenario in this piece, and it&#8217;s already ~44&#215; every chip on Earth, for one ordinary slice of working life. Every step up (bigger swarms, more people, faster tiers, larger models) multiplies the requirement further past everything that exists. Agents are multiplying from one toward the hundreds. The world has 19 million chips.</p><p>The chips aren&#8217;t overbuilt. They&#8217;re the floor of a new economy we&#8217;ve barely started building. The companies that make the packaging, the HBM, and the accelerators aren&#8217;t selling into a glut. They&#8217;re selling the first 1% of the substrate, against a demand curve set by the size of the workforce and the swarms it will command.</p><p>Every one of these companies is already among the most valuable on Earth. Every one of them is still too small for the floor it&#8217;s standing on. That&#8217;s the part the market hasn&#8217;t priced: not that the chips are overbuilt, but that the people who build them have barely begun.</p><p>A new economy is forming, one in which the foundations are in silicon.</p><p>Connect on LinkedIn: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p><p>Follow on X: <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><div><hr></div><p><a href="https://github.com/patrick-toulme/justabyte/blob/main/ai_compute_1_percent_post/compute_model.py">Model: </a><code>compute_model.py</code><em>, ~150 lines, every figure and number reproducible. Chip counts in FP8 H100-equivalents; supply figures SemiAnalysis-derived via IFP and Epoch; serving economics calibrated to SemiAnalysis InferenceMAX and LMSYS; training to SemiAnalysis cluster figures and Epoch. Every source is linked inline above. Change any number you doubt and rerun; the gap survives it.</em></p>]]></content:encoded></item><item><title><![CDATA[Portability Is a Myth: Why the Best AI Stacks Will Never Be Hardware-Agnostic]]></title><description><![CDATA[Hardware Vendors Should Build Their Own DSL &#8212; and Stop Chasing Portability]]></description><link>https://patricktoulme.substack.com/p/portability-is-a-myth-why-the-best</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/portability-is-a-myth-why-the-best</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sat, 16 May 2026 18:02:05 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!tNcN!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p>Connect on LinkedIn: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p><p>Follow on X: <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><div class="subscription-widget-wrap-editor" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe&quot;,&quot;language&quot;:&quot;en&quot;}" data-component-name="SubscribeWidgetToDOM"><div class="subscription-widget show-subscribe"><div class="preamble"><p class="cta-caption">Thanks for reading Just a Byte - AI Compilers, Silicon, and Systems! Subscribe for free to receive new posts and support my work.</p></div><form class="subscription-widget-subscribe"><input type="email" class="email-input" name="email" placeholder="Type your email&#8230;" tabindex="-1"><input type="submit" class="button primary" value="Subscribe"><div class="fake-input-wrapper"><div class="fake-input"></div><div class="fake-button"></div></div></form></div></div><h2><strong>Motivation</strong></h2><p><strong>Disclaimer:</strong> This blog is more opinionated than my prior blog posts.</p><p><strong>The AI industry says it wants portable kernels. It keeps building hardware-specific ones.</strong></p><p>TPUs have Pallas. Trainium has NKI. NVIDIA has CUDA C, CUTLASS, Triton, CuTile, CuTe &#8212; five DSLs and counting. AMD adopted Triton, the closest thing to a &#8220;portable&#8221; kernel DSL &#8212; then built FlyDSL anyway. Tenstorrent has tt-Metalium. Mojo targets NVIDIA and AMD.</p><p><strong>If portability worked, we&#8217;d have one universal DSL. We have many &#8212; and the number keeps growing.</strong></p><p>282 lines of Pallas Python on TPU. 4 million lines of generated CUDA C++ on Blackwell. Same operation &#8212; MoE grouped matrix multiplication. Zero shared code. <strong>Not because anyone failed at portability, but because the hardware is different enough that the optimal algorithms, tiling strategies, memory staging, and synchronization are all different. </strong><em><strong>Portability at the performance layer was never possible.</strong></em></p><p>And yet teams keep chasing it. Startups burn months building &#8220;hardware-agnostic&#8221; training stacks that run everywhere and run fast nowhere. <strong>Hardware vendors contort their silicon to fit someone else&#8217;s programming model.</strong> Engineers write portable kernels that achieve 30% MFU on two platforms instead of 90% MFU on one &#8212; and at frontier scale, that gap is <strong>millions of dollars in wasted compute.</strong></p><p><em><strong>This post argues that the best AI stacks of the future will not be portable &#8212; and that this is a good thing.</strong></em> It lays out why portability at the kernel level is a myth, why every hardware vendor will end up building their own DSL, and what those DSLs should look like to maximize both performance and usability.</p><div><hr></div><h2><strong>Three Layers of the AI Stack</strong></h2><p>Every AI training or inference stack has three layers. Understanding which layers are portable &#8212; and which aren&#8217;t &#8212; is the key to understanding why the industry keeps rebuilding infrastructure from scratch.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!tNcN!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 424w, /__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 848w, /__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 1272w, /__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!tNcN!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png" width="1456" height="927" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:927,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:397646,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/193263861?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 424w, /__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 848w, /__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 1272w, /__u/substackcdn.com/image/fetch/$s_!tNcN!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F88df8950-1bc3-4c95-ac69-aa684236b36d_5590x3560.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><h3><strong>Layer 1: The Math (Portable)</strong></h3><pre><code><code># This runs on TPU, GPU, CPU, Trainium. It's the same math everywhere.
y = jnp.dot(x, w)                    # matrix multiply
y = jax.nn.softmax(logits)            # softmax
loss = optax.softmax_cross_entropy(logits, labels)</code></code></pre><p>PyTorch and JAX are portable at this layer. <code>torch.matmul</code> and <code>jnp.dot</code> express a mathematical operation &#8212; multiply two matrices. The framework doesn&#8217;t specify how. <strong>This is real portability. It works, and it matters.</strong></p><p>But <code>jnp.dot</code> is not what runs on the hardware. What runs on the hardware is the output of a compiler that made thousands of hardware-specific decisions about tiling, memory placement, instruction selection, synchronization, and scheduling. The math is the same. Everything else is different.</p><h3><strong>Layer 2: The Compiler (The Portability Illusion)</strong></h3><p>XLA, torch.compile, Triton&#8217;s compiler &#8212; these sit between your portable math and the hardware. They take <code>jnp.dot</code> and produce TPU VLIW bundles, or CUDA PTX, or Trainium NeuronCore instructions. The API is portable. The output is not.</p><p>This is where the portability illusion lives. <strong>When someone says &#8220;JAX runs on TPU and GPU,&#8221; what they mean is: JAX&#8217;s tracing and HLO generation are shared, but XLA-TPU and XLA-GPU are </strong><em><strong>different compilers that produce fundamentally different code for fundamentally different hardware.</strong></em><strong> </strong>They share an IR format. They don&#8217;t share a backend.</p><p>The same is true for PyTorch. <code>torch.compile</code> with the Inductor backend targets NVIDIA GPUs. <code>torch.compile</code> with the XLA backend targets TPUs. The <code>torch.matmul</code> call is the same. Everything below it diverges.</p><h3><strong>Layer 3: The Hardware-Native Code (Not Portable)</strong></h3><p>This is where MFU is won or lost.</p><p>On TPU, a fused attention kernel is a Pallas program that manages VMEM explicitly, tiles across 8 sublanes &#215; 128 lanes, issues DMA operations for HBM&#8596;VMEM transfers, and produces VLIW bundles that pack MXU matmul, VPU elementwise, and XLU cross-lane ops into single cycles.</p><p>On Blackwell, the same logical operation is a CuTile program that allocates TMEM columns with atomic contention handling, elects a warp leader with <code>elect.sync</code>, issues tcgen05 MMA instructions, manages 20 mbarrier objects for async pipelining, and backs off with <code>NANOSLEEP</code> on TMEM allocation failure.</p><p>On Trainium, that operation is an NKI kernel that partitions work across NeuronCores, stages data through SBUF (on-chip SRAM) and PSUM (accumulation memory), and issues tensor engine instructions with explicit partition-dimension tiling.</p><p>These are not different implementations of the same algorithm. <strong>They are different algorithms driven by different hardware constraints.</strong> You cannot port one to another. You rewrite from scratch.</p><div><hr></div><h2><strong>The Evidence: Same Math, Zero Shared Code</strong></h2><p>Here&#8217;s the most concrete version of this argument. Mixture-of-Experts grouped matrix multiplication &#8212; the same operation compiled for two different hardware targets.</p><h3><strong>On TPU: MaxText MegaBlox (Pallas)</strong></h3><p>From my <a href="/__u/patricktoulme.substack.com/p/frontier-pretraining-infrastructure">last post</a>: MaxText&#8217;s MoE layer compiles into 29 GMM Pallas kernel calls. Each one is a <code>tpu_custom_call</code> that:</p><ul><li><p>Runs on a 3D grid <code>(tiles_n, num_active_tiles, tiles_k)</code></p></li><li><p>DMA-fetches <code>[512, 1024]</code> tiles of sorted tokens from HBM to VMEM</p></li><li><p>Selects the correct expert&#8217;s weight tile via <code>group_ids[grid_id]</code></p></li><li><p>Accumulates <code>dot(lhs_tile, rhs_tile)</code> in VMEM f32 scratch</p></li><li><p>Applies group boundary masks and stores results back to HBM</p></li></ul><p>The backward pass uses two separate kernels: <code>gmm</code> with <code>transpose_rhs=True</code> for input gradients, and <code>tgmm</code> &#8212; a structurally different kernel &#8212; for weight gradients. Tiling is <code>(512, 1024, 1024)</code>, tuned for the TPU MXU&#8217;s 256&#215;256 systolic arrays.</p><p>Total handwritten kernel code: 282 lines of Pallas Python targeting TPU memory hierarchy and MXU.</p><h3><strong>On NVIDIA: flashinfer MoE (CUDA)</strong></h3><p><a href="https://github.com/flashinfer-ai/flashinfer/pull/2917">flashinfer PR #2917</a>: 300 files of generated CUDA kernels for the same logical operation &#8212; MoE batched GEMM. Each file targets SM100 (Blackwell) and encodes its optimization parameters directly in the filename:</p><pre><code><code>batched_gemm_e2_sm100_s128x128x128_...bf16_tma_warpspecialized_cooperative_align.cu
batched_gemm_e2_sm100_s128x256x128_...fp8_tma_warpspecialized_pingpong_align.cu</code></code></pre><p>Tile sizes (128&#215;128&#215;128, 128&#215;256&#215;128). Data types (bf16, fp8). Memory access patterns (TMA). Scheduling strategies (cooperative, pingpong). Architecture target (SM100). All baked in. The PR adds approximately 4 million lines of generated CUDA &#8212; not because the developers enjoy writing CUDA, but because peak performance on SM100 requires SM100-specific code.</p><div><hr></div><h2><strong>Why a Portable DSL Can&#8217;t Fix This</strong></h2><p>The obvious response is: <strong>&#8220;What if we had one DSL that abstracts over these differences?&#8221;</strong></p><p>This is the promise of Triton, and to some extent the MLIR ecosystem. Write your kernel once in a hardware-agnostic DSL, and the compiler lowers it to TPU, GPU, or Trainium.</p><p><strong>The problem is that a portable DSL only works if the compiler is sufficiently complex to bridge the gap between abstract operations and hardware-specific instructions. </strong>And that compiler would need to do something no compiler has ever done.</p><h3><strong>What the Portable Compiler Would Need to Do</strong></h3><p>Take <code>tile_matmul(a, b)</code> in some abstract DSL. The compiler must decide:</p><p><strong>For TPU:</strong></p><ul><li><p>Allocate VMEM buffers for both operands</p></li><li><p>Issue DMA <code>copy-start</code>/<code>copy-done</code> pairs for HBM&#8594;VMEM transfers</p></li><li><p>Tile to (512, 1024, 1024) for the 256&#215;256 MXU systolic arrays</p></li><li><p>Pack MXU matmul + VPU elementwise + XLU cross-lane shuffle into VLIW bundles</p></li><li><p>Double-buffer DMA and compute across loop iterations</p></li></ul><p><strong>For Blackwell:</strong></p><ul><li><p>Allocate TMEM columns with atomic <code>UTCATOMSWS.FIND_AND_SET</code></p></li><li><p>Elect a warp leader with <code>elect.sync</code> before every MMA</p></li><li><p>Issue tcgen05 MMA instructions through the leader thread only</p></li><li><p>Manage 20 mbarrier objects with parity-based reuse for async pipelining</p></li><li><p>Back off with 100ns <code>NANOSLEEP</code> on TMEM allocation contention</p></li><li><p>Await results with 10ms barrier wait timeouts</p></li></ul><p><strong>For Trainium:</strong></p><ul><li><p>Stage data through SBUF (on-chip SRAM) with explicit partition-dimension tiling</p></li><li><p>Issue tensor engine instructions across NeuronCore partitions</p></li><li><p>Accumulate in PSUM memory space</p></li><li><p>Manage NeuronCore-level parallelism</p></li></ul><p>These aren&#8217;t different instruction selections for the same algorithm. They&#8217;re different algorithms. The double-buffered DMA approach on TPU, the leader-elected async MMA on Blackwell, and the partition-tiled tensor engine on Trainium are structurally different programs that solve the same mathematical problem through entirely different mechanisms.</p><h3><strong>The Union Problem</strong></h3><p>A portable DSL that exposes enough hardware detail to write high-performance kernels would need to contain <strong>the union of all hardware concepts:</strong></p><ul><li><p><strong>TPU:</strong> VMEM, HBM, DMA engines, MXU, VPU, XLU, sublanes/lanes, VLIW bundles</p></li><li><p><strong>Blackwell:</strong> SMEM, TMEM, L2 cache, tensor cores, tcgen05, mbarriers, TMA, warp specialization</p></li><li><p><strong>Trainium:</strong> SBUF, PSUM, HBM, tensor engine, NeuronCores, partition dimensions</p></li></ul><p>That&#8217;s not a DSL &#8212; it&#8217;s three DSLs wearing a trench coat. Every kernel written in it would need <code>if target == TPU: ... elif target == NVIDIA: ... elif target == TRAINIUM: ...</code> escape hatches that defeat the purpose of a shared abstraction.</p><p><strong>Alternatively, a portable DSL that hides these differences produces generic code that doesn&#8217;t exploit any of them</strong> &#8212; and you&#8217;re back at the compiler layer, hoping the compiler is smart enough to recover the performance you left on the table.</p><div><hr></div><h2><strong>The DSL Proliferation Is the Proof</strong></h2><p>If portability at the kernel level were achievable, the market would have converged on one DSL by now. Instead, every hardware vendor has independently built their own:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!a34p!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 424w, /__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 848w, /__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 1272w, /__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!a34p!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png" width="1456" height="671" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/a97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:671,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:303414,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/193263861?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 424w, /__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 848w, /__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 1272w, /__u/substackcdn.com/image/fetch/$s_!a34p!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa97a93dc-37fb-41aa-ba2b-c3dc976298b0_3192x1470.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p><strong>NVIDIA alone has </strong><em><strong>four</strong></em><strong> DSLs for writing kernels</strong> &#8212; CUDA C, CUTLASS/CuTe DSL, Triton, and CuTile &#8212; spanning from thread-level C to tile-level Python. Four abstraction levels for one vendor&#8217;s hardware, because even within a single architecture family, one DSL isn&#8217;t enough.</p><p><strong>AMD&#8217;s trajectory is particularly telling.</strong> They adopted Triton &#8212; the closest thing the industry has to a &#8220;portable&#8221; kernel DSL &#8212; and forked it for ROCm. But adopting Triton wasn&#8217;t enough. AMD still built FlyDSL: a Python DSL with explicit register-level layouts, lane-level operations, and CuTe-style tiling algebra, lowered through their own MLIR Fly dialect to ROCDL machine code. <strong>Even when you adopt a &#8220;portable&#8221; DSL, you still end up building a hardware native one to get the performance the portable one can&#8217;t deliver.</strong></p><p>This isn&#8217;t a failure of engineering. It&#8217;s evidence that hardware-specific kernel programming is irreducibly hardware-specific.</p><div><hr></div><h2><strong>For Highest MFU: Write the ISA in Python</strong></h2><p>Here&#8217;s the principle that unifies all the successful DSLs: <strong>the highest-performing kernel DSLs are thin Python wrappers around the hardware&#8217;s instruction set.</strong></p><p>Pallas doesn&#8217;t abstract away the MXU &#8212; it gives you <code>dot()</code> that maps directly to MXU matmul instructions, <code>load()</code>/<code>store()</code> that map to DMA operations, and Ref types that correspond to VMEM/HBM/SMEM memory spaces.</p><p>CuTile doesn&#8217;t abstract away tcgen05 &#8212; it gives you <code>tile_matmul()</code> that maps to tcgen05 MMA, <code>tile_load()</code> that maps to TMA descriptors, and automatic TMEM allocation that generates <code>UTCATOMSWS.FIND_AND_SET</code> instructions.</p><p>NKI doesn&#8217;t abstract away the tensor engine &#8212; it gives you <code>nisa.nc_matmul()</code> that maps to tensor engine instructions and explicit SBUF/PSUM memory staging. <strong>When you write NKI, you&#8217;re writing the Trainium ISA in Python.</strong></p><p>The value of these DSLs is not abstraction &#8212; it&#8217;s <em>ergonomics</em>. <strong>They let you write hardware-specific code in Python instead of C or assembly.</strong> The key is that they don&#8217;t hide the hardware. They expose it through a familiar syntax.</p><p>This is the right trade-off. C-level ISA access (CUDA C++, raw PTX, assembly) gives you maximum control but poor iteration speed and readability. An abstract portable DSL gives you ergonomics but takes away the hardware control that determines MFU. <strong>A Python DSL over the hardware ISA gives you both: hardware control with Python iteration speed.</strong></p><p>The flashinfer PR is 4 million lines of generated CUDA C++. MaxText&#8217;s MegaBlox kernel is 282 lines of Pallas Python. Both achieve high MFU on their respective hardware. The difference in developer experience is staggering &#8212; but the reason Pallas works isn&#8217;t that it&#8217;s more abstract than CUDA. <strong>It&#8217;s that it maps Python directly to TPU hardware concepts, so the 282 lines express the same level of hardware-specific intent as the generated CUDA, just more concisely.</strong></p><div><hr></div><h2><strong>Recommendations for Hardware Vendors</strong></h2><p>If you&#8217;re building AI accelerator hardware and want developers to achieve peak performance on your silicon, here&#8217;s what the evidence from every successful AI hardware platform tells us:</p><h3><strong>1. Build Your Own Python DSL</strong></h3><p><strong>Don&#8217;t wait for an open-source DSL to support your hardware.</strong> Google built Pallas. NVIDIA built CuTile. AWS built NKI. In every case, the vendor built DSL outperforms third-party alternatives because the DSL designers have deep knowledge of the hardware&#8217;s performance model.</p><p>Your DSL should let developers write kernels that map to your hardware&#8217;s execution model in Python. Not C. Not a custom language. Python &#8212; because the ML ecosystem lives in Python, and the fastest path from researcher insight to high-performance kernel is a Python file that directly expresses your hardware&#8217;s concepts.</p><p><strong>Build the DSL for the silicon. Don&#8217;t build the silicon for the DSL.</strong></p><h3><strong>2. Expose the ISA, Not Just High-Level Abstractions</strong></h3><p>Your DSL should operate at two levels:</p><p><strong>Low level: virtual ISA access.</strong> Let developers express the hardware&#8217;s instruction set &#8212; your matmul unit&#8217;s tile operations, your memory hierarchy&#8217;s explicit staging, your synchronization primitives. This is where peak MFU comes from. Developers who understand your hardware should be able to write near ISA level code in Python and know exactly what instructions will be generated.</p><p><strong>High level: compiler-lowered operations.</strong> Not all your users will want to write ISA &#8212; and that&#8217;s fine. For developers who don&#8217;t need peak performance or are prototyping, provide higher-level operations (<code>tile_matmul</code>, <code>tile_load</code>) that the compiler lowers to the ISA. This is where Pallas&#8217;s <code>dot()</code> and CuTile&#8217;s <code>tile_matmul()</code> live &#8212; high-level enough to be ergonomic, low-level enough that an expert can predict the generated code.</p><p><strong>Critically, the two levels must be mixable in the same kernel.</strong> Pallas GPU gets this right: you can write a kernel that uses high level <code>jnp</code> operations for the parts where the compiler is good enough, and drop into explicit tcgen05 instructions for the parts where you need direct hardware control &#8212; in the same function, in the same Python file. <strong>The developer chooses the abstraction level per-operation, not per-kernel.</strong></p><p>This is how you scale adoption. New users start with the high-level ops and get reasonable performance. Expert users drop into ISA-level control for the hot path. A DSL that only provides the high level interface leaves performance on the table for expert users. A DSL that only provides ISA access has too steep a learning curve for adoption. The DSLs that win will let developers slide between both levels seamlessly.</p><h3><strong>3. Do Not Constrain Your Hardware to an Open-Source DSL</strong></h3><p><strong>This is the most counterintuitive recommendation, but the evidence is clear.</strong></p><p>Triton is a good DSL for NVIDIA GPUs that are not Blackwell or Rubin. It is not a good target for designing your hardware around. If you build your chip to be &#8220;Triton-compatible,&#8221; you&#8217;ve constrained your hardware architecture to the abstractions Triton exposes &#8212; shared memory, thread blocks, warp-level operations. Those abstractions reflect NVIDIA&#8217;s hardware design choices. They may not reflect yours.</p><p><strong>NVIDIA itself doesn&#8217;t constrain Blackwell to Triton.</strong> CuTile exposes tcgen05, TMEM, and mbarrier &#8212; hardware features that don&#8217;t exist in Triton&#8217;s programming model. Triton will eventually support them, but CuTile had day one access. If NVIDIA doesn&#8217;t limit itself to the open-source DSL it helped create, you shouldn&#8217;t limit yourself to it either.</p><p>The risk of designing hardware for an existing DSL: <strong>your chip becomes a less-efficient implementation of someone else&#8217;s hardware model. </strong>The opportunity of designing a DSL for your hardware: your chip can exploit architectural innovations that no existing DSL anticipated.</p><p><em><strong>Build the DSL for the silicon. Don&#8217;t build the silicon for the DSL.</strong></em></p><h3><strong>4. Control Your Own IR</strong></h3><p>When you own your DSL, you own the intermediate representation &#8212; and <strong>the IR is where hardware specific semantics live.</strong> When you adopt someone else&#8217;s DSL, you inherit their IR&#8217;s assumptions about how hardware works.</p><p>If your chip ships a new memory space or a novel synchronization primitive, and you own the IR, you add an op and it&#8217;s available to users on day one. The compiler understands it, can optimize around it, can fuse across it, can schedule it relative to other operations. <strong>Your new hardware feature is a first class citizen in the optimization pipeline.</strong></p><p>If you&#8217;re lowering from someone else&#8217;s IR &#8212; say Triton&#8217;s TTIR &#8212; your options are worse. You either wait for upstream to add support, hack around it with extension intrinsics, or fork the IR and maintain it yourself.</p><p>Triton&#8217;s <code>tl.extra.nvidia</code> and <code>tl.extra.cuda</code> extensions illustrate the problem. They give you <em>access</em> to hardware-specific instructions from within Triton &#8212; but those calls are <strong>escape hatches,</strong> not integrated operations. The Triton compiler can&#8217;t reason about them, can&#8217;t fuse across them, can&#8217;t schedule them relative to surrounding Triton ops. <strong>You get the syntax of a high level DSL with the semantics of inline assembly.</strong></p><p>Compare that to CuTile, where tcgen05 is a first-class IR concept. The compiler understands it, generates TMEM allocation and mbarrier synchronization automatically, and schedules it as part of the overall kernel optimization. The difference isn&#8217;t access to the instruction &#8212; <em><strong>it&#8217;s whether the compiler understands the instruction.</strong></em></p><p>That&#8217;s the cost of building on someone else&#8217;s IR: your hardware-specific features will always be second-class citizens, bolted on through extensions rather than integrated into the optimization pipeline. <strong>This is why NVIDIA built CuTile instead of extending Triton for Blackwell &#8212; and it&#8217;s why you should own yours too.</strong></p><h3><strong>5. Build Your Compiler in MLIR &#8212; Open Source the Frontend, Keep the Backend</strong></h3><p>MLIR is often cited as the solution to the portability problem. It&#8217;s not &#8212; it&#8217;s a framework for building compilers, and the compilers people build with it are hardware specific. But that&#8217;s exactly why you should use it.</p><p><strong>MLIR gives you common infrastructure:</strong> parsing, verification, transformation utilities, and a pass pipeline framework that every compiler needs. Instead of building all of that from scratch, you define your own MLIR dialect with ops that match your hardware&#8217;s semantics, and plug into the existing ecosystem. Your DSL lowers to your dialect. Your dialect lowers to your backend. MLIR handles the plumbing.</p><p>The smart split is: <strong>open source your frontend compiler, keep your backend codegen closed.</strong> The frontend &#8212; the DSL to MLIR lowering, frontend optimization passes, the dialect definition &#8212; should be open. This is what developers interact with, debug against, and build tooling around. <strong>Transparency here builds trust and adoption.</strong></p><p>The backend &#8212; the final lowering from your MLIR dialect to machine code, the instruction selection, the register allocation, the hardware specific scheduling &#8212; <strong>can stay closed.</strong> This is where your proprietary IP lives, and keeping it closed doesn&#8217;t hurt developer experience because users don&#8217;t need to read your codegen to write high performance kernels.</p><p><strong>We&#8217;re seeing this pattern emerge across the industry.</strong> The Mosaic compiler in the JAX repository is open source &#8212; developers can read the TPU and GPU Pallas frontend passes, understand how their kernels are lowered, and debug compilation issues. The backend TPU codegen (libtpu) is closed. AWS&#8217;s NKI compiler is on the path to being open sourced at the frontend level, with the Trainium backend remaining proprietary.</p><p><strong>This is the right model.</strong> MLIR makes it cheap to build a non portable compiler. Open-sourcing the frontend makes it usable. Keeping the backend closed protects your IP. Everyone wins.</p><div><hr></div><h2><strong>The Implication for Frontier Labs</strong></h2><p>If you&#8217;re building a frontier training or inference stack, the evidence from every successful large scale deployment says the same thing:</p><p><strong>Pick your hardware. Go deep. Don&#8217;t plan to port.</strong></p><p>The &#8220;we&#8217;ll stay hardware agnostic and port later&#8221; strategy sounds prudent but costs you twice: you lose MFU today because you avoid hardware specific optimizations, and you spend engineering time later porting code <strong>that can&#8217;t actually be ported</strong> &#8212; it has to be rewritten.</p><p>MaxText achieves high performance on TPU v6e for MoE pretraining because it uses TPU native Pallas kernels, TPU native XLA fusion, and TPU native SPMD partitioning. The flashinfer MoE kernels will achieve peak MFU on SM100 because they&#8217;re 4 million lines of SM100 specific CUDA. Neither can run on the other&#8217;s hardware. Neither should.</p><p>The math is portable. Write your model in JAX or PyTorch &#8212; the <code>jnp.dot</code> and <code>torch.matmul</code> calls work everywhere. But the moment you need peak performance &#8212; the attention kernel, the MoE routing, the collective overlap, the memory scheduling &#8212; you&#8217;re writing for one hardware target. <strong>Accept it early, and you&#8217;ll ship faster than the team that&#8217;s still trying to write kernels that run everywhere and run fast nowhere.</strong></p><h3><strong>Inference: Where This Matters Even More</strong></h3><p>The portability argument applies to training, <strong>but it&#8217;s even more acute for inference.</strong> In training, you amortize kernel performance over weeks or months of runs. In inference, kernel performance <em>is</em> your cost per token &#8212; and cost per token is the metric that determines whether your product is economically viable at scale. Every percentage point of MFU you leave on the table by using a &#8220;portable&#8221; kernel instead of a hardware native one translates <strong>directly into higher serving costs</strong> &#8212; multiplied by every token you generate, forever.</p><h3><strong>The Cost of Kernel Development Is Collapsing</strong></h3><p>One historical argument for portability was that writing hardware-specific kernels was too expensive. If it takes a team of experts six months to write a custom attention kernel, you can&#8217;t afford to do it for every hardware target. Better to write it once and port.</p><p><strong>That calculus is changing. </strong>Agentic AI coding tools are <strong>dramatically reducing the cost </strong>of kernel development. An engineer with a Python DSL and an AI coding assistant can iterate on kernel implementations in hours instead of weeks. The DSLs themselves &#8212; Pallas, CuTile, NKI &#8212; already reduced the barrier from &#8220;CUDA C++ expert&#8221; to &#8220;Python developer who understands the hardware.&#8221; <strong>AI tooling reduces it further:</strong> generate a kernel skeleton, profile it, have the agent iterate on tiling and memory staging, converge on high MFU code in a fraction of the time.</p><p>When kernel development cost was high, portability was an economic argument: write once, amortize across targets. When kernel development cost is low, that argument collapses. <strong>It becomes cheaper to write two hardware native kernels that each achieve 90%+ MFU than to write one &#8220;portable&#8221; kernel that achieves 30% MFU everywhere and needs constant tuning to close the gap.</strong></p><h3><strong>Cost Per Token Is All That Matters</strong></h3><p><strong>At the end of the day, at frontier scale the economics are driven by one metric: cost per token. </strong>Frontier labs and organizations will use whatever silicon provides the best cost per token &#8212; and<em><strong> is usable.</strong></em> That second part is the bottleneck. The silicon that wins isn&#8217;t necessarily the silicon with the best theoretical FLOP/s. It&#8217;s the silicon where an engineer can write a high MFU kernel in a week instead of a quarter. A solid Python DSL that lets you write the ISA directly is a massive step in the usability direction &#8212; <strong>and usability is what determines whether your hardware gets adopted or sits in a rack collecting dust.</strong></p><div><hr></div><h2><strong>Conclusion</strong></h2><p><strong>The frameworks are portable. </strong><code>jnp.dot</code> runs on TPU, GPU, and Trainium. That&#8217;s a genuine achievement, and it matters &#8212; your model architecture, your training recipe, your evaluation harness all transfer across hardware.</p><p><strong>The kernels are not.</strong> 282 lines of Pallas and 4 million lines of CUDA solve the same math with zero shared code, and that gap isn&#8217;t closing &#8212; it&#8217;s widening with every generation of silicon. Blackwell introduced tcgen05 and TMEM. TPU v6e has a different MXU layout than v5e. Trainium&#8217;s tensor engine has its own tiling constraints. <strong>Each generation adds hardware specific concepts that reward hardware-specific code.</strong></p><p><strong>The industry should stop treating this as a problem to solve and start treating it as a reality to build around. Hardware vendors:</strong> build your own Python DSL, control your own IR, open source your frontend compiler. <strong>Frontier labs:</strong> pick your hardware and go deep &#8212; the cost of writing native kernels is collapsing, and the cost of not writing them is millions in wasted compute per year.</p><p><strong>The race isn&#8217;t to build the most portable stack. It&#8217;s to build the deepest one.</strong></p><div class="subscription-widget-wrap-editor" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe&quot;,&quot;language&quot;:&quot;en&quot;}" data-component-name="SubscribeWidgetToDOM"><div class="subscription-widget show-subscribe"><div class="preamble"><p class="cta-caption">Thanks for reading Just a Byte - AI Compilers, Silicon, and Systems! Subscribe for free to receive new posts and support my work.</p></div><form class="subscription-widget-subscribe"><input type="email" class="email-input" name="email" placeholder="Type your email&#8230;" tabindex="-1"><input type="submit" class="button primary" value="Subscribe"><div class="fake-input-wrapper"><div class="fake-input"></div><div class="fake-button"></div></div></form></div></div>]]></content:encoded></item><item><title><![CDATA[Launching pyptx — a Python DSL for writing NVIDIA PTX kernels directly.]]></title><description><![CDATA[The world's first open source Python DSL for NVIDIA PTX.]]></description><link>https://patricktoulme.substack.com/p/launching-pyptx-a-python-dsl-for</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/launching-pyptx-a-python-dsl-for</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sat, 25 Apr 2026 19:42:30 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!b8CV!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><p><strong>Repo:</strong> <a href="https://github.com/patrick-toulme/pyptx">https://github.com/patrick-toulme/pyptx</a></p><p><strong>Docs:</strong> <a href="https://pyptx.dev/">https://pyptx.dev/</a></p><p><em><strong>Connect on LinkedIn:</strong> <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></em></p><p><em><strong>Follow on X:</strong> <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></em></p><p>Launching pyptx &#8212; a Python DSL for writing NVIDIA PTX kernels directly.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!b8CV!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 424w, /__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 848w, /__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 1272w, /__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!b8CV!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png" width="1529" height="427" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/a95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:427,&quot;width&quot;:1529,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:165517,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:false,&quot;topImage&quot;:true,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/195468346?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F9ea1736f-ef97-48ee-8b06-cee8a9f486d4_1536x1024.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 424w, /__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 848w, /__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 1272w, /__u/substackcdn.com/image/fetch/$s_!b8CV!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fa95a91fa-8dbe-41e8-84cb-0d8e9c88db3d_1529x427.png 1456w" sizes="100vw" fetchpriority="high"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><div class="twitter-embed" data-attrs="{&quot;url&quot;:&quot;https://x.com/PatrickToulme/status/2048113408940159430?s=20&quot;,&quot;full_text&quot;:&quot;Launching pyptx &#8212; a Python DSL for writing NVIDIA PTX kernels.\n\nOne PTX instruction = one Python call. Write pure PTX in Python.\n\nDirect Hopper + Blackwell support: wgmma, TMA, tcgen05, mbarriers. JAX + PyTorch integration. \n\nIncludes GEMM, grouped GEMM, RMSNorm, SwiGLU, and a&quot;,&quot;username&quot;:&quot;PatrickToulme&quot;,&quot;name&quot;:&quot;Patrick C Toulme&quot;,&quot;profile_image_url&quot;:&quot;https://pbs.substack.com/profile_images/2013003105390981120/ydKai4dt_normal.jpg&quot;,&quot;date&quot;:&quot;2026-04-25T18:54:23.000Z&quot;,&quot;photos&quot;:[],&quot;quoted_tweet&quot;:{},&quot;reply_count&quot;:10,&quot;retweet_count&quot;:8,&quot;like_count&quot;:48,&quot;impression_count&quot;:1632,&quot;expanded_url&quot;:null,&quot;video_url&quot;:null,&quot;video_preview_media_key&quot;:null,&quot;belowTheFold&quot;:false}" data-component-name="Twitter2ToDOM"></div><p>Today I&#8217;m open-sourcing a project I&#8217;ve been building on personal time: pyptx, a Python DSL where the function body is the PTX instruction stream. One PTX instruction = one Python call. No optimizer, no autotuner, no tile IR between you and the hardware.</p><p>Why? Because the newest GPU features &#8212; Hopper&#8217;s wgmma, TMA multicast, mbarrier-based pipelines, Blackwell&#8217;s tcgen05.mma + TMEM + cooperative 2-SM MMA often only exist at the PTX level. For developers chasing peak performance, that has historically meant writing inline PTX inside CUDA C++.                                                                                        </p><p>pyptx brings that whole path into Python. Callable from JAX (via typed XLA FFI) and PyTorch (eager, torch.compile, and a C++ extension fast path).</p><p>A few numbers from real silicon:</p><p>&#8226; <strong>H100 bf16 GEMM</strong>: 815 TFLOPS, competitive with cuBLAS at matrix sizes &#8805; 6K</p><p>&#8226; <strong>B200 bf16 GEMM</strong>: 1240 TFLOPS on the 1SM kernel</p><p>&#8226; <strong>RMSNorm</strong>: 2.6 TB/s (88% of HBM3 peak, 3.9&#215; PyTorch eager)</p><p>&#8226; <strong>SwiGLU</strong>: 2.8 TB/s (94% HBM3)</p><p>The other half of the project is a transpiler in the opposite direction:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!k3YH!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 424w, /__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 848w, /__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 1272w, /__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!k3YH!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png" width="1456" height="422" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:422,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:316950,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/195468346?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 424w, /__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 848w, /__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 1272w, /__u/substackcdn.com/image/fetch/$s_!k3YH!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F41023d0d-b7f8-4a2d-be83-40ae51850097_3428x994.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p></p><p>  python -m pyptx.codegen kernel.ptx --sugar</p><p>takes PTX from anywhere &#8212; nvcc, Triton, CUTLASS output, DeepGEMM kernels &#8212; and emits editable pyptx Python. The parser/emitter round-trips byte-identical on 218+ real-world kernels. So you can read someone else&#8217;s kernel as Python, modify it, and ship the result.</p><p>Built end-to-end: parser, IR, emitter, transpiler, JAX integration, PyTorch integration, full Hopper + Blackwell ISA coverage, multi-arch wheels published to PyPI. ~17K lines of Python total.</p><p>Ships with maintained GEMM, grouped GEMM, RMSNorm, LayerNorm, and SwiGLU kernels for both Hopper and Blackwell, plus the PTX &#8594; Python transpiler.</p><p><code>  pip install pyptx[torch]    # for PyTorch</code></p><p><code>  pip install pyptx[jax]      # for JAX</code></p><p><code>  pip install pyptx[all]      # both</code></p><p><strong>Repo:</strong> <a href="https://github.com/patrick-toulme/pyptx">https://github.com/patrick-toulme/pyptx</a></p><p><strong>Docs:</strong> <a href="https://pyptx.dev">https://pyptx.dev</a></p><p>If you write GPU kernels &#8212; especially if you&#8217;ve ever wished Triton would let you express a specific wgmma pattern, or wanted to read a CUTLASS PTX dump as editable Python &#8212; try it. PRs welcome, especially Blackwell tuning and new ISA wrappers.</p><div class="subscription-widget-wrap-editor" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe&quot;,&quot;language&quot;:&quot;en&quot;}" data-component-name="SubscribeWidgetToDOM"><div class="subscription-widget show-subscribe"><div class="preamble"><p class="cta-caption">Thanks for reading Just a Byte - AI Compilers, Silicon, and Systems! Subscribe for free to receive new posts and support my work.</p></div><form class="subscription-widget-subscribe"><input type="email" class="email-input" name="email" placeholder="Type your email&#8230;" tabindex="-1"><input type="submit" class="button primary" value="Subscribe"><div class="fake-input-wrapper"><div class="fake-input"></div><div class="fake-button"></div></div></form></div></div>]]></content:encoded></item><item><title><![CDATA[Frontier Pretraining Infrastructure Is Already Open Source: GPT-OSS on TPU with MaxText]]></title><description><![CDATA[Tracing a Full MoE Training Step Through the XLA Compiler &#8212; and Why You Shouldn't Rebuild This From Scratch]]></description><link>https://patricktoulme.substack.com/p/frontier-pretraining-infrastructure</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/frontier-pretraining-infrastructure</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Mon, 06 Apr 2026 15:46:20 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!UwCR!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><p><strong>Full Code / IR Dumps:</strong> <a href="https://github.com/patrick-toulme/justabyte/tree/main/maxtext_pretraining">https://github.com/patrick-toulme/justabyte/tree/main/maxtext_pretraining</a></p><p><em><strong>Connect on LinkedIn:</strong> <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></em></p><p><em><strong>Follow on X:</strong> <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></em></p><h2>Motivation</h2><p>3.5 billion parameters. 16 MoE experts with top-2 routing. 2-way FSDP sharded across 2 chips, expert-parallel across the other 2 chips. 200 training steps at 22 TFLOP/s per device &#8212; every step generating 887 fused kernels, 104 async all-gathers, 12 ragged all-to-all collectives, 8 Pallas attention kernels, and 24 MegaBlox grouped matmuls. All compiled, scheduled, and overlapped automatically from a single <code>python3 -m maxtext.trainers.pre_train.train</code> command.</p><p>The cost of this entire experiment was under 10 dollars.</p><p><em><strong>I keep watching startups, neo-frontier labs and major organizations spend months rebuilding training infrastructure from scratch in PyTorch on NVIDIA GPUs. FSDP sharding, kernel fusion, collective overlap, MoE routing.</strong></em> Team after team duplicating the same work, and it's not clear most of them even achieve as high an MFU as MaxText gets out of the box. Meanwhile, I rarely hear MaxText mentioned as even an option.</p><p>I&#8217;m not claiming a frontier lab can take MaxText off the shelf and ship. But I am claiming they could get to a production training job dramatically faster by building on top of it &#8212; and this post shows why. I understand some organizations have bespoke requirements or privacy reasons that prevent them from using open source &#8212; this post isn't aimed at them.</p><p>This experiment runs on 4 chips, but MaxText and JAX's SPMD model scale with no code changes. You change <code>ici_fsdp_parallelism=2</code> to <code>ici_fsdp_parallelism=128</code> and the same code runs on 128 chips. The compilation pipeline I traced here &#8212; the fusions, the async collectives, the VMEM scheduling &#8212; is the same compilation pipeline at any scale. The only things that change are the mesh dimensions in a config flag.</p><p>I ran a GPT-OSS MoE pretraining job on 4 TPU v6e chips, dumped the IR, and traced the full compilation pipeline from Python down to fused TPU kernels. <em><strong>What I found: the XLA compiler does a staggering amount of work that would take months of manual kernel engineering. </strong></em>The fusion, the async scheduling, the memory management, the SPMD partitioning &#8212; it&#8217;s all generated automatically. <em><strong>This is the infrastructure teams are rebuilding by hand.</strong></em></p><div><hr></div><h2><strong>Setup</strong></h2><p><strong>Hardware:</strong> a single TPU v6e-4 node &#8211; 4 TPU v6e chips in a 2&#215;2 mesh, each with 32 GB HBM. <strong>Total:</strong> 128 GB HBM across the node.</p><pre><code><code>[TpuDevice(id=0, coords=(0,0,0)), TpuDevice(id=1, coords=(1,0,0)),
 TpuDevice(id=2, coords=(0,1,0)), TpuDevice(id=3, coords=(1,1,0))]</code></code></pre><p><strong>Software:</strong> MaxText cloned from <a href="https://github.com/AI-Hypercomputer/maxtext">AI-Hypercomputer/maxtext</a>, installed with <code>pip install -e ".[tpu]"</code>. JAX 0.9.0, libtpu 0.0.36. Nothing else.</p><p><strong>The model:</strong> a scaled variant of GPT-OSS (OpenAI&#8217;s recently open-sourced MoE architecture). The released gpt-oss-20b uses 32 experts with 4 active experts per token; I adjusted width/depth and used 16 experts with top-2 routing to fit on 4 chips&#8212;the point here is tracing the compiler + SPMD stack, not exact parity with the released checkpoints.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!zi9a!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 424w, /__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 848w, /__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 1272w, /__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!zi9a!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png" width="800" height="633" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/90009631-44cf-4510-8de2-c093d14cf847_800x633.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:633,&quot;width&quot;:800,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:31341,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/189578317?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 424w, /__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 848w, /__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 1272w, /__u/substackcdn.com/image/fetch/$s_!zi9a!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F90009631-44cf-4510-8de2-c093d14cf847_800x633.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>The parallelism: <strong>2-way FSDP &#215; Expert Parallelism-2</strong>. Weights are sharded across 2 chips (FSDP), and the 16 experts are partitioned across 2 chips (EP). This requires real multi-axis communication &#8211; all-gather for FSDP weight reconstruction and ragged all-to-all for MoE token dispatch.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!UwCR!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 424w, /__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 848w, /__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 1272w, /__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!UwCR!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png" width="1456" height="930" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/db2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:930,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:141097,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/189578317?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 424w, /__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 848w, /__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 1272w, /__u/substackcdn.com/image/fetch/$s_!UwCR!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdb2112e9-86be-4309-b590-300788a3f4e0_2160x1380.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><pre><code><code>python3 -m maxtext.trainers.pre_train.train src/maxtext/configs/base.yml \
  model_name="gpt-oss-20b" override_model_config=true \
  base_emb_dim=2048 base_num_decoder_layers=16 \
  num_experts=16 num_experts_per_tok=2 \
  megablox=true sparse_matmul=true capacity_factor=-1.0 \
  ici_fsdp_parallelism=2 ici_expert_parallelism=2 \
  attention='flash' remat_policy=full dtype=bfloat16 \
  dump_hlo=true dump_jaxpr=true ...</code></code></pre><p>That <code>ici_fsdp_parallelism=2 ici_expert_parallelism=2</code> line is the entire parallelism specification. <strong>The same codebase handles single chip debugging and multi thousand chip production training</strong> &#8211; the parallelism strategy is a configuration parameter, not code.</p><p><strong>Note:</strong> 2-way FSDP &#215; 2-way EP is not the most optimal parallelism configuration for this model and chip count. The point of this experiment is demonstrating the capability of the stack, not maximizing utilization or FLOP/s.</p><h2><strong>What is MaxText?</strong></h2><p>MaxText is Google&#8217;s open-source JAX/Flax training framework for large language models. It sits at <a href="https://github.com/AI-Hypercomputer/maxtext">AI-Hypercomputer/maxtext</a> and is the reference implementation for training on TPUs.</p><p>The architecture has four layers:</p><p><strong>Configuration.</strong> MaxText uses OmegaConf + Pydantic for model config. You pick a base architecture with <code>model_name="gpt-oss-20b"</code> (which selects the decoder layer class, attention pattern, MoE routing style) and then override any dimension with CLI flags: <code>base_emb_dim=2048 base_num_decoder_layers=16 num_experts=16</code>.</p><p><strong>Model layers (Flax NNX).</strong> The model definition is pure Python using Flax&#8217;s NNX API. <code>GptOssDecoderLayer</code> in <code>models/gpt_oss.py</code> is 289 lines &#8211; it composes an <code>Attention</code> module, an <code>RMSNorm</code>, and a <code>RoutedMoE</code> with residual connections and dropout. The scannable block wraps multiple decoder layers with alternating attention patterns (local sliding-window, global) for <code>nn.scan</code> compilation.</p><p><strong>Custom kernels (Pallas).</strong> Performance-critical operations &#8211; grouped matrix multiplication for MoE, flash attention &#8211; are hand-written Pallas kernels. Pallas is JAX&#8217;s kernel authoring language that compiles to <code>tpu_custom_call</code> ops. <strong>These are the only hand-written kernels in the stack; everything else is generated by XLA.</strong></p><p><strong>Parallelism (JAX SPMD).</strong> MaxText defines a 13-axis logical mesh (<code>fsdp</code>, <code>expert</code>, <code>tensor</code>, <code>data</code>, <code>sequence</code>, &#8230;) and annotates tensors with <code>nn.with_logical_constraint</code>. JAX&#8217;s Shardy partitioner reads these annotations and automatically inserts all-gathers, reduce-scatters, and ragged all-to-all collectives. </p><p>The key insight: <strong>MaxText is a </strong><em><strong>thin</strong></em><strong> layer.</strong> The model definition is high level Python. The kernels handle the MXU-level compute. Everything in between &#8211; fusion, scheduling, communication, memory management &#8211; is the compiler&#8217;s job.</p><h2><strong>Training Results</strong></h2><pre><code><code>step:   0, seconds: 30.309, TFLOP/s/device:  0.133, loss: 10.871  # compilation
step:   1, seconds:  0.531, TFLOP/s/device:  7.598, loss: 10.871  # warmup
step:   3, seconds:  0.182, TFLOP/s/device: 22.194, loss: 10.730  # steady state
step:  50, seconds:  0.180, TFLOP/s/device: 22.441, loss:  0.008
step: 100, seconds:  0.180, TFLOP/s/device: 22.446, loss:  0.002
step: 199, seconds:  0.180, TFLOP/s/device: 22.368, loss:  0.001</code></code></pre><p>Step 0 takes 30 seconds &#8211; XLA compilation. By step 3, we&#8217;re at steady state. Note: this run uses dataset_type=synthetic, so the rapid loss collapse is expected (fast memorization). The goal is validating throughput and end-to-end correctness of MoE + collectives + optimizer, not model quality. <strong>182 ms per step, 22.2 TFLOP/s per device, 5,640 tokens/s/device</strong>. The loss drops from 10.87 to 0.001 &#8211; synthetic data memorization confirming the full pipeline works: forward through 16 MoE layers, loss computation, backward through MegaBlox kernels and routing gradients, gradient reduce-scatter across devices, and AdamW optimizer update.</p><div><hr></div><h2><strong>One Binary Per Step</strong></h2><p>Here&#8217;s something that surprises people coming from PyTorch: <strong>the </strong><em><strong>entire</strong></em><strong> training step </strong>&#8211; forward pass through 16 MoE decoder layers, cross-entropy loss, backward pass through every layer including custom MegaBlox VJPs, gradient reduce-scatter across all 4 devices, and AdamW optimizer update &#8211; is compiled into <strong>a single XLA binary</strong>.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!ss3z!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 424w, /__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 848w, /__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 1272w, /__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!ss3z!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png" width="1440" height="572" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/bb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:572,&quot;width&quot;:1440,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:109010,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/189578317?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 424w, /__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 848w, /__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 1272w, /__u/substackcdn.com/image/fetch/$s_!ss3z!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fbb39085b-1b44-4f46-9e00-a9580b2ef2eb_1440x572.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>This is what <code>@jax.jit</code> does. When MaxText calls <code>train_step(state, batch)</code>, JAX traces the full Python function into a Jaxpr (18,079 lines of functional IR). XLA then compiles that Jaxpr into a single HLO module (22,636 lines after optimization). <strong>That module is compiled once into a TPU binary.</strong> Every subsequent step just re-executes the same binary with new inputs.</p><p>This is why step 0 takes 30 seconds and step 1 takes 0.5 seconds. Step 0 is compilation. Steps 1-199 are execution.</p><p>The consequence: <em><strong>XLA sees everything at compile time.</strong></em> That turns the training step into one schedulable program instead of a chain of compiled islands. XLA can fuse across the forward/backward/optimizer boundary, and it can overlap communication with compute by issuing async all-gathers early and consuming them later&#8212;e.g., pulling layer 5&#8217;s shards while layer 3&#8217;s backward pass runs. <strong>With separate compilation units, those boundaries act like barriers</strong>: more launches, more forced materialization, and less automatic comm/compute overlap.</p><h2><strong>The Compilation Stack: Four IR Layers</strong></h2><p>Between <code>python3 -m maxtext.trainers.pre_train.train</code> and the TPU executing 182 ms training steps, the code passes through four IR layers. MaxText&#8217;s <code>dump_hlo=true</code> and <code>dump_jaxpr=true</code> flags capture all of them.</p><h3><strong>Layer 1: Jaxpr</strong></h3><p>JAX traces the Python training step into Jaxpr &#8211; a functional IR that captures the computation graph. Our train step produces an <strong>18,079-line Jaxpr</strong> with 207 helper function definitions encoding the full MoE routing logic:</p><pre><code><code>train_step &#8594; jit &#8594; TransformerLinenPure.apply &#8594; Decoder.__call__
  &#8594; scan (16 layers) &#8594; GptOssScannableBlock.__call__
    &#8594; GptOssDecoderLayer.__call__
      &#8594; shard_map &#8594; Attention (splash_mha Pallas kernel)
      &#8594; shard_map &#8594; RoutedMoE.__call__
        &#8594; permute &#8594; ragged_all_to_all &#8594; gmm (MegaBlox Pallas kernel) &#8594; unpermute
  &#8594; cross_entropy_with_logits
  &#8594; TrainState.apply_gradients &#8594; adam &#8594; reduce_scatter</code></code></pre><p>The Jaxpr reveals the primitive inventory: <strong>112 </strong><code>dot_general</code> (matrix multiplies), <strong>120 </strong><code>sharding_constraint</code> (SPMD annotations), <strong>32 </strong><code>custom_vjp_call</code> (custom backward passes for GMM and attention), <strong>16 </strong><code>shard_map</code> (manual sharding regions), <strong>12 </strong><code>pallas_call</code> (custom TPU kernels), and <strong>16 </strong><code>scan</code> iterations.</p><p>The mesh configuration appears directly in the Jaxpr:</p><pre><code><code>ctx_mesh=Mesh('diloco': 1, 'data': 1, 'stage': 1, 'fsdp': 2,
              'fsdp_transpose': 1, 'sequence': 1, 'context': 1,
              'context_autoregressive': 1, 'tensor': 1,
              'tensor_transpose': 1, 'tensor_sequence': 1,
              'expert': 2, 'autoregressive': 1)</code></code></pre><p>Thirteen logical mesh axes, but only two are active: <code>fsdp=2</code> and <code>expert=2</code>. XLA&#8217;s SPMD partitioner reads these annotations and generates all the necessary collectives.</p><h3><strong>Layer 2: HLO Before Optimization</strong></h3><p>Jaxpr lowers to HLO (High-Level Operations), XLA&#8217;s main IR. Before optimization, our train step is <strong>14,541 lines of HLO</strong> with <strong>11,207 individual instructions</strong> and <strong>zero fusion blocks</strong>.</p><p>This is the key: <em>everything is separate</em>. Every add, multiply, convert, broadcast, reduce is its own HLO instruction. Let me show you what this actually looks like.</p><p><strong>RMSNorm &#8211; 10 separate instructions for one normalization:</strong></p><pre><code><code>%convert.105 = f32[4,1024,2048] convert(%input_bf16)           // bf16 &#8594; f32
%square.4   = f32[4,1024,2048] multiply(%convert.105, %convert.105)  // x&#178;
%reduce.92  = f32[4,1024]      reduce(%square.4, zero), dims={2}     // &#931;x&#178;
%reshape.1  = f32[4,1024,1]    reshape(%reduce.92)              // broadcast prep
%div.48     = f32[4,1024,1]    divide(%reshape.1, 2048.0)       // mean(x&#178;)
%add.217    = f32[4,1024,1]    add(%div.48, 1e-5)               // + &#949;
%rsqrt.4    = f32[4,1024,1]    rsqrt(%add.217)                  // 1/&#8730;(mean(x&#178;)+&#949;)
%bcast.1    = f32[4,1024,2048] broadcast(%rsqrt.4)              // expand
%mul.128    = f32[4,1024,2048] multiply(%convert.105, %bcast.1) // x * scale
%convert.106 = bf16[4,1024,2048] convert(%mul.128)              // f32 &#8594; bf16</code></code></pre><p>Ten HLO instructions. Each one reads its input from HBM and writes its output back to HBM. For a <code>f32[4,1024,2048]</code> tensor, that&#8217;s 32 MB per read/write. This RMSNorm alone would generate <strong>320 MB of memory traffic</strong> without fusion. There are 32 RMSNorm instances in this model (2 per layer &#215; 16 layers).</p><p><strong>Adam optimizer &#8211; 15 separate instructions per parameter update:</strong></p><pre><code><code>%mul.1073  = f32[2048,8] multiply(grad, 0.1)               // (1-&#946;&#8321;) * grad
%mul.1074  = f32[2048,8] multiply(mu_old, 0.9)             // &#946;&#8321; * &#956;
%add.1253  = f32[2048,8] add(%mul.1073, %mul.1074)         // &#956;_new
%div.614   = f32[2048,8] divide(%add.1253, bias_correction) // &#956;&#770;
%pow.59    = f32[2048,8] multiply(grad, grad)               // grad&#178;
%mul.1155  = f32[2048,8] multiply(%pow.59, 0.05)            // (1-&#946;&#8322;) * grad&#178;
%mul.1156  = f32[2048,8] multiply(nu_old, 0.95)             // &#946;&#8322; * &#957;
%add.1294  = f32[2048,8] add(%mul.1155, %mul.1156)          // &#957;_new
%div.696   = f32[2048,8] divide(%add.1294, bias_correction) // &#957;&#770;
%sqrt.64   = f32[2048,8] sqrt(%div.696)                     // &#8730;&#957;&#770;
%add.1336  = f32[2048,8] add(%sqrt.64, 1e-8)               // &#8730;&#957;&#770; + &#949;
%div.759   = f32[2048,8] divide(%div.614, %add.1336)       // &#956;&#770;/(&#8730;&#957;&#770;+&#949;)
%mul.1219  = f32[2048,8] multiply(param, 0.1)               // weight decay
%add.1377  = f32[2048,8] add(%div.759, %mul.1219)          // update + WD
%add.1422  = f32[2048,8] add(param, lr * %add.1377)        // &#952;_new</code></code></pre><p>This pattern repeats for <strong>40+ parameter tensors</strong>. That&#8217;s 600+ element-wise HLO instructions just for the optimizer.</p><h3><strong>Layer 3: HLO After Optimization</strong></h3><p>After XLA&#8217;s optimization passes, the picture changes dramatically:</p><p>The optimized IR is <em>longer</em> because XLA outlines each fused computation as a separate block, but those 11,207 individual instructions have been compressed into <strong>hundreds of fused kernels</strong>. Each fusion reads its inputs from HBM once, executes all operations in VMEM, and writes outputs once.</p><p>Let me show you what these fusions look like.</p><div><hr></div><h2><strong>Deep Dive: What XLA Actually Fuses</strong></h2><p>Before looking at individual fusions, here&#8217;s the full picture.</p><p>The backward pass dominates: <strong>492 fusions vs. 274 forward</strong>. This isn&#8217;t surprising &#8211; the backward pass through MoE requires separate <code>gmm</code> (input gradient) and <code>tgmm</code> (weight gradient) kernels for each forward GMM, plus chain-rule fusions through every activation, normalization, and residual connection. The 22 backward Pallas kernels vs. 10 forward kernels reflect the same 2:1+ ratio at the custom kernel level.</p><p>The optimizer is compact &#8211; just <strong>41 fusions</strong> &#8211; but each one is a monster. A single AdamW fusion handles gradient scaling, both moment EMAs, bias correction, weight decay, and the parameter update for an entire tensor. The L2 norm reduction for gradient clipping is fused in too.</p><p>Nearly a quarter of all fusions (188 out of 832) contain bf16&#8596;f32 type conversions. The model computes in bf16 for throughput but accumulates in f32 for numerical stability. XLA fuses these conversions into the surrounding operations so they never hit HBM as separate ops.</p><p>Now let me show you what these fusions look like in practice.</p><h3><strong>Fusion 1: Logits Matmul + Softmax Prep (kOutput fusion)</strong></h3><p>This is where XLA goes beyond element-wise fusion. The forward logits computation &#8211; RMSNorm scaling, the <code>[1024,2048] &#215; [2048,32768]</code> matmul into vocabulary space, <em>and</em> the <code>reduce_max</code> for softmax numerical stability &#8211; all in a single kOutput fusion:</p><pre><code><code>%fused_computation.1802 (weights: bf16[2048,32768],
    norm_weight: bf16[2048], rsqrt_scale: f32[1024],
    activation: bf16[1,1024,2048])
    -&gt; (bf16[1024], bf16[1024,32768]) {
  // RMSNorm: scale activation (nested kLoop fusion)
  %normed    = fusion(norm_weight, rsqrt_scale, activation)  // bf16[1024,2048]
  %w_reshaped = fusion(weights)                               // layout bitcast

  // THE MATMUL: [1024,2048] &#215; [2048,32768] &#8594; logits
  %logits    = bf16[1024,32768] convolution(%normed, %w_reshaped),
                 dim_labels=bf_io-&gt;bf                         // logits_dense forward

  // FUSED: reduce_max over vocab dim (softmax numerics)
  %row_max   = bf16[1024] reduce(%logits, -inf), dimensions={1}

  ROOT tuple(%row_max, %logits)
}</code></code></pre><p>The <code>reduce_max</code> consumes logits tiles as the convolution (matmul) produces them &#8211; the full <code>bf16[1024,32768]</code> tensor (64 MB) never needs a separate read pass. In the unfused version, XLA would write the logits to HBM, then read them back for the max. The <code>kind=kOutput</code> annotation means this fusion is anchored on the convolution &#8211; XLA built the fusion outward from the matmul, pulling in both its input preparation (RMSNorm) and its consumer (reduce_max).</p><h3><strong>Fusion 2: Backward Logits Matmul + RMSNorm Gradient (kOutput fusion)</strong></h3><p>The backward pass shows an even deeper kOutput fusion. The gradient flows backward through the logits projection and directly into the RMSNorm gradient &#8211; matmul and normalization backward fused into one kernel:</p><pre><code><code>%fused_computation.1488 (activation: bf16[1,1024,2048],
    norm_weight: bf16[2048], weights: bf16[2048,32768],
    dLogits: bf16[1024,32768], row_max: f32[1024],
    softmax_denom: f32[1024], labels: bf16[1024],
    label_indices: s32[1024], loss_scale: f32[1024])
    -&gt; (f32[1024], bf16[1024,2048]) {
  // Softmax gradient correction (nested kLoop fusion)
  %dLogits_corrected = fusion(dLogits, row_max, softmax_denom,
                              labels, label_indices, loss_scale)  // bf16[1024,32768]

  // THE MATMUL: dLogits &#215; W^T &#8594; dActivation
  %dAct      = bf16[1024,2048] convolution(%dLogits_corrected, weights),
                 dim_labels=bf_oi-&gt;bf                     // backward logits_dense

  // FUSED: RMSNorm backward -- no HBM round-trip after matmul
  %scaled    = bf16[1,1024,2048] multiply(bitcast(%dAct), broadcast(norm_weight))
  %f32_grad  = f32[1,1024,2048] convert(%scaled)         // bf16 &#8594; f32
  %f32_act   = f32[1,1024,2048] convert(activation)      // bf16 &#8594; f32
  %chain     = f32[1,1024,2048] multiply(%f32_act, %f32_grad)  // chain rule
  %grad_sum  = f32[1024] reduce(%chain, 0.0), dimensions={0,2} // RMSNorm weight grad

  ROOT tuple(%grad_sum, %dAct)
}</code></code></pre><p>The <code>[1024,32768] &#215; [32768,2048]</code> backward matmul produces <code>dAct</code>, which is immediately consumed by the RMSNorm gradient chain: scale by norm weight, convert bf16&#8594;f32, elementwise multiply with saved activations, reduce-sum over hidden dim. Seven operations after the matmul, all executing on the convolution&#8217;s output tiles without an HBM round-trip. That&#8217;s <strong>64 MB of avoided intermediate traffic</strong> for the matmul output alone.</p><h3><strong>Fusion 3: Grouped Attention Conv + Residual + RMSNorm Backward (kOutput fusion)</strong></h3><p>The deepest matmul fusion in the model. Inside the decoder backward pass, XLA fuses a 32-head grouped convolution with residual gradient accumulation and the full RMSNorm backward chain:</p><pre><code><code>%fused_computation.372 (saved_activations: bf16[8,1,1024,2048],
    layer_idx: s32[], norm_weight: bf16[2048],
    O_weight: bf16[1,2048,32,64],
    dQ_heads: bf16[1,1024,32,32], dK_heads: bf16[1,1024,32,32],
    residual_grad_1: bf16[1024,2048,1],
    residual_grad_2: bf16[1024,2048,1])
    -&gt; (f32[1024], bf16[1024,2048]) {
  // Extract this layer's activations from scan buffer
  %act_slice  = dynamic-slice(saved_activations, layer_idx, 0, 0, 0)  // bf16[1,1024,2048]
  %f32_act    = f32[1,1024,2048] convert(%act_slice)                   // bf16 &#8594; f32

  // Accumulate two residual gradient streams
  %residual   = bf16[1024,2048] add(residual_grad_1, residual_grad_2)

  // Prepare attention head gradients and weights
  %dHeads     = fusion(dQ_heads, dK_heads)         // bf16[1024,32,64]: merged Q+K grads
  %W_reshaped = fusion(O_weight)                   // bf16[2048,32,64]: reshaped O projection

  // THE MATMUL: 32-head grouped attention projection backward
  %dHidden    = bf16[1024,2048,1] convolution(%dHeads, %W_reshaped),
                  window={size=32}, dim_labels=b0f_o0i-&gt;bf0       // 32 attention heads

  // Add matmul result to residual chain
  %combined   = bf16[1024,2048] add(%residual, bitcast(%dHidden))

  // RMSNorm backward: scale, convert, chain-rule multiply, reduce
  %scaled     = bf16[1,1024,2048] multiply(bitcast(%combined), broadcast(norm_weight))
  %f32_grad   = f32[1,1024,2048] convert(%scaled)                 // bf16 &#8594; f32
  %chain      = f32[1,1024,2048] multiply(%f32_act, %f32_grad)    // chain rule
  %grad_sum   = f32[1024] reduce(%chain, 0.0), dimensions={0,2}   // RMSNorm weight grad

  ROOT tuple(%grad_sum, %combined)
}</code></code></pre><p>Count the operations: dynamic-slice from the scan buffer, bf16&#8594;f32 convert, two residual adds, a 32-head grouped convolution (<code>window={size=32}</code> &#8211; this is how TPU HLO represents grouped matmuls), broadcast, multiply, bf16&#8594;f32 convert, multiply, reduce-sum. <strong>~20 operations across three distinct algorithmic stages</strong> (residual accumulation, attention projection, normalization backward) in a single kernel. The <code>dynamic-slice</code> indexed by <code>layer_idx</code> is particularly notable &#8211; this fusion runs inside the scan while-loop, extracting per-layer activations from the recomputation buffer each iteration.</p><h3><strong>Fusion 4: Full AdamW Update (~40 ops &#8594; 1 kernel)</strong></h3><p>The entire AdamW optimizer step for a single MoE expert weight tensor &#8211; gradient clipping, both moment EMAs, bias correction, weight decay, parameter update, <em>and</em> L2 norm computation &#8211; all in one fusion:</p><pre><code><code>%fused_computation.1373 (weight: f32[8,8,2048,1024],
    lr: f32[], beta1_correction: f32[], beta2_correction: f32[],
    nu_prev: f32[8,8,2048,1024], grad_scale: f32[],
    mu_prev: f32[8,8,2048,1024], is_finite: pred[],
    gradient: f32[8,8,2048,1024])
    -&gt; (f32[], f32[...], f32[...], f32[...], f32[]) {
  // Gradient clipping: zero out if non-finite, else scale
  %scaled_grad = divide(gradient, grad_scale)
  %safe_grad   = select(is_finite, %scaled_grad, gradient)
  // First moment: &#956;_t = 0.9&#183;&#956;_{t-1} + 0.1&#183;g
  %mu_new      = add(multiply(%safe_grad, 0.1), multiply(mu_prev, 0.9))
  // Second moment: &#957;_t = 0.95&#183;&#957;_{t-1} + 0.05&#183;g&#178;
  %grad_sq     = multiply(%safe_grad, %safe_grad)
  %nu_new      = add(multiply(%grad_sq, 0.05), multiply(nu_prev, 0.95))
  // Bias-corrected update: &#956;&#770;/(&#8730;&#957;&#770; + &#949;)
  %nu_hat      = divide(%nu_new, beta2_correction)
  %update      = divide(%mu_new, multiply(beta1_correction, add(sqrt(%nu_hat), 1e-8)))
  // Weight decay + parameter update
  %new_w       = add(weight, multiply(lr, add(%update, multiply(weight, 0.1))))
  // L2 norms for gradient clipping and logging
  %w_norm      = reduce(multiply(%new_w, %new_w), dims={0,1,2,3})
  %g_norm      = reduce(%grad_sq, dims={0,1,2,3})

  ROOT tuple(%w_norm, %new_w, %nu_new, %mu_new, %g_norm)
}</code></code></pre><p>The tensor shape <code>f32[8,8,2048,1024]</code> is an MoE expert weight: <code>[num_experts=8, layers_per_scan=8, hidden=2048, ffn=1024]</code>. That&#8217;s a <strong>512 MiB</strong> tensor. Without fusion, the 15+ intermediate reads and writes would generate <strong>~8 GB of HBM traffic per parameter update</strong>. With fusion, it&#8217;s three reads (weight, gradient, both moments) and three writes (new weight, new &#956;, new &#957;). Two full reductions (weight norm, gradient norm) are computed in the same pass. This pattern repeats for every parameter in the model.</p><div><hr></div><h2><strong>Deep Dive: How MoE Compiles</strong></h2><p>This is the most complex part of the compilation pipeline. The MoE layer transforms from ~500 lines of Python into a choreographed sequence of sorts, collectives, and Pallas kernels that spans all 4 TPU chips simultaneously.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!xYAM!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 424w, /__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 848w, /__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 1272w, /__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!xYAM!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png" width="1440" height="848" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:848,&quot;width&quot;:1440,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:140249,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/189578317?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 424w, /__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 848w, /__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 1272w, /__u/substackcdn.com/image/fetch/$s_!xYAM!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F3e8904b5-eb22-4652-beb8-bd1b482942c1_1440x848.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><h3><strong>The Python: What the Developer Writes</strong></h3><p>The MoE forward pass in MaxText (<code>layers/moe.py</code>) follows this flow:</p><pre><code><code># 1. Gate: project to expert logits
gate_logits = self.gate(inputs)                    # [batch, seq, num_experts]
top_k_weights, top_k_indices = jax.lax.top_k(gate_logits, k=2)
top_k_weights = jax.nn.softmax(top_k_weights)

# 2. Permute: sort tokens by assigned expert
sorted_inputs = _sort_activations(inputs, argsort(experts))
group_sizes = jnp.bincount(experts, length=num_experts)

# 3. Dispatch: send tokens to expert-owning devices
x = jax.lax.ragged_all_to_all(sorted_inputs, offsets, sizes, ...)

# 4. Expert FFN: three grouped matrix multiplies
gate_out = gmm(x, w0, group_sizes, tiling=(512,1024,1024))  # gate proj
up_out   = gmm(x, w1, group_sizes, tiling=(512,1024,1024))  # up proj
hidden   = silu(gate_out) * (up_out + 1)                     # activation
output   = gmm(hidden, wo, group_sizes, tiling=(512,1024,1024))  # down proj

# 5. Combine: send results back and weight-sum
x = jax.lax.ragged_all_to_all(output, ...)         # reverse dispatch
output = unpermute(x, argsort(experts))
result = einsum("BKE,BK-&gt;BE", output, weights)     # weighted combination
</code></code></pre><p>The <code>gmm</code> function wraps a Pallas kernel with a custom VJP: <code>_gmm_fwd</code> calls the kernel for the forward pass, <code>_gmm_bwd</code> calls <code>gmm</code> with <code>transpose_rhs=True</code> for the input gradient and <code>tgmm</code> (a separate kernel) for the weight gradient. Three different tiling configurations optimize each independently.</p><p>The <code>shard_map</code> wrapper around this code gives each device a local view, and <code>ragged_all_to_all</code> &#8211; a JAX primitive &#8211; handles the variable-length token shuffle: each device sends different numbers of tokens to each peer based on the routing decisions.</p><h3><strong>The HLO: What XLA Compiles</strong></h3><p>Here&#8217;s what the MoE pipeline looks like in the optimized HLO. I&#8217;ve annotated each step:</p><p><strong>Step 1 &#8211; Token routing (sort by expert):</strong></p><pre><code><code>%sort.508 = (s32[2048], s32[2048]) sort(%experts, %iota), dimensions={0}, is_stable=true</code></code></pre><p><strong>Step 2 &#8211; Sort activations to match expert ordering:</strong></p><pre><code><code>%gather_custom_fusion.263 = bf16[2048,2048] fusion(%activations, %sort_indices)</code></code></pre><p><strong>Step 3 &#8211; Dispatch tokens via ragged all-to-all:</strong></p><pre><code><code>%ragged_all_to_all.295 = bf16[4096,2,1024] ragged-all-to-all(
    %sorted_tokens,       // bf16[2048,2,1024]: local tokens
    %output_buffer,       // bf16[4096,2,1024]: receive buffer (2x for worst case)
    %send_sizes,          // s32[2]: how many tokens to send to each peer
    %send_offsets,        // s32[2]: where they start in my buffer
    %recv_sizes,          // s32[2]: how many I'll receive from each peer
    %recv_offsets),       // s32[2]: where to put them
    channel_id=1,
    replica_groups={{0,1},{2,3}}</code></code></pre><p>The replica groups <code>{{0,1},{2,3}}</code> show the expert-parallel communication pattern: chips 0&#8596;1 and 2&#8596;3 exchange tokens along the expert axis of the 2&#215;2 mesh.</p><p><strong>Step 4 &#8211; GMM forward (Pallas custom call):</strong></p><pre><code><code>%gmm.1 = bf16[4096,2048] custom-call(
    %group_index,           // s32[]: scalar loop state
    %group_sizes,           // s32[9]: tokens per expert (+padding)
    %group_offsets_lhs,     // s32[15]: tile boundary metadata
    %group_offsets_rhs,     // s32[15]: tile boundary metadata
    %num_actual_groups,     // s32[1]: 8 experts per shard
    %dispatched_tokens,     // bf16[4096,2048]: sorted tokens (LHS)
    %expert_weights),       // bf16[8,2048,2048]: expert matrices (RHS)
    custom_call_target="tpu_custom_call",
    kernel_metadata={"tiling": {"tile_m": 512, "tile_k": 1024, "tile_n": 1024},
                     "cost_estimate": {"flops": "34359738368"}}</code></code></pre><p>Each GMM call does ~34.4 GFLOP &#8211; tiling <code>[512, 1024, 1024]</code> across the M (token), K (input), and N (output) dimensions. The group metadata tells the kernel where each expert&#8217;s tokens start and end in the sorted input.</p><p><strong>Step 5 &#8211; SiLU activation (fused into a single kernel):</strong></p><pre><code><code>%fusion.1232 = bf16[4096,2048] fusion(%gmm.0, %gmm_shuffled, %gmm.1, %weights),
    kind=kLoop, calls=%fused_computation.341   // &#8592; the SiLU fusion shown earlier</code></code></pre><p><strong>Step 6 &#8211; Combine via reverse ragged all-to-all</strong> back to original device mapping.</p><h3><strong>The Pallas Kernel: What Runs on the MXU</strong></h3><p>Inside the <code>tpu_custom_call</code>, the Pallas GMM kernel runs on a 3D grid <code>(tiles_n, num_active_tiles, tiles_k)</code>:</p><pre><code><code>grid=(tiles_n, num_active_tiles, tiles_k)
# tiles_n=2:      N-dimension tiles (parallel)
# num_active_tiles: M-tiles covering all groups (sequential)
# tiles_k=2:      K-dimension accumulation (sequential)</code></code></pre><p>The first dimension is marked <code>"parallel"</code> &#8211; independent axis. The kernel body:</p><ol><li><p><strong>Fetch</strong>: DMA a <code>[512, 1024]</code> tile of sorted tokens from HBM to VMEM</p></li><li><p><strong>Fetch</strong>: DMA the correct expert&#8217;s <code>[1024, 1024]</code> weight tile (selected by <code>group_ids[grid_id]</code>)</p></li><li><p><strong>Accumulate</strong>: <code>acc_scratch += dot(lhs_tile, rhs_tile)</code> in VMEM f32 scratch</p></li><li><p><strong>On last k-tile</strong>: Apply group boundary mask and store result to HBM</p></li></ol><p>The boundary mask handles the case where a single tile spans two experts: only rows belonging to the current expert are written.</p><h3><strong>The Backward Pass: Three Kernels, Not One</strong></h3><p>The backward pass through MoE requires custom VJPs at the <code>ops.py</code> level. For each of the three forward GMM calls:</p><ul><li><p><code>gmm</code> with <code>transpose_rhs=True</code> computes <code>&#8706;L/&#8706;input = grad @ W^T</code> &#8211; the input activation gradient</p></li><li><p><code>tgmm</code> computes <code>&#8706;L/&#8706;W = input^T @ grad</code> &#8211; the weight gradient</p></li></ul><p>The <code>tgmm</code> kernel has a fundamentally different structure: it produces <code>[num_experts, K, N]</code> output, detecting group boundaries to accumulate then store per-expert weight gradients. Its Pallas grid is <code>(tiles_n, tiles_k, num_active_tiles)</code> &#8211; note the reordered axes.</p><p>In the optimized HLO, you can see both:</p><pre><code><code>// Input gradient: gmm with transposed weights
%gmm.10 = bf16[4096,2048] custom-call(..., %expert_weights), target="tpu_custom_call"

// Weight gradient: tgmm (separate kernel body)
%tgmm.2 = bf16[8,2048,2048] custom-call(
    ..., %saved_activations, %upstream_grad), target="tpu_custom_call",
    kernel_metadata={"tiling": {"tile_m": 512, "tile_k": 1024, "tile_n": 1024},
                     "num_actual_groups": 8}</code></code></pre><div><hr></div><h2><strong>Deep Dive: Collectives and Async Overlap</strong></h2><p>The 2&#215;2 mesh creates three distinct communication topologies, all visible in the IR. In total: <strong>104 async all-gathers</strong>, <strong>12 ragged all-to-all</strong>, <strong>10 plain all-to-all</strong>, <strong>8 reduce-scatters</strong>, and <strong>22 all-reduces</strong> &#8211; 156 collective operations in a single training step.</p><h3><strong>Pattern 1: Async Collective Fusions (Compute + Comms Overlap)</strong></h3><p>This is the most sophisticated pattern. XLA doesn&#8217;t just make all-gathers async &#8211; it <strong>fuses them with compute into a single pipelined operation</strong>. While the current layer&#8217;s backward matmul runs, the all-gather for the next layer&#8217;s FSDP-sharded weights is in flight:</p><pre><code><code>%async_collective_fusion.1497 (
    shard: bf16[1,8,2048,1024],     // FSDP-sharded expert weight (half)
    full_weight: bf16[1,8,2048,2048],  // previous all-gather result
    semaphore: s32[2],              // collective sync state
    flags: u32[], u32[], ...,       // S(2) flag registers
    gate_bias: bf16[16],            // router bias
    gate_weight: bf16[1,2048,16],   // router projection
    activation: bf16[1024,2048])    // layer input
    -&gt; (bf16[1,1024,16], ..., bf16[1,8,2048,2048], ...) {

  // ===== COMPUTE: router backward matmul =====
  %dGate   = bf16[1024,16] convolution(activation, gate_weight),
               dim_labels=bf_io-&gt;bf                    // [1024,2048] &#215; [2048,16]
  %dGate   = add(bitcast(%dGate), broadcast(gate_bias))  // + bias

  // ===== COMMS: all-gather for NEXT layer (overlapped) =====
  %gathered = bf16[1,8,2048,2048] all-gather(shard),
               channel_id=38,
               replica_groups=[2,2]&lt;=[2,2]T(1,0),      // FSDP: {0,2},{1,3}
               dimensions={3},                          // double dim 3: 1024 &#8594; 2048
               frontend_attributes={chain_id="4"},      // pipeline ordering
               backend_config={async_collective_fusion_config:
                 {flag_start:"2", flag_end:"8"}}        // semaphore window

  ROOT tuple(%dGate, shard, full_weight, %gathered, semaphore, flags...)
}</code></code></pre><p>The <code>chain_id="4"</code> and <code>flag_start</code>/<code>flag_end</code> annotations reveal XLA&#8217;s <strong>double-buffered all-gather pipeline</strong>: chain IDs sequence the all-gathers across layers, and flag registers manage the handoff between pipeline stages. The semaphore state (<code>s32[2]</code> in memory space <code>S(4)</code>) and flag registers (<code>u32[]</code> in <code>S(2)</code>) are threaded through the tuple output as &#8220;continuation state&#8221; &#8211; passed from one async collective fusion to the next.</p><p>There are <strong>39 async collective fusions</strong> in total, wrapping <strong>104 individual all-gather operations</strong>. The <code>replica_groups=[2,2]&lt;=[2,2]T(1,0)</code> pattern (39 all-gathers) reconstructs FSDP-sharded weights across chips {0,2} and {1,3}. The <code>[1,4]&lt;=[4]</code> pattern (65 all-gathers) gathers across all 4 chips for activations and non-expert parameters.</p><h3><strong>Pattern 2: MoE Ragged All-to-All (Token Dispatch)</strong></h3><p>Expert parallelism uses <code>ragged_all_to_all</code> with <strong>variable length messages</strong> &#8211; each device sends different numbers of tokens to each peer depending on the routing decisions:</p><pre><code><code>%ragged_all_to_all.295 = bf16[4096,2,1024]
    ragged-all-to-all(
        %sorted_tokens,       // bf16[2048,2,1024]: local tokens (input)
        %output_buffer,       // bf16[4096,2,1024]: receive buffer (zeros)
        %send_sizes,          // s32[2]: how many tokens to send each peer
        %recv_sizes,          // s32[2]: how many I'll receive
        %send_offsets,        // s32[2]: where they start in my buffer
        %recv_offsets),       // s32[2]: where to put them
    channel_id=1, replica_groups={{0,1},{2,3}},
    barrier_config={"barrier_type":"CUSTOM","id":"3"}</code></code></pre><p>The replica groups <code>{{0,1},{2,3}}</code> show the expert-parallel topology: chips 0&#8596;1 and 2&#8596;3 exchange tokens along the horizontal axis of the 2&#215;2 mesh. There are <strong>12 ragged all-to-all</strong> operations (4 forward dispatch + 4 forward combine + 4 backward), plus <strong>8 plain all-to-all</strong> exchanging small <code>s32[2,1,1]</code> metadata tensors for token count coordination.</p><p>The output lands in <code>S(1)</code> (VMEM) &#8211; dispatched tokens go straight to on-chip memory, avoiding an HBM round-trip before the GMM kernels consume them.</p><h3><strong>Pattern 3: Batched Gradient Reduce-Scatter</strong></h3><p>After the backward pass, gradients are reduce-scattered back to FSDP-sharded form. XLA <strong>batches 10 parameter gradients from 2 transformer layers</strong> into a single collective:</p><pre><code><code>%all-reduce-scatter.17 (
    Q_grad:    bf16[2048,8,64],    // Q projection (layer A)
    K_grad:    bf16[32,64,2048],   // K projection (layer A)
    V_grad:    bf16[2048,32,64],   // V projection (layer A)
    O_grad:    bf16[2048,8,64],    // output projection (layer A)
    gate_grad: bf16[2048,16],      // router gate (layer A)
    Q_grad_B:  bf16[2048,8,64],    // Q projection (layer B)
    K_grad_B:  bf16[32,64,2048],   // K projection (layer B)
    V_grad_B:  bf16[2048,32,64],   // V projection (layer B)
    O_grad_B:  bf16[2048,8,64],    // output projection (layer B)
    gate_grad_B: bf16[2048,16])    // router gate (layer B)
    &#8594; (bf16[512,8,64], bf16[32,64,512], ...) {
  // All-reduce across all 4 devices, then scatter via partition-id indexing
  %all-reduce = all-reduce(inputs...), replica_groups={{0,1,2,3}}
  // Each device slices its 1/4 shard: dynamic-slice(..., partition_id * 512, ...)
}</code></code></pre><p>Ten gradient tensors, one collective launch. Each output is scattered from 2048 down to 512 along the FSDP dimension via <code>dynamic-slice</code> indexed by <code>partition-id</code>. For MoE expert weights, separate reduce-scatters operate along the FSDP axis only (<code>replica_groups={{0,2},{1,3}}</code>).</p><h3><strong>The Scheduling Trade-Off: Coalescing vs. Overlap</strong></h3><p>These patterns reveal a tension in the compiler&#8217;s scheduling strategy. <strong>Coalescing</strong> wants to batch collectives together &#8211; the 10 input reduce-scatter amortizes launch overhead by waiting until all 10 gradient tensors from 2 layers are ready, then issuing one large collective instead of 10 small ones. <strong>Overlap</strong> wants to start collectives as early as possible &#8211; the async collective fusions issue all-gathers for layer N+1 while layer N&#8217;s compute is still running, hiding latency behind useful work.</p><p>XLA makes different choices for different collective types. All-gathers are <strong>overlap optimized</strong>: each weight gather launches individually inside its own async collective fusion, pipelined with <code>chain_id</code> ordering, because waiting to batch them would stall the compute pipeline. Reduce-scatters are <strong>coalescing optimized</strong>: batching 10 gradients into one call is worth the delay because the gradients aren&#8217;t needed until the optimizer runs, so there&#8217;s no compute to overlap them with anyway. Ragged all-to-alls are <strong>neither</strong> &#8211; they&#8217;re synchronous barriers because the MoE routing decisions create data-dependent communication patterns that can&#8217;t be predicted at compile time.</p><p>The scheduler is balancing three constraints: minimize collective launch overhead (coalesce), maximize compute-comms overlap (pipeline early), and respect data dependencies (barrier where required). The fact that it makes different choices for all-gather vs. reduce-scatter vs. ragged all-to-all shows this isn&#8217;t a one size fits all heuristic &#8211; it&#8217;s a per collective scheduling decision informed by the dependency graph.</p><div><hr></div><h2><strong>What This Would Take to Hand-Fuse</strong></h2><p><strong>This is the infrastructure teams are reimplementing when they rebuild training stacks from scratch.</strong> Let me quantify it.</p><p><strong>Kernels you&#8217;d need to write:</strong> The hand-written kernel count alone is significant: 9 for MoE (gate projection, top-k selection, argsort, bincount, GMM forward, gated activation, GMM backward, TGMM for weight gradients, and unsort/combine) and 3 for attention (forward, dQ backward, dKV backward). Each needs its own tiling strategy, VMEM allocation, and DMA scheduling. But that's just the custom kernels &#8212; it doesn't include RMSNorm, softmax, embedding, cross-entropy, residual connections, or the AdamW optimizer. On TPU, XLA generates fused kernels for all of those automatically. On a from scratch PyTorch stack, you're writing or tuning those too.</p><p><strong>Memory management decisions:</strong> ~8 critical choices. VMEM scratch allocation for GMM accumulators. HBM buffer sizing for ragged all-to-all (worst case capacity). Double-buffering strategy for weight prefetch. Activation checkpointing (MaxText uses <code>remat_policy=full</code> &#8211; all activations recomputed in backward). Padding strategies for tile boundaries. When to materialize vs. recompute intermediate results.</p><p><strong>Communication patterns:</strong> ~6 distinct collective types. Ragged all-to-all for MoE dispatch (variable-length). Ragged all-to-all for MoE combine (reverse direction). FSDP all-gather for weight reconstruction. All-gather for routing metadata. Reduce-scatter for gradient sharding. All-reduce for loss aggregation. Each one needs to be overlapped with compute for efficiency.</p><p><strong>Custom backward passes:</strong> The MoE routing sort requires a custom VJP (<code>_sort_activations_custom_bwd</code>) because JAX&#8217;s automatic backward for indexing is inefficient. The GMM needs separate forward/backward kernel implementations. The attention uses three different Pallas kernels for forward, dQ, and dKV.</p><p><strong>Interacting parallelism dimensions:</strong> FSDP weight sharding, expert partitioning, and data parallelism interact in subtle ways. The <code>weight_gather_axes</code> logic in <code>moe.py</code> handles cases where weights are sharded across FSDP and must be all-gathered before GMM but reduce-scattered in backward.</p><p><strong>The MaxText/JAX/XLA stack automates all of this.</strong> The developer writes ~500 lines of high-level Python (the <code>RoutedMoE</code> class) and the GMM kernel interface (282 lines). The compiler generates the remaining equivalent of ~50,000+ lines of kernel code, memory management, and communication scheduling.</p><p><strong>The numbers tell the story:</strong> 11,207 individual HLO instructions compressed into 887 fused kernels. 39 async collective fusions containing 104 all-gathers with compute overlap. 12 ragged all-to-all collectives. 156 total collective operations. 8 Splash Attention calls and 24 GMM calls &#8211; the only handwritten code in the entire pipeline. <strong>XLA generated </strong><em><strong>everything else</strong></em><strong>.</strong></p><div><hr></div><h2><strong>What This Means</strong></h2><p><em><strong>This entire pretraining run &#8212; MoE architecture, Pallas kernels, XLA compilation, multi-axis SPMD sharding, async collective overlap &#8212; is open source.</strong></em> Every file is at <a href="https://github.com/AI-Hypercomputer/maxtext">AI-Hypercomputer/maxtext</a>. The model definition: 289 lines of Python. The MoE layer: layers/moe.py. The MegaBlox kernel: 282 lines.</p><p>Out of the 22,636 lines of optimized HLO, only the Splash Attention and MegaBlox GMM calls are handwritten kernels. Everything else &#8212; 887 fused kernels, 39 async collective fusions wrapping 104 all-gathers, ~960 prefetch operations across 4 memory spaces, the entire SPMD partitioning &#8212; is generated by XLA from high level JAX code.</p><p><em><strong>This is what I keep watching teams spend months rebuilding by hand.</strong></em> FSDP overlap, kernel fusion, MoE routing, collective scheduling &#8212; the compiler already does it, and it&#8217;s already open source. I&#8217;m not claiming a frontier lab can take MaxText off the shelf and ship tomorrow. But the distance between this codebase and a production training job is <em><strong>months shorter</strong></em> than starting from scratch on a custom PyTorch stack &#8212; <strong>and the MFU you get out of the box is likely higher than what most teams achieve after those months of work</strong>.</p><p><em><strong>MaxText isn&#8217;t a ceiling &#8212; it&#8217;s a floor, and the floor is already really good.</strong></em> If you&#8217;re a frontier lab, you could rent V7 pods, fork MaxText, and have an MoE training loop running at high MFU <strong>before your PyTorch team finishes writing their first custom kernel.</strong> And if XLA&#8217;s codegen isn&#8217;t enough for a specific op, you write a Pallas kernel for that op &#8212; the same way Splash Attention and MegaBlox already exist in the stack. <strong>You&#8217;re not replacing the compiler; you&#8217;re surgically overriding it where it matters.</strong></p><p>The architecture, the kernels, the compiler flags, the SPMD partitioning &#8212; it&#8217;s all the same code. <strong>The only difference between this experiment and a production training run is the number of chips.</strong></p><h2><strong>Reproduce This</strong></h2><pre><code><code># Clone and install (requires Python 3.12+)
git clone https://github.com/AI-Hypercomputer/maxtext.git
cd maxtext &amp;&amp; pip install -e ".[tpu]"

# Set dump flags
export DECOUPLE_GCLOUD=TRUE
export XLA_FLAGS="--xla_dump_to=/tmp/xla_dump \
  --xla_dump_hlo_module_re=jit_train_step \
  --xla_dump_hlo_as_text --xla_dump_hlo_as_proto"
export LIBTPU_INIT_ARGS=" \
  --xla_tpu_enable_async_collective_fusion=true \
  --xla_tpu_enable_async_collective_fusion_fuse_all_gather=true \
  --xla_tpu_enable_async_collective_fusion_multiple_steps=true \
  --xla_enable_async_all_gather=true \
  --xla_tpu_overlap_compute_collective_tc=true \
  --xla_tpu_scoped_vmem_limit_kib=98304 \
  --xla_tpu_enable_data_parallel_all_reduce_opt=true \
  --xla_tpu_data_parallel_opt_different_sized_ops=true"

# Train (adjust for your chip count)
python3 -m maxtext.trainers.pre_train.train \
  src/maxtext/configs/base.yml \
  model_name="gpt-oss-20b" override_model_config=true \
  dataset_type=synthetic steps=200 \
  base_emb_dim=2048 base_num_decoder_layers=16 \
  base_num_query_heads=32 base_num_kv_heads=8 head_dim=64 \
  base_mlp_dim=2048 base_moe_mlp_dim=2048 \
  num_experts=16 num_experts_per_tok=2 \
  megablox=true sparse_matmul=true capacity_factor=-1.0 \
  per_device_batch_size=1 max_target_length=1024 vocab_size=32768 \
  attention='flash' sa_block_q=512 \
  ici_fsdp_parallelism=2 ici_expert_parallelism=2 \
  remat_policy=full enable_checkpointing=false reuse_example_batch=1 \
  base_output_directory=/tmp/maxtext_output run_name=gpt_oss_3b \
  dump_hlo=true dump_hlo_local_dir=/tmp/hlo_dumps \
  dump_jaxpr=true dump_jaxpr_local_dir=/tmp/jaxpr_dumps \
  gcs_metrics=false

# Examine the IR
ls /tmp/xla_dump/  # HLO before/after optimization
wc -l /tmp/xla_dump/*after_optimizations.txt
grep -c 'fusion' /tmp/xla_dump/*after_optimizations.txt</code></code></pre><p>The IR dumps, training scripts, and full logs from this experiment are at <a href="https://github.com/patrick-toulme/justabyte">https://github.com/patrick-toulme/justabyte</a>.</p><div><hr></div><p><em>If you found this useful, subscribe to <a href="/__u/patricktoulme.substack.com/">Just a Byte</a> for more deep dives into AI compilers, silicon, and systems.</em></p><p><em>Connect with me on <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">LinkedIn</a>.</em></p>]]></content:encoded></item><item><title><![CDATA[My Custom TPU Now Runs Llama — One Giant Fused Megakernel]]></title><description><![CDATA[No CUDA. No handwritten kernels. The entire model compiled into a single binary.]]></description><link>https://patricktoulme.substack.com/p/my-custom-tpu-now-runs-llama-one</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/my-custom-tpu-now-runs-llama-one</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sun, 15 Feb 2026 19:41:09 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!f8bB!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p><strong>One Model. One Giant Megakernel. Zero handwritten code.</strong></p><p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><p>My fully custom TPU now runs a Llama model. If you haven&#8217;t seen the original post, <a href="/__u/patricktoulme.substack.com/p/i-built-a-tpu-from-scratch-rtl-mlir">I built a TPU from scratch </a>&#8212; custom Verilog RTL, custom ISA, custom compiler &#8212; with a full PJRT backend so JAX treats it as a real device. The compiler takes the entire model and fuses it into a single megakernel binary. No op-by-op dispatch. No handwritten kernels. Pure compiler codegen.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!f8bB!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 424w, /__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 848w, /__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 1272w, /__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!f8bB!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png" width="1456" height="642" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:642,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:448001,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:false,&quot;topImage&quot;:true,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/188057660?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 424w, /__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 848w, /__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 1272w, /__u/substackcdn.com/image/fetch/$s_!f8bB!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F078086e3-e02c-4476-be7e-1615d154e562_2684x1184.png 1456w" sizes="100vw" fetchpriority="high"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p><strong>Full Video of JAX Execution on X: </strong></p><div class="twitter-embed" data-attrs="{&quot;url&quot;:&quot;https://x.com/PatrickToulme/status/2023095643699335219?s=20&quot;,&quot;full_text&quot;:&quot;One Model. One Giant Megakernel. Zero handwritten code.\n\nMy fully custom TPU now runs a Llama model.\n\nThe compiler takes the entire model &#8212; attention, FFN, norms, everything &#8212; and fuses it into a single megakernel binary. No op-by-op dispatch. No kernel launch overhead. One &quot;,&quot;username&quot;:&quot;PatrickToulme&quot;,&quot;name&quot;:&quot;Patrick C Toulme&quot;,&quot;profile_image_url&quot;:&quot;https://pbs.substack.com/profile_images/2013003105390981120/ydKai4dt_normal.jpg&quot;,&quot;date&quot;:&quot;2026-02-15T18:02:43.000Z&quot;,&quot;photos&quot;:[{&quot;img_url&quot;:&quot;https://substackcdn.com/image/upload/w_1028,c_limit,q_auto:best/l_twitter_play_button_rvaygk,w_88/k05stf1svut0ecxdrjmt&quot;,&quot;link_url&quot;:&quot;https://t.co/Km2CEpd00Q&quot;}],&quot;quoted_tweet&quot;:{},&quot;reply_count&quot;:13,&quot;retweet_count&quot;:9,&quot;like_count&quot;:98,&quot;impression_count&quot;:3835,&quot;expanded_url&quot;:null,&quot;video_url&quot;:&quot;https://video.twimg.com/amplify_video/2023094710466912256/vid/avc1/1582x720/XcLK3ABxfEjmD3TF.mp4&quot;,&quot;video_preview_media_key&quot;:null,&quot;belowTheFold&quot;:false}" data-component-name="Twitter2ToDOM"></div><p>Getting here wasn&#8217;t straightforward. <strong>The biggest challenge was accuracy.</strong> When your logits don&#8217;t match the reference implementation, the bug could be anywhere &#8212; the RTL, the compiler, the register allocation, the instruction scheduling. Debugging this through waveform traces alone was brutal, so <strong>I built an ISA simulator in Python.</strong> It executes the same binary the hardware runs, instruction by instruction, which made it far easier to isolate where values started diverging.</p><p>Once accuracy was solid, I shifted to performance. The main optimizations were <strong>DMA/compute overlap</strong> so the machine isn&#8217;t stalling on memory fetches, and enhancements to the VLIW scheduler to pack instructions more tightly. The result is a pipeline that actually keeps the functional units busy instead of sitting idle between operations.</p><p>Same stack as before: JAX &#8594; HLO &#8594; MLIR &#8594; ASM &#8594; VLIW &#8594; Binary. Still no CUDA. Still running on my custom Verilog RTL.</p><p><strong>Questions? Comments?</strong></p><p>Reach out on X or Linkedin:</p><p>X: <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><p>Linkedin: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p><p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p>]]></content:encoded></item><item><title><![CDATA[I Built a TPU from scratch – RTL, MLIR compiler, PJRT runtime, runs JAX]]></title><description><![CDATA[No CUDA. No Kernels. 100% compiler codegeneration.]]></description><link>https://patricktoulme.substack.com/p/i-built-a-tpu-from-scratch-rtl-mlir</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/i-built-a-tpu-from-scratch-rtl-mlir</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sat, 31 Jan 2026 21:45:55 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!ZgYC!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><p>I built a TPU from scratch. Custom Verilog RTL. JAX&#8594;HLO&#8594;MLIR&#8594;ASM&#8594;VLIW&#8594;Binary. Real PJRT backend. Runs on Verilog simulator. No CUDA. No handwritten kernels. Pure compiler codegen. Below is a transformer layer compiling and executing from JAX.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!ZgYC!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 424w, /__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 848w, /__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 1272w, /__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!ZgYC!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp" width="1456" height="806" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:806,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:77710,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/webp&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:false,&quot;topImage&quot;:true,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/186447871?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 424w, /__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 848w, /__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 1272w, /__u/substackcdn.com/image/fetch/$s_!ZgYC!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F067a7697-9da1-4be8-b83a-640f726524e3_1467x812.webp 1456w" sizes="100vw" fetchpriority="high"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>Full Video of JAX Execution on X: <a href="https://x.com/PatrickToulme/status/2017706776288719194?s=20">https://x.com/PatrickToulme/status/2017706776288719194?s=20</a></p><div class="twitter-embed" data-attrs="{&quot;url&quot;:&quot;https://x.com/PatrickToulme/status/2017706776288719194?s=20&quot;,&quot;full_text&quot;:&quot;I built a TPU from scratch. Custom Verilog RTL. JAX&#8594;HLO&#8594;MLIR&#8594;ASM&#8594;VLIW&#8594;Binary. Real PJRT backend. Runs on Verilog simulator. No CUDA. No handwritten kernels. Pure compiler codegen. Below is a transformer layer compiling and executing from JAX. &quot;,&quot;username&quot;:&quot;PatrickToulme&quot;,&quot;name&quot;:&quot;Patrick C Toulme&quot;,&quot;profile_image_url&quot;:&quot;https://pbs.substack.com/profile_images/2013003105390981120/ydKai4dt_normal.jpg&quot;,&quot;date&quot;:&quot;2026-01-31T21:09:17.000Z&quot;,&quot;photos&quot;:[{&quot;img_url&quot;:&quot;https://substackcdn.com/image/upload/w_1028,c_limit,q_auto:best/l_twitter_play_button_rvaygk,w_88/wcnln0azsicls5yitj1n&quot;,&quot;link_url&quot;:&quot;https://t.co/xe2VS1FDXR&quot;}],&quot;quoted_tweet&quot;:{},&quot;reply_count&quot;:1,&quot;retweet_count&quot;:0,&quot;like_count&quot;:16,&quot;impression_count&quot;:185,&quot;expanded_url&quot;:null,&quot;video_url&quot;:&quot;https://video.twimg.com/amplify_video/2017706528094969857/vid/avc1/1302x720/1cPO-YtttGjYXhg2.mp4&quot;,&quot;video_preview_media_key&quot;:null,&quot;belowTheFold&quot;:false}" data-component-name="Twitter2ToDOM"></div><p></p><p>Questions? Comments?</p><p>Reach out on X or Linkedin:</p><p>X: <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><p>Linkedin: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p>]]></content:encoded></item><item><title><![CDATA[CuTile on Blackwell: NVIDIA's Compiler Moat Is Already Built ]]></title><description><![CDATA[Tracing a fused MoE kernel from CuTile IR to PTX/SASS on Blackwell.]]></description><link>https://patricktoulme.substack.com/p/cutile-on-blackwell-nvidias-compiler</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/cutile-on-blackwell-nvidias-compiler</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sun, 25 Jan 2026 21:47:03 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!8Ogn!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p><strong>Full Code / IR Dumps</strong>: <a href="https://github.com/patrick-toulme/justabyte/tree/main/cutile_blackwell_post">https://github.com/patrick-toulme/justabyte/tree/main/cutile_blackwell_post</a></p><p class="button-wrapper" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe now&quot;,&quot;action&quot;:null,&quot;class&quot;:null}" data-component-name="ButtonCreateButton"><a class="button primary" href="/__u/patricktoulme.substack.com/subscribe"><span>Subscribe now</span></a></p><p><em><strong>Connect on LinkedIn:</strong> <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></em></p><p><em><strong>Follow on X:</strong> <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></em></p><h2>Motivation</h2><p>86 lines of Python. 180KB of shared memory. 46 tcgen05 instructions. A contention handling retry loop with nanosleep backoff.</p><p><strong>NVIDIA is building a compiler moat, and making a play for the GPU kernel DSL market.</strong></p><p>CuTile compiled a Mixture of Experts kernel for Blackwell, and the gap between what I wrote and what the compiler generated is the evidence. Last post, I traced Pallas through the TPU compiler stack. Pallas exists because XLA can&#8217;t rewrite your algorithm &#8212; it optimizes what you give it, but it can&#8217;t infer that you never need to materialize the full attention matrix. <strong>CuTile is NVIDIA&#8217;s answer</strong>: a tile level DSL that lets you express algorithms the compiler can&#8217;t find automatically, while hiding the thread level complexity that makes GPU programming painful.</p><p>The difference is what CuTile targets. Blackwell&#8217;s tcgen05 instruction family is publicly documented. Triton, Pallas, and CUTLASS all support it to some degree. But high performance requires <strong>patterns that are hard to discover</strong> &#8212; single thread MMA issuance via leader election, multi barrier pipelining, TMEM allocation with contention handling (retry/backoff). CuTile generates these automatically. Triton built Gluon as an escape hatch because the high level abstractions couldn&#8217;t find the right codegen. <strong>CuTile doesn&#8217;t need an escape hatch</strong>.</p><div class="captioned-image-container"><figure><a class="image-link image2" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!6fsO!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 424w, /__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 848w, /__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 1272w, /__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!6fsO!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png" width="1200" height="254" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/b083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:254,&quot;width&quot;:1200,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:1221499,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/185139919?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 424w, /__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 848w, /__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 1272w, /__u/substackcdn.com/image/fetch/$s_!6fsO!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb083b40a-bde8-44ed-b21b-d27bc379841b_1200x254.png 1456w" sizes="100vw" loading="lazy"></picture><div></div></div></a></figure></div><p>This post traces a real MoE kernel through the compilation stack: Python &#8594; CuTile IR &#8594; [MLIR black box] &#8594; PTX &#8594; SASS. The compiler is a black box, but the outputs tell you everything.</p><div><hr></div><h2>Setup</h2><p>All experiments were performed on a <strong>B200 GPU</strong> rented from Lambda Labs. The B200 is NVIDIA&#8217;s Blackwell architecture (sm_100a), requiring:</p><ul><li><p><strong>CUDA 13.1+</strong> (CuTile is not available on earlier versions)</p></li><li><p><strong>NVIDIA Driver 580+</strong> with <strong>open kernel modules</strong> (<code>nvidia-driver-580-open</code>)</p></li><li><p><strong>CuTile Python package</strong> from NVIDIA (<a href="https://github.com/NVIDIA/cutile-python">github source</a>)</p></li></ul><h3>Extracting Compilation Artifacts</h3><p>For the CuTile IR, I used the library&#8217;s built in <code>print_ir()</code> functionality. For PTX and SASS, I extracted them from the compiled <code>.cubin</code> files using <code>cuobjdump</code>.</p><p>The MLIR intermediate passes are <strong>not publicly accessible</strong> &#8212; they require an internal extension (<code>cuda.tile_internal._internal_cext</code>) that isn&#8217;t available with the public release. I extracted the pass names by running <code>strings</code> on the <code>tileiras</code> compiler binary.</p><div><hr></div><h2>The Kernel: NVIDIA&#8217;s Fused MoE</h2><p><a href="https://github.com/NVIDIA/cutile-python/blob/main/samples/MoE.py">The kernel comes from NVIDIA&#8217;s CuTile examples</a> &#8212; a production quality Mixture of Experts implementation. Here&#8217;s the core computation:</p><pre><code>@ct.kernel
def fused_moe_kernel(
    A,                     # Input tokens, shape (batch, K)
    B,                     # Expert weights, shape (num_experts, N, K)
    C,                     # Output tensor, shape (num_tokens * topk, N)
    topk_weights,          # Router weights for each token-expert pair
    sorted_token_ids,      # Token indices sorted by expert assignment
    sorted_expert_ids,     # Expert index for each TILE_M
    num_token_replicas: int,
    mul_routed_weight: ConstBool,
    TILE_M: ConstInt,
    TILE_N: ConstInt,
    TILE_K: ConstInt,
):
    M = sorted_token_ids.shape[0]
    N = B.shape[1]
    K = B.shape[2]

    GROUP_SIZE_M = 8
    bid_m, bid_n = swizzle_2d(M, N, TILE_M, TILE_N, GROUP_SIZE_M)

    # Gather token indices for this block
    token_id_indices = bid_m * TILE_M + ct.arange(TILE_M, dtype=ct.int32)
    token_ids = ct.gather(sorted_token_ids, token_id_indices)
    a_row_indices = token_ids // num_token_replicas

    # Each TILE_M block processes one expert
    expert_id = ct.load(sorted_expert_ids, index=bid_m, shape=())

    # Main matmul loop
    accumulator = ct.full((TILE_M, TILE_N), 0.0, dtype=ct.float32)
    for k in range(0, ct.cdiv(K, TILE_K)):
        a_col_indices = k * TILE_K + ct.arange(TILE_K, dtype=ct.int32)
        a = ct.gather(A, (a_row_indices[:, None], a_col_indices[None, :]))

        b = ct.load(B, (expert_id, k, bid_n), shape=(1, TILE_K, TILE_N),
                    order=(0, 2, 1), padding_mode=ct.PaddingMode.ZERO)
        b = b.reshape((TILE_K, TILE_N))

        accumulator = ct.mma(a, b, accumulator)

    if mul_routed_weight:
        moe_weight = ct.gather(topk_weights, token_ids)
        accumulator = accumulator * moe_weight[:, None]

    # Scatter results back
    c_col_indices = bid_n * TILE_N + ct.arange(TILE_N, dtype=ct.int32)
    accumulator = ct.astype(accumulator, C.dtype)
    ct.scatter(C, (token_ids[:, None], c_col_indices[None, :]), accumulator)</code></pre><h3>What&#8217;s in the Kernel</h3><ul><li><p><strong>Tile level operations</strong>: <code>ct.gather</code>, <code>ct.load</code>, <code>ct.mma</code>, <code>ct.scatter</code></p></li><li><p><strong>2D block swizzling</strong> for L2 cache locality</p></li><li><p><strong>Irregular memory access</strong> patterns (gather/scatter for token routing)</p></li><li><p><strong>Compile time constants</strong> (<code>ConstInt</code>, <code>ConstBool</code>) for specialization</p></li></ul><h3>What&#8217;s NOT in the Kernel</h3><ul><li><p>Thread indices (<code>threadIdx.x</code>)</p></li><li><p>Shared memory allocation</p></li><li><p>Synchronization barriers</p></li><li><p>Async copy operations</p></li><li><p>Register allocation hints</p></li><li><p>The <code>tcgen05</code> instruction family</p></li></ul><p>The kernel is ~86 lines of tile level Python. The compiler generates ~1,900 lines of PTX.</p><div><hr></div><h2>CuTile IR: The First Layer</h2><p>CuTile traces the kernel into a typed intermediate representation. Here&#8217;s the full IR for <code>fused_moe_kernel</code>:</p><p><a href="https://github.com/patrick-toulme/justabyte/blob/main/cutile_blackwell_post/moe_dumps/01_cutile_ir/cutile_ir_output.txt">CuTile IR</a></p><h3>Tensor Views and Memory Layout</h3><p>The IR creates explicit tensor views with stride information:</p><pre><code>B{$177, $179, $181, $183, $185, $187, B_6}: Array[bfloat16,(?,?,?):(?,?,1)] =
    make_tensor_view(base_ptr=$177, shape=($179, $181, $183), dynamic_strides=($185, $187))

sorted_expert_ids{$201, $203, sorted_expert_ids_2}: Array[int32,(?):(1)] =
    make_tensor_view(base_ptr=$201, shape=($203), dynamic_strides=())</code></pre><p>The <code>(?,?,?):(?,?,1)</code> notation encodes shape and strides &#8212; the last dimension has stride 1 (contiguous), while outer dimensions have dynamic strides.</p><h3>The Swizzle Computation</h3><p>The 2D swizzle becomes explicit integer arithmetic:</p><pre><code>$261: int32 = tile_bid(axis=0)
$307: Tile[int32,()] = raw_binary_arith(lhs=$200, rhs=$357, fn="cdiv", ...)
$312: Tile[int32,()] = raw_binary_arith(lhs=$181, rhs=$371, fn="cdiv", ...)
$315: Tile[int32,()] = raw_binary_arith(lhs=$380, rhs=$312, fn="mul", ...)
$318: Tile[int32,()] = raw_binary_arith(lhs=$389, rhs=$315, fn="floordiv", ...)
...
$331: Tile[int32,()] = raw_where(cond=$431, x=$432, y=$425)
$332: Tile[int32,()] = raw_binary_arith(lhs=$321, rhs=$331, fn="add", ...)</code></pre><p>The <code>swizzle_2d()</code> helper expands into ~20 IR operations for group-based block assignment.</p><h3>The Token System</h3><p>CuTile uses tokens for memory ordering &#8212; similar to Pallas&#8217;s implicit ordering through Refs:</p><pre><code>$40: Tile[int32,(128)], $token.0: Token =
    load_pointer_tko(pointer=$546, mask=$543, padding_value=$548, token=$token, latency=None)

$578: Tile[int32,(1)], $token.2: Token =
    tile_load_token_ordered(array=sorted_expert_ids{...}, index=($332), token=$token,
                            order=(0,), padding_mode=PaddingMode.UNDETERMINED, latency=None)</code></pre><p>The <code>_tko</code> suffix means &#8220;token-keyed ordering&#8221; &#8212; loads and stores carry tokens that establish memory dependencies.</p><h3>The Main Loop</h3><p>The k-loop becomes explicit with accumulator threading:</p><pre><code>$792: Tile[float32,(128,128)] = for k in range($60, $65, $624)
    (with accumulator.0: Tile[float32,(128,128)] = $58)
do (k: int32, accumulator.0: Tile[float32,(128,128)]):
    $69: int32 = raw_binary_arith(lhs=k, rhs=$169, fn="mul", ...)
    ...
    $91: Tile[bfloat16,(128,64)], $token.8: Token =
        load_pointer_tko(pointer=$717, mask=$713, padding_value=$720, token=$token, latency=None)

    $108: Tile[bfloat16,(1,64,128)], $token.10: Token =
        tile_load_token_ordered(array=B{...}, index=($49, k, $337), token=$token,
                                order=(0, 2, 1), padding_mode=PaddingMode.ZERO, latency=None)

    $113: Tile[bfloat16,(64,128)] = tile_reshape(x=$108)
    $119: Tile[float32,(128,128)] = tile_mma(x=$91, y=$113, acc=accumulator.0)
    continue $119</code></pre><p>The <code>tile_mma</code> operation is the matrix multiply-accumulate &#8212; this single IR operation will expand into <strong>Blackwell's </strong><code>tcgen05.mma</code><strong> instructions.</strong></p><h2>The Compiler Wall</h2><p>Everything after CuTile IR is a <em><strong>black box</strong></em>. The MLIR passes, the layout transformations, the async scheduling  etc. are not visible.</p><h3>The MLIR Black Box: What I Found</h3><p>I ran <code>strings</code> on the <code>tileiras</code> compiler binary <code>/usr/local/cuda-13.1/bin/tileiras</code>) and extracted the pass pipeline. Here are the 30+ MLIR passes that transform CuTile IR into PTX:</p><h3>Layout Assignment</h3><pre><code>tileas-assign-dot-layouts         # Assign layouts for MMA operations
tileas-assign-load-store-layouts  # Assign layouts for memory operations
tileas-assign-pipeline-layouts    # Assign layouts for async pipeline</code></pre><h3>Kernel Planning</h3><pre><code>tileas-plan-cta                   # Plan CTA (thread block) structure
tileas-generate-schedule          # Generate execution schedule</code></pre><h3>Optimization Passes</h3><pre><code>tileas-optimize-alloc-tensor         # Optimize tensor allocations
tileas-optimize-dot-accumulation     # Optimize MMA accumulation
tileas-optimize-reduce               # Optimize reduction operations
tileas-refine-atom-by-resource       # Refine atomics based on resources</code></pre><h3>Pipelining</h3><pre><code>tileas-dynamic-persistent         # Dynamic persistent thread scheduling
tileas-materialize-async          # Materialize async operations
tileas-materialize-schedule       # Materialize the schedule
tileas-unspecialized-pipeline     # Create unspecialized pipeline
tileas-prepare-for-scheduling     # Prepare IR for scheduling</code></pre><h3>Loop Transformations</h3><pre><code>tileas-slice-and-fuse                  # Slice and fuse loops
tileas-slicing                         # Loop slicing
tileas-unroll-register-loops           # Unroll register-bound loops</code></pre><h3>Final Lowering</h3><pre><code>tileas-attach-tma-desc-args      # Attach TMA descriptor arguments
tileas-insert-OCG-knobs          # Insert optimization hints for backend
convert-nv-tileas-to-llvm        # Convert to LLVM IR</code></pre><h3>The Dialect Hierarchy</h3><p>Based on the pass names, this is what I think the dialect progression looks like:</p><pre><code>ccuda_tile (Python IR)
    &#8595;
nv_tileaa (NVIDIA address-annotated)
    &#8595;
nv_tileas (NVIDIA address-specialized)
    &#8595;
LLVM IR
    &#8595;
PTX
    &#8595;
SASS</code></pre><p>The <code>cuda_tile</code> dialect has 100+ operations. The <code>nv_tileaa</code> and <code>nv_tileas</code> dialects are NVIDIA internal. We can see the inputs and outputs, but the transformations are proprietary.</p><h3>What the Passes Reveal</h3><p>From error messages in the binary:</p><pre><code>"TileASRemoveBufferAliasPass failed to converge"
"failed to assign pipelineLayout in assignPipelineLayouts pass"
"Materialize the pipeline schedule to generate warp-specialized or unspecialized IR"</code></pre><p>The compiler can generate <strong>warp-specialized IR</strong> &#8212; different code paths for different warps within a thread block. This is how it implements the leader election pattern we'll see in the PTX.</p><h2>PTX: Where Blackwell Shows Up</h2><p>The compiled PTX reveals what the MLIR passes generated. Let&#8217;s trace each kernel operation to its PTX implementation.</p><h3><strong>Kernel Header</strong></h3><pre><code><code>.version 9.1
.target sm_100a
.address_size 64

.visible .entry fused_moe_kernel(
    .param .u64 .ptr .global .align 1 fused_moe_kernel_param_0,
    .param .u32 fused_moe_kernel_param_1,
    ...
)
.reqntid 256
.minnctapersm 1
.reg .pred  %p&lt;4315&gt;;
.reg .b16   %rs&lt;5&gt;;
.reg .b32   %r&lt;1178&gt;;
.reg .b64   %rd&lt;1202&gt;;
.shared .align 128 .b8 global_smem[180424];
</code></code></pre><p>Key details:</p><ul><li><p><strong>PTX 9.1</strong></p></li><li><p><strong>sm_100a</strong> &#8212; Blackwell architecture target</p></li><li><p><strong>256 threads</strong> per block (<code>.reqntid 256</code>)</p></li><li><p><strong>4,315 predicates</strong></p></li><li><p><strong>1,202 b64 PTX virtual registers declared</strong></p></li><li><p><strong>180KB shared memory</strong></p></li></ul><h3><strong>The tcgen05 Allocation System</strong></h3><p>The first thing the kernel does  is allocate <strong>TMEM columns</strong> (and write the TMEM address into shared memory).:</p><pre><code><code>// Get warp ID
mov.u32     %r1, %tid.x;
shr.u32     %r2, %r1, 5;                    // warp_id = tid / 32
shfl.sync.idx.b32  %r3, %r2, 0, 31, -1;     // Broadcast warp 0's ID

// Elect a single leader thread
elect.sync  %r205|%p3937, -1;
bar.sync    0;

// Only warp 0, thread 0 executes allocation
setp.ne.s32 %p3938, %r3, 0;
@%p3938 bra $L__BB0_2;

// Allocate 128 TMEM columns (write TMEM base address into shared memory)
mov.b32     %r206, 128;
tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [global_smem+180416], %r206;
tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;
</code></code></pre><h2><strong>What is tcgen05?</strong></h2><p><code>tcgen05</code> is Blackwell&#8217;s new instruction family that changes how threads cooperate on matrix operations. This is introduced on Blackwell (sm_100a) and represents NVIDIA&#8217;s new approach to tile level execution.</p><p>The tcgen05 instruction family includes:</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!8Ogn!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 424w, /__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 848w, /__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 1272w, /__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!8Ogn!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png" width="2288" height="1253" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/b097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:1253,&quot;width&quot;:2288,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:235163,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/185139919?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F74506fae-5532-40ef-be79-c5498acfadf1_2344x1466.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 424w, /__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 848w, /__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 1272w, /__u/substackcdn.com/image/fetch/$s_!8Ogn!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fb097184d-6f5a-450b-8049-9fb4425dd0b6_2288x1253.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>The typical tcgen05 lifecycle follows this pattern:</p><pre><code><code>// 1. Allocate TMEM columns (TMEM base address written into shared memory)

tcgen05.alloc.cta_group::1.sync.aligned.shared::cta.b32 [smem+offset], %r;

// 2. Release the allocation permit
tcgen05.relinquish_alloc_permit.cta_group::1.sync.aligned;

// 3. Perform matrix operations (repeated in main loop)
tcgen05.mma.cta_group::1.kind::f16 [col], A, B, desc, {...}, pred;

// 4. Commit results and signal barrier
tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [barrier];

// 5. Deallocate columns (once at kernel end)
tcgen05.dealloc.cta_group::1.sync.aligned.b32 col, count;
</code></code></pre><p>The <code>.cta_group::1</code> suffix indicates thread block group scope, while <code>.shared::cluster</code> enables cluster-wide synchronization across multiple CTAs.</p><p>This is <strong>not</strong> in my kernel. The compiler determined that the MMA operations need this allocation and inserted it automatically.</p><h3><strong>Barrier Initialization</strong></h3><p>After allocation, the kernel initializes <strong>20 mbarrier objects</strong>:</p><pre><code><code>// Only thread 0 initializes barriers
@%p3939 bra $L__BB0_4;

mov.b32     %r207, 128;
mbarrier.init.shared.b64 [global_smem+180400], %r207;
mov.b32     %r208, 1;
mbarrier.init.shared.b64 [global_smem+180408], %r208;
mbarrier.init.shared.b64 [global_smem+180304], %r207;
mbarrier.init.shared.b64 [global_smem+180336], %r208;
mbarrier.init.shared.b64 [global_smem+180312], %r207;
mbarrier.init.shared.b64 [global_smem+180344], %r208;
mbarrier.init.shared.b64 [global_smem+180320], %r207;
mbarrier.init.shared.b64 [global_smem+180352], %r208;
mbarrier.init.shared.b64 [global_smem+180328], %r207;
mbarrier.init.shared.b64 [global_smem+180360], %r208;
mbarrier.init.shared.b64 [global_smem+180224], %r208;
mbarrier.init.shared.b64 [global_smem+180264], %r208;
mbarrier.init.shared.b64 [global_smem+180232], %r208;
mbarrier.init.shared.b64 [global_smem+180272], %r208;
mbarrier.init.shared.b64 [global_smem+180240], %r208;
mbarrier.init.shared.b64 [global_smem+180280], %r208;
mbarrier.init.shared.b64 [global_smem+180248], %r208;
mbarrier.init.shared.b64 [global_smem+180288], %r208;
mbarrier.init.shared.b64 [global_smem+180256], %r208;
mbarrier.init.shared.b64 [global_smem+180296], %r208;

$L__BB0_4:
fence.mbarrier_init.release.cluster;
bar.sync    0;
</code></code></pre><p>The barriers are initialized with different expected arrival counts:</p><ul><li><p>Some expect 128 arrivals</p></li><li><p>Some expect 1 arrival (leader-only operations)</p></li></ul><p>The <code>fence.mbarrier_init.release.cluster</code> ensures barrier visibility across the cluster &#8212; cluster-scoped synchronization.</p><h3><strong>Mapping: </strong><code>swizzle_2d()</code><strong> &#8594; Integer Arithmetic</strong></h3><p><strong>Kernel:</strong></p><pre><code>bid_m, bid_n = swizzle_2d(M, N, TILE_M, TILE_N, GROUP_SIZE_M)</code></pre><p><strong>PTX (60+ instructions):</strong></p><pre><code><code>mov.u32     %r5, %clusterid.x;              // Get cluster ID (not just CTA ID!)
shr.s32     %r211, %r1112, 31;              // Sign extension for ceiling div
shr.u32     %r212, %r211, 25;
add.s32     %r213, %r1112, %r212;
shr.s32     %r214, %r213, 7;                // &#247; 128 (TILE_M)
and.b32     %r215, %r213, -128;
setp.ne.s32 %p3945, %r1112, %r215;
setp.gt.s32 %p3946, %r1112, -1;
and.pred    %p3947, %p3946, %p3945;
selp.b32    %r216, 1, 0, %p3947;
add.s32     %r7, %r214, %r216;              // num_bid_m = cdiv(M, TILE_M)
...
div.s32     %r219, %r5, %r16;               // group_id = bid / num_bid_in_group
mul.lo.s32  %r220, %r219, %r16;
...
rem.s32     %r231, %r5, %r18;               // bid % group_size_m
...
shl.b32     %r20, %r241, 7;                 // bid_m * TILE_M
</code></code></pre><p>The swizzle pattern interleaves blocks across the N dimension within groups of 8, <strong>improving L2 cache reuse</strong>. The compiler expands this into explicit integer operations.</p><h3><strong>Mapping: </strong><code>ct.gather()</code><strong> &#8594; Predicated Global Loads</strong></h3><p><strong>Kernel:</strong></p><pre><code>token_ids = ct.gather(sorted_token_ids, token_id_indices)</code></pre><p><strong>PTX:</strong></p><pre><code><code>// Compute index and check bounds
setp.ge.u32   %p3959, %r22, %r1121;         // index &gt;= array_length?
mov.b32       %r1152, 0;                     // Default value if out-of-bounds
@%p3959 bra   $L__BB0_8;                     // Skip load if OOB
ld.global.b32 %r1152, [%rd1];                // Actual global load

// Repeat for each element (8x unrolled)
add.s32       %r243, %r22, 16;
setp.ge.u32   %p3960, %r243, %r1122;
mov.b32       %r1153, 0;
@%p3960 bra   $L__BB0_10;
ld.global.b32 %r1153, [%rd2];

add.s32       %r245, %r22, 32;
setp.ge.u32   %p3961, %r245, %r1123;
mov.b32       %r1154, 0;
@%p3961 bra   $L__BB0_12;
ld.global.b32 %r1154, [%rd3];
// ... continues for 8 elements
</code></code></pre><p>The compiler generates <strong>predicated loads</strong> with bounds checking. Each gather element becomes a separate load with its own bounds check.</p><h3><strong>Mapping: K-Loop &#8594; Async Pipeline with Double Buffering</strong></h3><p><strong>Kernel:</strong></p><pre><code>for k in range(0, ct.cdiv(K, TILE_K)):
    a = ct.gather(A, ...)
    b = ct.load(B, ...)
    accumulator = ct.mma(a, b, accumulator)</code></pre><p><strong>PTX (software-pipelined):</strong></p><pre><code><code>$L__BB0_33:
// Wait for previous iteration's data to arrive
mov.b32     %r386, 10000000;                // 10ms timeout
mbarrier.try_wait.parity.acquire.cta.shared::cta.b64 %p4141, [%r385], %r1165, %r386;
not.pred    %p4142, %p4141;
@%p4142 bra $L__BB0_33;                     // Spin until barrier ready

// Issue async copies for NEXT iteration (double buffering!)
selp.b32    %r415, 16, 0, %p4298;           // Copy size or 0 if masked
cp.async.cg.shared.global [%r414], [%rd1170], 16, %r415;
cp.async.cg.shared.global [%r419], [%rd1171], 16, %r420;
cp.async.cg.shared.global [%r424], [%rd1172], 16, %r425;
cp.async.cg.shared.global [%r429], [%rd1173], 16, %r430;
cp.async.cg.shared.global [%r434], [%rd1174], 16, %r435;
cp.async.cg.shared.global [%r439], [%rd1175], 16, %r440;
cp.async.cg.shared.global [%r444], [%rd1176], 16, %r445;
cp.async.cg.shared.global [%r449], [%rd1177], 16, %r450;

// Signal that async copies are issued
cp.async.mbarrier.arrive.noinc.shared.b64 [%r452+180304];
</code></code></pre><p>My simple for loop became a <strong>software pipelined async copy</strong> with:</p><ul><li><p><code>mbarrier.try_wait</code> for synchronization</p></li><li><p><code>cp.async.cg.shared.global</code> for async global&#8594;shared memory copies</p></li><li><p>Double buffering (issue copies for iteration N+1 while computing iteration N)</p></li><li><p>Parity based barrier reuse</p></li></ul><h3><strong>Mapping: </strong><code>ct.mma()</code><strong> &#8594; tcgen05.mma</strong></h3><p><strong>Kernel:</strong></p><pre><code>accumulator = ct.mma(a, b, accumulator)</code></pre><p><strong>PTX:</strong></p><pre><code><code>// Elect a leader thread for this MMA
elect.sync    %r1087|%p4269, -1;
not.pred      %p4270, %p4269;
@%p4270 bra   $L__BB0_133;                   // Non-leaders skip

// Encode shared memory addresses
or.b64        %rd1130, %rd516, 4611756662049538048;
or.b64        %rd1131, %rd515, 4611756662049538048;
mov.b32       %r1088, 0;
mov.b32       %r1089, 136316048;             // MMA descriptor

// Issue the tcgen05 MMA &#8212; only the leader executes this!
tcgen05.mma.cta_group::1.kind::f16 [%r4], %rd1131, %rd1130, %r1089,
                                   {%r1088, %r1088, %r1088, %r1088}, %p4314;

$L__BB0_133:
// Second MMA tile (unrolled)
elect.sync    %r1090|%p4271, -1;
not.pred      %p4272, %p4271;
@%p4272 bra   $L__BB0_135;
add.s64       %rd1132, %rd516, 4611756662049538050;
add.s64       %rd1133, %rd515, 4611756662049538050;
tcgen05.mma.cta_group::1.kind::f16 [%r4], %rd1133, %rd1132, %r1092, {...}, %p4273;

$L__BB0_135:
// Third MMA tile...
</code></code></pre><p><strong>Key insight</strong>: The <code>tcgen05.mma</code> is issued by a <strong>single elected thread</strong>, not all threads. The hardware handles distributing the matrix multiply across the warp. This is <strong>fundamentally different from previous architectures</strong> where each thread contributed to the MMA.</p><p>The <code>elect.sync</code> instruction selects one thread per warp to issue the instruction. Other threads skip directly to the next synchronization point.</p><h3><strong>Mapping: Loop Commit &#8594; Barrier Arrival</strong></h3><p>After each iteration&#8217;s MMAs complete:</p><pre><code><code>// Leader commits results and signals barrier
elect.sync    %r1099|%p4280, -1;
not.pred      %p4281, %p4280;
@%p4281 bra   $L__BB0_141;
tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%r179+40];

$L__BB0_141:
elect.sync    %r1100|%p4282, -1;
not.pred      %p4283, %p4282;
@%p4283 bra   $L__BB0_143;
tcgen05.commit.cta_group::1.mbarrier::arrive::one.shared::cluster.b64 [%r178+32];
</code></code></pre><p>The <code>tcgen05.commit</code> instruction:</p><ol><li><p>Commits the MMA results to shared memory</p></li><li><p>Signals an mbarrier arrival (<code>.mbarrier::arrive::one</code>)</p></li><li><p>Has <strong>cluster scope</strong> (<code>.shared::cluster</code>) &#8212; visible across CTAs in a cluster</p></li></ol><h3><strong>Kernel Cleanup: Deallocation</strong></h3><p>At the end of the kernel:</p><pre><code><code>// Only warp 0 deallocates
setp.ne.s32 %p4289, %r3, 0;
bar.sync    0;
@%p4289 bra $L__BB0_150;

mov.b32     %r1103, 128;
tcgen05.dealloc.cta_group::1.sync.aligned.b32 %r4, %r1103;

$L__BB0_150:
ret;
</code></code></pre><p>The <code>tcgen05.dealloc</code> releases the allocated TMEM columns allocated at kernel start.</p><div><hr></div><h2><strong>SASS: The Final Layer</strong></h2><p>The SASS (Shader ASSembly) is the actual machine code that runs on the GPU. Here&#8217;s what the PTX becomes.</p><h3><strong>The TMEM allocation: contention-handling retry pattern</strong></h3><p>The PTX <code>tcgen05.alloc</code> lowers into a contention-handling atomic retry loop in SASS:</p><pre><code>/*0100*/  ELECT P0, URZ, PT ;                    // Elect leader
/*0110*/  @!P0 BRA 0x2a0 ;                       // Non-leaders skip

/*0120*/  UMOV UR4, 0x4 ;
/*0130*/  IMAD.MOV.U32 R3, RZ, RZ, 0xf ;
/*0140*/  IMAD.MOV.U32 R5, RZ, RZ, 0x1 ;
/*0150*/  DEPBAR.LE SB0, 0x36 ;

// Atomic find-and-set with retry loop
/*0160*/  UTCATOMSWS.FIND_AND_SET.ALIGN UP0, UR4, UR4 ;
/*0170*/  IMAD.U32 R2, RZ, RZ, UR4 ;
/*0180*/  BRA.U UP0, 0x1f0 ;                     // If acquired, continue
/*0190*/  NANOSLEEP 0x64 ;                       // 100ns backoff!
/*01a0*/  UMOV UR4, 0x4 ;
/*01b0*/  DEPBAR.LE SB0, 0x36 ;
/*01c0*/  UTCATOMSWS.FIND_AND_SET.ALIGN UP0, UR4, UR4 ;
/*01d0*/  IMAD.U32 R2, RZ, RZ, UR4 ;
/*01e0*/  BRA.U !UP0, 0x190 ;                    // Retry if not acquired</code></pre><p><strong>New Blackwell instructions:</strong></p><ul><li><p><code>UTCATOMSWS.FIND_AND_SET.ALIGN</code> &#8212; Atomic find-and-set on the thread coalescing state</p></li><li><p><code>NANOSLEEP 0x64</code> &#8212; Sleep for ~100ns (bounded [0, 2t])  (backoff for contention)</p></li></ul><p>This is a <strong>contention handling retry loop with nanosleep backoff</strong> &#8212; the compiler generated a sophisticated synchronization primitive from my simple kernel.</p><h3><strong>Barrier Initialization in SASS</strong></h3><pre><code>/*0420*/  FENCE.VIEW.ASYNC.S ;                   // System-scope fence
/*0430*/  SYNCS.EXCH.64 URZ, [UR14+0x2c0b0], UR4 ;
/*0440*/  SYNCS.EXCH.64 URZ, [UR14+0x2c0b8], UR6 ;
/*0450*/  SYNCS.EXCH.64 URZ, [UR14+0x2c050], UR4 ;
/*0460*/  SYNCS.EXCH.64 URZ, [UR14+0x2c070], UR6 ;
// ... 20 total barrier initializations</code></pre><p>The <code>SYNCS.EXCH.64</code> is Blackwell&#8217;s barrier exchange instruction &#8212; it atomically swaps a value into the barrier object.</p><h3><strong>Virtual Memory Accounting</strong></h3><pre><code>/*02f0*/  UVIRTCOUNT.DEALLOC.SMPOOL 0x80 ;</code></pre><p><code>UVIRTCOUNT.DEALLOC.SMPOOL</code> deallocates from the shared memory pool&#8217;s virtual counter. This is part of Blackwell&#8217;s resource accounting system &#8212; tracking how much shared memory is &#8220;logically&#8221; allocated even when the physical allocation is different.</p><h3><strong>Optimized Barrier Waiting</strong></h3><pre><code>// Try to acquire barrier (non-blocking)
SYNCS.PHASECHK.TRANS64.TRYWAIT P4, [UR10+0x2c070], R14 ;

// If not ready, sleep and retry
@!P4 NANOSLEEP.SYNCS 0x989680 ;               // Sleep 10ms while waiting
@!P4 SYNCS.PHASECHK.TRANS64 P4, [R15+URZ+0x2c070], R14 ;  // Blocking check</code></pre><p>The <code>NANOSLEEP.SYNCS</code> variant sleeps while waiting for a synchronization operation &#8212; more power efficient than a busy-wait loop.</p><h3><strong>tcgen05 in SASS: Under the Hood</strong></h3><p>At the SASS level, the tcgen05 PTX instructions expand into sequences that reveal more about the hardware implementation. The allocation uses <code>UTCATOMSWS</code> (Unified Thread Coalescing Atomic on Shared Workspace):</p><pre><code>// tcgen05.alloc expands to:
/*0160*/  UTCATOMSWS.FIND_AND_SET.ALIGN UP0, UR4, UR4 ;  // Find free column slot
/*0170*/  IMAD.U32 R2, RZ, RZ, UR4 ;                      // Extract result
/*0180*/  BRA.U UP0, 0x1f0 ;                              // Branch if acquired
/*0190*/  NANOSLEEP 0x64 ;                                // 100ns backoff on contention</code></pre><p>The <code>UTCATOMSWS.FIND_AND_SET.ALIGN</code> atomically searches a bitmap for a free slot and sets it &#8212; this is how the hardware manages dynamic column allocation. The <code>.ALIGN</code> suffix ensures the operation respects alignment constraints for the tcgen05 column structure.</p><p>For deallocation, the compiler generates:</p><pre><code>// tcgen05.dealloc expands to:
/*xxxx*/  UTCATOMSWS.AND URZ, UR4 ;                      // Clear allocation bits
/*xxxx*/  UVIRTCOUNT.DEALLOC.SMPOOL 0x80 ;              // Return 128 bytes to pool</code></pre><p>The <code>UVIRTCOUNT.DEALLOC.SMPOOL</code> instruction is part of Blackwell&#8217;s virtual memory accounting system &#8212; it tracks &#8220;logical&#8221; allocations separately from physical shared memory.</p><div><hr></div><h2>What CuTile Generated That I Didn&#8217;t Write</h2><p>Let&#8217;s be explicit about the gap between the kernel and the compiled code.</p><h3>Memory Management</h3><p><strong>What I wrote &#8594; What the compiler generated:</strong></p><ul><li><p>No shared memory allocation &#8594; 180KB <code>global_smem[180424]</code></p></li><li><p>No explicit barriers &#8594; 20 mbarrier objects</p></li><li><p>Simple for-loop &#8594; Double-buffered async pipeline</p></li><li><p>No register hints &#8594; 1,202 b64 PTX virtual registers declared</p></li></ul><h3>Synchronization</h3><p><strong>What I wrote &#8594; What the compiler generated:</strong></p><ul><li><p>No thread coordination &#8594; <code>elect.sync</code> for leader election</p></li><li><p>No barriers &#8594; <code>mbarrier.try_wait.parity</code> with 10ms timeout</p></li><li><p>No atomics &#8594; <code>UTCATOMSWS.FIND_AND_SET</code> with NANOSLEEP backoff</p></li><li><p>No cluster operations &#8594; <code>fence.mbarrier_init.release.cluster</code></p></li></ul><h3>Instruction Selection</h3><p><strong>What I wrote &#8594; What the compiler generated:</strong></p><ul><li><p><code>ct.mma(a, b, acc)</code> &#8594; <code>tcgen05.mma.cta_group::1.kind::f16</code></p></li><li><p><code>ct.gather()</code> &#8594; Predicated loads with bounds checks</p></li><li><p><code>ct.scatter()</code> &#8594; Predicated stores with bounds checks</p></li><li><p><code>for k in range()</code> &#8594; Software pipelined loop with async copies</p></li></ul><div><hr></div><div><hr></div><h2><strong>Where This Is Heading</strong></h2><p><strong>The technical evidence points to an uncomfortable conclusion: NVIDIA&#8217;s first party compiler advantage isn&#8217;t just real &#8212; it&#8217;s structural and widening.</strong></p><p><strong>The moat isn&#8217;t just documentation.</strong> Yes, tcgen05 is in PTX ISA 9.1. Yes, <code>elect.sync</code> and <code>nanosleep</code> are public. But NVIDIA has a history of undocumented PTX and SASS instructions that show up in cuBLAS and cuDNN before they appear in any ISA reference &#8212; sometimes they never do. I can&#8217;t claim CuTile is using undocumented instructions in this kernel, but <strong>it would be naive to assume the public PTX ISA is the complete picture of what NVIDIA&#8217;s own compilers can target</strong>. Even setting that aside, the orchestration problem remains: knowing which patterns, with what parameters, for which kernel shapes. My 86 line kernel became 20 barriers with specific arrival counts, 100ns backoff on TMEM allocation, 10ms timeouts on barrier waits, leader election before every MMA. You can find pieces in Flash Attention 4, in CUTLASS examples, in blog posts. The specific combination? That&#8217;s encoded in CuTile&#8217;s MLIR passes.</p><p><strong>The gap compounds with each generation.</strong> Blackwell introduced tcgen05, cluster-scope barriers, TMEM, CTA pairs. Each feature multiplies the configuration space. The Triton team struggled enough with Blackwell attention that they built Gluon as an escape hatch &#8212; a lower level frontend that exposes more control because the high level abstractions couldn&#8217;t find the right codegen. CuTile doesn&#8217;t need an escape hatch. <strong>The compiler already knows</strong>.</p><p><strong>Rubin will repeat this pattern.</strong> Whatever instruction family Rubin introduces &#8212; tcgen06, new cluster primitives, new memory hierarchies &#8212; NVIDIA&#8217;s compiler team has been generating code for it since before the silicon taped out. CuTile will ship Rubin support <strong>on day one with tuned codegen.</strong> Triton and Pallas will start reverse engineering patterns from publicly released kernels and CUTLASS examples. This isn&#8217;t a criticism of those projects; it&#8217;s the structural reality of targeting hardware you don&#8217;t design.</p><p><strong>Frameworks will notice.</strong> If CuTile consistently generates faster code for complex kernels &#8212; MoE, attention, whatever the next bottleneck is &#8212; how long before PyTorch&#8217;s Inductor or JAX&#8217;s Pallas add CuTile as a backend option? Vendor lock-in concerns are real, but frameworks already have NVIDIA specific paths through cuDNN and cuBLAS. A CuTile backend is the same trade-off with better codegen for custom ops.</p><p><strong>The black box is the product.</strong> CuTile&#8217;s closed MLIR passes let NVIDIA iterate on compiler internals without breaking user code. Your tile level kernel stays stable while the codegen improves underneath. This is the opposite of Triton&#8217;s model, where you sometimes need to understand the compiler to get good performance. For NVIDIA, hiding the complexity is the point &#8212; they want you writing tile level Python, not debugging PTX.</p><p><strong>My predictions:</strong></p><ul><li><p>Within two years, at least <em><strong>one major framework</strong></em> adds CuTile as an optional backend for NVIDIA GPUs.</p></li><li><p>Rubin ships with CuTile support day one. Triton gets functional Rubin support within 6 months, competitive codegen within 18.</p></li></ul><div><hr></div><h2><strong>Conclusion</strong></h2><p><strong>NVIDIA is building a compiler moat.</strong></p><p>Not just through documentation &#8212; though the gap between public PTX and what their own tools can emit is its own question. <strong>The deeper moat is orchestration: </strong>which patterns to combine, which parameters to use, which configurations work for which workloads. My 86 line kernel became 1,900 lines of PTX with 20 barriers, contention handling TMEM allocation with nanosleep backoff, and software pipelined async copies. Every one of those primitives is discoverable. The specific combination that makes this MoE kernel fast is not.</p><p>Triton can emit tcgen05 instructions. Matching CuTile&#8217;s codegen quality &#8212; for this kernel, for attention, for whatever comes next &#8212; requires <strong>reproducing years of tuning encoded in proprietary MLIR passes.</strong> And that&#8217;s assuming the public ISA is the whole story, which history suggests it isn&#8217;t.</p><p>Each architecture adds complexity. <strong>Each generation widens the gap.</strong></p><p>If you want to see this yourself, the artifacts are in the repo. Run <code>strings</code> on <code>tileiras</code>, extract the PTX from your <code>.cubin</code> files, trace your operations through the stack. The compiler is a black box, but the outputs tell you everything you need to know about where GPU programming is heading.</p><div><hr></div><p><em>Full Code / IR Dumps: <a href="https://github.com/patrick-toulme/justabyte/tree/main/cutile_blackwell_post">https://github.com/patrick-toulme/justabyte/tree/main/cutile_blackwell_post</a></em></p><p><em>Connect on LinkedIn: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></em></p><p><em>Follow on X: <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></em></p>]]></content:encoded></item><item><title><![CDATA[When XLA Isn't Enough: From Pallas to VLIW with Splash Attention on TPU]]></title><description><![CDATA[When does XLA hit its limits? How do you write the TPU Pallas kernel that the compiler cannot automatically find? Why can't XLA generate Splash Attention?]]></description><link>https://patricktoulme.substack.com/p/when-xla-isnt-enough-from-pallas</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/when-xla-isnt-enough-from-pallas</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sun, 11 Jan 2026 21:18:48 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!fqnv!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p><strong>Full Code / IR Dumps: </strong><a href="https://github.com/patrick-toulme/justabyte/tree/main/tpu_pallas_post">https://github.com/patrick-toulme/justabyte/tree/main/tpu_pallas_post</a></p><p><strong>Connect On Linkedin:</strong> <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p><div class="subscription-widget-wrap-editor" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe&quot;,&quot;language&quot;:&quot;en&quot;}" data-component-name="SubscribeWidgetToDOM"><div class="subscription-widget show-subscribe"><div class="preamble"><p class="cta-caption">Thanks for reading Just a Byte - AI Compilers, Silicon, and Systems! Subscribe for free to receive new posts and support my work.</p></div><form class="subscription-widget-subscribe"><input type="email" class="email-input" name="email" placeholder="Type your email&#8230;" tabindex="-1"><input type="submit" class="button primary" value="Subscribe"><div class="fake-input-wrapper"><div class="fake-input"></div><div class="fake-button"></div></div></form></div></div><p><strong>Follow on X: </strong><a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><h2>Motivation</h2><p><a href="/__u/patricktoulme.substack.com/p/from-jax-to-vliw-tracing-a-computation">Last post, the TPU compiler did everything</a>: 8 lines of JAX became 250 VLIW bundles across 5 fused kernels, with automatic memory orchestration, dual-MXU scheduling, and async DMA overlap. The thesis was simple&#8212;write high-level code, let the compiler figure out the rest.</p><p><strong>This post is about where that breaks down.</strong></p><p>Same hardware, same compiler: naive attention compiles to 10,245 VLIW bundles and moves 526MB through HBM. Splash Attention (<strong>block_size=512</strong>) compiles to 1,651 bundles and moves 14MB. <strong>XLA optimized both aggressively</strong>. The 6&#215; gap isn&#8217;t a compiler failure&#8212;<strong>it&#8217;s a compiler limitation.</strong> XLA can fuse ops and schedule instructions, but it can&#8217;t rewrite your algorithm. <strong>It can&#8217;t infer the streaming/online-softmax reformulation</strong>, that you never need to materialize the full attention matrix, that there&#8217;s a numerically-stable streaming formulation hiding inside your einsum.</p><p><strong>In principle, a compiler could learn this</strong>. But no one has taught one yet. FlashAttention exists because a human figured out the trick and wrote a kernel. Until compilers get smarter, we need an escape hatch.</p><p><strong>Pallas is that escape hatch.</strong> This post traces Splash Attention&#8212;<a href="https://github.com/jax-ml/jax/blob/main/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py">JAX&#8217;s FlashAttention for TPUs</a>&#8212;through the <a href="https://github.com/jax-ml/jax/tree/d40b4c951a9b53207974330c620130f664a0383e/jaxlib/mosaic/dialect/tpu">TPU Pallas Compiler</a>, and compares it to naive attention at every layer of IR.</p><h2>Setup: </h2><p>All experiments were performed on a TPU V6e Trillium from Google Cloud. We use the same dump flags from the previous post with the addition of xla_mosaic flags which dumps the Mosaic compiler IR for Pallas.</p><pre><code><code>os.environ["XLA_FLAGS"] = (
    f"--xla_dump_hlo_as_text "
    f"--xla_dump_to={HLO_PATH} "
    f"--xla_dump_hlo_pass_re=.* "
)
os.environ["LIBTPU_INIT_ARGS"] = (
    f"--xla_jf_dump_to={LLO_PATH} "
    f"--xla_jf_dump_hlo_text=true "
    f"--xla_jf_dump_llo_text=true "
    f"--xla_jf_dump_llo_html=false "
    f"--xla_jf_dump_llo_static_gaps=true "
    f"--xla_jf_emit_annotations=true "
    f"--xla_jf_debug_level=2 "
    f"--xla_mosaic_dump_to={MOSAIC_PATH} "
    f"--xla_mosaic_enable_dump_debug_info=true "
    f"--xla_mosaic_enable_llo_source_annotations=true"
)</code></code></pre><p><strong>Full Code / IR Dumps: </strong><a href="https://github.com/patrick-toulme/justabyte/tree/main/tpu_pallas_post">https://github.com/patrick-toulme/justabyte/tree/main/tpu_pallas_post</a></p><h2>What this post covers:</h2><ul><li><p><strong>The Compiler&#8217;s Limitation</strong>: Why XLA can&#8217;t fuse across the attention matrix</p></li><li><p><strong>Splash Attention Kernel</strong>: How Pallas expresses online softmax with grid, BlockSpec, and scratch</p></li><li><p><strong>Pallas Compiler Pipeline</strong>: The path from Pallas code &#8594; Mosaic &#8594; LLO &#8594; VLIW</p></li><li><p><strong>HLO/Mosaic/LLO Walkthrough</strong>: What each IR layer reveals</p></li><li><p><strong>Reference Attention Comparison</strong>: Same algorithm through XLA&#8217;s standard path</p></li><li><p><strong>Analysis</strong>: Bundle counts, HBM traffic, block size tradeoffs</p></li></ul><h3>The Compiler&#8217;s Limitation</h3><p>Why can&#8217;t XLA automatically achieve Pallas-level performance for every case? Looking at what XLA <em>does</em> fuse for reference JAX attention:</p><pre><code><code>fusion.5: Q @ K^T + mask + reduce_max  &#8594;  (max, scores)
fusion.2: exp(scores - max) + reduce_sum  &#8594;  sum
fusion:   normalize + S @ V  &#8594;  output</code></code></pre><p>XLA fuses the matmul with the mask and max reduction. It fuses the exp with the sum reduction. It&#8217;s doing real work here. But the 128MB attention matrix still gets written to HBM between <code>fusion.5</code> and <code>fusion.2</code>.</p><p>The problem is that standard softmax is inherently multi-pass: you need the max over <em>all</em> values before you can compute any exp. <strong>XLA can fuse along dataflow edges, but it can&#8217;t restructure the algorithm</strong>. The attention matrix has to exist somewhere because multiple operations need to read it.</p><p>Online softmax sidesteps this by maintaining running statistics &#8212; you compute a tile&#8217;s contribution to max/sum/output, then <em>update</em> your estimates as you process the next tile. <strong>This is an algorithmic transformation, not a fusion pattern</strong>. XLA would need to:</p><ol><li><p>Recognize that the softmax + matmul pattern permits streaming</p></li><li><p>Prove the online reformulation is numerically equivalent</p></li><li><p>Restructure the computation to process tiles incrementally</p></li></ol><p><strong>Concretely:</strong> it would have to replace <code>scores = QK&#7488;; p = softmax(scores); o = pV</code> with a loop that carries (m, l, o) and updates them per KV tile.</p><p><strong>No production compiler does this today (to my knowledge). </strong>Pallas lets you write the streaming algorithm directly.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!fqnv!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 424w, /__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 848w, /__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 1272w, /__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!fqnv!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png" width="1456" height="809" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:809,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:130502,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182968804?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 424w, /__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 848w, /__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 1272w, /__u/substackcdn.com/image/fetch/$s_!fqnv!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F5c33bcb2-9c9b-4c73-a7b1-591561ac1676_1800x1000.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><h2>Splash Attention Kernel</h2><p>Splash Attention is JAX&#8217;s FlashAttention for TPUs, built on Pallas. <a href="https://arxiv.org/abs/2205.14135">The algorithm itself is well-documented elsewhere</a> &#8212; here we focus on how Pallas expresses it.</p><pre><code><code>// Standard softmax (requires full matrix):
max = reduce_max(scores)        // Need ALL scores first
exp_scores = exp(scores - max)  // Then process all
sum = reduce_sum(exp_scores)    // Need ALL exp values
output = exp_scores / sum

// Online softmax (streaming):
for each KV_block:
    m_new = max(m_prev, local_max)           // Update running max
    correction = exp(m_prev - m_new)         // Rescale factor
    l_new = correction * l_prev + local_sum  // Update running sum
    o_new = correction * o_prev + local_out  // Update running output</code></code></pre><h3>Pallas Kernel Structure</h3><p><a href="https://docs.jax.dev/en/latest/pallas/index.html">A Pallas kernel</a> is a function that operates on <em>Refs</em> &#8212; mutable views into memory regions. The runtime calls your kernel once per grid point, with Refs pointing to the appropriate tiles:</p><pre><code><code>def flash_attention_kernel(
    # Prefetched scalar inputs (live in SMEM)
    data_next_ref,      # Which KV block to load
    block_mask_ref,     # Skip mask
    mask_next_ref,      # Partial mask index
    # Tiled inputs (sliced per grid point)
    q_ref,              # [block_q, head_dim] slice of Q
    k_ref,              # [block_kv, head_dim] slice of K  
    v_ref,              # [block_kv, head_dim] slice of V
    ...
    # Scratch space (persists across grid iterations)
    m_scratch_ref,      # [block_q, 128] running max
    l_scratch_ref,      # [block_q, 128] running sum
    o_scratch_ref,      # [block_q, head_dim] running output
    # Output tile
    o_ref,
    *,
    # Static config (compiled into kernel)
    bq: int,
    bkv: int,
    ...
):</code></code></pre><p>Inside the kernel, you read/write Refs with <code>ref[...]</code> syntax. Pallas traces these operations and compiles them to Mosaic IR.</p><h3>BlockSpec: Mapping Grid to Memory</h3><p><a href="https://docs.jax.dev/en/latest/pallas/grid_blockspec.html#blockspec-a-k-a-how-to-chunk-up-inputs">BlockSpec</a> defines how each grid point maps to a tile of the input/output arrays:</p><pre><code><code>pl.BlockSpec(
    block_shape=(None, bq, head_dim),  # None = full dim, bq = tiled
    index_map=lambda h, i, j, *_: (h, i, 0)  # grid coords &#8594; tile origin
)
</code></code></pre><p>For Q, the index map is simple &#8212; grid point <code>(h, i, j)</code> reads Q block <code>(h, i, :)</code>. But K and V need indirection for sparse attention:</p><pre><code><code>def k_index_map(h, i, j, data_next_ref, block_mask_ref, mask_next_ref):
    # Look up which KV block to actually load (enables block skipping)
    next_j, *_ = _next_nonzero(h, i, j, data_next_ref, block_mask_ref, mask_next_ref)
    return (h // q_heads_per_kv_head, next_j, 0)
</code></code></pre><p>The <code>data_next_ref</code> indirection is how Splash skips masked blocks &#8212; if block <code>j</code> is fully masked, <code>data_next[h,i,j]</code> points to the next valid block instead.</p><h3><a href="https://docs.jax.dev/en/latest/pallas/grid_blockspec.html#grid-a-k-a-kernels-in-a-loop">Grid and Dimension Semantics</a></h3><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!gYVO!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 424w, /__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 848w, /__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 1272w, /__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!gYVO!,w_2400,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png" width="1200" height="656.8681318681319" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:false,&quot;imageSize&quot;:&quot;large&quot;,&quot;height&quot;:797,&quot;width&quot;:1456,&quot;resizeWidth&quot;:1200,&quot;bytes&quot;:221614,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182968804?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:&quot;center&quot;,&quot;offset&quot;:false}" class="sizing-large" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 424w, /__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 848w, /__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 1272w, /__u/substackcdn.com/image/fetch/$s_!gYVO!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F4dced65d-e0bb-48e6-86fc-7bb312185b98_1900x1040.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><pre><code><code>grid = (num_q_heads, q_seq_len // block_q, grid_width)

compiler_params = pltpu.CompilerParams(
    dimension_semantics=("parallel", "arbitrary", "arbitrary")
)
</code></code></pre><ul><li><p><strong>Dimension 0 (heads)</strong>: <code>"parallel"</code> &#8212; no cross-iteration dependencies, can execute in any order</p></li><li><p><strong>Dimension 1 (Q blocks)</strong>: <code>"arbitrary"</code> &#8212; independent because scratch is <em>per-(head, q_block)</em></p></li><li><p><strong>Dimension 2 (KV blocks)</strong>: <code>"arbitrary"</code> &#8212; dependent because scratch accumulates across KV tiles</p></li></ul><p>The <code>"parallel"</code> hint lets the compiler vectorize across heads. <code>"arbitrary"</code> means &#8220;don&#8217;t assume anything&#8221; &#8212; safe but conservative.</p><h3><a href="https://docs.jax.dev/en/latest/pallas/grid_blockspec.html#grid-a-k-a-kernels-in-a-loop">PrefetchScalarGridSpec</a></h3><p>Splash uses <code>PrefetchScalarGridSpec</code> to overlap data loading with compute:</p><pre><code><code>pltpu.PrefetchScalarGridSpec(
    num_scalar_prefetch=3,  # First 3 inputs go to SMEM, prefetched ahead
    in_specs=[...],
    out_specs=[...],
    grid=grid,
)</code></code></pre><p>The first 3 arguments (<code>data_next</code>, <code>block_mask</code>, <code>mask_next</code>) are small index arrays. Putting them in SMEM with prefetch means the kernel can look up the next block address while the current block is still computing.</p><h3>Memory Spaces</h3><p>Pallas exposes TPU memory hierarchy:</p><pre><code><code># VMEM (on-chip SRAM) &#8212; default for BlockSpec refs
pl.BlockSpec((bq, head_dim), index_map)

# SMEM (scalar memory) &#8212; for small indexing data
pl.BlockSpec((num_heads,), lambda *_: (0,), memory_space=pltpu.SMEM)

# Scratch (persists across grid iterations)
jax.ShapeDtypeStruct((bq, NUM_LANES), jnp.float32)  # m_scratch, l_scratch</code></code></pre><p>The scratch refs (<code>m_scratch_ref</code>, <code>l_scratch_ref</code>, <code>o_scratch_ref</code>) are the key to online softmax &#8212; they accumulate across the KV dimension without spilling to HBM.</p><h2>Pallas Compiler Pipeline</h2><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!_OPO!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 424w, /__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 848w, /__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 1272w, /__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!_OPO!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png" width="1456" height="453" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:453,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:119546,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182968804?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 424w, /__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 848w, /__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 1272w, /__u/substackcdn.com/image/fetch/$s_!_OPO!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F2eb88d4a-b684-44e3-957a-8fedd7c4513d_1800x560.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>When building a DSL compiler, you want to reuse as much of the existing compiler infrastructure as possible rather than writing everything from scratch. Pallas shows up in HLO as an opaque <code>tpu_custom_call</code>. XLA still does its usual whole-graph work (layout, scheduling, surrounding fusions), but it <strong>can&#8217;t see or rewrite the kernel body</strong>. The payload is serialized MLIR that goes straight to Mosaic &#8594; LLO &#8594; VLIW. The key insight is that both XLA's fusion region codegen and Pallas share the same LLO backend. Whether the TPU is running a compiler-generated fusion or a hand-written Pallas kernel, the final stages&#8212;MXU scheduling, DMA overlap, VLIW packing&#8212;are identical. Pallas just enters the pipeline at a different point.</p><h2>What XLA Sees</h2><p>The whole Pallas kernel becomes a <strong>single opaque HLO op:</strong></p><pre><code><code>%custom-call = bf16[8,2048,128] custom-call(
    bf16[8,2048,128] %q, bf16[8,2048,128] %k, bf16[8,2048,128] %v, ...
), custom_call_target="tpu_custom_call", backend_config="..."
</code></code></pre><p>XLA&#8217;s fusion passes can&#8217;t look inside &#8212; the kernel goes directly to Mosaic, which compiles it to LLO &#8594; VLIW bundles. The <code>backend_config</code> encodes the full Pallas program as serialized MLIR IR.</p><h2>HLO</h2><p>Here&#8217;s what the Pallas kernel looks like in HLO:</p><div class="github-gist" data-attrs="{&quot;innerHTML&quot;:&quot;<div id=\&quot;gist144398218\&quot; class=\&quot;gist\&quot;>\n    <div class=\&quot;gist-file\&quot; translate=\&quot;no\&quot; data-color-mode=\&quot;light\&quot; data-light-theme=\&quot;light\&quot;>\n      <div class=\&quot;gist-data\&quot;>\n        <div class=\&quot;js-gist-file-update-container js-task-list-container\&quot;>\n  <div id=\&quot;file-pallas_attn-md\&quot; class=\&quot;file my-2\&quot;>\n      <div id=\&quot;file-pallas_attn-md-readme\&quot; class=\&quot;Box-body readme blob p-5 p-xl-6 \&quot;\n    style=\&quot;overflow: auto\&quot; tabindex=\&quot;0\&quot; role=\&quot;region\&quot;\n    aria-label=\&quot;pallas_attn.md content, created by patrick-toulme on 07:31PM today.\&quot;\n  >\n    <article class=\&quot;markdown-body entry-content container-lg\&quot; itemprop=\&quot;text\&quot;><pre><code>HloModule jit_splash_attention_kernel, entry_computation_layout={(bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)})-&amp;gt;bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}\n\nENTRY %main.28 (Arg_0.1: bf16[8,2048,128], Arg_1.2: bf16[8,2048,128], Arg_2.3: bf16[8,2048,128]) -&amp;gt; bf16[8,2048,128] {\n  %constant.4 = s8[1,2,2]{2,1,0} constant({ { { 0, 0 }, { 0, 1 } } })\n  %constant.5 = s8[1,2,2]{2,1,0} constant({ { { 1, 0 }, { 2, 1 } } })\n  %Arg_0.1 = bf16[8,2048,128]{2,1,0} parameter(0), metadata={op_name=\&quot;q\&quot;}\n  %Arg_1.2 = bf16[8,2048,128]{2,1,0} parameter(1), metadata={op_name=\&quot;k\&quot;}\n  %Arg_2.3 = bf16[8,2048,128]{2,1,0} parameter(2), metadata={op_name=\&quot;v\&quot;}\n  %constant.6 = s32[2048]{0} constant({...})\n  %broadcast.0 = s32[2048,128]{1,0} broadcast(%constant.6), dimensions={0}, metadata={op_name=\&quot;jit(splash_attention_kernel)/jit(main)/splash_attention_kernel/splash_kernel_b1024_h8_s2048/jit(_splash_attention)/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=1045}\n  %custom-call.0 = (f32[1024,128]{1,0}, f32[1024,128]{1,0}, f32[1024,128]{1,0}, bf16[8,2048,128]{2,1,0}) custom-call(%constant.4, %constant.5, %Arg_0.1, %Arg_1.2, %Arg_2.3, /*index=5*/%broadcast.0), custom_call_target=\&quot;tpu_custom_call\&quot;, operand_layout_constraints={s8[1,2,2]{2,1,0}, s8[1,2,2]{2,1,0}, bf16[8,2048,128]{2,1,0}, bf16[8,2048,128]{2,1,0}, bf16[8,2048,128]{2,1,0}, s32[2048,128]{1,0}}, metadata={op_name=\&quot;jit(splash_attention_kernel)/jit(main)/splash_attention_kernel/splash_kernel_b1024_h8_s2048/jit(_splash_attention)/splash_mha_fwd_c602b16d/pallas_call\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=1100}, backend_config={\&quot;custom_call_config\&quot;: {\&quot;body\&quot;: \&quot;TUzvUgFNTElSMjEuMC4wZ2l0AAFRCwEDBQcJAQMLAzsNDxETFRcZGx0fISMlJykrLS8xMzU3OTs9P0FDRQMKCWYILwH7BwsXCxMXCxMLCwsbCwsLCysLCxMTExcLCwsLCw8PDw8XGxMTFxNzUwsLExcXFxcXCxMLExMTExMTCxMTExcPCxMTExMTEwsTDxcTCxMTExMPCxcPDxMXCxMXExMTExMPpYUPCwsLCwsLCwsLFxMPExcPExMLExMPExMTGwsFD2GBkY1hKgIqAgFGBgsLExMTExcTFwsLCwsTExMTHxcXCwszFwsLFxMXExMTIwsLawsbCwtzCw8LC0sfCx8LHwsfCycLJwsnCx8LGxsbGxsbGxsPDw8fCycLJxMLJxMLJxMLJxMLGxMLGxMfFwsfExMLHw8fExMTEw8PHw8fExMPHw8LJxMTDw8PExMPExMfDxMfDw8nExMPDw8TEw8PHxMTHw8THxMTHxcfCxMTHwsTExMfExMfExMTHxMTEx8TEx8rUw8TEx8TEx9TExMfExMfExMfExMfCw8LHxMTEx8PCxMTHxMTHxMTHxMTHxcfExMfMxMTHxMTHxMTEx8TFwsTEycfExMTExMTDxMTJxMTHw8TExMnFw8XCxcXCyMLFxMfFxMTExMTDx8TCxMfDxMLFx8LExMfKxMTEx8LExMfExMTHxcTEx8XExMfFx8LExMfKxMTHxcTEx8XExMfExMTHxMTEx8TEx8TEx8TEx8TExMfFx8TEx8zExMTHxMTHxMTEx8TExMfExMfExMfExMfExMfExMfExMfExMfBwVZWQkFXUkBLw8HHycHCw8rHwsnFx8jHwcnG0MvKx8fAlovHwVHAwMTLgMFSRV7XgUDAxMqAwVLFW4EJwVNBU8FUQMDmgRWCAVTBVUFVwVZAwUSAgoFFgIaAgVbBV0VV0oDFSIEJx0b2gQdogOmAwVfBWEFYwVlBWcdoykdo6kRAQUdow8DAxPuBAMDYgVaCB0NfgUdDYoFFU4C8gUdDZIHIw0HMQEAAAAAAAAAAAQAAAAAAACAAAAAAAAAACMNBSEABAAAAAAAAIAAAAAAAAAADScNKR0NMgMdQgNGAx1OA1IDHVoDXgMdZgNqAx1yA3YDBWkVV7IDBWsdDc4DHQ3mAx0NAgQdDS4EHQ1GBB0vWgQFbR0b/gQdGxoFHRsyBR1GBUoFEQsABW8dG3oHHRuGBx0bAggdGxoIHRsyCB0bSggFcR03NgMV2ScdfgOCAx03rgMFcx03ygMDA2W/HTfSAx034gMV4ScFdQMDZSoEFe0nFXEnHQ12BB2KBI4EBXcdDa4EAwMTTgUdDWYFHQ1yBR0NGgYVkgZJHQ3WBxENEWFmZmluZV9tYXA8KGQwLCBkMSwgZDIpIC0+IChkMCwgZDEsIGQyKT4AYWZmaW5lX21hcDwoZDAsIGQxKSAtPiAoZDAsIGQxKT4AEQ0BBXkFewV9BX8FgQWDBYUFhwWJHToDPgMdx4oDFdljHZfGAx3qA+4DFeFjHZf+Ax3pEgQFiwMDZcUdLzIEFe1jHZdCBAMDEz0d6V4EAwWvPfk9BY0jdHB1Lm1lbW9yeV9zcGFjZTx2bWVtPgAjdHB1LnBpcGVsaW5lX21vZGU8c3luY2hyb25vdXM+ACN0cHUuZGltZW5zaW9uX3NlbWFudGljczxhcmJpdHJhcnk+ACN0cHUuZGltZW5zaW9uX3NlbWFudGljczxwYXJhbGxlbD4AI3RwdS5tZW1vcnlfc3BhY2U8c21lbT4AI3RwdS5kb3RfZGltZW5zaW9uX251bWJlcnM8WzFdLCBbMV0sIFswXSwgWzBdLCBbMCwgMCwgMSwgMF0sIFtdLCBbXT4AI3RwdS5kb3RfZGltZW5zaW9uX251bWJlcnM8WzFdLCBbMF0sIFswXSwgWzFdLCBbMCwgMCwgMSwgMV0sIFtdLCBbXT4ABY8FkSMBAQEdH/IEHR8OBR0fJgUdPgVCBR0xUgUdNgKeBQWTBZUFlwWZHTGuBR0xugUdB+4FHX8OBgMFrz35JgYDAxOGBh2uBrIGBZsFnSMNAxEBAAAAAAAAAB3SBtYGBZ8FoR0aBx4HHR8uBx02ArYHHR8OCB0fJggdHz4IAwWWAr8ZmgIFowWlAw+iAqYCHaoCrgKyArYCugK+AsUZx8ICxgIFpwEHAgL//w0lBakjDQcxCAAAAAAAAAACAAAAAAAAAAIAAAAAAAAABasRDQkFrQWvARHKAtIC2gLiAuoC8gL6AgIDAwUjzgIlTQnJAwUj1gIlTQnLAwUj3gIlTQnNAwUj5gIlTwnPAweN/SPuAiVPCdEDB439I/YCJU8J0wMHjf0j/gIlTwnVAwUjBgMlTQnXAwUdURnJAwUdURnLAwUdURnNAwUdUxnPAwUdUxnRAwUdUxnTAwUdUxnVAwUdURnXEQEBEQMBFY+RLQMH5ggtWwWxLQMJRg8jTg8LBbMtAwkyER3KEQsVWVYDBbUtAwmiEhPaEgcVW2IDBbctAwkiJBNiJAcVXW4DBbktAwmuJBfKJAsVX3oDBbstYQeBJ08Vk4YDBb0tYQeNI0EV244DLWEHZgJBvR2SA5YDBb8tYQf+AgkVHRWeAxUtqgMFwS0DB94IK1EVld0tAwfmCB9dFVm2AxVbugMVXb4DFV/CAxWT2xWZkS0DB+YIH2UVnZEtAwcCCSVRHRXaAxUt3gMVn90tAwcCCRdTFY+hBcMtAwkmDyMuDwsdFfYDFS36AxWV4xWZoRWdoR0VCgQVLQ4EFZ/jFRYEJx0vGgQtAwdKCxcjHRUpHS8mBC0DB0oLByURDQUVj6ctAwliC2l+CwcdFToEFS0+BBWV7xWZpxWdpx0VTgQVLVIEFZ/vHRWpLQMHmgwHLRViBCcdL2YELQMHugwXPR0VDx0vcgQtAwe6DAc/FXoEDx0RfgQtAwfCDBE1AwMThgQTCRAAAOAPBcUVkgQPHRGWBC0DB8YMM0EFxx1zogQVpgQPHRGqBC0DB8YMGXsVsgQPHRG2BC0DB8oMJUkdMb4EFcIEDx0RxgQtAwfKDCVZHRXOBBXSBA8dEdYELQMHygwjgxXeBA8dEeIELQMHygwJHQMFEgLqBBYCGgIjAQkhAQAAAAEAAAADAAAAAAAAABMJARX2BA8dEfoELQMH5gwzbRUCBQ8dEQYFLQMH5gwJLSMBCSEBAAAAAQAAAAIAAAAAAAAAFRIFDx0RFgUtAwfqDDNtFR4FDx0RIgUtAwfqDAktFSoFDx0RLgUtAwfuDDNtFTYFDx0ROgUtAwfuDAktBckVe6kFyy0DB7IMCXERAQIgFVYFCR0HWgUtAweKCyllFXFjBc0VagUJHQduBS0DB44LK08VdgUJHQd6BS0DB44LU3cVggUJHQeGBS0DB54LESUVjgUJHQeSBS0DB6oLFTcDAxOaBSURCQAAAAAVogUJHQemBS0DB7YLE48DBzoCCgI+An1CAn0VsgUJHQe2BS0DB/4LI00VvgUJHQfCBS0DB/4LU48df8oFFc4FCR0H0gUtAwf+CyOPAwOvPR3eBeIFBc8V5gVJHTPqBS0DCaoJPbIJDy0DBxIME0UVe/YFFXH6BRVX/gUVWQIGFVsGBhVdCgYVX5MVEgZJHTMWBi0DCaoJJ7IJDxUeBkkdMyIGLQMHxgkVOxEBIR1zLgYVMgZJHTM2Bi0DCcIJJ8oJDwMDZT4GEQ0VHUYGSgYF0RVOBl4GHVIGVgYF0y1aBgcCBR89BdUVYgZqBh0zZgYtAwf2CSlzFU4CbgYVe3IGFXF2BhVXegYVWX4GFVuCBhVdXxMJkMzMzD8djga7BdcdM5YGLQMHhgoTUR0fux2iBrsF2QMDE6oGJRcJAACA/wXbFbYGCR0HugYtAwcaDBs5AwViAl4IZgJqAh0fxgYVygYJHQfOBi0DBxoMG0sF3RXaBgkdB94GLQMHIgwbUR1z5gYV6gYJHQfuBi0DB0YMNYcdcgL2BhX6BgkdB/4GLQMHRgwrhx12AgYHFQoHCR0HDgctAwdGDBuJAwMTFgclFwkAAAAABd8VIgcJHQcmBy0DB1IMTXMDBWICYghmAmoCFTIHCR0HNgctAwdSDBudHXICPgcVQgcJHQdGBy0DB14MKUcddgJOBxVSBwkdB1YHLQMHXgwZSR0xXgcVYgcJHQdmBy0DB2IMLUkdf24HFXIHCR0HdgctAwdiDBtJFX4HCR0HggctAwdmDAktFYoHCR0HjgctAwdmDDFVFZYHCR0HmgctAwd2DBU3HRWiBxWmBwkdB6oHLQMHggwRMwMDE7IHJQUJAAAAABW6BwkdB74HLQMHhgwbYQMHOgIOAj4CfUICfR1zygcVzgcJHQfSBy0DB44MHXsV2gcJHQfeBy0DB5IMQ2MdMeYHFeoHCR0H7gctAweSDC9jHX/2BxX6BwkdB/4HLQMHkgwvdRUGCAkdBwoILQMHkgwJKRUSCCkdNRYILQMHUgszbRUeCCkdNSIILQMHUgsJLRUqCCkdNS4ILQMHVgszgxU2CCkdNToILQMHVgsJLRVCCCkdNUYILQMHWgszbRVOCCkdNVIILQMHWgsJLSNhcml0aC5mYXN0bWF0aDxub25lPgAjYXJpdGgub3ZlcmZsb3c8bm9uZT4AI3ZlY3Rvci5raW5kPG1heGltdW1mPgAjdmVjdG9yLmtpbmQ8YWRkPgABAgIDJwUCIAIECRcGAgcFCQkTwQsBCQECBBf7BwUCIAIEH8EnBQIgAiAJAUEX+wUCIAIECcMnAwIgCScFAiACBB8nBwUCIAIEHycFAiACIAEHF/sFAiACBAHDJwUCIAUJBRsBAQEHBw8PDyEVFRUPAQULAQEBBwcHAQEBBQsBAQEHBwUBAScFAiACBAEnBQIgAiALBEIdBQERAZICBwMBJQ0RAZ4CBwNNdxsBAQEBAQEHAQcBDwEPAQ8BIQEVARUBFQEPAQMD5wsDAREH5+sDCwUFGx0GHgQDAQMdAwM5CwMBEQc5pQMLBR8hHxQ5AyMJAx9NAwOGAkEDCQkGhgIDBQNNAwOHBQMDAwOHBQMDBQaHAwUHF1FTCwWHIQlPF1FTAwOKAloCAwkJBooCAwUDVwMDiQUDAwMDiQUDAwUGiQMFBxNbXQsFiSEJWRNbXQMDjgJBAwkJBo4CAwUDYQMDiwUDAwMDiwUDAwUGiwMFBxVlZwsFiyEJYxVlZxkAOQMBBRkAOQMDbQUDAwcGbQMDAwMHBm0DAwMFFQZtAxMJCSUnKRcGNgQDAQMrAwPxCwMBEQfxmwMLBS0vAwNvBQMDBwZvAwMDAwcGbwMDAwUVBm8DEwkHMzU3FwZKBAMBAzkdBlYEAwEDMQMDOwsDAREHO6UDCwU9Px8UOwNBCQObigIDAyoCCwMBAwMuArMDASMHLgJDAwEFTU8DA7UFAwMDA7UFAwMFBrUDBQcTU1UDA7cFAwMDA7cFAwMFBrcDBQcVWVsDA0UFAwMDA0UFAwMDA0UFAwMFBkUDGwkLX2FjEwZFAxkDZQMDRwUDAwcGRwMDA1EDA0cFAwMFBkcDGwkNaWttEwZHAxkDbwMDMgKWBQMRJQcyAqoFAxEHZ3FzAwNGArMDASMHRgJDAwEFO3cDA0oCswMBIwdKAkMDAQVNeycHxgVDAwEFeX01A9oF1gUDHQkGUgIDHQN/JwdSAkMDHQWDgQMDuQUDAwMDuQUDAwUGuQMrBxGHiRsHKgZWAgMdA4sRB0IGOgYDLQWNhQMDigZaAgMJCQaaBgMRA5E3Bp4GAxEHj3WTAwNeAqYGAxcpB14CvgYDFwWVlxMGwgYDIwOZCQZuAgMFA5s5B24CFwMFBVedGwfiBlYCAxEDnysH8gYXAxEFlaEtBwIHFwMRA6MDA3oCEgcDFykHegIqBwMXBaWnEwZ+AgMjA6kJBn4CAwUDqysHOgcXAwUFV58tB0oHFwMFA68hB1oHFwMFBbFdLwdqBxcDBQWtswMDgQUDAwMDgQUDAwUGgQMFBxO3uQsFgSEJnxO3uQMDgwUDAwMDgwUDAwUGgwMFBxW9vwsFgyEJtRW9vwMDSwUDAwcGSwMDA1EDA0sFAwMFBksDGwkPw8XHEwZLAxkDyTsGngcDBQPLAwOCAq4HAwUlB4ICwgcDBQelzc8bB8YH9wMFA7EDA70FAwMDA70FAwMFBr0DBQcX1dchB+IHFwMFBdPZLwfyBxcDBQXb0QMDhQUDAwMDhQUDAwUGhQMFBxff4QsFhSEJ3Rff4QMDKgLzAwEZADsDAQUZADsDA/XzAwERB/XrAwsFBUMdBmoEAwEDRQMDPwsDAREHP6UDCwVHSR8UPwNLCQNDmQMDqwUDAwMDqwUDAwUGqwMFBxVNTwMDrYIEAwkJBq0DBQNTMQetFwMFBVVRGweeBPcDBQNXAwOxBQMDAwOxBQMDBQaxAwUHF1tdIQe6BBcDBQVfWTMGygQDGQNhAwMrBQMDAwMrBQMDAwMrBQMDBQYrAxsJGWVnaRMGKwMZA2sTBisDGwNjCwUr5gQLbxllZ2kDAx4CQQMJCQYeAgMFA3EDA3UFAwMDA3UFAwMFBnUDBQcTdXcLBXUhCXMTdXcDAyICQQMJCQYiAgMFA3sDA3cFAwMDA3cFAwMFBncDBQcVf4ELBXchCX0Vf4EDAyYCQQMJCQYmAgMFA4UDA3kFAwMDA3kFAwMFBnkDBQcXiYsLBXkhCYcXiYsZAD8DAQUZAD8PAAENEQEKAwcDDw8LAQEBAQEBBwEHAQMDAQsDAQMDAQsDAQ8EAQcBAwsNEQEOAwcDJz8LAQEBAQEBBwEHAQMDaQUDAwcGaQMDAwMHBmkDAwMFFQZpAxMJCQsNDxcG8gMDAQMRAwPlCwMBEQflmwMLBRMVAwNrBQMDBwZrAwMDAwcGawMDAwUVBmsDEwkHGRsdFwYGBAMBAx8DAwELAwEDAwELAwEPBAEHASEjDREBEgMHAyc/CwEBAQEBAQcBBwEDA1UFAwMHBlUDAwMDBwZVAwMDBRUGVQMTCQkLDQ8XBpoDAwEDEQMD3wsDAREH35sDCwUTFQMDZwUDAwcGZwMDAwMHBmcDAwMFFQZnAxMJBxkbHRcG1gMDAQMfAwMBCwMBAwMBCwMBDwQBBwEhIw0RARYDBwMPDwsBAQEBAQEHAQcBAwMBCwMBAwMBCwMBDwQBBQMLDREBGgMHAxETCwEBAQEBAQcBBwEDAwELAwEDAwELAwEDAwELAwEPBAEFCw0NEQEeAwcDERMLAQEBAQEBBwEHAQMDAQsDAQMDAQsDAQMDAQsDAQ8EAQULDQ0RASIDBwMREwsBAQEBAQEHAQcBAwMBCwMBAwMBCwMBAwMBCwMBDwQBBQsNDREBJgMHAw8PCwEBAQEBAQcBBwEDAwELAwEDAwELAwEPBAEHAQMLBgMBBQEA6h/hGQsZFQ0qAmUJDR1JDRMLX0ETI2U/JTM1Xx0jISMpMS0LCx8LHR0lGxEpDQkZGRkZGRkZGQsVDQkdCxEVdx1LMwsvHSUlHQ0TLQ1JC0syAhcfGxMbFxcTFy8XFxcXDxkXFRkZJRcZFSMjIxkfDw8NCR0RYnVpbHRpbgBzdGFibGVfbW9zYWljAHRwdQBhcml0aAB2ZWN0b3IAbW9kdWxlAGFyaXRoLmNvbnN0YW50AHZlY3Rvci5sb2FkAGFyaXRoLmluZGV4X2Nhc3QAdmVjdG9yLmJyb2FkY2FzdAB0cHUudmVjdG9yX3N0b3JlAGZ1bmMuZnVuYwBmdW5jLnJldHVybgBhcml0aC5jbXBpAHZlY3Rvci5zaGFwZV9jYXN0AG1lbXJlZi5sb2FkAGFyaXRoLmV4dHNpAHNjZi55aWVsZAB0cHUucmVwZWF0AGFyaXRoLmV4dHVpAHNjZi5pZgBhcml0aC5tdWxmAGFyaXRoLm11bGkAdHB1Lm1hdG11bABhcml0aC5hZGRpAHZlY3Rvci5tdWx0aV9yZWR1Y3Rpb24AYXJpdGguc3ViZgBtYXRoLmV4cABhcml0aC5hZGRmAGFyaXRoLmRpdmYAYXJpdGgudHJ1bmNmAHRwdS5pb3RhAGFyaXRoLnNlbGVjdABhcml0aC5tYXhpbXVtZgBhcml0aC5leHRmAC9ob21lL3B0b3VsbWUvbWluaWNvbmRhMy9lbnZzL3ZsbG0vbGliL3B5dGhvbjMuMTIvc2l0ZS1wYWNrYWdlcy9qYXgvZXhwZXJpbWVudGFsL3BhbGxhcy9vcHMvdHB1L3NwbGFzaF9hdHRlbnRpb24vc3BsYXNoX2F0dGVudGlvbl9rZXJuZWwucHkAZmxhc2hfYXR0ZW50aW9uX2tlcm5lbC48bG9jYWxzPi5ib2R5AC9nZXQAZmxhc2hfYXR0ZW50aW9uX2tlcm5lbC48bG9jYWxzPi5lbmQAdmFsdWUAL2NvbnZlcnRfZWxlbWVudF90eXBlAHN5bV9uYW1lAC9zd2FwAGZ1bmN0aW9uX3R5cGUAL2Jyb2FkY2FzdF9pbl9kaW0AdHJhbnNmb3JtX2luZGljZXMAd2luZG93X2JvdW5kcwBmbGFzaF9hdHRlbnRpb25fa2VybmVsAC9tdWwAX2FwcGx5X21hc2tfYW5kX3NvZnRfY2FwAGZsYXNoX2F0dGVudGlvbl9rZXJuZWwuPGxvY2Fscz4uaW5pdABfbmV4dF9ub256ZXJvAC9ob21lL3B0b3VsbWUvanVzdGFieXRlL3RwdV9wYWxsYXNfcG9zdC8uL3BhbGxhc19rZXJuZWwucHkAcHJlZGljYXRlAC9yZXBlYXQAL2FkZABwaXBlbGluZV9tb2RlAC9ndAAvY29uZABkaW1lbnNpb24AbWFpbgB0cmFuc2Zvcm1fMAB0cmFuc2Zvcm1fMQB0cmFuc2Zvcm1fMgB0cmFuc2Zvcm1fMwB0cmFuc2Zvcm1fNAB0cmFuc2Zvcm1fNQB0cmFuc2Zvcm1fNgB0cmFuc2Zvcm1fNwAvZXEAdGltZXMAb3BlcmFuZFNlZ21lbnRTaXplcwBzdHJpZGVzAC9kb3RfZ2VuZXJhbABkaW1lbnNpb25fbnVtYmVycwB0cmFuc3Bvc2VfbGhzAHRyYW5zcG9zZV9yaHMAa2luZAByZWR1Y3Rpb25fZGltcwAvc3ViAC9leHAAc3RhYmxlX21vc2FpYy52ZXJzaW9uAHNwbGFzaF9taGFfZndkX2M2MDJiMTZkAGRpbWVuc2lvbl9zZW1hbnRpY3MAaXRlcmF0aW9uX2JvdW5kcwBzY2FsYXJfcHJlZmV0Y2gAc2NyYXRjaF9vcGVyYW5kcwB3aW5kb3dfcGFyYW1zAF9zcGxhc2hfYXR0ZW50aW9uX2ZvcndhcmQuPGxvY2Fscz4udl9pbmRleF9tYXAAX3NwbGFzaF9hdHRlbnRpb25fZm9yd2FyZABfc3BsYXNoX2F0dGVudGlvbl9jdXN0b20AX3NwbGFzaF9hdHRlbnRpb24AU3BsYXNoQXR0ZW50aW9uS2VybmVsLl9fY2FsbF9fAGJlbmNobWFya19rZXJuZWwuPGxvY2Fscz4uc3BsYXNoX2F0dGVudGlvbl9rZXJuZWwAYmVuY2htYXJrX2tlcm5lbAA8bW9kdWxlPgBfbmV4dF9ub256ZXJvLjxsb2NhbHM+LjxsYW1iZGE+AF9zcGxhc2hfYXR0ZW50aW9uX2ZvcndhcmQuPGxvY2Fscz4ua19pbmRleF9tYXAAL2RpdgBmYXN0bWF0aAAvc2NhbgBmbGFzaF9hdHRlbnRpb25fa2VybmVsLjxsb2NhbHM+LnJ1bgBvdmVyZmxvd0ZsYWdzAC9pb3RhAC9nZQBDYXVzYWxNYXNrLl9faW5pdF9fLjxsb2NhbHM+LmNhdXNhbF9tYXNrX2Z1bmN0aW9uAC9ob21lL3B0b3VsbWUvbWluaWNvbmRhMy9lbnZzL3ZsbG0vbGliL3B5dGhvbjMuMTIvc2l0ZS1wYWNrYWdlcy9qYXgvZXhwZXJpbWVudGFsL3BhbGxhcy9vcHMvdHB1L3NwbGFzaF9hdHRlbnRpb24vc3BsYXNoX2F0dGVudGlvbl9tYXNrLnB5AC9waml0AC9zZWxlY3RfbgAvcmVkdWNlX21heAAvbWF4AC9yZWR1Y2Vfc3VtAA==\&quot;, \&quot;serialization_format\&quot;: 1, \&quot;needs_layout_passes\&quot;: true}}\n  ROOT %get-tuple-element.0 = bf16[8,2048,128]{2,1,0} get-tuple-element(%custom-call.0), index=3, metadata={op_name=\&quot;jit(splash_attention_kernel)/jit(main)/splash_attention_kernel/splash_kernel_b1024_h8_s2048/jit(_splash_attention)/splash_mha_fwd_c602b16d/pallas_call\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=1100}\n}\n</code></pre>\n</article>\n  </div>\n\n  </div>\n</div>\n\n      </div>\n      <div class=\&quot;gist-meta\&quot;>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/d9d15b515f57fe61773ac44a17b45973/raw/d23a342dbe066aa131ba40c3a1ec61583e2503c9/pallas_attn.md/%22 style=\&quot;float:right\&quot; class=\&quot;Link--inTextBlock\&quot;>view raw</a>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/d9d15b515f57fe61773ac44a17b45973#file-pallas_attn-md\%22 class=\&quot;Link--inTextBlock\&quot;>\n          pallas_attn.md\n        </a>\n        hosted with &amp;#10084; by <a class=\&quot;Link--inTextBlock\&quot; href=/__u/patricktoulme.substack.com/%22https://github.com/%22>GitHub</a>\n      </div>\n    </div>\n</div>\n&quot;,&quot;stylesheet&quot;:&quot;https://github.githubassets.com/assets/gist-embed-68783a026c0c.css&quot;}" data-component-name="GitgistToDOM"><link rel="stylesheet" href="https://github.githubassets.com/assets/gist-embed-68783a026c0c.css"><div id="gist144398218" class="gist">
    <div class="gist-file" data-color-mode="light" data-light-theme="light">
      <div class="gist-data">
        <div class="js-gist-file-update-container js-task-list-container">
  <div id="file-pallas_attn-md" class="file my-2">
      <div id="file-pallas_attn-md-readme" class="Box-body readme blob p-5 p-xl-6 " style="overflow:auto">
    <article class="markdown-body entry-content container-lg" itemprop="text"><pre><code>HloModule jit_splash_attention_kernel, entry_computation_layout={(bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)})-&gt;bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}

ENTRY %main.28 (Arg_0.1: bf16[8,2048,128], Arg_1.2: bf16[8,2048,128], Arg_2.3: bf16[8,2048,128]) -&gt; bf16[8,2048,128] {
  %constant.4 = s8[1,2,2]{2,1,0} constant({ { { 0, 0 }, { 0, 1 } } })
  %constant.5 = s8[1,2,2]{2,1,0} constant({ { { 1, 0 }, { 2, 1 } } })
  %Arg_0.1 = bf16[8,2048,128]{2,1,0} parameter(0), metadata={op_name="q"}
  %Arg_1.2 = bf16[8,2048,128]{2,1,0} parameter(1), metadata={op_name="k"}
  %Arg_2.3 = bf16[8,2048,128]{2,1,0} parameter(2), metadata={op_name="v"}
  %constant.6 = s32[2048]{0} constant({...})
  %broadcast.0 = s32[2048,128]{1,0} broadcast(%constant.6), dimensions={0}, metadata={op_name="jit(splash_attention_kernel)/jit(main)/splash_attention_kernel/splash_kernel_b1024_h8_s2048/jit(_splash_attention)/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=1045}
  %custom-call.0 = (f32[1024,128]{1,0}, f32[1024,128]{1,0}, f32[1024,128]{1,0}, bf16[8,2048,128]{2,1,0}) custom-call(%constant.4, %constant.5, %Arg_0.1, %Arg_1.2, %Arg_2.3, /*index=5*/%broadcast.0), custom_call_target="tpu_custom_call", operand_layout_constraints={s8[1,2,2]{2,1,0}, s8[1,2,2]{2,1,0}, bf16[8,2048,128]{2,1,0}, bf16[8,2048,128]{2,1,0}, bf16[8,2048,128]{2,1,0}, s32[2048,128]{1,0}}, metadata={op_name="jit(splash_attention_kernel)/jit(main)/splash_attention_kernel/splash_kernel_b1024_h8_s2048/jit(_splash_attention)/splash_mha_fwd_c602b16d/pallas_call" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=1100}, backend_config={"custom_call_config": {"body": "TUzvUgFNTElSMjEuMC4wZ2l0AAFRCwEDBQcJAQMLAzsNDxETFRcZGx0fISMlJykrLS8xMzU3OTs9P0FDRQMKCWYILwH7BwsXCxMXCxMLCwsbCwsLCysLCxMTExcLCwsLCw8PDw8XGxMTFxNzUwsLExcXFxcXCxMLExMTExMTCxMTExcPCxMTExMTEwsTDxcTCxMTExMPCxcPDxMXCxMXExMTExMPpYUPCwsLCwsLCwsLFxMPExcPExMLExMPExMTGwsFD2GBkY1hKgIqAgFGBgsLExMTExcTFwsLCwsTExMTHxcXCwszFwsLFxMXExMTIwsLawsbCwtzCw8LC0sfCx8LHwsfCycLJwsnCx8LGxsbGxsbGxsPDw8fCycLJxMLJxMLJxMLJxMLGxMLGxMfFwsfExMLHw8fExMTEw8PHw8fExMPHw8LJxMTDw8PExMPExMfDxMfDw8nExMPDw8TEw8PHxMTHw8THxMTHxcfCxMTHwsTExMfExMfExMTHxMTEx8TEx8rUw8TEx8TEx9TExMfExMfExMfExMfCw8LHxMTEx8PCxMTHxMTHxMTHxMTHxcfExMfMxMTHxMTHxMTEx8TFwsTEycfExMTExMTDxMTJxMTHw8TExMnFw8XCxcXCyMLFxMfFxMTExMTDx8TCxMfDxMLFx8LExMfKxMTEx8LExMfExMTHxcTEx8XExMfFx8LExMfKxMTHxcTEx8XExMfExMTHxMTEx8TEx8TEx8TEx8TExMfFx8TEx8zExMTHxMTHxMTEx8TExMfExMfExMfExMfExMfExMfExMfExMfBwVZWQkFXUkBLw8HHycHCw8rHwsnFx8jHwcnG0MvKx8fAlovHwVHAwMTLgMFSRV7XgUDAxMqAwVLFW4EJwVNBU8FUQMDmgRWCAVTBVUFVwVZAwUSAgoFFgIaAgVbBV0VV0oDFSIEJx0b2gQdogOmAwVfBWEFYwVlBWcdoykdo6kRAQUdow8DAxPuBAMDYgVaCB0NfgUdDYoFFU4C8gUdDZIHIw0HMQEAAAAAAAAAAAQAAAAAAACAAAAAAAAAACMNBSEABAAAAAAAAIAAAAAAAAAADScNKR0NMgMdQgNGAx1OA1IDHVoDXgMdZgNqAx1yA3YDBWkVV7IDBWsdDc4DHQ3mAx0NAgQdDS4EHQ1GBB0vWgQFbR0b/gQdGxoFHRsyBR1GBUoFEQsABW8dG3oHHRuGBx0bAggdGxoIHRsyCB0bSggFcR03NgMV2ScdfgOCAx03rgMFcx03ygMDA2W/HTfSAx034gMV4ScFdQMDZSoEFe0nFXEnHQ12BB2KBI4EBXcdDa4EAwMTTgUdDWYFHQ1yBR0NGgYVkgZJHQ3WBxENEWFmZmluZV9tYXA8KGQwLCBkMSwgZDIpIC0+IChkMCwgZDEsIGQyKT4AYWZmaW5lX21hcDwoZDAsIGQxKSAtPiAoZDAsIGQxKT4AEQ0BBXkFewV9BX8FgQWDBYUFhwWJHToDPgMdx4oDFdljHZfGAx3qA+4DFeFjHZf+Ax3pEgQFiwMDZcUdLzIEFe1jHZdCBAMDEz0d6V4EAwWvPfk9BY0jdHB1Lm1lbW9yeV9zcGFjZTx2bWVtPgAjdHB1LnBpcGVsaW5lX21vZGU8c3luY2hyb25vdXM+ACN0cHUuZGltZW5zaW9uX3NlbWFudGljczxhcmJpdHJhcnk+ACN0cHUuZGltZW5zaW9uX3NlbWFudGljczxwYXJhbGxlbD4AI3RwdS5tZW1vcnlfc3BhY2U8c21lbT4AI3RwdS5kb3RfZGltZW5zaW9uX251bWJlcnM8WzFdLCBbMV0sIFswXSwgWzBdLCBbMCwgMCwgMSwgMF0sIFtdLCBbXT4AI3RwdS5kb3RfZGltZW5zaW9uX251bWJlcnM8WzFdLCBbMF0sIFswXSwgWzFdLCBbMCwgMCwgMSwgMV0sIFtdLCBbXT4ABY8FkSMBAQEdH/IEHR8OBR0fJgUdPgVCBR0xUgUdNgKeBQWTBZUFlwWZHTGuBR0xugUdB+4FHX8OBgMFrz35JgYDAxOGBh2uBrIGBZsFnSMNAxEBAAAAAAAAAB3SBtYGBZ8FoR0aBx4HHR8uBx02ArYHHR8OCB0fJggdHz4IAwWWAr8ZmgIFowWlAw+iAqYCHaoCrgKyArYCugK+AsUZx8ICxgIFpwEHAgL//w0lBakjDQcxCAAAAAAAAAACAAAAAAAAAAIAAAAAAAAABasRDQkFrQWvARHKAtIC2gLiAuoC8gL6AgIDAwUjzgIlTQnJAwUj1gIlTQnLAwUj3gIlTQnNAwUj5gIlTwnPAweN/SPuAiVPCdEDB439I/YCJU8J0wMHjf0j/gIlTwnVAwUjBgMlTQnXAwUdURnJAwUdURnLAwUdURnNAwUdUxnPAwUdUxnRAwUdUxnTAwUdUxnVAwUdURnXEQEBEQMBFY+RLQMH5ggtWwWxLQMJRg8jTg8LBbMtAwkyER3KEQsVWVYDBbUtAwmiEhPaEgcVW2IDBbctAwkiJBNiJAcVXW4DBbktAwmuJBfKJAsVX3oDBbstYQeBJ08Vk4YDBb0tYQeNI0EV244DLWEHZgJBvR2SA5YDBb8tYQf+AgkVHRWeAxUtqgMFwS0DB94IK1EVld0tAwfmCB9dFVm2AxVbugMVXb4DFV/CAxWT2xWZkS0DB+YIH2UVnZEtAwcCCSVRHRXaAxUt3gMVn90tAwcCCRdTFY+hBcMtAwkmDyMuDwsdFfYDFS36AxWV4xWZoRWdoR0VCgQVLQ4EFZ/jFRYEJx0vGgQtAwdKCxcjHRUpHS8mBC0DB0oLByURDQUVj6ctAwliC2l+CwcdFToEFS0+BBWV7xWZpxWdpx0VTgQVLVIEFZ/vHRWpLQMHmgwHLRViBCcdL2YELQMHugwXPR0VDx0vcgQtAwe6DAc/FXoEDx0RfgQtAwfCDBE1AwMThgQTCRAAAOAPBcUVkgQPHRGWBC0DB8YMM0EFxx1zogQVpgQPHRGqBC0DB8YMGXsVsgQPHRG2BC0DB8oMJUkdMb4EFcIEDx0RxgQtAwfKDCVZHRXOBBXSBA8dEdYELQMHygwjgxXeBA8dEeIELQMHygwJHQMFEgLqBBYCGgIjAQkhAQAAAAEAAAADAAAAAAAAABMJARX2BA8dEfoELQMH5gwzbRUCBQ8dEQYFLQMH5gwJLSMBCSEBAAAAAQAAAAIAAAAAAAAAFRIFDx0RFgUtAwfqDDNtFR4FDx0RIgUtAwfqDAktFSoFDx0RLgUtAwfuDDNtFTYFDx0ROgUtAwfuDAktBckVe6kFyy0DB7IMCXERAQIgFVYFCR0HWgUtAweKCyllFXFjBc0VagUJHQduBS0DB44LK08VdgUJHQd6BS0DB44LU3cVggUJHQeGBS0DB54LESUVjgUJHQeSBS0DB6oLFTcDAxOaBSURCQAAAAAVogUJHQemBS0DB7YLE48DBzoCCgI+An1CAn0VsgUJHQe2BS0DB/4LI00VvgUJHQfCBS0DB/4LU48df8oFFc4FCR0H0gUtAwf+CyOPAwOvPR3eBeIFBc8V5gVJHTPqBS0DCaoJPbIJDy0DBxIME0UVe/YFFXH6BRVX/gUVWQIGFVsGBhVdCgYVX5MVEgZJHTMWBi0DCaoJJ7IJDxUeBkkdMyIGLQMHxgkVOxEBIR1zLgYVMgZJHTM2Bi0DCcIJJ8oJDwMDZT4GEQ0VHUYGSgYF0RVOBl4GHVIGVgYF0y1aBgcCBR89BdUVYgZqBh0zZgYtAwf2CSlzFU4CbgYVe3IGFXF2BhVXegYVWX4GFVuCBhVdXxMJkMzMzD8djga7BdcdM5YGLQMHhgoTUR0fux2iBrsF2QMDE6oGJRcJAACA/wXbFbYGCR0HugYtAwcaDBs5AwViAl4IZgJqAh0fxgYVygYJHQfOBi0DBxoMG0sF3RXaBgkdB94GLQMHIgwbUR1z5gYV6gYJHQfuBi0DB0YMNYcdcgL2BhX6BgkdB/4GLQMHRgwrhx12AgYHFQoHCR0HDgctAwdGDBuJAwMTFgclFwkAAAAABd8VIgcJHQcmBy0DB1IMTXMDBWICYghmAmoCFTIHCR0HNgctAwdSDBudHXICPgcVQgcJHQdGBy0DB14MKUcddgJOBxVSBwkdB1YHLQMHXgwZSR0xXgcVYgcJHQdmBy0DB2IMLUkdf24HFXIHCR0HdgctAwdiDBtJFX4HCR0HggctAwdmDAktFYoHCR0HjgctAwdmDDFVFZYHCR0HmgctAwd2DBU3HRWiBxWmBwkdB6oHLQMHggwRMwMDE7IHJQUJAAAAABW6BwkdB74HLQMHhgwbYQMHOgIOAj4CfUICfR1zygcVzgcJHQfSBy0DB44MHXsV2gcJHQfeBy0DB5IMQ2MdMeYHFeoHCR0H7gctAweSDC9jHX/2BxX6BwkdB/4HLQMHkgwvdRUGCAkdBwoILQMHkgwJKRUSCCkdNRYILQMHUgszbRUeCCkdNSIILQMHUgsJLRUqCCkdNS4ILQMHVgszgxU2CCkdNToILQMHVgsJLRVCCCkdNUYILQMHWgszbRVOCCkdNVIILQMHWgsJLSNhcml0aC5mYXN0bWF0aDxub25lPgAjYXJpdGgub3ZlcmZsb3c8bm9uZT4AI3ZlY3Rvci5raW5kPG1heGltdW1mPgAjdmVjdG9yLmtpbmQ8YWRkPgABAgIDJwUCIAIECRcGAgcFCQkTwQsBCQECBBf7BwUCIAIEH8EnBQIgAiAJAUEX+wUCIAIECcMnAwIgCScFAiACBB8nBwUCIAIEHycFAiACIAEHF/sFAiACBAHDJwUCIAUJBRsBAQEHBw8PDyEVFRUPAQULAQEBBwcHAQEBBQsBAQEHBwUBAScFAiACBAEnBQIgAiALBEIdBQERAZICBwMBJQ0RAZ4CBwNNdxsBAQEBAQEHAQcBDwEPAQ8BIQEVARUBFQEPAQMD5wsDAREH5+sDCwUFGx0GHgQDAQMdAwM5CwMBEQc5pQMLBR8hHxQ5AyMJAx9NAwOGAkEDCQkGhgIDBQNNAwOHBQMDAwOHBQMDBQaHAwUHF1FTCwWHIQlPF1FTAwOKAloCAwkJBooCAwUDVwMDiQUDAwMDiQUDAwUGiQMFBxNbXQsFiSEJWRNbXQMDjgJBAwkJBo4CAwUDYQMDiwUDAwMDiwUDAwUGiwMFBxVlZwsFiyEJYxVlZxkAOQMBBRkAOQMDbQUDAwcGbQMDAwMHBm0DAwMFFQZtAxMJCSUnKRcGNgQDAQMrAwPxCwMBEQfxmwMLBS0vAwNvBQMDBwZvAwMDAwcGbwMDAwUVBm8DEwkHMzU3FwZKBAMBAzkdBlYEAwEDMQMDOwsDAREHO6UDCwU9Px8UOwNBCQObigIDAyoCCwMBAwMuArMDASMHLgJDAwEFTU8DA7UFAwMDA7UFAwMFBrUDBQcTU1UDA7cFAwMDA7cFAwMFBrcDBQcVWVsDA0UFAwMDA0UFAwMDA0UFAwMFBkUDGwkLX2FjEwZFAxkDZQMDRwUDAwcGRwMDA1EDA0cFAwMFBkcDGwkNaWttEwZHAxkDbwMDMgKWBQMRJQcyAqoFAxEHZ3FzAwNGArMDASMHRgJDAwEFO3cDA0oCswMBIwdKAkMDAQVNeycHxgVDAwEFeX01A9oF1gUDHQkGUgIDHQN/JwdSAkMDHQWDgQMDuQUDAwMDuQUDAwUGuQMrBxGHiRsHKgZWAgMdA4sRB0IGOgYDLQWNhQMDigZaAgMJCQaaBgMRA5E3Bp4GAxEHj3WTAwNeAqYGAxcpB14CvgYDFwWVlxMGwgYDIwOZCQZuAgMFA5s5B24CFwMFBVedGwfiBlYCAxEDnysH8gYXAxEFlaEtBwIHFwMRA6MDA3oCEgcDFykHegIqBwMXBaWnEwZ+AgMjA6kJBn4CAwUDqysHOgcXAwUFV58tB0oHFwMFA68hB1oHFwMFBbFdLwdqBxcDBQWtswMDgQUDAwMDgQUDAwUGgQMFBxO3uQsFgSEJnxO3uQMDgwUDAwMDgwUDAwUGgwMFBxW9vwsFgyEJtRW9vwMDSwUDAwcGSwMDA1EDA0sFAwMFBksDGwkPw8XHEwZLAxkDyTsGngcDBQPLAwOCAq4HAwUlB4ICwgcDBQelzc8bB8YH9wMFA7EDA70FAwMDA70FAwMFBr0DBQcX1dchB+IHFwMFBdPZLwfyBxcDBQXb0QMDhQUDAwMDhQUDAwUGhQMFBxff4QsFhSEJ3Rff4QMDKgLzAwEZADsDAQUZADsDA/XzAwERB/XrAwsFBUMdBmoEAwEDRQMDPwsDAREHP6UDCwVHSR8UPwNLCQNDmQMDqwUDAwMDqwUDAwUGqwMFBxVNTwMDrYIEAwkJBq0DBQNTMQetFwMFBVVRGweeBPcDBQNXAwOxBQMDAwOxBQMDBQaxAwUHF1tdIQe6BBcDBQVfWTMGygQDGQNhAwMrBQMDAwMrBQMDAwMrBQMDBQYrAxsJGWVnaRMGKwMZA2sTBisDGwNjCwUr5gQLbxllZ2kDAx4CQQMJCQYeAgMFA3EDA3UFAwMDA3UFAwMFBnUDBQcTdXcLBXUhCXMTdXcDAyICQQMJCQYiAgMFA3sDA3cFAwMDA3cFAwMFBncDBQcVf4ELBXchCX0Vf4EDAyYCQQMJCQYmAgMFA4UDA3kFAwMDA3kFAwMFBnkDBQcXiYsLBXkhCYcXiYsZAD8DAQUZAD8PAAENEQEKAwcDDw8LAQEBAQEBBwEHAQMDAQsDAQMDAQsDAQ8EAQcBAwsNEQEOAwcDJz8LAQEBAQEBBwEHAQMDaQUDAwcGaQMDAwMHBmkDAwMFFQZpAxMJCQsNDxcG8gMDAQMRAwPlCwMBEQflmwMLBRMVAwNrBQMDBwZrAwMDAwcGawMDAwUVBmsDEwkHGRsdFwYGBAMBAx8DAwELAwEDAwELAwEPBAEHASEjDREBEgMHAyc/CwEBAQEBAQcBBwEDA1UFAwMHBlUDAwMDBwZVAwMDBRUGVQMTCQkLDQ8XBpoDAwEDEQMD3wsDAREH35sDCwUTFQMDZwUDAwcGZwMDAwMHBmcDAwMFFQZnAxMJBxkbHRcG1gMDAQMfAwMBCwMBAwMBCwMBDwQBBwEhIw0RARYDBwMPDwsBAQEBAQEHAQcBAwMBCwMBAwMBCwMBDwQBBQMLDREBGgMHAxETCwEBAQEBAQcBBwEDAwELAwEDAwELAwEDAwELAwEPBAEFCw0NEQEeAwcDERMLAQEBAQEBBwEHAQMDAQsDAQMDAQsDAQMDAQsDAQ8EAQULDQ0RASIDBwMREwsBAQEBAQEHAQcBAwMBCwMBAwMBCwMBAwMBCwMBDwQBBQsNDREBJgMHAw8PCwEBAQEBAQcBBwEDAwELAwEDAwELAwEPBAEHAQMLBgMBBQEA6h/hGQsZFQ0qAmUJDR1JDRMLX0ETI2U/JTM1Xx0jISMpMS0LCx8LHR0lGxEpDQkZGRkZGRkZGQsVDQkdCxEVdx1LMwsvHSUlHQ0TLQ1JC0syAhcfGxMbFxcTFy8XFxcXDxkXFRkZJRcZFSMjIxkfDw8NCR0RYnVpbHRpbgBzdGFibGVfbW9zYWljAHRwdQBhcml0aAB2ZWN0b3IAbW9kdWxlAGFyaXRoLmNvbnN0YW50AHZlY3Rvci5sb2FkAGFyaXRoLmluZGV4X2Nhc3QAdmVjdG9yLmJyb2FkY2FzdAB0cHUudmVjdG9yX3N0b3JlAGZ1bmMuZnVuYwBmdW5jLnJldHVybgBhcml0aC5jbXBpAHZlY3Rvci5zaGFwZV9jYXN0AG1lbXJlZi5sb2FkAGFyaXRoLmV4dHNpAHNjZi55aWVsZAB0cHUucmVwZWF0AGFyaXRoLmV4dHVpAHNjZi5pZgBhcml0aC5tdWxmAGFyaXRoLm11bGkAdHB1Lm1hdG11bABhcml0aC5hZGRpAHZlY3Rvci5tdWx0aV9yZWR1Y3Rpb24AYXJpdGguc3ViZgBtYXRoLmV4cABhcml0aC5hZGRmAGFyaXRoLmRpdmYAYXJpdGgudHJ1bmNmAHRwdS5pb3RhAGFyaXRoLnNlbGVjdABhcml0aC5tYXhpbXVtZgBhcml0aC5leHRmAC9ob21lL3B0b3VsbWUvbWluaWNvbmRhMy9lbnZzL3ZsbG0vbGliL3B5dGhvbjMuMTIvc2l0ZS1wYWNrYWdlcy9qYXgvZXhwZXJpbWVudGFsL3BhbGxhcy9vcHMvdHB1L3NwbGFzaF9hdHRlbnRpb24vc3BsYXNoX2F0dGVudGlvbl9rZXJuZWwucHkAZmxhc2hfYXR0ZW50aW9uX2tlcm5lbC48bG9jYWxzPi5ib2R5AC9nZXQAZmxhc2hfYXR0ZW50aW9uX2tlcm5lbC48bG9jYWxzPi5lbmQAdmFsdWUAL2NvbnZlcnRfZWxlbWVudF90eXBlAHN5bV9uYW1lAC9zd2FwAGZ1bmN0aW9uX3R5cGUAL2Jyb2FkY2FzdF9pbl9kaW0AdHJhbnNmb3JtX2luZGljZXMAd2luZG93X2JvdW5kcwBmbGFzaF9hdHRlbnRpb25fa2VybmVsAC9tdWwAX2FwcGx5X21hc2tfYW5kX3NvZnRfY2FwAGZsYXNoX2F0dGVudGlvbl9rZXJuZWwuPGxvY2Fscz4uaW5pdABfbmV4dF9ub256ZXJvAC9ob21lL3B0b3VsbWUvanVzdGFieXRlL3RwdV9wYWxsYXNfcG9zdC8uL3BhbGxhc19rZXJuZWwucHkAcHJlZGljYXRlAC9yZXBlYXQAL2FkZABwaXBlbGluZV9tb2RlAC9ndAAvY29uZABkaW1lbnNpb24AbWFpbgB0cmFuc2Zvcm1fMAB0cmFuc2Zvcm1fMQB0cmFuc2Zvcm1fMgB0cmFuc2Zvcm1fMwB0cmFuc2Zvcm1fNAB0cmFuc2Zvcm1fNQB0cmFuc2Zvcm1fNgB0cmFuc2Zvcm1fNwAvZXEAdGltZXMAb3BlcmFuZFNlZ21lbnRTaXplcwBzdHJpZGVzAC9kb3RfZ2VuZXJhbABkaW1lbnNpb25fbnVtYmVycwB0cmFuc3Bvc2VfbGhzAHRyYW5zcG9zZV9yaHMAa2luZAByZWR1Y3Rpb25fZGltcwAvc3ViAC9leHAAc3RhYmxlX21vc2FpYy52ZXJzaW9uAHNwbGFzaF9taGFfZndkX2M2MDJiMTZkAGRpbWVuc2lvbl9zZW1hbnRpY3MAaXRlcmF0aW9uX2JvdW5kcwBzY2FsYXJfcHJlZmV0Y2gAc2NyYXRjaF9vcGVyYW5kcwB3aW5kb3dfcGFyYW1zAF9zcGxhc2hfYXR0ZW50aW9uX2ZvcndhcmQuPGxvY2Fscz4udl9pbmRleF9tYXAAX3NwbGFzaF9hdHRlbnRpb25fZm9yd2FyZABfc3BsYXNoX2F0dGVudGlvbl9jdXN0b20AX3NwbGFzaF9hdHRlbnRpb24AU3BsYXNoQXR0ZW50aW9uS2VybmVsLl9fY2FsbF9fAGJlbmNobWFya19rZXJuZWwuPGxvY2Fscz4uc3BsYXNoX2F0dGVudGlvbl9rZXJuZWwAYmVuY2htYXJrX2tlcm5lbAA8bW9kdWxlPgBfbmV4dF9ub256ZXJvLjxsb2NhbHM+LjxsYW1iZGE+AF9zcGxhc2hfYXR0ZW50aW9uX2ZvcndhcmQuPGxvY2Fscz4ua19pbmRleF9tYXAAL2RpdgBmYXN0bWF0aAAvc2NhbgBmbGFzaF9hdHRlbnRpb25fa2VybmVsLjxsb2NhbHM+LnJ1bgBvdmVyZmxvd0ZsYWdzAC9pb3RhAC9nZQBDYXVzYWxNYXNrLl9faW5pdF9fLjxsb2NhbHM+LmNhdXNhbF9tYXNrX2Z1bmN0aW9uAC9ob21lL3B0b3VsbWUvbWluaWNvbmRhMy9lbnZzL3ZsbG0vbGliL3B5dGhvbjMuMTIvc2l0ZS1wYWNrYWdlcy9qYXgvZXhwZXJpbWVudGFsL3BhbGxhcy9vcHMvdHB1L3NwbGFzaF9hdHRlbnRpb24vc3BsYXNoX2F0dGVudGlvbl9tYXNrLnB5AC9waml0AC9zZWxlY3RfbgAvcmVkdWNlX21heAAvbWF4AC9yZWR1Y2Vfc3VtAA==", "serialization_format": 1, "needs_layout_passes": true}}
  ROOT %get-tuple-element.0 = bf16[8,2048,128]{2,1,0} get-tuple-element(%custom-call.0), index=3, metadata={op_name="jit(splash_attention_kernel)/jit(main)/splash_attention_kernel/splash_kernel_b1024_h8_s2048/jit(_splash_attention)/splash_mha_fwd_c602b16d/pallas_call" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=1100}
}
</code></pre>
</article>
  </div>

  </div>
</div>

      </div>
      <div class="gist-meta">
        <a href="https://gist.github.com/patrick-toulme/d9d15b515f57fe61773ac44a17b45973/raw/d23a342dbe066aa131ba40c3a1ec61583e2503c9/pallas_attn.md" style="float:right" class="Link--inTextBlock">view raw</a>
        <a href="https://gist.github.com/patrick-toulme/d9d15b515f57fe61773ac44a17b45973#file-pallas_attn-md" class="Link--inTextBlock">
          pallas_attn.md
        </a>
        hosted with &#10084; by <a class="Link--inTextBlock" href="https://github.com">GitHub</a>
      </div>
    </div>
</div>
</div><p>The key insight: XLA sees the entire Splash Attention kernel as one opaque custom-call. The <code>backend_config</code> contains the entire kernel body as base64-encoded MLIR (that <code>"TUzvUg..."</code> blob). XLA can&#8217;t optimize inside it &#8212; it just hands the blob to Mosaic for TPU-specific compilation.</p><p>Notice the scratch buffers in the return tuple: <code>f32[1024,128]</code> appears three times. These are the online softmax accumulators (<code>m_scratch</code>, <code>l_scratch</code>, <code>o_scratch</code>). The fourth element &#8212; <code>bf16[8,2048,128]</code> &#8212; is the actual output.</p><p>The <code>s8[1,2,2]</code> block masks encode which blocks to process. With <code>seq_len=2048</code> and <code>block_size=1024</code>, we have <code>2048/1024 = 2</code> blocks per dimension, hence the 2&#215;2 shape.</p><h2>Mosaic Compilation Pipeline</h2><p>When the HLO custom-call with <code>custom_call_target="tpu_custom_call"</code> reaches the TPU backend, the base64-encoded kernel body gets deserialized and fed to Mosaic. The compilation proceeds through 12 passes:</p><h3>Stage 1: Original MLIR</h3><p>The kernel arrives as high-level MLIR with explicit memory spaces and TPU-specific operations:</p><pre><code><code>func.func @main(
  %arg0: i32,  // grid index: head
  %arg1: i32,  // grid index: q_block
  %arg2: i32,  // grid index: kv_block
  %arg3: memref&lt;1x2x2xi8, #tpu.memory_space&lt;smem&gt;&gt;,   // block_mask
  %arg4: memref&lt;1x2x2xi8, #tpu.memory_space&lt;smem&gt;&gt;,   // data_next
  %arg5: memref&lt;1x1024x128xbf16, #tpu.memory_space&lt;vmem&gt;&gt;,  // Q block
  %arg6: memref&lt;1x1024x128xbf16, #tpu.memory_space&lt;vmem&gt;&gt;,  // K block
  %arg7: memref&lt;1x1024x128xbf16, #tpu.memory_space&lt;vmem&gt;&gt;,  // V block
  %arg8: memref&lt;1024x128xi32, #tpu.memory_space&lt;vmem&gt;&gt;,     // q_sequence (for causal)
  %arg9:  memref&lt;1024x128xf32, #tpu.memory_space&lt;vmem&gt;&gt;,    // m_scratch
  %arg10: memref&lt;1024x128xf32, #tpu.memory_space&lt;vmem&gt;&gt;,    // l_scratch
  %arg11: memref&lt;1024x128xf32, #tpu.memory_space&lt;vmem&gt;&gt;,    // o_scratch
  %arg12: memref&lt;1x1024x128xbf16, #tpu.memory_space&lt;vmem&gt;&gt;  // output
) attributes {
  dimension_semantics = [#tpu.dimension_semantics&lt;parallel&gt;,   // heads
                         #tpu.dimension_semantics&lt;arbitrary&gt;,  // q_blocks
                         #tpu.dimension_semantics&lt;arbitrary&gt;], // kv_blocks
  iteration_bounds = array&lt;i64: 8, 2, 2&gt;,  // 8 heads, 2 q_blocks, 2 kv_blocks
  scalar_prefetch = 2 : i64,  // prefetch first 2 args (mask info)
  window_params = [...]
}
</code></code></pre><p>The function attributes encode Pallas&#8217;s grid specification. <code>iteration_bounds = [8, 2, 2]</code> means the kernel runs 8 &#215; 2 &#215; 2 = 32 iterations. The first dimension is <code>parallel</code> (heads can run independently), while Q and KV blocks are <code>arbitrary</code> (have data dependencies through scratch buffers).</p><p>The kernel body shows the attention computation:</p><pre><code><code>// Load Q block, K block
%q = vector.load %arg5[%c0, %c0, %c0] : vector&lt;1x1024x128xbf16&gt;
%k = vector.load %arg6[%c0, %kv_idx, %c0] : vector&lt;1x1024x128xbf16&gt;

// QK^T matmul -&gt; [1024, 1024] attention scores
%qk = tpu.matmul %21, %24, %cst {
  dimension_numbers = #tpu.dot_dimension_numbers&lt;[1], [1], [0], [0], [0, 0, 1, 0], [], []&gt;
} : vector&lt;1024x128xbf16&gt;, vector&lt;1024x128xbf16&gt;, vector&lt;1024x1024xf32&gt;

// Apply causal mask using iota comparison
%29 = tpu.iota {dimension = 1 : i32} : vector&lt;1024x1024xi32&gt;
%34 = arith.cmpi sge, %33, %31 : vector&lt;1024x1024xi32&gt;
%36 = arith.select %34, %25, %35 : vector&lt;1024x1024xi1&gt;, vector&lt;1024x1024xf32&gt;

// Online softmax: max reduction
%37 = vector.multi_reduction &lt;maximumf&gt;, %36, %cst_20 [1] : vector&lt;1024x1024xf32&gt; to vector&lt;1024xf32&gt;
%40 = arith.maximumf %18, %39 : vector&lt;1024x128xf32&gt;  // update running max

// exp(scores - max) and sum reduction
%42 = arith.subf %36, %41 : vector&lt;1024x1024xf32&gt;
%43 = math.exp %42 : vector&lt;1024x1024xf32&gt;
%44 = vector.multi_reduction &lt;add&gt;, %43, %cst_21 [1] : vector&lt;1024x1024xf32&gt; to vector&lt;1024xf32&gt;

// Rescale previous accumulator: alpha = exp(m_prev - m_next)
%47 = arith.subf %18, %40 : vector&lt;1024x128xf32&gt;
%48 = math.exp %47 : vector&lt;1024x128xf32&gt;
%49 = arith.mulf %48, %19 : vector&lt;1024x128xf32&gt;  // alpha * l_prev
%50 = arith.addf %46, %49 : vector&lt;1024x128xf32&gt;  // l_next = l_curr + alpha * l_prev

// S @ V matmul for current output contribution
%57 = tpu.matmul %43, %56, %cst_28 {
  dimension_numbers = #tpu.dot_dimension_numbers&lt;[1], [0], [0], [1], [0, 0, 1, 1], [], []&gt;
} : vector&lt;1024x1024xf32&gt;, vector&lt;1024x128xf32&gt;, vector&lt;1024x128xf32&gt;

// Update output accumulator: o_next = alpha * o_prev + o_curr
%60 = arith.mulf %58, %59 : vector&lt;1024x128xf32&gt;  // alpha * o_prev
%61 = arith.addf %60, %57 : vector&lt;1024x128xf32&gt;  // o_next

// Store updated accumulators
tpu.vector_store %arg9[...], %40   // m_scratch
tpu.vector_store %arg10[...], %50  // l_scratch
tpu.vector_store %arg11[...], %61  // o_scratch
</code></code></pre><h3>Stage 2-3: Layout and Tiling</h3><p>The <code>infer-vector-layout</code> and tiling passes map abstract vectors onto the VPU&#8217;s physical structure (8 sublanes &#215; 128 lanes):</p><pre><code><code>// Before: abstract vector and memref
%qk = tpu.matmul %q, %k, %zeros : vector&lt;1024x1024xf32&gt;
%arg5: memref&lt;1x1024x128xbf16, #tpu.memory_space&lt;vmem&gt;&gt;

// After: explicit layouts matching hardware
%qk = tpu.matmul %q, %k, %zeros {
  in_layout = [#tpu.vpad&lt;"16,{0,0},(16,128)"&gt;,   // bf16: 16 sublanes
               #tpu.vpad&lt;"16,{0,0},(16,128)"&gt;],
  out_layout = [#tpu.vpad&lt;"32,{0,0},(8,128)"&gt;]   // f32: 8 sublanes
} : ...

%arg5: memref&lt;1x1024x128xbf16, 
  #tpu.tiled&lt;(8,128)(2,1)&gt;,  // 8&#215;128 tiles, 2 bf16 per 32-bit word
  #tpu.memory_space&lt;vmem&gt;&gt;</code></code></pre><p>The <code>(8,128)</code> tile shape matches the VPU dimensions. When layouts don&#8217;t match between operations, the compiler inserts explicit relayouts.</p><h3>Stage 4: Lower to LLO</h3><p>The <code>lower-to-llo</code> pass transforms everything to Low-Level Operations &#8212; the final MLIR dialect before machine code:</p><pre><code><code>// High-level vector store
tpu.vector_store %arg11[%c0, %c0], %result : memref&lt;1024x128xf32&gt;, vector&lt;1024x128xf32&gt;

// Becomes LLO with explicit addressing
%addr = "llo.saddr_scaled"(%arg11, %offset) &lt;{multiplier_in_bytes = 512}&gt;
llo.vector_store %result into %addr : vector&lt;8x128xf32&gt; into i32
llo.vector_store %result into %addr + %8 : vector&lt;8x128xf32&gt; into i32
// ... (128 stores for 1024 rows at 8 rows per tile)
</code></code></pre><p>The 1024&#215;128 vector gets broken into 128 tiles of 8&#215;128, each stored with computed VMEM addresses.</p><h2>Pallas LLO Compilation: From Mosaic to Machine Code</h2><p>Next, the Mosaic IR is converted to native TPU LLO IR. We will now trace the computation through the LLO compiler.</p><p>For this LLO walkthrough, we use <code>block_size=512</code> (grid 8&#215;4&#215;4) to show more loop iterations and better demonstrate software pipelining. The HLO and Mosaic sections above use <code>block_size=1024</code>.<br><br><strong>Kernel Structure (Pass 02: Original)</strong><br><br>The initial LLO reveals the kernel's memory layout and loop structure:</p><pre><code><code>$region0: #{splash_mha_fwd_e9701070.1}
  // Memory allocations
  #allocation1 [shape = 'u32[144,128]', space=vmem, size = 0x12000, tag = 'internal scratch']
  #allocation3 [shape = 'u8[512]', space=smem, size = 0x200, tag = 'prefetched SMEM operand 0']
  #allocation4 [shape = 'u8[512]', space=smem, size = 0x200, tag = 'prefetched SMEM operand 1']
  
  // Kernel inputs/outputs
  %s0 = inlined_call_operand.hbm [shape: s8[1,4,4], index: 0]   // block_mask
  %s1 = inlined_call_operand.hbm [shape: s8[1,4,4], index: 1]   // data_next
  %s2 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 2]  // Q
  %s3 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 3]  // K
  %s4 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 4]  // V
  %s5 = inlined_call_operand.vmem [shape: s32[2048,128], index: 5]    // q_sequence (causal)
  
  // Scratch buffers for online softmax
  %s6 = inlined_call_operand.hbm [shape: f32[512,128], index: 6]  // m_scratch output
  %s7 = inlined_call_operand.hbm [shape: f32[512,128], index: 7]  // l_scratch output
  %s8 = inlined_call_operand.hbm [shape: f32[512,128], index: 8]  // o_scratch output
  %s9 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 9]  // final output
  
  // Prefetch mask data to SMEM before main loop
  %16 = dma.hbm_to_smem /*hbm=*/%s0, /*size=*/16, /*smem=*/[#allocation3]
  %18 = dma.hbm_to_smem /*hbm=*/%s1, /*size=*/16, /*smem=*/[#allocation4]
  %19 = dma.done [#allocation2], 32
</code></code></pre><p>The kernel operand layout reflects Pallas&#8217;s <code>PrefetchScalarGridSpec</code>:</p><ul><li><p>Operands 0-1: Block mask indices prefetched to SMEM</p></li><li><p>Operands 2-4: Q, K, V matrices streaming from HBM</p></li><li><p>Operand 5: Causal sequence indices (pre-loaded to VMEM)</p></li><li><p>Operands 6-8: Online softmax scratch buffers</p></li></ul><h3>Loop Structure with Phi Nodes</h3><p>With <code>block_size=512</code> and <code>seq_len=2048</code>, the grid is <code>(8 heads, 4 q_blocks, 4 kv_blocks)</code>. Software pipelining adds 2 iterations for warmup/cooldown:</p><pre><code><code>loop: start=0, step=1, limit=130   // 8&#215;4&#215;4 = 128 + 2 for pipeline

$region107: #{splash_mha_fwd_e9701070.1} parent=1 // loop_body
  // Iteration tracking across 3 dimensions
  %s6473 = sphi 0, %s37    /* iteration index, stage = 0 */
  %s6475 = sphi 0, %s59    /* iter bound = 0 (head index) */
  %s6477 = sphi 0, %s55    /* iter bound = 1 (q_block index) */
  %s6479 = sphi 0, %s51    /* iter bound = 2 (kv_block index) */
  
  // Pipelined stage tracking (stage 1 lags stage 0)
  %s6481 = sphi 0, %s6475  /* stage = 1 iter bound = 0 */
  %s6483 = sphi 0, %s6477  /* stage = 1 iter bound = 1 */
  %s6485 = sphi 0, %s6479  /* stage = 1 iter bound = 2 */
  
  // Online softmax state (carries across KV blocks)
  %s6487 = sphi 0, %s66    /* running max state */
  %s6493 = sphi 0, %s128   /* running sum state */
  %s6499 = sphi 0, %s190   /* running output state */
  
  // Dimension wrap-around logic
  %s49 = sadd.s32 1, %s6479               /* kv_block + 1 */
  %p50 = scmp.ge.s32.totalorder %s49, 4   /* kv_block &gt;= 4? */
  %s51 = scalar_select %p50, 0, %s49      /* wrap to 0 */
  
  %s52 = sadd.s32 1, %s6477               /* q_block + 1 (if kv wrapped) */
  %s53 = scalar_select %p50, %s52, %s6477 /* conditional increment */
  %p54 = scmp.ge.s32.totalorder %s53, 4   /* q_block &gt;= 4? */
  %s55 = scalar_select %p54, 0, %s53      /* wrap to 0 */
  
  %s56 = sadd.s32 1, %s6475               /* head + 1 (if q wrapped) */
  %p58 = scmp.ge.s32.totalorder %s57, 8   /* head &gt;= 8? */
  %s59 = scalar_select %p58, 0, %s57      /* wrap to 0 */
</code></code></pre><p>The phi nodes track iteration state across the flattened 3D grid. The <code>stage = 1</code> variables implement software pipelining&#8212;while stage 0 loads the next tile, stage 1 processes the current tile.</p><h3>DMA Operations with Double Buffering</h3><p>The kernel uses double buffering for async DMA overlap:</p><pre><code><code>// Compute buffer slot for double buffering
%s14824_s0 = sand.u32 1, %s8466_s28   /* slot = iteration % 2 */
%s344_s11 = scalar_lea.sflag [#allocation6], %s15337_s19  /* sync flag */

// Compute HBM source address
%s6544_s15 = sshll.u32 %s353_s16, 6   /* address offset */
%s355_s3 = scalar_lea.hbm %s15333_s2, %s6544_s15

// Compute VMEM destination
%s347_s24 = scalar_lea.vmem [#allocation5], %s6541_s30
%s358_s1 = int_to_ptr.vmem %s357_s1

// Issue async DMA: HBM &#8594; VMEM
%7188 = dma.hbm_to_vmem [thread:$0] (!%p8805_p8), 
        /*hbm=*/%s355_s3, 
        /*size_in_granules=*/4096,   /* 256KB for 512&#215;128 bf16 */
        /*vmem=*/%s358_s1, 
        /*dst_syncflagno=*/%s344_s11
/* metadata:
   window_bounds: (1, 64, 1)    // tile shape
   iteration_bounds: (8, 4, 4)  // grid dimensions
   element_size: 2048 bytes */

// Wait for previous DMA before consuming buffer
%8381 = dma.done.wait (%p15353_p5), %s454_s2, 4096
</code></code></pre><p>The <code>(!%p8805_p8)</code> predicate enables block skipping&#8212;if the current tile is fully masked, the DMA is skipped entirely.</p><h3>Matmul Operations (Pass 13: Post-MXU-Assigner)</h3><p>After MXU assignment, all matmuls are initially on <code>mxu0</code>:</p><pre><code><code>// Q @ K^T for attention scores (line 749 in splash_attention_kernel.py)
%1237 = vmatprep.subr.mxu0 0
%1238 = vmatpush1.bf16.xpose.msra.mxu0 %v1164
%1269 = vmatprep.mubr.bf16.mxu0 0
%1270 = vmatmul.mubr.bf16.gmra.mxu0 %v1047

// S @ V for output accumulation (line 801)
%1279 = vmatprep.mubr.bf16.mxu0 0
%1280 = vmatmul.mubr.bf16.gmra.mxu0 %v1050
</code></code></pre><p>The operations decode as:</p><ul><li><p><code>vmatprep.subr</code>: Prepare RHS with subtraction (accumulator init)</p></li><li><p><code>vmatpush1.bf16.xpose</code>: Push bf16 tile to MXU, transposed</p></li><li><p><code>vmatmul.mubr.bf16.gmra</code>: Multiply-accumulate with bf16 inputs</p></li></ul><h3>Online Softmax in LLO</h3><p>The online softmax pattern appears in three stages:</p><p><strong>1. Max Reduction Tree:</strong></p><pre><code><code>// Tree reduction: max across 4 tiles
%v2530 = vmax.f32 %v9322, %v9326   /* max(tile0, tile1) */
%v2531 = vmax.f32 %v2530, %v9324   /* max(result, tile2) */
%v2532 = vmax.f32 %v2531, %v9343   /* max(result, tile3) */

// Cross-lane reduction using XLU
%2533 = vmax.xlane.f32.xlu0 %v2532   /* reduce across 128 lanes */
%v2534 = vpop.xlane.xlu0 %2533       /* pop result */
</code></code></pre><p><strong>2. Exponential via Base Conversion:</strong></p><pre><code><code>// exp(x) = 2^(x &#215; log&#8322;(e)) where log&#8322;(e) = 1.442695
%v3170 = vmul.f32 1.442695, %v2914   /* x &#215; log&#8322;(e) */
%7348 = vpow2.f32 %v3170             /* 2^result = exp(x) */
%v10295 = vpop.eup %7348             /* pop from EUP */
</code></code></pre><p><strong>3. Running Max Update (the key online softmax insight):</strong></p><pre><code><code>// m_next = max(m_prev, m_curr)
%v10156 = vmax.f32 %v10130, %v2539   /* running max update */

// scores - max (for numerical stability)
%v2918 = vsub.f32 %v9345, %v10156
%v2919 = vsub.f32 %v9347, %v10156
</code></code></pre><h3>Final VLIW Bundles (Pass 79)</h3><p>The final bundles show aggressive instruction-level parallelism. Here&#8217;s the entry bundle:</p><pre><code><code>0x0 : { 
  %s8468_s30 = smov [#allocation3]  ;;  // Load SMEM constant addresses
  %s8469_s12 = smov [#allocation4]  ;;
  %s14794_s0 = inlined_call_operand.hbm [shape: s8[1,4,4], index: 0]  ;;  // block_mask
  %s14795_s2 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 2]  ;;  // Q
  %s14796_s3 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 3]  ;;  // K
  %s14797_s4 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 4]  ;;  // V
  ...
} /* entry bundle */
</code></code></pre><h3>The Online Softmax Interleaving</h3><p><strong>The most impressive bundles show matmul, softmax, and mask operations executing in parallel:</strong></p><pre><code><code>0x20b : &gt; { 
  %v2535_v40 = vmax.f32 %v9345, %v9347  ;;        // Tree reduction: max
  %v1281_v41 = vpop.f32.mrf.mxu0  ;;               // Pop Q@K^T result from MXU0
  %7068 = vmatmul.mubr.bf16.gmra.mxu0 %v6597  ;;   // Start next matmul on MXU0
  %v2532_v42 = vmax.f32 %v2531, %v9343  ;;        // Continue max reduction
  %v9365_v48 = vsel %vm2025, %v1630, -2.38e+38  ;; // Apply causal mask
  %v1960_v34 = vld [vmem:[%s8897_s21 + $0x30]]  ;; // Load next Q sequence index
  %vm2041 = vcmp.ge.s32 %v1959, %v9314  ;;        // Compute next mask predicate
}
</code></code></pre><p><strong>In a single VLIW bundle:</strong></p><ol><li><p>Pop matmul result from MXU</p></li><li><p>Start next matmul</p></li><li><p>Tree reduction for running max</p></li><li><p>Apply causal mask via <code>vsel</code></p></li><li><p>Load next tile</p></li><li><p>Compute next mask predicate</p></li></ol><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!Pgpr!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 424w, /__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 848w, /__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 1272w, /__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!Pgpr!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png" width="1456" height="777" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:777,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:139681,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182968804?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 424w, /__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 848w, /__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 1272w, /__u/substackcdn.com/image/fetch/$s_!Pgpr!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F86f1d299-249a-42b5-9c72-b545623fa8db_1800x960.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p><strong>This is the power of Pallas</strong>&#8212;the online softmax operations interleave with matmul, eliminating the HBM round-trips that plague reference attention.</p><h3>Dual-MXU Distribution (Final Bundles)</h3><p>The VLIW scheduler distributes work across both MXUs:</p><pre><code><code>0x20c : &gt; { 
  %v1634_v45 = vpop.f32.mrf.mxu1  ;;              // Pop from MXU1
  %7148 = vmatmul.mubr.bf16.gmra.mxu1 %v6597  ;;  // Matmul on MXU1
  %7069 = vmatprep.mubr.bf16.mxu0 %v9334  ;;      // Prepare MXU0
  %v9367_v49 = vsel %vm2026, %v1281, -2.38e+38  ;; // Mask application
  %vm2042 = vcmp.ge.s32 %v1960, %v9305  ;;
}
</code></code></pre><h3>What Makes This Possible</h3><p>The key insight from the LLO is that <strong>everything happens in one loop</strong>:</p><ol><li><p><strong>DMA prefetch</strong> loads next Q/K/V tiles while current tiles compute</p></li><li><p><strong>Matmul</strong> computes attention scores (Q @ K^T)</p></li><li><p><strong>Online softmax</strong> updates running max/sum/output <em>in the same iteration</em></p></li><li><p><strong>S @ V</strong> accumulates output using rescaled softmax weights</p></li><li><p><strong>Scratch buffers</strong> in VMEM carry state across KV blocks</p></li></ol><p>Reference attention requires three separate fusions with HBM materialization between them. Pallas keeps everything in VMEM&#8212;the 128MB attention matrix never exists.</p><h3>The Compiler&#8217;s Role</h3><p>The TPU Pallas compiler handled:</p><ul><li><p><strong>MXU assignment</strong> (pass 13): Initially all on mxu0</p></li><li><p><strong>VLIW bundle scheduling</strong> (pass 29-33): Distributes across both MXUs</p></li><li><p><strong>DMA scheduling</strong> (pass 35-37): Overlaps memory access with compute</p></li><li><p><strong>Register allocation</strong> (pass 45-55): Manages VMEM pressure</p></li><li><p><strong>Final bundle packing</strong> (pass 73-79): Maximizes ILP</p></li></ul><p><strong>The Pallas programmer writes the algorithm; the compiler handles the hardware mapping.</strong> But the algorithm itself&#8212;online softmax with tiled accumulation&#8212;is something no compiler can currently invent. That&#8217;s why Pallas exists.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!v67X!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 424w, /__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 848w, /__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 1272w, /__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!v67X!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png" width="1456" height="890" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:890,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:153444,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182968804?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 424w, /__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 848w, /__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 1272w, /__u/substackcdn.com/image/fetch/$s_!v67X!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F04621d5c-6db2-478e-af7f-9df6ce54c191_1800x1100.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><h2>Reference Attention Jax</h2><p>The reference implementation lives in JAX&#8217;s Splash Attention library. It&#8217;s the naive O(n&#178;) memory baseline that materializes the full attention matrix:</p><pre><code><code>def _attention_reference_default(mask, q, k, v, segment_ids, mask_value, ...):
    # Q @ K^T -&gt; full [seq, seq] attention matrix
    logits = jnp.einsum("sd,td-&gt;st", q.astype(jnp.float32), k.astype(jnp.float32))
    
    # Apply causal mask - masked positions get a large negative value
    logits = jnp.where(mask, logits, mask_value)  # mask_value = -2.38e+38
    
    # Numerically stable softmax
    m = logits.max(axis=-1)              # reduce_max over key dimension
    s = jnp.exp(logits - m[..., None])      # subtract max for stability
    l = s.sum(axis=-1)                    # reduce_sum for normalization
    s = s / l[..., None]                   # normalize to probabilities
    
    # Output projection: S @ V
    o = jnp.einsum("st,td-&gt;sd", s, v.astype(jnp.float32))
    return o
</code></code></pre><h2>Reference Attention HLO</h2><h3>Initial HLO (Before Optimization)</h3><p>XLA traces the reference implementation into standard HLO ops. Here&#8217;s the complete structure:</p><div class="github-gist" data-attrs="{&quot;innerHTML&quot;:&quot;<div id=\&quot;gist144398154\&quot; class=\&quot;gist\&quot;>\n    <div class=\&quot;gist-file\&quot; translate=\&quot;no\&quot; data-color-mode=\&quot;light\&quot; data-light-theme=\&quot;light\&quot;>\n      <div class=\&quot;gist-data\&quot;>\n        <div class=\&quot;js-gist-file-update-container js-task-list-container\&quot;>\n  <div id=\&quot;file-ref_attn-md\&quot; class=\&quot;file my-2\&quot;>\n      <div id=\&quot;file-ref_attn-md-readme\&quot; class=\&quot;Box-body readme blob p-5 p-xl-6 \&quot;\n    style=\&quot;overflow: auto\&quot; tabindex=\&quot;0\&quot; role=\&quot;region\&quot;\n    aria-label=\&quot;ref_attn.md content, created by patrick-toulme on 07:26PM today.\&quot;\n  >\n    <article class=\&quot;markdown-body entry-content container-lg\&quot; itemprop=\&quot;text\&quot;><pre><code>HloModule jit_reference_attention, entry_computation_layout={(bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)})-&amp;gt;f32[8,2048,128]{2,1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}\n\n%region_0.25 (Arg_0.22: f32[], Arg_1.23: f32[]) -&amp;gt; f32[] {\n  %Arg_0.22 = f32[] parameter(0), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot;}\n  %Arg_1.23 = f32[] parameter(1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot;}\n  ROOT %maximum.24 = f32[] maximum(%Arg_0.22, %Arg_1.23), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=196}\n}\n\n%region_1.36 (Arg_0.33: f32[], Arg_1.34: f32[]) -&amp;gt; f32[] {\n  %Arg_0.33 = f32[] parameter(0), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot;}\n  %Arg_1.34 = f32[] parameter(1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot;}\n  ROOT %add.35 = f32[] add(%Arg_0.33, %Arg_1.34), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=198}\n}\n\nENTRY %main.47 (Arg_0.1: bf16[8,2048,128], Arg_1.2: bf16[8,2048,128], Arg_2.3: bf16[8,2048,128]) -&amp;gt; f32[8,2048,128] {\n  %constant.4 = pred[8,2048,2048]{2,1,0} constant({...})\n  %Arg_0.1 = bf16[8,2048,128]{2,1,0} parameter(0), metadata={op_name=\&quot;q\&quot;}\n  %convert.0 = f32[8,2048,128]{2,1,0} convert(%Arg_0.1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/convert_element_type\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=184}\n  %Arg_1.2 = bf16[8,2048,128]{2,1,0} parameter(1), metadata={op_name=\&quot;k\&quot;}\n  %convert.1 = f32[8,2048,128]{2,1,0} convert(%Arg_1.2), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/convert_element_type\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=184}\n  %dot.0 = f32[8,2048,2048]{2,1,0} dot(%convert.0, %convert.1), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={2}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(sd,td-&amp;gt;st)/dot_general\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=184}\n  %constant.0 = f32[] constant(-2.38197633e+38)\n  %broadcast.1 = f32[8,2048,2048]{2,1,0} broadcast(%constant.0), dimensions={}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}\n  %select.1 = f32[8,2048,2048]{2,1,0} select(%constant.4, %dot.0, %broadcast.1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}\n  %constant.1 = f32[] constant(-inf)\n  %reduce.0 = f32[8,2048]{1,0} reduce(%select.1, %constant.1), dimensions={2}, to_apply=%region_0.25, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=196}\n  %reshape.0 = f32[8,2048,1]{2,1,0} reshape(%reduce.0), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %broadcast.2 = f32[8,2048,1]{2,1,0} broadcast(%reshape.0), dimensions={0,1,2}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %reshape.1 = f32[8,2048]{1,0} reshape(%broadcast.2), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %broadcast.3 = f32[8,2048,2048]{2,1,0} broadcast(%reshape.1), dimensions={0,1}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %subtract.0 = f32[8,2048,2048]{2,1,0} subtract(%select.1, %broadcast.3), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %exponential.0 = f32[8,2048,2048]{2,1,0} exponential(%subtract.0), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/exp\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %constant.2 = f32[] constant(0)\n  %reduce.1 = f32[8,2048]{1,0} reduce(%exponential.0, %constant.2), dimensions={2}, to_apply=%region_1.36, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=198}\n  %reshape.2 = f32[8,2048,1]{2,1,0} reshape(%reduce.1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n  %broadcast.4 = f32[8,2048,1]{2,1,0} broadcast(%reshape.2), dimensions={0,1,2}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n  %reshape.3 = f32[8,2048]{1,0} reshape(%broadcast.4), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n  %broadcast.5 = f32[8,2048,2048]{2,1,0} broadcast(%reshape.3), dimensions={0,1}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n  %divide.0 = f32[8,2048,2048]{2,1,0} divide(%exponential.0, %broadcast.5), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n  %Arg_2.3 = bf16[8,2048,128]{2,1,0} parameter(2), metadata={op_name=\&quot;v\&quot;}\n  %convert.2 = f32[8,2048,128]{2,1,0} convert(%Arg_2.3), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/convert_element_type\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=201}\n  ROOT %dot.1 = f32[8,2048,128]{2,1,0} dot(%divide.0, %convert.2), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={1}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(st,td-&amp;gt;sd)/dot_general\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=201}\n}\n</code></pre>\n</article>\n  </div>\n\n  </div>\n</div>\n\n      </div>\n      <div class=\&quot;gist-meta\&quot;>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/a5fef642ee7d20146abf66ca2afd1932/raw/f6a827060241b4d48deb387fa0b8ee24dea73ffe/ref_attn.md/%22 style=\&quot;float:right\&quot; class=\&quot;Link--inTextBlock\&quot;>view raw</a>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/a5fef642ee7d20146abf66ca2afd1932#file-ref_attn-md\%22 class=\&quot;Link--inTextBlock\&quot;>\n          ref_attn.md\n        </a>\n        hosted with &amp;#10084; by <a class=\&quot;Link--inTextBlock\&quot; href=/__u/patricktoulme.substack.com/%22https://github.com/%22>GitHub</a>\n      </div>\n    </div>\n</div>\n&quot;,&quot;stylesheet&quot;:&quot;https://github.githubassets.com/assets/gist-embed-68783a026c0c.css&quot;}" data-component-name="GitgistToDOM"><link rel="stylesheet" href="https://github.githubassets.com/assets/gist-embed-68783a026c0c.css"><div id="gist144398154" class="gist">
    <div class="gist-file" data-color-mode="light" data-light-theme="light">
      <div class="gist-data">
        <div class="js-gist-file-update-container js-task-list-container">
  <div id="file-ref_attn-md" class="file my-2">
      <div id="file-ref_attn-md-readme" class="Box-body readme blob p-5 p-xl-6 " style="overflow:auto">
    <article class="markdown-body entry-content container-lg" itemprop="text"><pre><code>HloModule jit_reference_attention, entry_computation_layout={(bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)})-&gt;f32[8,2048,128]{2,1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}

%region_0.25 (Arg_0.22: f32[], Arg_1.23: f32[]) -&gt; f32[] {
  %Arg_0.22 = f32[] parameter(0), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max"}
  %Arg_1.23 = f32[] parameter(1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max"}
  ROOT %maximum.24 = f32[] maximum(%Arg_0.22, %Arg_1.23), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=196}
}

%region_1.36 (Arg_0.33: f32[], Arg_1.34: f32[]) -&gt; f32[] {
  %Arg_0.33 = f32[] parameter(0), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum"}
  %Arg_1.34 = f32[] parameter(1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum"}
  ROOT %add.35 = f32[] add(%Arg_0.33, %Arg_1.34), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=198}
}

ENTRY %main.47 (Arg_0.1: bf16[8,2048,128], Arg_1.2: bf16[8,2048,128], Arg_2.3: bf16[8,2048,128]) -&gt; f32[8,2048,128] {
  %constant.4 = pred[8,2048,2048]{2,1,0} constant({...})
  %Arg_0.1 = bf16[8,2048,128]{2,1,0} parameter(0), metadata={op_name="q"}
  %convert.0 = f32[8,2048,128]{2,1,0} convert(%Arg_0.1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/convert_element_type" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=184}
  %Arg_1.2 = bf16[8,2048,128]{2,1,0} parameter(1), metadata={op_name="k"}
  %convert.1 = f32[8,2048,128]{2,1,0} convert(%Arg_1.2), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/convert_element_type" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=184}
  %dot.0 = f32[8,2048,2048]{2,1,0} dot(%convert.0, %convert.1), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={2}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(sd,td-&gt;st)/dot_general" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=184}
  %constant.0 = f32[] constant(-2.38197633e+38)
  %broadcast.1 = f32[8,2048,2048]{2,1,0} broadcast(%constant.0), dimensions={}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}
  %select.1 = f32[8,2048,2048]{2,1,0} select(%constant.4, %dot.0, %broadcast.1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}
  %constant.1 = f32[] constant(-inf)
  %reduce.0 = f32[8,2048]{1,0} reduce(%select.1, %constant.1), dimensions={2}, to_apply=%region_0.25, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=196}
  %reshape.0 = f32[8,2048,1]{2,1,0} reshape(%reduce.0), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %broadcast.2 = f32[8,2048,1]{2,1,0} broadcast(%reshape.0), dimensions={0,1,2}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %reshape.1 = f32[8,2048]{1,0} reshape(%broadcast.2), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %broadcast.3 = f32[8,2048,2048]{2,1,0} broadcast(%reshape.1), dimensions={0,1}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %subtract.0 = f32[8,2048,2048]{2,1,0} subtract(%select.1, %broadcast.3), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %exponential.0 = f32[8,2048,2048]{2,1,0} exponential(%subtract.0), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/exp" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %constant.2 = f32[] constant(0)
  %reduce.1 = f32[8,2048]{1,0} reduce(%exponential.0, %constant.2), dimensions={2}, to_apply=%region_1.36, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=198}
  %reshape.2 = f32[8,2048,1]{2,1,0} reshape(%reduce.1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
  %broadcast.4 = f32[8,2048,1]{2,1,0} broadcast(%reshape.2), dimensions={0,1,2}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
  %reshape.3 = f32[8,2048]{1,0} reshape(%broadcast.4), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
  %broadcast.5 = f32[8,2048,2048]{2,1,0} broadcast(%reshape.3), dimensions={0,1}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
  %divide.0 = f32[8,2048,2048]{2,1,0} divide(%exponential.0, %broadcast.5), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
  %Arg_2.3 = bf16[8,2048,128]{2,1,0} parameter(2), metadata={op_name="v"}
  %convert.2 = f32[8,2048,128]{2,1,0} convert(%Arg_2.3), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/convert_element_type" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=201}
  ROOT %dot.1 = f32[8,2048,128]{2,1,0} dot(%divide.0, %convert.2), lhs_batch_dims={0}, lhs_contracting_dims={2}, rhs_batch_dims={0}, rhs_contracting_dims={1}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(st,td-&gt;sd)/dot_general" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=201}
}
</code></pre>
</article>
  </div>

  </div>
</div>

      </div>
      <div class="gist-meta">
        <a href="https://gist.github.com/patrick-toulme/a5fef642ee7d20146abf66ca2afd1932/raw/f6a827060241b4d48deb387fa0b8ee24dea73ffe/ref_attn.md" style="float:right" class="Link--inTextBlock">view raw</a>
        <a href="https://gist.github.com/patrick-toulme/a5fef642ee7d20146abf66ca2afd1932#file-ref_attn-md" class="Link--inTextBlock">
          ref_attn.md
        </a>
        hosted with &#10084; by <a class="Link--inTextBlock" href="https://github.com">GitHub</a>
      </div>
    </div>
</div>
</div><p>Notice: XLA sees every operation&#8212;two dots, a select, two reductions, elementwise ops. <strong>This is fundamentally different from Pallas</strong>, where XLA sees one opaque <code>custom-call</code>.</p><h3>Optimized HLO (After Fusion)</h3><p>After optimization passes, XLA fuses operations into three main kernels plus helper fusions. Here&#8217;s the complete optimized HLO:</p><div class="github-gist" data-attrs="{&quot;innerHTML&quot;:&quot;<div id=\&quot;gist144398175\&quot; class=\&quot;gist\&quot;>\n    <div class=\&quot;gist-file\&quot; translate=\&quot;no\&quot; data-color-mode=\&quot;light\&quot; data-light-theme=\&quot;light\&quot;>\n      <div class=\&quot;gist-data\&quot;>\n        <div class=\&quot;js-gist-file-update-container js-task-list-container\&quot;>\n  <div id=\&quot;file-final_ref_attn-md\&quot; class=\&quot;file my-2\&quot;>\n      <div id=\&quot;file-final_ref_attn-md-readme\&quot; class=\&quot;Box-body readme blob p-5 p-xl-6 \&quot;\n    style=\&quot;overflow: auto\&quot; tabindex=\&quot;0\&quot; role=\&quot;region\&quot;\n    aria-label=\&quot;final_ref_attn.md content, created by patrick-toulme on 07:28PM today.\&quot;\n  >\n    <article class=\&quot;markdown-body entry-content container-lg\&quot; itemprop=\&quot;text\&quot;><pre><code>HloModule jit_reference_attention, is_scheduled=true, entry_computation_layout={(bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)})-&amp;gt;f32[8,2048,128]{2,1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}\n\n%copy_fusion.2 (input.2: pred[8,2048,2048]) -&amp;gt; pred[8,2048,2048] {\n  %input.2 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(0)\n  ROOT %copy.4 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} copy(%input.2)\n}\n\n%fused_computation.1 (param_0.19: f32[8,2048], param_1.20: f32[8,2048], param_2.14: pred[8,2048,2048], param_3.4: f32[8,2048,2048]) -&amp;gt; f32[8,2048,2048] {\n  %param_2.14 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(2)\n  %fusion.12 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} fusion(%param_2.14), kind=kLoop, output_to_operand_aliasing={{}: (0, {})}, calls=%copy_fusion.2\n  %param_3.4 = f32[8,2048,2048]{1,2,0:T(8,128)} parameter(3)\n  %constant.21 = f32[]{:T(128)} constant(-2.38197633e+38)\n  %broadcast.21 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%constant.21), dimensions={}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[]}}\n  %select.5 = f32[8,2048,2048]{1,2,0:T(8,128)} select(%fusion.12, %param_3.4, %broadcast.21), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}\n  %param_1.20 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(1)\n  %broadcast.15 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%param_1.20), dimensions={0,1}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;1\&quot;,\&quot;128\&quot;]}}\n  %subtract.4 = f32[8,2048,2048]{1,2,0:T(8,128)} subtract(%select.5, %broadcast.15), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %exponential.4 = f32[8,2048,2048]{1,2,0:T(8,128)} exponential(%subtract.4), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/exp\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %param_0.19 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(0)\n  %broadcast.11 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%param_0.19), dimensions={0,1}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;1\&quot;,\&quot;128\&quot;]}}\n  ROOT %divide.2 = f32[8,2048,2048]{1,2,0:T(8,128)} divide(%exponential.4, %broadcast.11), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n}\n\n%bitcast_fusion.1 (bitcast_input.1: bf16[8,2048,128]) -&amp;gt; bf16[8,2048,128] {\n  %bitcast_input.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(0)\n  ROOT %bitcast.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} bitcast(%bitcast_input.1)\n}\n\n%fused_computation (param_0.1: bf16[8,2048,128], param_1.18: f32[8,2048], param_2.12: f32[8,2048], param_3.3: pred[8,2048,2048], param_4: f32[8,2048,2048]) -&amp;gt; f32[8,2048,128] {\n  %param_1.18 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(1)\n  %param_2.12 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(2)\n  %param_3.3 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(3)\n  %param_4 = f32[8,2048,2048]{1,2,0:T(8,128)} parameter(4)\n  %fusion.1 = f32[8,2048,2048]{1,2,0:T(8,128)} fusion(%param_1.18, %param_2.12, %param_3.3, %param_4), kind=kLoop, calls=%fused_computation.1, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=199}\n  %param_0.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(0)\n  %fusion.8 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} fusion(%param_0.1), kind=kLoop, calls=%bitcast_fusion.1\n  ROOT %convolution-base-dilated.2 = f32[8,2048,128]{2,1,0:T(8,128)} convolution(%fusion.1, %fusion.8), window={size=8 stride=7 lhs_dilate=8}, dim_labels=0bf_0io-&amp;gt;0bf, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(st,td-&amp;gt;sd)/dot_general\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=201}\n}\n\n%region_1.36 (Arg_0.33: f32[], Arg_1.34: f32[]) -&amp;gt; f32[] {\n  %Arg_1.34 = f32[]{:T(128)} parameter(1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot;}\n  %Arg_0.33 = f32[]{:T(128)} parameter(0), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot;}\n  ROOT %add.35 = f32[]{:T(128)} add(%Arg_0.33, %Arg_1.34), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=198}\n}\n\n%copy_fusion.1 (input.1: pred[8,2048,2048]) -&amp;gt; pred[8,2048,2048] {\n  %input.1 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(0)\n  ROOT %copy.3 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} copy(%input.1)\n}\n\n%fused_computation.3 (param_0.24: f32[8,2048], param_1.26: pred[8,2048,2048], param_2.19: f32[8,2048,2048]) -&amp;gt; f32[8,2048] {\n  %param_1.26 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(1)\n  %fusion.11 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} fusion(%param_1.26), kind=kLoop, output_to_operand_aliasing={{}: (0, {})}, calls=%copy_fusion.1\n  %param_2.19 = f32[8,2048,2048]{1,2,0:T(8,128)} parameter(2)\n  %constant.16 = f32[]{:T(128)} constant(-2.38197633e+38)\n  %broadcast.23 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%constant.16), dimensions={}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[]}}\n  %select.7 = f32[8,2048,2048]{1,2,0:T(8,128)} select(%fusion.11, %param_2.19, %broadcast.23), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}\n  %param_0.24 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(0)\n  %broadcast.16 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%param_0.24), dimensions={0,1}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;1\&quot;,\&quot;128\&quot;]}}\n  %subtract.6 = f32[8,2048,2048]{1,2,0:T(8,128)} subtract(%select.7, %broadcast.16), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %exponential.6 = f32[8,2048,2048]{1,2,0:T(8,128)} exponential(%subtract.6), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/exp\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=197}\n  %constant.14 = f32[]{:T(128)} constant(0)\n  ROOT %reduce.2 = f32[8,2048]{1,0:T(8,128)S(1)} reduce(%exponential.6, %constant.14), dimensions={2}, to_apply=%region_1.36, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=198}\n}\n\n%region_0.25 (Arg_0.22: f32[], Arg_1.23: f32[]) -&amp;gt; f32[] {\n  %Arg_1.23 = f32[]{:T(128)} parameter(1), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot;}\n  %Arg_0.22 = f32[]{:T(128)} parameter(0), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot;}\n  ROOT %maximum.24 = f32[]{:T(128)} maximum(%Arg_0.22, %Arg_1.23), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=196}\n}\n\n%bitcast_fusion (bitcast_input: bf16[8,2048,128]) -&amp;gt; bf16[8,2048,128] {\n  %bitcast_input = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(0)\n  ROOT %bitcast = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} bitcast(%bitcast_input)\n}\n\n%bitcast_fusion.2 (bitcast_input.2: bf16[8,2048,128]) -&amp;gt; bf16[8,2048,128] {\n  %bitcast_input.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(0)\n  ROOT %bitcast.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} bitcast(%bitcast_input.2)\n}\n\n%copy_fusion (input: pred[8,2048,2048]) -&amp;gt; pred[8,2048,2048] {\n  %input = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)} parameter(0)\n  ROOT %copy.2 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} copy(%input)\n}\n\n%fused_computation.7 (param_0.25: pred[8,2048,2048], param_1.28: bf16[8,2048,128], param_2.21: bf16[8,2048,128]) -&amp;gt; (f32[8,2048], f32[8,2048,2048]) {\n  %param_0.25 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)} parameter(0)\n  %fusion.10 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} fusion(%param_0.25), kind=kLoop, output_to_operand_aliasing={{}: (0, {})}, calls=%copy_fusion\n  %param_1.28 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(1)\n  %fusion.7 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} fusion(%param_1.28), kind=kLoop, calls=%bitcast_fusion\n  %param_2.21 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(2)\n  %fusion.9 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} fusion(%param_2.21), kind=kLoop, calls=%bitcast_fusion.2\n  %convolution-base-dilated.3 = f32[8,2048,2048]{1,2,0:T(8,128)} convolution(%fusion.7, %fusion.9), window={size=8 stride=7 lhs_dilate=8}, dim_labels=0bf_0oi-&amp;gt;0bf, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(sd,td-&amp;gt;st)/dot_general\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=184}\n  %constant.22 = f32[]{:T(128)} constant(-2.38197633e+38)\n  %broadcast.25 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%constant.22), dimensions={}, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[]}}\n  %select.9 = f32[8,2048,2048]{1,2,0:T(8,128)} select(%fusion.10, %convolution-base-dilated.3, %broadcast.25), metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=195}\n  %constant.15 = f32[]{:T(128)} constant(-inf)\n  %reduce.3 = f32[8,2048]{1,0:T(8,128)S(1)} reduce(%select.9, %constant.15), dimensions={2}, to_apply=%region_0.25, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=196}\n  ROOT %tuple = (f32[8,2048]{1,0:T(8,128)S(1)}, f32[8,2048,2048]{1,2,0:T(8,128)}) tuple(%reduce.3, %convolution-base-dilated.3)\n}\n\nENTRY %main.47 (Arg_0.1: bf16[8,2048,128], Arg_1.2: bf16[8,2048,128], Arg_2.3: bf16[8,2048,128]) -&amp;gt; f32[8,2048,128] {\n  %Arg_0.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(0), metadata={op_name=\&quot;q\&quot;}\n  %copy-start = (bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, u32[]{:S(2)}) copy-start(%Arg_0.1), cross_program_prefetch_index=0\n  %constant.4 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)} constant({...})\n  %Arg_2.3 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(2), metadata={op_name=\&quot;v\&quot;}\n  %Arg_1.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(1), metadata={op_name=\&quot;k\&quot;}\n  %copy-start.1 = (pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)}, pred[8,2048,2048]{1,2,0:T(32,128)(4,1)}, u32[]{:S(2)}) copy-start(%constant.4)\n  %copy-done = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} copy-done(%copy-start)\n  %fusion.5 = (f32[8,2048]{1,0:T(8,128)S(1)}, f32[8,2048,2048]{1,2,0:T(8,128)}) fusion(%constant.4, %copy-done, %Arg_1.2), kind=kOutput, calls=%fused_computation.7, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(sd,td-&amp;gt;st)/dot_general\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=184}\n  %get-tuple-element.1 = f32[8,2048,2048]{1,2,0:T(8,128)} get-tuple-element(%fusion.5), index=1, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=196}\n  %get-tuple-element = f32[8,2048]{1,0:T(8,128)S(1)} get-tuple-element(%fusion.5), index=0, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=196}\n  %copy-start.2 = (bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, u32[]{:S(2)}) copy-start(%Arg_2.3)\n  %copy-done.1 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} copy-done(%copy-start.1)\n  %fusion.2 = f32[8,2048]{1,0:T(8,128)S(1)} fusion(%get-tuple-element, %copy-done.1, %get-tuple-element.1), kind=kLoop, calls=%fused_computation.3, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=198}\n  %copy-done.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} copy-done(%copy-start.2)\n  ROOT %fusion = f32[8,2048,128]{2,1,0:T(8,128)} fusion(%copy-done.2, %fusion.2, %get-tuple-element, %copy-done.1, %get-tuple-element.1), kind=kOutput, calls=%fused_computation, metadata={op_name=\&quot;jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(st,td-&amp;gt;sd)/dot_general\&quot; source_file=\&quot;/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py\&quot; source_line=201}\n}\n</code></pre>\n</article>\n  </div>\n\n  </div>\n</div>\n\n      </div>\n      <div class=\&quot;gist-meta\&quot;>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/cc0ea7a1c5c8f73476d0aafbd7f14b82/raw/62ea137bb12b5e81d9fbd1b3043be0e1d0c716f9/final_ref_attn.md/%22 style=\&quot;float:right\&quot; class=\&quot;Link--inTextBlock\&quot;>view raw</a>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/cc0ea7a1c5c8f73476d0aafbd7f14b82#file-final_ref_attn-md\%22 class=\&quot;Link--inTextBlock\&quot;>\n          final_ref_attn.md\n        </a>\n        hosted with &amp;#10084; by <a class=\&quot;Link--inTextBlock\&quot; href=/__u/patricktoulme.substack.com/%22https://github.com/%22>GitHub</a>\n      </div>\n    </div>\n</div>\n&quot;,&quot;stylesheet&quot;:&quot;https://github.githubassets.com/assets/gist-embed-68783a026c0c.css&quot;}" data-component-name="GitgistToDOM"><link rel="stylesheet" href="https://github.githubassets.com/assets/gist-embed-68783a026c0c.css"><div id="gist144398175" class="gist">
    <div class="gist-file" data-color-mode="light" data-light-theme="light">
      <div class="gist-data">
        <div class="js-gist-file-update-container js-task-list-container">
  <div id="file-final_ref_attn-md" class="file my-2">
      <div id="file-final_ref_attn-md-readme" class="Box-body readme blob p-5 p-xl-6 " style="overflow:auto">
    <article class="markdown-body entry-content container-lg" itemprop="text"><pre><code>HloModule jit_reference_attention, is_scheduled=true, entry_computation_layout={(bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)})-&gt;f32[8,2048,128]{2,1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}

%copy_fusion.2 (input.2: pred[8,2048,2048]) -&gt; pred[8,2048,2048] {
  %input.2 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(0)
  ROOT %copy.4 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} copy(%input.2)
}

%fused_computation.1 (param_0.19: f32[8,2048], param_1.20: f32[8,2048], param_2.14: pred[8,2048,2048], param_3.4: f32[8,2048,2048]) -&gt; f32[8,2048,2048] {
  %param_2.14 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(2)
  %fusion.12 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} fusion(%param_2.14), kind=kLoop, output_to_operand_aliasing={{}: (0, {})}, calls=%copy_fusion.2
  %param_3.4 = f32[8,2048,2048]{1,2,0:T(8,128)} parameter(3)
  %constant.21 = f32[]{:T(128)} constant(-2.38197633e+38)
  %broadcast.21 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%constant.21), dimensions={}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":[]}}
  %select.5 = f32[8,2048,2048]{1,2,0:T(8,128)} select(%fusion.12, %param_3.4, %broadcast.21), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}
  %param_1.20 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(1)
  %broadcast.15 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%param_1.20), dimensions={0,1}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["1","128"]}}
  %subtract.4 = f32[8,2048,2048]{1,2,0:T(8,128)} subtract(%select.5, %broadcast.15), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %exponential.4 = f32[8,2048,2048]{1,2,0:T(8,128)} exponential(%subtract.4), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/exp" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %param_0.19 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(0)
  %broadcast.11 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%param_0.19), dimensions={0,1}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["1","128"]}}
  ROOT %divide.2 = f32[8,2048,2048]{1,2,0:T(8,128)} divide(%exponential.4, %broadcast.11), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
}

%bitcast_fusion.1 (bitcast_input.1: bf16[8,2048,128]) -&gt; bf16[8,2048,128] {
  %bitcast_input.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(0)
  ROOT %bitcast.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} bitcast(%bitcast_input.1)
}

%fused_computation (param_0.1: bf16[8,2048,128], param_1.18: f32[8,2048], param_2.12: f32[8,2048], param_3.3: pred[8,2048,2048], param_4: f32[8,2048,2048]) -&gt; f32[8,2048,128] {
  %param_1.18 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(1)
  %param_2.12 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(2)
  %param_3.3 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(3)
  %param_4 = f32[8,2048,2048]{1,2,0:T(8,128)} parameter(4)
  %fusion.1 = f32[8,2048,2048]{1,2,0:T(8,128)} fusion(%param_1.18, %param_2.12, %param_3.3, %param_4), kind=kLoop, calls=%fused_computation.1, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/div" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=199}
  %param_0.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(0)
  %fusion.8 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} fusion(%param_0.1), kind=kLoop, calls=%bitcast_fusion.1
  ROOT %convolution-base-dilated.2 = f32[8,2048,128]{2,1,0:T(8,128)} convolution(%fusion.1, %fusion.8), window={size=8 stride=7 lhs_dilate=8}, dim_labels=0bf_0io-&gt;0bf, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(st,td-&gt;sd)/dot_general" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=201}
}

%region_1.36 (Arg_0.33: f32[], Arg_1.34: f32[]) -&gt; f32[] {
  %Arg_1.34 = f32[]{:T(128)} parameter(1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum"}
  %Arg_0.33 = f32[]{:T(128)} parameter(0), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum"}
  ROOT %add.35 = f32[]{:T(128)} add(%Arg_0.33, %Arg_1.34), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=198}
}

%copy_fusion.1 (input.1: pred[8,2048,2048]) -&gt; pred[8,2048,2048] {
  %input.1 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(0)
  ROOT %copy.3 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} copy(%input.1)
}

%fused_computation.3 (param_0.24: f32[8,2048], param_1.26: pred[8,2048,2048], param_2.19: f32[8,2048,2048]) -&gt; f32[8,2048] {
  %param_1.26 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} parameter(1)
  %fusion.11 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} fusion(%param_1.26), kind=kLoop, output_to_operand_aliasing={{}: (0, {})}, calls=%copy_fusion.1
  %param_2.19 = f32[8,2048,2048]{1,2,0:T(8,128)} parameter(2)
  %constant.16 = f32[]{:T(128)} constant(-2.38197633e+38)
  %broadcast.23 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%constant.16), dimensions={}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":[]}}
  %select.7 = f32[8,2048,2048]{1,2,0:T(8,128)} select(%fusion.11, %param_2.19, %broadcast.23), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}
  %param_0.24 = f32[8,2048]{1,0:T(8,128)S(1)} parameter(0)
  %broadcast.16 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%param_0.24), dimensions={0,1}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["1","128"]}}
  %subtract.6 = f32[8,2048,2048]{1,2,0:T(8,128)} subtract(%select.7, %broadcast.16), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/sub" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %exponential.6 = f32[8,2048,2048]{1,2,0:T(8,128)} exponential(%subtract.6), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/exp" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=197}
  %constant.14 = f32[]{:T(128)} constant(0)
  ROOT %reduce.2 = f32[8,2048]{1,0:T(8,128)S(1)} reduce(%exponential.6, %constant.14), dimensions={2}, to_apply=%region_1.36, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=198}
}

%region_0.25 (Arg_0.22: f32[], Arg_1.23: f32[]) -&gt; f32[] {
  %Arg_1.23 = f32[]{:T(128)} parameter(1), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max"}
  %Arg_0.22 = f32[]{:T(128)} parameter(0), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max"}
  ROOT %maximum.24 = f32[]{:T(128)} maximum(%Arg_0.22, %Arg_1.23), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=196}
}

%bitcast_fusion (bitcast_input: bf16[8,2048,128]) -&gt; bf16[8,2048,128] {
  %bitcast_input = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(0)
  ROOT %bitcast = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} bitcast(%bitcast_input)
}

%bitcast_fusion.2 (bitcast_input.2: bf16[8,2048,128]) -&gt; bf16[8,2048,128] {
  %bitcast_input.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(0)
  ROOT %bitcast.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} bitcast(%bitcast_input.2)
}

%copy_fusion (input: pred[8,2048,2048]) -&gt; pred[8,2048,2048] {
  %input = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)} parameter(0)
  ROOT %copy.2 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} copy(%input)
}

%fused_computation.7 (param_0.25: pred[8,2048,2048], param_1.28: bf16[8,2048,128], param_2.21: bf16[8,2048,128]) -&gt; (f32[8,2048], f32[8,2048,2048]) {
  %param_0.25 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)} parameter(0)
  %fusion.10 = pred[8,2048,2048]{1,2,0:T(8,128)(4,1)} fusion(%param_0.25), kind=kLoop, output_to_operand_aliasing={{}: (0, {})}, calls=%copy_fusion
  %param_1.28 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} parameter(1)
  %fusion.7 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} fusion(%param_1.28), kind=kLoop, calls=%bitcast_fusion
  %param_2.21 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(2)
  %fusion.9 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} fusion(%param_2.21), kind=kLoop, calls=%bitcast_fusion.2
  %convolution-base-dilated.3 = f32[8,2048,2048]{1,2,0:T(8,128)} convolution(%fusion.7, %fusion.9), window={size=8 stride=7 lhs_dilate=8}, dim_labels=0bf_0oi-&gt;0bf, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(sd,td-&gt;st)/dot_general" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=184}
  %constant.22 = f32[]{:T(128)} constant(-2.38197633e+38)
  %broadcast.25 = f32[8,2048,2048]{1,2,0:T(8,128)} broadcast(%constant.22), dimensions={}, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/broadcast_in_dim" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":[]}}
  %select.9 = f32[8,2048,2048]{1,2,0:T(8,128)} select(%fusion.10, %convolution-base-dilated.3, %broadcast.25), metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(jit(_where))/select_n" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=195}
  %constant.15 = f32[]{:T(128)} constant(-inf)
  %reduce.3 = f32[8,2048]{1,0:T(8,128)S(1)} reduce(%select.9, %constant.15), dimensions={2}, to_apply=%region_0.25, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=196}
  ROOT %tuple = (f32[8,2048]{1,0:T(8,128)S(1)}, f32[8,2048,2048]{1,2,0:T(8,128)}) tuple(%reduce.3, %convolution-base-dilated.3)
}

ENTRY %main.47 (Arg_0.1: bf16[8,2048,128], Arg_1.2: bf16[8,2048,128], Arg_2.3: bf16[8,2048,128]) -&gt; f32[8,2048,128] {
  %Arg_0.1 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(0), metadata={op_name="q"}
  %copy-start = (bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, u32[]{:S(2)}) copy-start(%Arg_0.1), cross_program_prefetch_index=0
  %constant.4 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)} constant({...})
  %Arg_2.3 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(2), metadata={op_name="v"}
  %Arg_1.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)} parameter(1), metadata={op_name="k"}
  %copy-start.1 = (pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)}, pred[8,2048,2048]{1,2,0:T(32,128)(4,1)}, u32[]{:S(2)}) copy-start(%constant.4)
  %copy-done = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} copy-done(%copy-start)
  %fusion.5 = (f32[8,2048]{1,0:T(8,128)S(1)}, f32[8,2048,2048]{1,2,0:T(8,128)}) fusion(%constant.4, %copy-done, %Arg_1.2), kind=kOutput, calls=%fused_computation.7, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(sd,td-&gt;st)/dot_general" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=184}
  %get-tuple-element.1 = f32[8,2048,2048]{1,2,0:T(8,128)} get-tuple-element(%fusion.5), index=1, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=196}
  %get-tuple-element = f32[8,2048]{1,0:T(8,128)S(1)} get-tuple-element(%fusion.5), index=0, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_max" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=196}
  %copy-start.2 = (bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)}, bf16[8,2048,128]{2,1,0:T(8,128)(2,1)}, u32[]{:S(2)}) copy-start(%Arg_2.3)
  %copy-done.1 = pred[8,2048,2048]{1,2,0:T(32,128)(4,1)S(1)} copy-done(%copy-start.1)
  %fusion.2 = f32[8,2048]{1,0:T(8,128)S(1)} fusion(%get-tuple-element, %copy-done.1, %get-tuple-element.1), kind=kLoop, calls=%fused_computation.3, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/reduce_sum" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=198}
  %copy-done.2 = bf16[8,2048,128]{2,1,0:T(8,128)(2,1)S(1)} copy-done(%copy-start.2)
  ROOT %fusion = f32[8,2048,128]{2,1,0:T(8,128)} fusion(%copy-done.2, %fusion.2, %get-tuple-element, %copy-done.1, %get-tuple-element.1), kind=kOutput, calls=%fused_computation, metadata={op_name="jit(reference_attention)/jit(main)/reference_attention/reference_attention_h8_s2048/jit(_wrapped)/vmap(st,td-&gt;sd)/dot_general" source_file="/home/ptoulme/miniconda3/envs/vllm/lib/python3.12/site-packages/jax/experimental/pallas/ops/tpu/splash_attention/splash_attention_kernel.py" source_line=201}
}
</code></pre>
</article>
  </div>

  </div>
</div>

      </div>
      <div class="gist-meta">
        <a href="https://gist.github.com/patrick-toulme/cc0ea7a1c5c8f73476d0aafbd7f14b82/raw/62ea137bb12b5e81d9fbd1b3043be0e1d0c716f9/final_ref_attn.md" style="float:right" class="Link--inTextBlock">view raw</a>
        <a href="https://gist.github.com/patrick-toulme/cc0ea7a1c5c8f73476d0aafbd7f14b82#file-final_ref_attn-md" class="Link--inTextBlock">
          final_ref_attn.md
        </a>
        hosted with &#10084; by <a class="Link--inTextBlock" href="https://github.com">GitHub</a>
      </div>
    </div>
</div>
</div><h3>Key Optimization Patterns</h3><p><strong>1. Multi-Output Fusion (fusion.5)</strong></p><p>The first matmul fuses with mask application and reduce_max. It returns a <strong>tuple</strong> with both outputs:</p><ul><li><p><code>f32[8,2048]</code> &#8212; the row-wise maximum values</p></li><li><p><code>f32[8,2048,2048]</code> &#8212; the full attention scores</p></li></ul><p><strong>2. Recomputation vs Storage Trade-off</strong></p><p>Notice that <code>exp(scores - max)</code> is computed <strong>three times</strong>:</p><ul><li><p>In <code>fusion.2</code> to compute the sum</p></li><li><p>In <code>fused_computation.1</code> (inside <code>fusion</code>) to compute normalized weights</p></li><li><p>The mask is also re-applied multiple times</p></li></ul><p>This is intentional: recomputing is cheaper than storing and reloading the full <code>[8,2048,2048]</code> tensor.</p><p><strong>The O(n&#178;) Problem is Visible:</strong></p><pre><code><code>%get-tuple-element.1 = f32[8,2048,2048]{1,2,0:T(8,128)} get-tuple-element(%fusion.5), index=1</code></code></pre><p>This 128MB tensor (<code>8 &#215; 2048 &#215; 2048 &#215; 4</code> bytes) lives in HBM between fusion.5 and downstream fusions. <strong>Pallas avoids this entirely with online softmax.</strong></p><div><hr></div><h2>Reference Attention LLO</h2><p>After HLO optimization, each fusion compiles through the LLO (Low-Level Operations) pipeline. The fusion.5 kernel alone goes through <strong>79 optimization passes</strong> before producing final VLIW bundles.</p><h3>Original LLO: fusion.5 (Pass 02)</h3><p>The initial LLO reveals the kernel&#8217;s structure before optimization:</p><pre><code><code>$region0: #{fusion.5}
  // VMEM allocations for intermediate results
  #allocation0 [shape = 'f32[262144]{0}', space=vmem, size = 0x100000]     // 1MB scratch
  #allocation7 [shape = 'f32[131072]{0}', space=vmem, size = 0x80000]      // 512KB scratch
  #allocation8 [shape = 's32[1]{0}', space=sflag, size = 0x4]              // sync flag
  
  // Kernel inputs/outputs (from HLO fusion interface)
  %s0 = inlined_call_operand.hbm [shape: pred[8,2048,2048], index: 0]  // mask from HBM
  %s1 = inlined_call_operand.vmem [shape: bf16[8,2048,128], index: 1]  // Q in VMEM
  %s2 = inlined_call_operand.hbm [shape: bf16[8,2048,128], index: 2]   // K from HBM
  %s3 = inlined_call_operand.vmem [shape: f32[8,2048], index: 3]       // max output (VMEM)
  %s4 = inlined_call_operand.hbm [shape: f32[8,2048,2048], index: 4]   // scores output (HBM!)
  
  // Initialize max output to -inf (16 vector stores for 8&#215;2048 shape)
  %7 = vst [vmem:[%s3] sm:$0xff] /*vst_source=*/-inf
  %s8 = scalar_lea.vmem %s3, 8
  %9 = vst [vmem:[%s8] sm:$0xff] /*vst_source=*/-inf
  // ... 14 more stores to initialize max buffer
</code></code></pre><p><strong>Key observations:</strong></p><ul><li><p><strong>Q lives in VMEM</strong> but <strong>K streams from HBM</strong> via DMA</p></li><li><p><strong>The full attention scores write to HBM</strong> &#8212; 128MB output!</p></li><li><p><strong>The causal mask also streams from HBM</strong> &#8212; 32MB of predicates</p></li></ul><h3>Loop Structure (Pass 02)</h3><pre><code><code>  $region2: #{fusion.5} parent=0
    // Allocations for double-buffered DMA
    #allocation1 [shape = 'u8[524288]{0}', space=vmem, tag = 'operand span for K']  // 512KB
    #allocation2 [shape = 's32[2]{0}', space=sflag]                                  // sync flags
    #allocation4 [shape = 'u8[4194304]{0}', space=vmem, tag = 'operand span for mask'] // 4MB
    #allocation6 [shape = 'u8[16777216]{0}', space=vmem, tag = 'operand span for output'] // 16MB
    
    // Initialize sync flags for double buffering
    %38 = vsyncpa [#allocation2], 0
    %40 = vsyncpa [#allocation2 + $0x1], 0
    
    loop: start=0, step=1, limit=18
    
    $region4: #{fusion.5} parent=2 // loop_header
      // Phi nodes track iteration state across loop iterations
      %s48 = sphi 0, %s52 /* iteration index, stage = 0 */
      %p49 = scmp.ge.s32.totalorder %s48, 18 /* loop exit test */
      
      // Multi-dimensional iteration bounds (8 heads &#215; 2 K-blocks)
      %s55 = sphi 0, %s88 /* iter bound = 0 (head index) */
      %s56 = sphi 0, %s84 /* iter bound = 1 (K block index) */
      // ... more phi nodes for software pipelining stages
</code></code></pre><p><strong>Loop dimensions:</strong> The kernel tiles K into 2 blocks of 1024 rows each:</p><ul><li><p>8 heads &#215; 2 K-blocks = 16 main iterations</p></li><li><p>Plus 2 for software pipelining = 18 total iterations</p></li></ul><h3>Matmul Operations (Pass 02 - Original)</h3><p>Before optimization, each matmul tile uses explicit load-store patterns:</p><pre><code><code>// Load Q slice, unpack bf16, push to MXU
%v322 = vld [vmem:[%s321] sm:$0xf]           // Load Q tile (bf16)
%v323 = vunpack.c.l.bf16 %v322               // Unpack low bf16 to f32
%325 = vst [vmem:[%s320] sm:$0xff] /*vst_source=*/%v323  // Store unpacked
%v326 = vld [vmem:[%s320] sm:$0xff]          // Reload (wasteful!)
%327 = vmatpush1.xpose.msra.mxu0 %v326       // Push to MXU systolic array

// Prepare RHS with zeros (accumulator init)
%319 = vmatprep.subr.mxu0 0.0                // Prepare zero for subtraction

// Repeat for all 16 tiles of the 128-dim K dimension...
%328 = vmatprep.subr.mxu0 0.0
%v331 = vld [vmem:[%s330] sm:$0xf]
%v332 = vunpack.c.l.bf16 %v331
%334 = vst [vmem:[%s329] sm:$0xff] /*vst_source=*/%v332
%v335 = vld [vmem:[%s329] sm:$0xff]
%336 = vmatpush1.xpose.msra.mxu0 %v335
// ... continues for all K tiles
</code></code></pre><p>This is verbose and inefficient &#8212; later passes eliminate redundant loads/stores.</p><h3>VLIW Bundle Packing (Pass 29)</h3><p>The bundle packer groups independent operations into VLIW bundles:</p><pre><code><code>0x1   :  { %7 = vst [vmem:[%s3] sm:$0xff] /*vst_source=*/%v57179  ;;  
           %51239 = vst [vmem:[%s3 + $0x8] sm:$0xff] /*vst_source=*/%v57179  ;;  
           %51240 = vst [vmem:[%s3 + $0x10] sm:$0xff] /*vst_source=*/%v57179  ;; 
           // ... 16 parallel stores in one bundle!
         }

// Later: DMA + address computation in parallel
0x1e  : &gt; { %57143 = dma.hbm_to_vmem [thread:$0] (!%p57141), /*hbm=*/%s192, 
                     /*size_in_granules=*/8192, /*vmem=*/%s195, /*dst_syncflagno=*/%s180 }
0x1f  : &gt; { %p222 = pnand %p51265, %p221  ;;  
            %s202 = scalar_lea.vmem [#allocation4], %s51260 }
</code></code></pre><h3>Final VLIW Bundles (Pass 79) - Matmul Section</h3><p>After all optimizations, both MXUs execute in parallel:</p><pre><code><code>// MXU operations now distributed across both units
0xa   : &gt; { %55856 = vmatprep.subr.mxu0 %v465  ;;  
            %56016 = vmatprep.subr.mxu1 %v7203 }

0xb   : &gt; { %55857 = vmatpush3.xpose.msra.mxu0 %v323  ;;  
            %56017 = vmatpush3.xpose.msra.mxu1 %v7059  ;;  
            %v332 = vunpack.c.h.bf16 %v51273  ;;  
            %v7068 = vunpack.c.h.bf16 %v52005 }

0xc   : &gt; { %55858 = vmatprep.subr.mxu0 %v474  ;;  
            %56018 = vmatprep.subr.mxu1 %v7212  ;;  
            %v483 = vunpack.c.l.bf16 %v51288  ;;  
            %v7221 = vunpack.c.h.bf16 %v52022 }
</code></code></pre><p><strong>Key optimization:</strong> The VLIW scheduler distributes work across both MXUs:</p><ul><li><p><strong>mxu0</strong>: 1,396 operations (Q @ K^T lower tiles)</p></li><li><p><strong>mxu1</strong>: 1,383 operations (Q @ K^T upper tiles)</p></li></ul><p>This effectively <strong>doubles throughput</strong> compared to single-MXU execution.</p><h3>Final VLIW Bundles - Reduce_Max + Mask Application</h3><p>The most complex bundles overlap matmul results with mask lookup and max reduction:</p><pre><code><code>0xf4  : &gt; { %v618_v36 = vpop.f32.mrf.mxu0  ;;                    // Pop matmul result
            %55937 = vmatmul.mubr.bf16.gmra.mxu0 %v57834_v53 }   // Start next matmul

0xf5  : &gt; { %52037 = vst [vmem:[%s57903_s21 + $0x10] sm:$0xff] /*vst_source=*/%v7347_v37  ;;  
            %56098 = vmatprep.mubr.bf16.mxu1 %v57863_v16  ;;     // Prepare mxu1
            %v631_v54 = vpop.f32.mrf.mxu0  ;;                    // Pop another result
            %v7358_v6 = vsel %vm, %v7347, -2.38e+38 }            // Apply mask!

0xf9  : &gt; { %v7392_v33 = vsel %vm, %v7381, -2.38e+38  ;;         // Mask application
            %52041 = vst [vmem:[%s57903_s21 + $0x90] sm:$0xff]  ;;  // Store to HBM buffer
            %v665_v37 = vmax.f32 %v627_v12, %v659_v20  ;;        // Tree reduction max
            %v669_v44 = vpop.f32.mrf.mxu0  ;;                    // Pop matmul result
            %vm58016_vm2 = vcmp.ne.s32.totalorder %v854, 0 }     // Prepare next mask

0xfa  : &gt; { %v7399_v43 = vmax.f32 %v7358_v6, %v7392_v33  ;;      // Continue max tree
            %v7403_v53 = vpop.f32.mrf.mxu1  ;;                   // Pop from mxu1
            %v680_v59 = vsel %vm, %v669, -2.38e+38  ;;           // Mask more tiles
            %51310 = vst ... }                                    // Store result
</code></code></pre><p><strong>In a single VLIW bundle (0xf9):</strong></p><ul><li><p>Pop matmul result from MXU</p></li><li><p>Apply mask via <code>vsel</code></p></li><li><p>Store to HBM buffer</p></li><li><p>Compute tree reduction <code>vmax.f32</code></p></li><li><p>Prepare next mask predicate</p></li></ul><p>This is <strong>5 independent operations in parallel</strong> &#8212; the power of VLIW scheduling.</p><h2>Comparison: Pallas vs Reference Jax</h2><p>Now we can directly compare the two approaches using our test configuration: <strong>(8 heads, 2048 seq_len, 128 head_dim)</strong>.</p><h3>Algorithm: Online vs Naive Softmax</h3><p>The fundamental difference lies in the softmax computation strategy:</p><p>The reference implementation computes the full <code>[8, 2048, 2048]</code> attention matrix (128MB in f32), stores it to HBM, then reads it back multiple times for softmax operations. Pallas processes small tiles (e.g., <code>[512, 512]</code> or <code>[1024, 1024]</code>) that fit in VMEM, computing and consuming each tile before moving to the next.</p><h3>Kernel Structure</h3><p><strong>Reference: 3 Separate Fusions</strong></p><pre><code><code>fusion.5: Q @ K^T + mask + reduce_max  &#8594;  2,559 bundles
fusion.2: exp + reduce_sum             &#8594;  1,925 bundles  
fusion:   normalize + S @ V            &#8594;  5,761 bundles
&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;&#9472;
Total:                                   ~10,245 bundles</code></code></pre><p>Each fusion requires:</p><ol><li><p>Reading inputs from HBM</p></li><li><p>Computing results</p></li><li><p>Writing outputs back to HBM</p></li><li><p>Synchronization before next fusion</p></li></ol><p><strong>Pallas: 1 Custom Kernel (varies by block size)</strong></p><pre><code><code>splash_mha_fwd (block_size=512)   &#8594;  1,651 bundles
splash_mha_fwd (block_size=1024)  &#8594;  4,302 bundles
splash_mha_fwd (block_size=2048)  &#8594; 15,153 bundles</code></code></pre><p>Everything happens in a single kernel with explicit VMEM management. No intermediate HBM round-trips.</p><h3>VLIW Bundle Comparison</h3><p>The block size tradeoff is fascinating:</p><ul><li><p><strong>block_size=512</strong>: More iterations (4&#215;4=16 per head), but each iteration is simple. Fewest total bundles.</p></li><li><p><strong>block_size=1024</strong>: Balanced (2&#215;2=4 per head). Still significantly fewer bundles than reference.</p></li><li><p><strong>block_size=2048</strong>: One iteration covers entire sequence (1&#215;1=1 per head). More bundles than reference because the large tile operations can&#8217;t amortize overhead.</p></li></ul><h3>HBM Bandwidth Analysis</h3><p>For our test case <code>(8 heads, 2048 seq_len, 128 head_dim)</code>:</p><p><strong>Reference HBM Traffic:</strong></p><pre><code><code>Inputs:
  Q, K, V: 3 &#215; 8 &#215; 2048 &#215; 128 &#215; 2 bytes = 6 MB

Intermediates (written then read):
  S matrix after fusion.5: 8 &#215; 2048 &#215; 2048 &#215; 4 = 128 MB (write)
  S matrix for fusion.2:   8 &#215; 2048 &#215; 2048 &#215; 4 = 128 MB (read)
  exp(S) after fusion.2:   8 &#215; 2048 &#215; 2048 &#215; 4 = 128 MB (write)
  exp(S) for fusion:       8 &#215; 2048 &#215; 2048 &#215; 4 = 128 MB (read)
  
Output:
  O: 8 &#215; 2048 &#215; 128 &#215; 4 = 8 MB

Total: ~6 + 512 + 8 = ~526 MB HBM traffic
</code></code></pre><p><strong>Pallas HBM Traffic:</strong></p><pre><code><code>Inputs:
  Q, K, V: 3 &#215; 8 &#215; 2048 &#215; 128 &#215; 2 bytes = 6 MB

Intermediates:
  None! Everything stays in VMEM

Output:
  O: 8 &#215; 2048 &#215; 128 &#215; 4 = 8 MB

Total: ~14 MB HBM traffic</code></code></pre><p>That&#8217;s a <strong>37&#215; reduction</strong> in HBM bandwidth. Since TPU performance is often memory-bound, this directly translates to speedup.</p><h3>The Block Size Sweet Spot</h3><p>Why does block_size=512 produce fewer bundles than block_size=1024, and why does block_size=2048 produce <em>more</em> bundles than reference attention?</p><p>At <strong>block_size=2048</strong> (equal to seq_len), there&#8217;s no tiling&#8212;you compute the full [2048, 2048] attention matrix in one iteration. You&#8217;ve lost the memory benefit of online softmax but kept all its bookkeeping. It&#8217;s reference attention with extra overhead.</p><p>At <strong>block_size=512 vs 1024</strong>, the total math is the same&#8212;only the schedule changes. Smaller tiles reduce VMEM footprint, which improves double-buffering and DMA/compute overlap. Larger tiles reduce loop count but increase register pressure, making it harder for the backend to keep the MXUs busy while hiding memory latency.</p><h2>Why Pallas Matters</h2><p>Last post, the takeaway was: trust the compiler. This post&#8217;s takeaway is: know when not to.</p><p><strong>XLA is doing real work on reference attention</strong> &#8212; it fuses the matmul with the mask and max reduction, it overlaps DMA with compute, it balances both MXUs. The 10,245 VLIW bundles represent a well-optimized implementation of the algorithm you wrote. <strong>The problem is the algorithm itself</strong>.</p><p><strong>Pallas doesn&#8217;t replace XLA</strong>. For most operations, the automatic path is the right path. But when you&#8217;re hitting a fundamental algorithmic limit &#8212; when you know there&#8217;s a streaming formulation the compiler can&#8217;t discover &#8212; <strong>Pallas is how you escape</strong>. You write the tiled, memory-conscious kernel; Mosaic and LLO handle the rest.</p><p>The practical insights:</p><ul><li><p><strong>Block size matters, but depends on your shapes.</strong> For our test case (8 heads, seq_len=2048, head_dim=128), block_size=512 produced 6&#215; fewer bundles than reference attention. block_size=2048 produced <em>more</em> bundles. Smaller tiles pipeline better and reduce register pressure &#8212; but the optimal block size depends on your sequence length, head count, and head dimension. Profile your actual workload.</p></li><li><p><strong>HBM bandwidth is often the bottleneck.</strong> Reference attention moved 526MB through HBM; Splash moved 14MB. The 37&#215; reduction in memory traffic is why online softmax wins, not the bundle count.</p></li><li><p><strong>The compiler still does heavy lifting.</strong> Pallas gives you control over algorithm-level tiling; Mosaic handles vector layouts, MXU scheduling, and VLIW packing automatically. You&#8217;re not writing assembly.</p></li></ul><p>The dump flags for tracing Pallas kernels:</p><pre><code><code>os.environ["XLA_FLAGS"] = (
    f"--xla_dump_hlo_as_text "
    f"--xla_dump_to={HLO_PATH} "
    f"--xla_dump_hlo_pass_re=.* "
)
os.environ["LIBTPU_INIT_ARGS"] = (
    f"--xla_jf_dump_to={LLO_PATH} "
    f"--xla_jf_dump_hlo_text=true "
    f"--xla_jf_dump_llo_text=true "
    f"--xla_jf_dump_llo_html=false "
    f"--xla_jf_dump_llo_static_gaps=true "
    f"--xla_jf_emit_annotations=true "
    f"--xla_jf_debug_level=2 "
    f"--xla_mosaic_dump_to={MOSAIC_PATH} "
    f"--xla_mosaic_enable_dump_debug_info=true "
    f"--xla_mosaic_enable_llo_source_annotations=true"
)</code></code></pre><p>The HLO shows an opaque <code>custom-call</code> &#8212; the interesting stuff is in the Mosaic MLIR dumps (<code>0001-original.txt</code> through <code>0012-post-finalize-llo.txt</code>) and the final LLO bundles.</p><p>In principle, <strong>a compiler could learn online softmax</strong>. Someone could teach XLA to recognize the attention pattern, prove the streaming reformulation is valid, and generate the tiled kernel automatically. Until then, <strong>we write Pallas</strong>.</p><p>Questions? Message me on LinkedIn: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a> or follow me on X: <a href="https://x.com/PatrickToulme">https://x.com/PatrickToulme</a></p><div class="subscription-widget-wrap-editor" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe&quot;,&quot;language&quot;:&quot;en&quot;}" data-component-name="SubscribeWidgetToDOM"><div class="subscription-widget show-subscribe"><div class="preamble"><p class="cta-caption">Thanks for reading Just a Byte - AI Compilers, Silicon, and Systems! Subscribe for free to receive new posts and support my work.</p></div><form class="subscription-widget-subscribe"><input type="email" class="email-input" name="email" placeholder="Type your email&#8230;" tabindex="-1"><input type="submit" class="button primary" value="Subscribe"><div class="fake-input-wrapper"><div class="fake-input"></div><div class="fake-button"></div></div></form></div></div>]]></content:encoded></item><item><title><![CDATA[From JAX to VLIW: Tracing a Computation Through the TPU Compiler Stack]]></title><description><![CDATA[What actually happens when you compile Jax code on TPU? How does Jax compile to TPU assembly?]]></description><link>https://patricktoulme.substack.com/p/from-jax-to-vliw-tracing-a-computation</link><guid isPermaLink="false">https://patricktoulme.substack.com/p/from-jax-to-vliw-tracing-a-computation</guid><dc:creator><![CDATA[Patrick C. Toulme]]></dc:creator><pubDate>Sat, 27 Dec 2025 22:41:51 GMT</pubDate><enclosure url="https://substackcdn.com/image/fetch/$s_!cbPG!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin" length="0" type="image/jpeg"/><content:encoded><![CDATA[<p><strong>Full Code and IR Dumps to Follow Along</strong> - <a href="https://github.com/patrick-toulme/justabyte/tree/main/tpu_compiler_post">https://github.com/patrick-toulme/justabyte/tree/main/tpu_compiler_post</a></p><h2>Motivation</h2><p><br>Eight lines of JAX code in Python become 250 VLIW bundles across 5 fused kernels. A matrix multiply, RMS normalization, softmax, and another matrix multiply &#8212; the kind of operation that runs billions of times inside every transformer. <em>Here&#8217;s what happens between </em><code>jax.jit(f)(x)</code><em> and electrons moving through a TPU.</em></p><p>There&#8217;s surprisingly little public information about this &#8212; <em>Google&#8217;s TPU compiler is closed-source</em>, and the internal IRs are undocumented. I rented a TPU v6e for under a dollar and traced a small computation through four layers of compiler IR.</p><p><em><strong>The key insight: TPUs reward experimentation.</strong></em> GPUs have automatic codegen too &#8212; Inductor, XLA, Triton &#8212; but in practice, peak performance often still requires hand-tuned kernels. FlashAttention exists because no compiler found that optimization automatically. The TPU compiler gets you closer to the ceiling without manual intervention: it fuses operations, schedules hardware, and orchestrates memory in ways that would require expert kernel engineering on GPU. The performance ceiling might be lower than a perfect custom kernel, but the floor is much higher. This post shows exactly how that works.</p><p>Whether you&#8217;re debugging why a custom op is slow, deciding between TPUs and GPUs for a new workload, or just curious what <code>jax.jit</code> actually does &#8212; this post traces the full compilation path with real IR dumps.</p><h2>Setup</h2><p>I have rented a <a href="https://docs.cloud.google.com/tpu/docs/v6e">TPU V6e Trillium</a> on Google Cloud for these experiments. The cost is pretty minimal. I think this entire experiment was under one dollar.</p><p></p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!cbPG!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 424w, /__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 848w, /__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 1272w, /__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!cbPG!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin" width="600" height="337" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:337,&quot;width&quot;:600,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:null,&quot;alt&quot;:&quot;Google Cloud announces Trillium TPUs now available&quot;,&quot;title&quot;:null,&quot;type&quot;:null,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:null,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="Google Cloud announces Trillium TPUs now available" title="Google Cloud announces Trillium TPUs now available" srcset="/__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 424w, /__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 848w, /__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 1272w, /__u/substackcdn.com/image/fetch/$s_!cbPG!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F707c98c4-9819-47af-bd0a-f214d1c314c5_600x337.bin 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a><figcaption class="image-caption">Google TPU V6e Trillium - <a href="https://blog.google/feed/trillium-tpus/">Link</a></figcaption></figure></div><h2>Jax Code</h2><p>I have written some simple Jax code that we can use to trace the TPU&#8217;s compilation.</p><div class="github-gist" data-attrs="{&quot;innerHTML&quot;:&quot;<div id=\&quot;gist143984797\&quot; class=\&quot;gist\&quot;>\n    <div class=\&quot;gist-file\&quot; translate=\&quot;no\&quot; data-color-mode=\&quot;light\&quot; data-light-theme=\&quot;light\&quot;>\n      <div class=\&quot;gist-data\&quot;>\n        <div class=\&quot;js-gist-file-update-container js-task-list-container\&quot;>\n  <div id=\&quot;file-jax-md\&quot; class=\&quot;file my-2\&quot;>\n      <div id=\&quot;file-jax-md-readme\&quot; class=\&quot;Box-body readme blob p-5 p-xl-6 \&quot;\n    style=\&quot;overflow: auto\&quot; tabindex=\&quot;0\&quot; role=\&quot;region\&quot;\n    aria-label=\&quot;jax.md content, created by patrick-toulme on 10:24PM today.\&quot;\n  >\n    <article class=\&quot;markdown-body entry-content container-lg\&quot; itemprop=\&quot;text\&quot;><pre><code>import os\n\n# Create dump directories\nDUMP_ROOT = \&quot;compiler_dump/\&quot;\nHLO_DUMP_PATH = os.path.join(DUMP_ROOT, \&quot;hlo\&quot;)\nLLO_DUMP_PATH = os.path.join(DUMP_ROOT, \&quot;llo\&quot;)\n\nos.makedirs(HLO_DUMP_PATH, exist_ok=True)\nos.makedirs(LLO_DUMP_PATH, exist_ok=True)\n\nos.environ[\&quot;XLA_FLAGS\&quot;] = (\n    f\&quot;--xla_dump_hlo_as_text \&quot;\n    f\&quot;--xla_dump_to={HLO_DUMP_PATH} \&quot;\n    f\&quot;--xla_dump_hlo_pass_re=.* \&quot;\n)\n\nos.environ[\&quot;LIBTPU_INIT_ARGS\&quot;] = (\n    f\&quot;--xla_jf_dump_to={LLO_DUMP_PATH} \&quot;\n    f\&quot;--xla_jf_dump_hlo_text=true \&quot;\n    f\&quot;--xla_jf_dump_llo_text=true \&quot;\n    f\&quot;--xla_jf_dump_llo_html=false \&quot;\n    f\&quot;--xla_jf_dump_llo_static_gaps=true \&quot;\n    f\&quot;--xla_jf_emit_annotations=true \&quot;\n    f\&quot;--xla_jf_debug_level=2\&quot;\n)\n\n# Import JAX after setting env vars\nimport jax\nimport jax.numpy as jnp\n\n\n@jax.named_call\ndef matmul_1(x, w1):\n    \&quot;\&quot;\&quot;Stage 1: Linear projection (like Q @ K^T)\&quot;\&quot;\&quot;\n    return x @ w1\n\n\n@jax.named_call\ndef rms_norm(h):\n    \&quot;\&quot;\&quot;Stage 2: RMS Normalization\&quot;\&quot;\&quot;\n    rms = jnp.sqrt(jnp.mean(h ** 2, axis=-1, keepdims=True) + 1e-6)\n    return h / rms\n\n\n@jax.named_call\ndef softmax(h):\n    \&quot;\&quot;\&quot;Stage 3: Softmax (row-wise, numerically stable)\&quot;\&quot;\&quot;\n    h_max = jnp.max(h, axis=-1, keepdims=True)\n    exp_h = jnp.exp(h - h_max)\n    return exp_h / jnp.sum(exp_h, axis=-1, keepdims=True)\n\n\n@jax.named_call\ndef matmul_2(h, w2):\n    \&quot;\&quot;\&quot;Stage 4: Output projection (like attention @ V)\&quot;\&quot;\&quot;\n    return h @ w2\n\n\ndef mini_attention(x, w1, w2):\n    \&quot;\&quot;\&quot;\n    A minimal attention-like block:\n    matmul &#8594; rms_norm &#8594; softmax &#8594; matmul\n    \n    \&quot;\&quot;\&quot;\n    h = matmul_1(x, w1)\n    h = rms_norm(h)\n    h = softmax(h)\n    out = matmul_2(h, w2)\n    return out\n\n\ndef main():\n    # Small shapes to keep IR readable\n    batch, d_in, d_mid, d_out = 16, 64, 64, 32\n    \n    # Create inputs\n    key = jax.random.PRNGKey(42)\n    k1, k2, k3 = jax.random.split(key, 3)\n    \n    x = jax.random.normal(k1, (batch, d_in))\n    w1 = jax.random.normal(k2, (d_in, d_mid)) * 0.02\n    w2 = jax.random.normal(k3, (d_mid, d_out)) * 0.02\n    \n    # JIT compile and run\n    jitted_fn = jax.jit(mini_attention)\n    \n    # First call triggers compilation (and IR dump)\n    result = jitted_fn(x, w1, w2)\n    \n    # Block until computation is done\n    result.block_until_ready()\n    \n    print(f\&quot;Input shape:  {x.shape}\&quot;)\n    print(f\&quot;Output shape: {result.shape}\&quot;)\n    print(f\&quot;Output sample: {result[0, :5]}\&quot;)\n    print(f\&quot;\\nDumps written to:\&quot;)\n    print(f\&quot;  HLO: {HLO_DUMP_PATH}\&quot;)\n    print(f\&quot;  LLO: {LLO_DUMP_PATH}\&quot;)\n\n\nif __name__ == \&quot;__main__\&quot;:\n    main()\n    \n</code></pre>\n</article>\n  </div>\n\n  </div>\n</div>\n\n      </div>\n      <div class=\&quot;gist-meta\&quot;>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/6ddecf105b05236654cb89e13373457f/raw/534270b164357a7cb6c90457873e4d33479d76c8/jax.md/%22 style=\&quot;float:right\&quot; class=\&quot;Link--inTextBlock\&quot;>view raw</a>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/6ddecf105b05236654cb89e13373457f#file-jax-md\%22 class=\&quot;Link--inTextBlock\&quot;>\n          jax.md\n        </a>\n        hosted with &amp;#10084; by <a class=\&quot;Link--inTextBlock\&quot; href=/__u/patricktoulme.substack.com/%22https://github.com/%22>GitHub</a>\n      </div>\n    </div>\n</div>\n&quot;,&quot;stylesheet&quot;:&quot;https://github.githubassets.com/assets/gist-embed-ed91f9610ae6.css&quot;}" data-component-name="GitgistToDOM"><link rel="stylesheet" href="https://github.githubassets.com/assets/gist-embed-ed91f9610ae6.css"><div id="gist143984797" class="gist">
    <div class="gist-file" data-color-mode="light" data-light-theme="light">
      <div class="gist-data">
        <div class="js-gist-file-update-container js-task-list-container">
  <div id="file-jax-md" class="file my-2">
      <div id="file-jax-md-readme" class="Box-body readme blob p-5 p-xl-6 " style="overflow:auto">
    <article class="markdown-body entry-content container-lg" itemprop="text"><pre><code>import os

# Create dump directories
DUMP_ROOT = "compiler_dump/"
HLO_DUMP_PATH = os.path.join(DUMP_ROOT, "hlo")
LLO_DUMP_PATH = os.path.join(DUMP_ROOT, "llo")

os.makedirs(HLO_DUMP_PATH, exist_ok=True)
os.makedirs(LLO_DUMP_PATH, exist_ok=True)

os.environ["XLA_FLAGS"] = (
    f"--xla_dump_hlo_as_text "
    f"--xla_dump_to={HLO_DUMP_PATH} "
    f"--xla_dump_hlo_pass_re=.* "
)

os.environ["LIBTPU_INIT_ARGS"] = (
    f"--xla_jf_dump_to={LLO_DUMP_PATH} "
    f"--xla_jf_dump_hlo_text=true "
    f"--xla_jf_dump_llo_text=true "
    f"--xla_jf_dump_llo_html=false "
    f"--xla_jf_dump_llo_static_gaps=true "
    f"--xla_jf_emit_annotations=true "
    f"--xla_jf_debug_level=2"
)

# Import JAX after setting env vars
import jax
import jax.numpy as jnp


@jax.named_call
def matmul_1(x, w1):
    """Stage 1: Linear projection (like Q @ K^T)"""
    return x @ w1


@jax.named_call
def rms_norm(h):
    """Stage 2: RMS Normalization"""
    rms = jnp.sqrt(jnp.mean(h ** 2, axis=-1, keepdims=True) + 1e-6)
    return h / rms


@jax.named_call
def softmax(h):
    """Stage 3: Softmax (row-wise, numerically stable)"""
    h_max = jnp.max(h, axis=-1, keepdims=True)
    exp_h = jnp.exp(h - h_max)
    return exp_h / jnp.sum(exp_h, axis=-1, keepdims=True)


@jax.named_call
def matmul_2(h, w2):
    """Stage 4: Output projection (like attention @ V)"""
    return h @ w2


def mini_attention(x, w1, w2):
    """
    A minimal attention-like block:
    matmul &#8594; rms_norm &#8594; softmax &#8594; matmul
    
    """
    h = matmul_1(x, w1)
    h = rms_norm(h)
    h = softmax(h)
    out = matmul_2(h, w2)
    return out


def main():
    # Small shapes to keep IR readable
    batch, d_in, d_mid, d_out = 16, 64, 64, 32
    
    # Create inputs
    key = jax.random.PRNGKey(42)
    k1, k2, k3 = jax.random.split(key, 3)
    
    x = jax.random.normal(k1, (batch, d_in))
    w1 = jax.random.normal(k2, (d_in, d_mid)) * 0.02
    w2 = jax.random.normal(k3, (d_mid, d_out)) * 0.02
    
    # JIT compile and run
    jitted_fn = jax.jit(mini_attention)
    
    # First call triggers compilation (and IR dump)
    result = jitted_fn(x, w1, w2)
    
    # Block until computation is done
    result.block_until_ready()
    
    print(f"Input shape:  {x.shape}")
    print(f"Output shape: {result.shape}")
    print(f"Output sample: {result[0, :5]}")
    print(f"\nDumps written to:")
    print(f"  HLO: {HLO_DUMP_PATH}")
    print(f"  LLO: {LLO_DUMP_PATH}")


if __name__ == "__main__":
    main()
    
</code></pre>
</article>
  </div>

  </div>
</div>

      </div>
      <div class="gist-meta">
        <a href="https://gist.github.com/patrick-toulme/6ddecf105b05236654cb89e13373457f/raw/534270b164357a7cb6c90457873e4d33479d76c8/jax.md" style="float:right" class="Link--inTextBlock">view raw</a>
        <a href="https://gist.github.com/patrick-toulme/6ddecf105b05236654cb89e13373457f#file-jax-md" class="Link--inTextBlock">
          jax.md
        </a>
        hosted with &#10084; by <a class="Link--inTextBlock" href="https://github.com">GitHub</a>
      </div>
    </div>
</div>
</div><p>The above Jax code is really just a compiled matmul + rms_norm + softmax + matmul. We jit the computation with jax.jit. Jax then traces this computation into HLO (High Level Operations) IR. HLO is heavily open source - <a href="https://github.com/openxla/xla/tree/main/xla/hlo">HLO.</a></p><h2>TPU Compiler Top Level View</h2><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!UY9Y!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 424w, /__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 848w, /__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 1272w, /__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!UY9Y!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png" width="1456" height="985" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:985,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:317013,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182379451?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 424w, /__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 848w, /__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 1272w, /__u/substackcdn.com/image/fetch/$s_!UY9Y!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F06d702bb-6489-448c-9ccc-86447e18aa17_2042x1382.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!SiN4!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 424w, /__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 848w, /__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 1272w, /__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!SiN4!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png" width="1456" height="895" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/dd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:895,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:207454,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182379451?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 424w, /__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 848w, /__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 1272w, /__u/substackcdn.com/image/fetch/$s_!SiN4!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2Fdd1275f5-bc17-47d4-959c-8bd49f383649_1810x1112.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>At a high level the TPU compiler pipeline is Jax&#8594;HLO&#8594;LLO&#8594;VLIW bundles. Technically there are other stages, such as StableHLO, Jaxpr, TPU TLP etc., but for our purposes we can follow this diagram.</p><h2>HLO IR</h2><p>The frontend TPU compiler uses XLA as it&#8217;s primary infrastructure. XLA works off an IR called HLO (High Level Operations), which is an SSA graph IR.</p><p>There is some information publicly available on the frontend TPU HLO compiler and some open sourcing at <a href="https://github.com/openxla/xla">OpenXLA</a>. </p><p>This is similar to PyTorch FX IR. </p><div class="github-gist" data-attrs="{&quot;innerHTML&quot;:&quot;<div id=\&quot;gist143984817\&quot; class=\&quot;gist\&quot;>\n    <div class=\&quot;gist-file\&quot; translate=\&quot;no\&quot; data-color-mode=\&quot;light\&quot; data-light-theme=\&quot;light\&quot;>\n      <div class=\&quot;gist-data\&quot;>\n        <div class=\&quot;js-gist-file-update-container js-task-list-container\&quot;>\n  <div id=\&quot;file-hlo-md\&quot; class=\&quot;file my-2\&quot;>\n      <div id=\&quot;file-hlo-md-readme\&quot; class=\&quot;Box-body readme blob p-5 p-xl-6 \&quot;\n    style=\&quot;overflow: auto\&quot; tabindex=\&quot;0\&quot; role=\&quot;region\&quot;\n    aria-label=\&quot;hlo.md content, created by patrick-toulme on 10:25PM today.\&quot;\n  >\n    <article class=\&quot;markdown-body entry-content container-lg\&quot; itemprop=\&quot;text\&quot;><pre><code>HloModule jit_mini_attention, entry_computation_layout={(f32[16,64]{1,0:T(8,128)}, f32[64,64]{1,0:T(8,128)}, f32[64,32]{0,1:T(8,128)})-&amp;gt;f32[16,32]{1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}\n\n%region_0.15 (Arg_0.12: f32[], Arg_1.13: f32[]) -&amp;gt; f32[] {\n  %Arg_0.12 = f32[] parameter(0), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot;}\n  %Arg_1.13 = f32[] parameter(1), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot;}\n  ROOT %add.14 = f32[] add(%Arg_0.12, %Arg_1.13), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n}\n\n%region_1.28 (Arg_0.25: f32[], Arg_1.26: f32[]) -&amp;gt; f32[] {\n  %Arg_0.25 = f32[] parameter(0), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot;}\n  %Arg_1.26 = f32[] parameter(1), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot;}\n  ROOT %maximum.27 = f32[] maximum(%Arg_0.25, %Arg_1.26), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=48}\n}\n\n%region_2.39 (Arg_0.36: f32[], Arg_1.37: f32[]) -&amp;gt; f32[] {\n  %Arg_0.36 = f32[] parameter(0), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot;}\n  %Arg_1.37 = f32[] parameter(1), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot;}\n  ROOT %add.38 = f32[] add(%Arg_0.36, %Arg_1.37), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n}\n\nENTRY %main.47 (Arg_0.1: f32[16,64], Arg_1.2: f32[64,64], Arg_2.3: f32[64,32]) -&amp;gt; f32[16,32] {\n  %Arg_0.1 = f32[16,64]{1,0} parameter(0), metadata={op_name=\&quot;x\&quot;}\n  %Arg_1.2 = f32[64,64]{1,0} parameter(1), metadata={op_name=\&quot;w1\&quot;}\n  %dot.10 = f32[16,64]{1,0} dot(%Arg_0.1, %Arg_1.2), lhs_contracting_dims={1}, rhs_contracting_dims={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/matmul_1/dot_general\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=35}\n  %multiply.11 = f32[16,64]{1,0} multiply(%dot.10, %dot.10), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/integer_pow\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %constant.9 = f32[] constant(0)\n  %reduce.16 = f32[16]{0} reduce(%multiply.11, %constant.9), dimensions={1}, to_apply=%region_0.15, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %reshape.17 = f32[16,1]{1,0} reshape(%reduce.16), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %constant.6 = f32[] constant(64)\n  %broadcast.7 = f32[16,1]{1,0} broadcast(%constant.6), dimensions={}\n  %divide.18 = f32[16,1]{1,0} divide(%reshape.17, %broadcast.7), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %constant.4 = f32[] constant(1e-06)\n  %broadcast.5 = f32[16,1]{1,0} broadcast(%constant.4), dimensions={}\n  %add.19 = f32[16,1]{1,0} add(%divide.18, %broadcast.5), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/add\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %sqrt.20 = f32[16,1]{1,0} sqrt(%add.19), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/sqrt\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %broadcast.21 = f32[16,1]{1,0} broadcast(%sqrt.20), dimensions={0,1}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %reshape.22 = f32[16]{0} reshape(%broadcast.21), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %broadcast.23 = f32[16,64]{1,0} broadcast(%reshape.22), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %divide.24 = f32[16,64]{1,0} divide(%dot.10, %broadcast.23), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %constant.8 = f32[] constant(-inf)\n  %reduce.29 = f32[16]{0} reduce(%divide.24, %constant.8), dimensions={1}, to_apply=%region_1.28, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=48}\n  %reshape.30 = f32[16,1]{1,0} reshape(%reduce.29), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=48}\n  %broadcast.31 = f32[16,1]{1,0} broadcast(%reshape.30), dimensions={0,1}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %reshape.32 = f32[16]{0} reshape(%broadcast.31), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %broadcast.33 = f32[16,64]{1,0} broadcast(%reshape.32), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %subtract.34 = f32[16,64]{1,0} subtract(%divide.24, %broadcast.33), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %exponential.35 = f32[16,64]{1,0} exponential(%subtract.34), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/exp\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %reduce.40 = f32[16]{0} reduce(%exponential.35, %constant.9), dimensions={1}, to_apply=%region_2.39, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %reshape.41 = f32[16,1]{1,0} reshape(%reduce.40), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/broadcast_in_dim\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %broadcast.42 = f32[16,1]{1,0} broadcast(%reshape.41), dimensions={0,1}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %reshape.43 = f32[16]{0} reshape(%broadcast.42), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %broadcast.44 = f32[16,64]{1,0} broadcast(%reshape.43), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %divide.45 = f32[16,64]{1,0} divide(%exponential.35, %broadcast.44), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %Arg_2.3 = f32[64,32]{1,0} parameter(2), metadata={op_name=\&quot;w2\&quot;}\n  ROOT %dot.46 = f32[16,32]{1,0} dot(%divide.45, %Arg_2.3), lhs_contracting_dims={1}, rhs_contracting_dims={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/matmul_2/dot_general\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=56}\n}\n</code></pre>\n</article>\n  </div>\n\n  </div>\n</div>\n\n      </div>\n      <div class=\&quot;gist-meta\&quot;>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/5aad7bfd90ba3046aa0dfe8ae5cb5468/raw/738cdb7591be460913ca2c7b3edc9fd635561006/hlo.md/%22 style=\&quot;float:right\&quot; class=\&quot;Link--inTextBlock\&quot;>view raw</a>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/5aad7bfd90ba3046aa0dfe8ae5cb5468#file-hlo-md\%22 class=\&quot;Link--inTextBlock\&quot;>\n          hlo.md\n        </a>\n        hosted with &amp;#10084; by <a class=\&quot;Link--inTextBlock\&quot; href=/__u/patricktoulme.substack.com/%22https://github.com/%22>GitHub</a>\n      </div>\n    </div>\n</div>\n&quot;,&quot;stylesheet&quot;:&quot;https://github.githubassets.com/assets/gist-embed-ed91f9610ae6.css&quot;}" data-component-name="GitgistToDOM"><link rel="stylesheet" href="https://github.githubassets.com/assets/gist-embed-ed91f9610ae6.css"><div id="gist143984817" class="gist">
    <div class="gist-file" data-color-mode="light" data-light-theme="light">
      <div class="gist-data">
        <div class="js-gist-file-update-container js-task-list-container">
  <div id="file-hlo-md" class="file my-2">
      <div id="file-hlo-md-readme" class="Box-body readme blob p-5 p-xl-6 " style="overflow:auto">
    <article class="markdown-body entry-content container-lg" itemprop="text"><pre><code>HloModule jit_mini_attention, entry_computation_layout={(f32[16,64]{1,0:T(8,128)}, f32[64,64]{1,0:T(8,128)}, f32[64,32]{0,1:T(8,128)})-&gt;f32[16,32]{1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}

%region_0.15 (Arg_0.12: f32[], Arg_1.13: f32[]) -&gt; f32[] {
  %Arg_0.12 = f32[] parameter(0), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum"}
  %Arg_1.13 = f32[] parameter(1), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum"}
  ROOT %add.14 = f32[] add(%Arg_0.12, %Arg_1.13), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=41}
}

%region_1.28 (Arg_0.25: f32[], Arg_1.26: f32[]) -&gt; f32[] {
  %Arg_0.25 = f32[] parameter(0), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max"}
  %Arg_1.26 = f32[] parameter(1), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max"}
  ROOT %maximum.27 = f32[] maximum(%Arg_0.25, %Arg_1.26), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max" source_file="/home/ptoulme/tpu.py" source_line=48}
}

%region_2.39 (Arg_0.36: f32[], Arg_1.37: f32[]) -&gt; f32[] {
  %Arg_0.36 = f32[] parameter(0), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum"}
  %Arg_1.37 = f32[] parameter(1), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum"}
  ROOT %add.38 = f32[] add(%Arg_0.36, %Arg_1.37), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=50}
}

ENTRY %main.47 (Arg_0.1: f32[16,64], Arg_1.2: f32[64,64], Arg_2.3: f32[64,32]) -&gt; f32[16,32] {
  %Arg_0.1 = f32[16,64]{1,0} parameter(0), metadata={op_name="x"}
  %Arg_1.2 = f32[64,64]{1,0} parameter(1), metadata={op_name="w1"}
  %dot.10 = f32[16,64]{1,0} dot(%Arg_0.1, %Arg_1.2), lhs_contracting_dims={1}, rhs_contracting_dims={0}, metadata={op_name="jit(mini_attention)/jit(main)/matmul_1/dot_general" source_file="/home/ptoulme/tpu.py" source_line=35}
  %multiply.11 = f32[16,64]{1,0} multiply(%dot.10, %dot.10), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/integer_pow" source_file="/home/ptoulme/tpu.py" source_line=41}
  %constant.9 = f32[] constant(0)
  %reduce.16 = f32[16]{0} reduce(%multiply.11, %constant.9), dimensions={1}, to_apply=%region_0.15, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=41}
  %reshape.17 = f32[16,1]{1,0} reshape(%reduce.16), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/broadcast_in_dim" source_file="/home/ptoulme/tpu.py" source_line=41}
  %constant.6 = f32[] constant(64)
  %broadcast.7 = f32[16,1]{1,0} broadcast(%constant.6), dimensions={}
  %divide.18 = f32[16,1]{1,0} divide(%reshape.17, %broadcast.7), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=41}
  %constant.4 = f32[] constant(1e-06)
  %broadcast.5 = f32[16,1]{1,0} broadcast(%constant.4), dimensions={}
  %add.19 = f32[16,1]{1,0} add(%divide.18, %broadcast.5), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/add" source_file="/home/ptoulme/tpu.py" source_line=41}
  %sqrt.20 = f32[16,1]{1,0} sqrt(%add.19), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/sqrt" source_file="/home/ptoulme/tpu.py" source_line=41}
  %broadcast.21 = f32[16,1]{1,0} broadcast(%sqrt.20), dimensions={0,1}, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %reshape.22 = f32[16]{0} reshape(%broadcast.21), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %broadcast.23 = f32[16,64]{1,0} broadcast(%reshape.22), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %divide.24 = f32[16,64]{1,0} divide(%dot.10, %broadcast.23), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %constant.8 = f32[] constant(-inf)
  %reduce.29 = f32[16]{0} reduce(%divide.24, %constant.8), dimensions={1}, to_apply=%region_1.28, metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max" source_file="/home/ptoulme/tpu.py" source_line=48}
  %reshape.30 = f32[16,1]{1,0} reshape(%reduce.29), metadata={op_name="jit(mini_attention)/jit(main)/softmax/broadcast_in_dim" source_file="/home/ptoulme/tpu.py" source_line=48}
  %broadcast.31 = f32[16,1]{1,0} broadcast(%reshape.30), dimensions={0,1}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}
  %reshape.32 = f32[16]{0} reshape(%broadcast.31), metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}
  %broadcast.33 = f32[16,64]{1,0} broadcast(%reshape.32), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}
  %subtract.34 = f32[16,64]{1,0} subtract(%divide.24, %broadcast.33), metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}
  %exponential.35 = f32[16,64]{1,0} exponential(%subtract.34), metadata={op_name="jit(mini_attention)/jit(main)/softmax/exp" source_file="/home/ptoulme/tpu.py" source_line=49}
  %reduce.40 = f32[16]{0} reduce(%exponential.35, %constant.9), dimensions={1}, to_apply=%region_2.39, metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=50}
  %reshape.41 = f32[16,1]{1,0} reshape(%reduce.40), metadata={op_name="jit(mini_attention)/jit(main)/softmax/broadcast_in_dim" source_file="/home/ptoulme/tpu.py" source_line=50}
  %broadcast.42 = f32[16,1]{1,0} broadcast(%reshape.41), dimensions={0,1}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}
  %reshape.43 = f32[16]{0} reshape(%broadcast.42), metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}
  %broadcast.44 = f32[16,64]{1,0} broadcast(%reshape.43), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}
  %divide.45 = f32[16,64]{1,0} divide(%exponential.35, %broadcast.44), metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}
  %Arg_2.3 = f32[64,32]{1,0} parameter(2), metadata={op_name="w2"}
  ROOT %dot.46 = f32[16,32]{1,0} dot(%divide.45, %Arg_2.3), lhs_contracting_dims={1}, rhs_contracting_dims={0}, metadata={op_name="jit(mini_attention)/jit(main)/matmul_2/dot_general" source_file="/home/ptoulme/tpu.py" source_line=56}
}
</code></pre>
</article>
  </div>

  </div>
</div>

      </div>
      <div class="gist-meta">
        <a href="https://gist.github.com/patrick-toulme/5aad7bfd90ba3046aa0dfe8ae5cb5468/raw/738cdb7591be460913ca2c7b3edc9fd635561006/hlo.md" style="float:right" class="Link--inTextBlock">view raw</a>
        <a href="https://gist.github.com/patrick-toulme/5aad7bfd90ba3046aa0dfe8ae5cb5468#file-hlo-md" class="Link--inTextBlock">
          hlo.md
        </a>
        hosted with &#10084; by <a class="Link--inTextBlock" href="https://github.com">GitHub</a>
      </div>
    </div>
</div>
</div><p>The above HLO is the output of the Jax tracer after converting from StableHLO to HLO. No optimization passes have been performed yet. </p><p>We can clearly see - instructions, shapes, dtypes and metadata.</p><h2>HLO Optimization Passes</h2><p>The XLA compiler ran 71 optimization passes on this module. Here are the key transformations for our small toy program.</p><h3>Algebraic Simplifier</h3><p>The algebraic simplifier runs early in the pipeline (pass #6). The Algebraic Simplifier is open source at OpenXLA -  <a href="https://github.com/openxla/xla/blob/80e2510a29014026abb93042280cddba1285620b/xla/hlo/transforms/simplifiers/algebraic_simplifier.cc">algebraic_simplifier.cc</a> The pass is nearly 10k lines of CPP code. One optimization is converting division by constants to multiplication:</p><pre><code><code>// Before (pass #5)
%constant.6 = f32[] constant(64)
%broadcast.7 = f32[16,1]{1,0} broadcast(%constant.6), dimensions={}
%divide.18 = f32[16,1]{1,0} divide(%reshape.17, %broadcast.7)

// After (pass #6)
%constant = f32[] constant(0.015625)
%broadcast = f32[16,1]{1,0} broadcast(%constant), dimensions={}
%multiply = f32[16,1]{1,0} multiply(%reshape.17, %broadcast)
</code></code></pre><p>Division is expensive on most hardware. The compiler precomputes <code>1/64 = 0.015625</code> and converts the divide to a multiply. We see this in our toy program.</p><h3>Layout Assignment</h3><p>Layout assignment and tiling (pass #25 &#8594; #29) adds TPU tile annotations to all shapes.</p><p>Layout assignment is partially open source at OpenXLA - <a href="https://github.com/openxla/xla/blob/9caad7b3520548142ccd6a2d528a06be6c474de1/xla/service/layout_assignment.cc">layout_assignment.cc</a></p><pre><code><code>// Before
%dot.10 = f32[16,64]{1,0} dot(...)
%reduce.16 = f32[16]{0} reduce(...)

// After
%dot.10 = f32[16,64]{1,0:T(8,128)} dot(...)
%reduce.16 = f32[16]{0:T(128)} reduce(...)
</code></code></pre><p>The <code>{1,0}</code> specifies row-major layout (the last dimension is contiguous in memory). The <code>:T(8,128)</code> is TPU-specific tiling &#8212; the tensor is divided into tiles of 8 rows &#215; 128 columns, matching the VPU's 8 sublanes &#215; 128 lanes.</p><h3>Fusion</h3><p>The fusion pass (pass #38-39) combines independent HLO operations into fused kernels.</p><div class="captioned-image-container"><figure><a class="image-link image2 is-viewable-img" target="_blank" href="/__u/substackcdn.com/image/fetch/$s_!EpOw!,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png" data-component-name="Image2ToDOM"><div class="image2-inset"><picture><source type="image/webp" srcset="/__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 424w, /__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 848w, /__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 1272w, /__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_webp, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 1456w" sizes="100vw"><img src="/__u/substackcdn.com/image/fetch/$s_!EpOw!,w_1456,c_limit,f_auto,q_auto:good,fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png" width="1456" height="1236" data-attrs="{&quot;src&quot;:&quot;https://substack-post-media.s3.amazonaws.com/public/images/05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png&quot;,&quot;srcNoWatermark&quot;:null,&quot;fullscreen&quot;:null,&quot;imageSize&quot;:null,&quot;height&quot;:1236,&quot;width&quot;:1456,&quot;resizeWidth&quot;:null,&quot;bytes&quot;:296400,&quot;alt&quot;:null,&quot;title&quot;:null,&quot;type&quot;:&quot;image/png&quot;,&quot;href&quot;:null,&quot;belowTheFold&quot;:true,&quot;topImage&quot;:false,&quot;internalRedirect&quot;:&quot;https://patricktoulme.substack.com/i/182379451?img=https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png&quot;,&quot;isProcessing&quot;:false,&quot;align&quot;:null,&quot;offset&quot;:false}" class="sizing-normal" alt="" srcset="/__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_424, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 424w, /__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_848, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 848w, /__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_1272, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 1272w, /__u/substackcdn.com/image/fetch/$s_!EpOw!, /__u/patricktoulme.substack.com/w_1456, /__u/patricktoulme.substack.com/c_limit, /__u/patricktoulme.substack.com/f_auto, /__u/patricktoulme.substack.com/q_auto:good, /__u/patricktoulme.substack.com/fl_progressive:steep/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F05311629-d22a-45ec-8a98-d4734a62d116_1880x1596.png 1456w" sizes="100vw" loading="lazy"></picture><div class="image-link-expand"><div class="pencraft pc-display-flex pc-gap-8 pc-reset"><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container restack-image"><svg aria-hidden="true" width="20" height="20" viewBox="0 0 20 20" fill="none" stroke-width="1.5" stroke="var(--color-fg-primary)" stroke-linecap="round" stroke-linejoin="round" xmlns="http://www.w3.org/2000/svg"><g><path d="M2.53001 7.81595C3.49179 4.73911 6.43281 2.5 9.91173 2.5C13.1684 2.5 15.9537 4.46214 17.0852 7.23684L17.6179 8.67647M17.6179 8.67647L18.5002 4.26471M17.6179 8.67647L13.6473 6.91176M17.4995 12.1841C16.5378 15.2609 13.5967 17.5 10.1178 17.5C6.86118 17.5 4.07589 15.5379 2.94432 12.7632L2.41165 11.3235M2.41165 11.3235L1.5293 15.7353M2.41165 11.3235L6.38224 13.0882"></path></g></svg></button><button tabindex="0" type="button" class="pencraft pc-reset pencraft icon-container view-image"><svg xmlns="http://www.w3.org/2000/svg" width="20" height="20" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round" class="lucide lucide-maximize2 lucide-maximize-2"><polyline points="15 3 21 3 21 9"></polyline><polyline points="9 21 3 21 3 15"></polyline><line x1="21" x2="14" y1="3" y2="10"></line><line x1="3" x2="10" y1="21" y2="14"></line></svg></button></div></div></div></a></figure></div><p>The TPU HLO fusion pass is also open source at OpenXLA - <a href="https://github.com/openxla/xla/blob/9caad7b3520548142ccd6a2d528a06be6c474de1/xla/service/instruction_fusion.cc">instruction_fusion.cc</a></p><pre><code><code>ENTRY %main.47 (Arg_0.1: f32[16,64], Arg_1.2: f32[64,64], Arg_2.3: f32[64,32]) -&gt; f32[16,32] {
  %Arg_0.1 = f32[16,64]{1,0:T(8,128)} parameter(0), metadata={op_name="x"}
  %Arg_1.2 = f32[64,64]{1,0:T(8,128)} parameter(1), metadata={op_name="w1"}
  %Arg_2.3 = f32[64,32]{0,1:T(8,128)} parameter(2), metadata={op_name="w2"}
  %convolution = f32[16,64]{1,0:T(8,128)} convolution(%Arg_0.1, %Arg_1.2), dim_labels=bf_io-&gt;bf
  %multiply_reduce_fusion = f32[16]{0:T(128)} fusion(%convolution), kind=kLoop, calls=%fused_computation.3
  %add_sqrt_fusion = f32[16]{0:T(128)} fusion(%multiply_reduce_fusion), kind=kLoop, calls=%fused_computation.9
  %fusion.5 = f32[16]{0:T(128)} fusion(%convolution, %add_sqrt_fusion), kind=kLoop, calls=%fused_computation.8
  %fusion.2 = f32[16]{0:T(128)} fusion(%fusion.5, %convolution, %add_sqrt_fusion), kind=kLoop, calls=%fused_computation.4
  ROOT %fusion = f32[16,32]{1,0:T(8,128)} fusion(%Arg_2.3, %fusion.2, %fusion.5, %convolution, %add_sqrt_fusion), kind=kOutput
}
</code></code></pre><p>The original ~25 HLO ops collapse into 6 operations. Each <code>fusion</code> wraps multiple operations into a single kernel. The <code>kind=kLoop</code> indicates elementwise fusion; <code>kind=kOutput</code> indicates producer-consumer fusion (the producer&#8217;s output is immediately consumed).</p><p>Note that <code>dot</code> became <code>convolution</code> &#8212; the TPU compiler canonicalizes matrix multiplications this way, likely because both operations map to the same MXU hardware.</p><p>Each fusion has a body. The <code>add_sqrt_fusion</code> contains:</p><pre><code><code>%fused_computation.9 (param_0.29: f32[16]) -&gt; f32[16] {
  %param_0.29 = f32[16]{0:T(128)} parameter(0)
  %constant.12 = f32[] constant(0.015625)
  %broadcast.37 = f32[16]{0:T(128)} broadcast(%constant.12), dimensions={}
  %multiply.7 = f32[16]{0:T(128)} multiply(%param_0.29, %broadcast.37)
  %constant.16 = f32[] constant(1e-06)
  %broadcast.36 = f32[16]{0:T(128)} broadcast(%constant.16), dimensions={}
  %add.5 = f32[16]{0:T(128)} add(%multiply.7, %broadcast.36)
  ROOT %sqrt.5 = f32[16]{0:T(128)} sqrt(%add.5)
}</code></code></pre><p>This is the RMS norm denominator: <code>sqrt(mean + epsilon)</code>.</p><h3>Multi-Output Fusion (pass #43)</h3><p>When multiple fusions share a common operand, they can be merged. Multi-output fusion (pass #43) identifies sibling fusions with shared inputs and combines them, returning a tuple:</p><pre><code><code>%fused_computation.3 (param_0.33: f32[16,64], param_1.35: f32[64,64]) -&gt; (f32[16], f32[16,64]) {
  %convolution.3 = f32[16,64]{1,0:T(8,128)} convolution(%param_0.33, %param_1.35), dim_labels=bf_io-&gt;bf
  %multiply.6 = f32[16,64]{1,0:T(8,128)} multiply(%convolution.3, %convolution.3)
  %constant.14 = f32[] constant(0)
  %reduce.0 = f32[16]{0:T(128)} reduce(%multiply.6, %constant.14), dimensions={1}, to_apply=%region_0.15
  ROOT %tuple = (f32[16]{0:T(128)}, f32[16,64]{1,0:T(8,128)}) tuple(%reduce.0, %convolution.3)
}</code></code></pre><p>Before this pass, the matmul feeding into both the normalization path and the RMS reduction path would have been computed in separate fusions. Multi-output fusion merges them &#8212; the matmul happens once, and both the normalized result (<code>convolution.3</code>) and the squared sum (<code>reduce.0</code>) are returned together. The alternative would be either recomputing the matmul or spilling it to HBM between fusions.</p><p>The caller extracts both values:</p><pre><code><code>%multiply_reduce_fusion = (f32[16]{0:T(128)}, f32[16,64]{1,0:T(8,128)}) fusion(%Arg_0.1, %Arg_1.2), kind=kOutput
%get-tuple-element = f32[16]{0:T(128)} get-tuple-element(%multiply_reduce_fusion), index=0
%get-tuple-element.1 = f32[16,64]{1,0:T(8,128)} get-tuple-element(%multiply_reduce_fusion), index=1</code></code></pre><h3>Memory Space Assignment &amp; Async Scheduling</h3><p>The final passes assign tensors to specific memory spaces and insert async memory operations. In HLO, memory spaces are annotated as <code>S(n)</code>:</p><ul><li><p><code>S(0)</code> (often omitted) &#8212; HBM (high bandwidth memory, off-chip)</p></li><li><p><code>S(1)</code> &#8212; VMEM (on-chip SRAM)</p></li><li><p><code>S(2)</code>, <code>S(3)</code>, etc. &#8212; additional device-specific memory spaces</p></li></ul><p>The <code>after_codegen.txt</code> shows the scheduled program:</p><pre><code><code>ENTRY %main.47 (Arg_0.1: f32[16,64], Arg_1.2: f32[64,64], Arg_2.3: f32[64,32]) -&gt; f32[16,32] {
  %Arg_1.2 = f32[64,64]{1,0:T(8,128)} parameter(1), metadata={op_name="w1"}
  %copy-start = (f32[64,64]{1,0:T(8,128)S(1)}, f32[64,64]{1,0:T(8,128)}, u32[]{:S(2)}) copy-start(%Arg_1.2), cross_program_prefetch_index=0
  %Arg_2.3 = f32[64,32]{0,1:T(8,128)} parameter(2), metadata={op_name="w2"}
  %Arg_0.1 = f32[16,64]{1,0:T(8,128)} parameter(0), metadata={op_name="x"}
  %copy-done = f32[64,64]{1,0:T(8,128)S(1)} copy-done(%copy-start)
  %multiply_reduce_fusion = (f32[16]{0:T(128)S(1)}, f32[16,64]{1,0:T(8,128)S(1)}) fusion(%Arg_0.1, %copy-done), kind=kOutput, ...
  %copy-start.1 = (f32[64,32]{0,1:T(8,128)S(1)}, f32[64,32]{0,1:T(8,128)}, u32[]{:S(2)}) copy-start(%Arg_2.3)
  %get-tuple-element.1 = f32[16,64]{1,0:T(8,128)S(1)} get-tuple-element(%multiply_reduce_fusion), index=1
  %get-tuple-element = f32[16]{0:T(128)S(1)} get-tuple-element(%multiply_reduce_fusion), index=0
  %add_sqrt_fusion = f32[16]{0:T(128)S(1)} fusion(%get-tuple-element), kind=kLoop, ...
  %fusion.5 = f32[16]{0:T(128)S(1)} fusion(%get-tuple-element.1, %add_sqrt_fusion), kind=kLoop, ...
  %fusion.2 = f32[16]{0:T(128)S(1)} fusion(%fusion.5, %get-tuple-element.1, %add_sqrt_fusion), kind=kLoop, ...
  %copy-done.1 = f32[64,32]{0,1:T(8,128)S(1)} copy-done(%copy-start.1)
  ROOT %fusion = f32[16,32]{1,0:T(8,128)} fusion(%copy-done.1, %fusion.2, %fusion.5, %get-tuple-element.1, %add_sqrt_fusion), kind=kOutput, ...
}
</code></code></pre><p>Notice the parameters start without <code>S(1)</code> &#8212; they live in HBM. The <code>copy-start</code>/<code>copy-done</code> pairs are async DMA operations that move data to VMEM. The tuple returned by <code>copy-start</code> contains the destination buffer (in <code>S(1)</code>), the source reference, and a sync token (in <code>S(2)</code>).</p><p>The compiler overlaps memory transfers with computation:</p><ol><li><p><code>copy-start(w1)</code> &#8212; initiate DMA of w1 from HBM &#8594; VMEM</p></li><li><p><code>copy-done(w1)</code> &#8212; wait for transfer, then use in matmul&#8321;</p></li><li><p><code>copy-start(w2)</code> &#8212; initiate DMA of w2 <em>while</em> RMS norm and softmax execute</p></li><li><p><code>copy-done(w2)</code> &#8212; wait for transfer, then use in matmul&#8322;</p></li></ol><p>The <code>backend_config</code> also shows estimated cycle counts per fusion:</p><pre><code><code>"estimated_cycles":"2248"  // multiply_reduce_fusion
"estimated_cycles":"2120"  // add_sqrt_fusion  
"estimated_cycles":"2143"  // fusion.5 (reduce_max)
"estimated_cycles":"2162"  // fusion.2 (reduce_sum)
"estimated_cycles":"3140"  // final fusion (matmul_2)
</code></code></pre><div><hr></div><p>Below is the final HLO that the frontend TPU compiler emits. </p><div class="github-gist" data-attrs="{&quot;innerHTML&quot;:&quot;<div id=\&quot;gist143984843\&quot; class=\&quot;gist\&quot;>\n    <div class=\&quot;gist-file\&quot; translate=\&quot;no\&quot; data-color-mode=\&quot;light\&quot; data-light-theme=\&quot;light\&quot;>\n      <div class=\&quot;gist-data\&quot;>\n        <div class=\&quot;js-gist-file-update-container js-task-list-container\&quot;>\n  <div id=\&quot;file-final_hlo-md\&quot; class=\&quot;file my-2\&quot;>\n      <div id=\&quot;file-final_hlo-md-readme\&quot; class=\&quot;Box-body readme blob p-5 p-xl-6 \&quot;\n    style=\&quot;overflow: auto\&quot; tabindex=\&quot;0\&quot; role=\&quot;region\&quot;\n    aria-label=\&quot;final_hlo.md content, created by patrick-toulme on 10:28PM today.\&quot;\n  >\n    <article class=\&quot;markdown-body entry-content container-lg\&quot; itemprop=\&quot;text\&quot;><pre><code>HloModule jit_mini_attention, is_scheduled=true, entry_computation_layout={(f32[16,64]{1,0:T(8,128)}, f32[64,64]{1,0:T(8,128)}, f32[64,32]{0,1:T(8,128)})-&amp;gt;f32[16,32]{1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}\n\n%fused_computation.1 (param_0.21: f32[16], param_1.23: f32[16], param_2.10: f32[16,64], param_3.5: f32[16]) -&amp;gt; f32[16,64] {\n  %param_2.10 = f32[16,64]{1,0:T(8,128)S(1)} parameter(2)\n  %param_3.5 = f32[16]{0:T(128)S(1)} parameter(3)\n  %broadcast.29 = f32[16,64]{1,0:T(8,128)} broadcast(%param_3.5), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;8\&quot;]}}\n  %divide.5 = f32[16,64]{1,0:T(8,128)} divide(%param_2.10, %broadcast.29), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %param_1.23 = f32[16]{0:T(128)S(1)} parameter(1)\n  %broadcast.24 = f32[16,64]{1,0:T(8,128)} broadcast(%param_1.23), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;8\&quot;]}}\n  %subtract.3 = f32[16,64]{1,0:T(8,128)} subtract(%divide.5, %broadcast.24), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %exponential.3 = f32[16,64]{1,0:T(8,128)} exponential(%subtract.3), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/exp\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %param_0.21 = f32[16]{0:T(128)S(1)} parameter(0)\n  %broadcast.18 = f32[16,64]{1,0:T(8,128)} broadcast(%param_0.21), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;8\&quot;]}}\n  ROOT %divide.1 = f32[16,64]{1,0:T(8,128)} divide(%exponential.3, %broadcast.18), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n}\n\n%bitcast_fusion.1 (bitcast_input.1: f32[64,32]) -&amp;gt; f32[64,32] {\n  %bitcast_input.1 = f32[64,32]{0,1:T(8,128)S(1)} parameter(0)\n  ROOT %bitcast.1 = f32[64,32]{0,1:T(8,128)} bitcast(%bitcast_input.1)\n}\n\n%fused_computation (param_0.1: f32[64,32], param_1.21: f32[16], param_2.9: f32[16], param_3.3: f32[16,64], param_4: f32[16]) -&amp;gt; f32[16,32] {\n  %param_1.21 = f32[16]{0:T(128)S(1)} parameter(1)\n  %param_2.9 = f32[16]{0:T(128)S(1)} parameter(2)\n  %param_3.3 = f32[16,64]{1,0:T(8,128)S(1)} parameter(3)\n  %param_4 = f32[16]{0:T(128)S(1)} parameter(4)\n  %fusion.1 = f32[16,64]{1,0:T(8,128)} fusion(%param_1.21, %param_2.9, %param_3.3, %param_4), kind=kLoop, calls=%fused_computation.1, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %param_0.1 = f32[64,32]{0,1:T(8,128)S(1)} parameter(0)\n  %fusion.7 = f32[64,32]{0,1:T(8,128)} fusion(%param_0.1), kind=kLoop, calls=%bitcast_fusion.1\n  ROOT %convolution.2 = f32[16,32]{1,0:T(8,128)} convolution(%fusion.1, %fusion.7), dim_labels=bf_io-&amp;gt;bf, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/matmul_2/dot_general\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=56}\n}\n\n%region_0.15 (Arg_0.12: f32[], Arg_1.13: f32[]) -&amp;gt; f32[] {\n  %Arg_1.13 = f32[]{:T(128)} parameter(1), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot;}\n  %Arg_0.12 = f32[]{:T(128)} parameter(0), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot;}\n  ROOT %add.14 = f32[]{:T(128)} add(%Arg_0.12, %Arg_1.13), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n}\n\n%bitcast_fusion (bitcast_input: f32[16,64]) -&amp;gt; f32[16,64] {\n  %bitcast_input = f32[16,64]{1,0:T(8,128)} parameter(0)\n  ROOT %bitcast = f32[16,64]{1,0:T(8,128)} bitcast(%bitcast_input)\n}\n\n%bitcast_fusion.2 (bitcast_input.2: f32[64,64]) -&amp;gt; f32[64,64] {\n  %bitcast_input.2 = f32[64,64]{1,0:T(8,128)S(1)} parameter(0)\n  ROOT %bitcast.2 = f32[64,64]{1,0:T(8,128)} bitcast(%bitcast_input.2)\n}\n\n%fused_computation.3 (param_0.33: f32[16,64], param_1.35: f32[64,64]) -&amp;gt; (f32[16], f32[16,64]) {\n  %param_0.33 = f32[16,64]{1,0:T(8,128)} parameter(0)\n  %fusion.6 = f32[16,64]{1,0:T(8,128)} fusion(%param_0.33), kind=kLoop, calls=%bitcast_fusion\n  %param_1.35 = f32[64,64]{1,0:T(8,128)S(1)} parameter(1)\n  %fusion.8 = f32[64,64]{1,0:T(8,128)} fusion(%param_1.35), kind=kLoop, calls=%bitcast_fusion.2\n  %convolution.3 = f32[16,64]{1,0:T(8,128)S(1)} convolution(%fusion.6, %fusion.8), dim_labels=bf_io-&amp;gt;bf, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/matmul_1/dot_general\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=35}\n  %multiply.6 = f32[16,64]{1,0:T(8,128)} multiply(%convolution.3, %convolution.3), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/integer_pow\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %constant.14 = f32[]{:T(128)} constant(0)\n  %reduce.0 = f32[16]{0:T(128)S(1)} reduce(%multiply.6, %constant.14), dimensions={1}, to_apply=%region_0.15, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  ROOT %tuple = (f32[16]{0:T(128)S(1)}, f32[16,64]{1,0:T(8,128)S(1)}) tuple(%reduce.0, %convolution.3)\n}\n\n%region_2.39 (Arg_0.36: f32[], Arg_1.37: f32[]) -&amp;gt; f32[] {\n  %Arg_1.37 = f32[]{:T(128)} parameter(1), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot;}\n  %Arg_0.36 = f32[]{:T(128)} parameter(0), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot;}\n  ROOT %add.38 = f32[]{:T(128)} add(%Arg_0.36, %Arg_1.37), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n}\n\n%fused_computation.4 (param_0.30: f32[16], param_1.33: f32[16,64], param_2.17: f32[16]) -&amp;gt; f32[16] {\n  %param_1.33 = f32[16,64]{1,0:T(8,128)S(1)} parameter(1)\n  %param_2.17 = f32[16]{0:T(128)S(1)} parameter(2)\n  %broadcast.32 = f32[16,64]{1,0:T(8,128)} broadcast(%param_2.17), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;8\&quot;]}}\n  %divide.7 = f32[16,64]{1,0:T(8,128)} divide(%param_1.33, %broadcast.32), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %param_0.30 = f32[16]{0:T(128)S(1)} parameter(0)\n  %broadcast.25 = f32[16,64]{1,0:T(8,128)} broadcast(%param_0.30), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;8\&quot;]}}\n  %subtract.5 = f32[16,64]{1,0:T(8,128)} subtract(%divide.7, %broadcast.25), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/sub\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %exponential.5 = f32[16,64]{1,0:T(8,128)} exponential(%subtract.5), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/exp\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=49}\n  %constant.13 = f32[]{:T(128)} constant(0)\n  ROOT %reduce.1 = f32[16]{0:T(128)S(1)} reduce(%exponential.5, %constant.13), dimensions={1}, to_apply=%region_2.39, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n}\n\n%region_1.28 (Arg_0.25: f32[], Arg_1.26: f32[]) -&amp;gt; f32[] {\n  %Arg_1.26 = f32[]{:T(128)} parameter(1), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot;}\n  %Arg_0.25 = f32[]{:T(128)} parameter(0), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot;}\n  ROOT %maximum.27 = f32[]{:T(128)} maximum(%Arg_0.25, %Arg_1.26), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=48}\n}\n\n%fused_computation.8 (param_0.32: f32[16,64], param_1.34: f32[16]) -&amp;gt; f32[16] {\n  %param_0.32 = f32[16,64]{1,0:T(8,128)S(1)} parameter(0)\n  %param_1.34 = f32[16]{0:T(128)S(1)} parameter(1)\n  %broadcast.35 = f32[16,64]{1,0:T(8,128)} broadcast(%param_1.34), dimensions={0}, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[\&quot;8\&quot;]}}\n  %divide.9 = f32[16,64]{1,0:T(8,128)} divide(%param_0.32, %broadcast.35), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=42}\n  %constant.15 = f32[]{:T(128)} constant(-inf)\n  ROOT %reduce.2 = f32[16]{0:T(128)S(1)} reduce(%divide.9, %constant.15), dimensions={1}, to_apply=%region_1.28, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=48}\n}\n\n%fused_computation.9 (param_0.29: f32[16]) -&amp;gt; f32[16] {\n  %param_0.29 = f32[16]{0:T(128)S(1)} parameter(0)\n  %constant.12 = f32[]{:T(128)} constant(0.015625)\n  %broadcast.37 = f32[16]{0:T(128)} broadcast(%constant.12), dimensions={}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[]}}\n  %multiply.7 = f32[16]{0:T(128)} multiply(%param_0.29, %broadcast.37), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/div\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %constant.16 = f32[]{:T(128)} constant(1e-06)\n  %broadcast.36 = f32[16]{0:T(128)} broadcast(%constant.16), dimensions={}, backend_config={\&quot;flag_configs\&quot;:[],\&quot;scoped_memory_configs\&quot;:[],\&quot;used_scoped_memory_configs\&quot;:[],\&quot;output_chunk_bound_config\&quot;:{\&quot;output_chunk_bound\&quot;:[]}}\n  %add.5 = f32[16]{0:T(128)} add(%multiply.7, %broadcast.36), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/add\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  ROOT %sqrt.5 = f32[16]{0:T(128)S(1)} sqrt(%add.5), metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/sqrt\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n}\n\nENTRY %main.47 (Arg_0.1: f32[16,64], Arg_1.2: f32[64,64], Arg_2.3: f32[64,32]) -&amp;gt; f32[16,32] {\n  %Arg_1.2 = f32[64,64]{1,0:T(8,128)} parameter(1), metadata={op_name=\&quot;w1\&quot;}\n  %copy-start = (f32[64,64]{1,0:T(8,128)S(1)}, f32[64,64]{1,0:T(8,128)}, u32[]{:S(2)}) copy-start(%Arg_1.2), cross_program_prefetch_index=0\n  %Arg_2.3 = f32[64,32]{0,1:T(8,128)} parameter(2), metadata={op_name=\&quot;w2\&quot;}\n  %Arg_0.1 = f32[16,64]{1,0:T(8,128)} parameter(0), metadata={op_name=\&quot;x\&quot;}\n  %copy-done = f32[64,64]{1,0:T(8,128)S(1)} copy-done(%copy-start)\n  %multiply_reduce_fusion = (f32[16]{0:T(128)S(1)}, f32[16,64]{1,0:T(8,128)S(1)}) fusion(%Arg_0.1, %copy-done), kind=kOutput, calls=%fused_computation.3, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/matmul_1/dot_general\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=35}\n  %copy-start.1 = (f32[64,32]{0,1:T(8,128)S(1)}, f32[64,32]{0,1:T(8,128)}, u32[]{:S(2)}) copy-start(%Arg_2.3)\n  %get-tuple-element.1 = f32[16,64]{1,0:T(8,128)S(1)} get-tuple-element(%multiply_reduce_fusion), index=1, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %get-tuple-element = f32[16]{0:T(128)S(1)} get-tuple-element(%multiply_reduce_fusion), index=0, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %add_sqrt_fusion = f32[16]{0:T(128)S(1)} fusion(%get-tuple-element), kind=kLoop, calls=%fused_computation.9, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/rms_norm/sqrt\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=41}\n  %fusion.5 = f32[16]{0:T(128)S(1)} fusion(%get-tuple-element.1, %add_sqrt_fusion), kind=kLoop, calls=%fused_computation.8, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_max\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=48}\n  %fusion.2 = f32[16]{0:T(128)S(1)} fusion(%fusion.5, %get-tuple-element.1, %add_sqrt_fusion), kind=kLoop, calls=%fused_computation.4, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/softmax/reduce_sum\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=50}\n  %copy-done.1 = f32[64,32]{0,1:T(8,128)S(1)} copy-done(%copy-start.1)\n  ROOT %fusion = f32[16,32]{1,0:T(8,128)} fusion(%copy-done.1, %fusion.2, %fusion.5, %get-tuple-element.1, %add_sqrt_fusion), kind=kOutput, calls=%fused_computation, metadata={op_name=\&quot;jit(mini_attention)/jit(main)/matmul_2/dot_general\&quot; source_file=\&quot;/home/ptoulme/tpu.py\&quot; source_line=56}\n}\n</code></pre>\n</article>\n  </div>\n\n  </div>\n</div>\n\n      </div>\n      <div class=\&quot;gist-meta\&quot;>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/23183f93a42054bd6fb5cdcd4efb72ab/raw/8e48a6298bdadbf8f7b50cc2b50e504a1bace21d/final_hlo.md/%22 style=\&quot;float:right\&quot; class=\&quot;Link--inTextBlock\&quot;>view raw</a>\n        <a href=/__u/patricktoulme.substack.com/%22https://gist.github.com/patrick-toulme/23183f93a42054bd6fb5cdcd4efb72ab#file-final_hlo-md\%22 class=\&quot;Link--inTextBlock\&quot;>\n          final_hlo.md\n        </a>\n        hosted with &amp;#10084; by <a class=\&quot;Link--inTextBlock\&quot; href=/__u/patricktoulme.substack.com/%22https://github.com/%22>GitHub</a>\n      </div>\n    </div>\n</div>\n&quot;,&quot;stylesheet&quot;:&quot;https://github.githubassets.com/assets/gist-embed-ed91f9610ae6.css&quot;}" data-component-name="GitgistToDOM"><link rel="stylesheet" href="https://github.githubassets.com/assets/gist-embed-ed91f9610ae6.css"><div id="gist143984843" class="gist">
    <div class="gist-file" data-color-mode="light" data-light-theme="light">
      <div class="gist-data">
        <div class="js-gist-file-update-container js-task-list-container">
  <div id="file-final_hlo-md" class="file my-2">
      <div id="file-final_hlo-md-readme" class="Box-body readme blob p-5 p-xl-6 " style="overflow:auto">
    <article class="markdown-body entry-content container-lg" itemprop="text"><pre><code>HloModule jit_mini_attention, is_scheduled=true, entry_computation_layout={(f32[16,64]{1,0:T(8,128)}, f32[64,64]{1,0:T(8,128)}, f32[64,32]{0,1:T(8,128)})-&gt;f32[16,32]{1,0:T(8,128)}}, allow_spmd_sharding_propagation_to_parameters={true,true,true}, allow_spmd_sharding_propagation_to_output={true}

%fused_computation.1 (param_0.21: f32[16], param_1.23: f32[16], param_2.10: f32[16,64], param_3.5: f32[16]) -&gt; f32[16,64] {
  %param_2.10 = f32[16,64]{1,0:T(8,128)S(1)} parameter(2)
  %param_3.5 = f32[16]{0:T(128)S(1)} parameter(3)
  %broadcast.29 = f32[16,64]{1,0:T(8,128)} broadcast(%param_3.5), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["8"]}}
  %divide.5 = f32[16,64]{1,0:T(8,128)} divide(%param_2.10, %broadcast.29), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %param_1.23 = f32[16]{0:T(128)S(1)} parameter(1)
  %broadcast.24 = f32[16,64]{1,0:T(8,128)} broadcast(%param_1.23), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["8"]}}
  %subtract.3 = f32[16,64]{1,0:T(8,128)} subtract(%divide.5, %broadcast.24), metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}
  %exponential.3 = f32[16,64]{1,0:T(8,128)} exponential(%subtract.3), metadata={op_name="jit(mini_attention)/jit(main)/softmax/exp" source_file="/home/ptoulme/tpu.py" source_line=49}
  %param_0.21 = f32[16]{0:T(128)S(1)} parameter(0)
  %broadcast.18 = f32[16,64]{1,0:T(8,128)} broadcast(%param_0.21), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["8"]}}
  ROOT %divide.1 = f32[16,64]{1,0:T(8,128)} divide(%exponential.3, %broadcast.18), metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}
}

%bitcast_fusion.1 (bitcast_input.1: f32[64,32]) -&gt; f32[64,32] {
  %bitcast_input.1 = f32[64,32]{0,1:T(8,128)S(1)} parameter(0)
  ROOT %bitcast.1 = f32[64,32]{0,1:T(8,128)} bitcast(%bitcast_input.1)
}

%fused_computation (param_0.1: f32[64,32], param_1.21: f32[16], param_2.9: f32[16], param_3.3: f32[16,64], param_4: f32[16]) -&gt; f32[16,32] {
  %param_1.21 = f32[16]{0:T(128)S(1)} parameter(1)
  %param_2.9 = f32[16]{0:T(128)S(1)} parameter(2)
  %param_3.3 = f32[16,64]{1,0:T(8,128)S(1)} parameter(3)
  %param_4 = f32[16]{0:T(128)S(1)} parameter(4)
  %fusion.1 = f32[16,64]{1,0:T(8,128)} fusion(%param_1.21, %param_2.9, %param_3.3, %param_4), kind=kLoop, calls=%fused_computation.1, metadata={op_name="jit(mini_attention)/jit(main)/softmax/div" source_file="/home/ptoulme/tpu.py" source_line=50}
  %param_0.1 = f32[64,32]{0,1:T(8,128)S(1)} parameter(0)
  %fusion.7 = f32[64,32]{0,1:T(8,128)} fusion(%param_0.1), kind=kLoop, calls=%bitcast_fusion.1
  ROOT %convolution.2 = f32[16,32]{1,0:T(8,128)} convolution(%fusion.1, %fusion.7), dim_labels=bf_io-&gt;bf, metadata={op_name="jit(mini_attention)/jit(main)/matmul_2/dot_general" source_file="/home/ptoulme/tpu.py" source_line=56}
}

%region_0.15 (Arg_0.12: f32[], Arg_1.13: f32[]) -&gt; f32[] {
  %Arg_1.13 = f32[]{:T(128)} parameter(1), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum"}
  %Arg_0.12 = f32[]{:T(128)} parameter(0), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum"}
  ROOT %add.14 = f32[]{:T(128)} add(%Arg_0.12, %Arg_1.13), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=41}
}

%bitcast_fusion (bitcast_input: f32[16,64]) -&gt; f32[16,64] {
  %bitcast_input = f32[16,64]{1,0:T(8,128)} parameter(0)
  ROOT %bitcast = f32[16,64]{1,0:T(8,128)} bitcast(%bitcast_input)
}

%bitcast_fusion.2 (bitcast_input.2: f32[64,64]) -&gt; f32[64,64] {
  %bitcast_input.2 = f32[64,64]{1,0:T(8,128)S(1)} parameter(0)
  ROOT %bitcast.2 = f32[64,64]{1,0:T(8,128)} bitcast(%bitcast_input.2)
}

%fused_computation.3 (param_0.33: f32[16,64], param_1.35: f32[64,64]) -&gt; (f32[16], f32[16,64]) {
  %param_0.33 = f32[16,64]{1,0:T(8,128)} parameter(0)
  %fusion.6 = f32[16,64]{1,0:T(8,128)} fusion(%param_0.33), kind=kLoop, calls=%bitcast_fusion
  %param_1.35 = f32[64,64]{1,0:T(8,128)S(1)} parameter(1)
  %fusion.8 = f32[64,64]{1,0:T(8,128)} fusion(%param_1.35), kind=kLoop, calls=%bitcast_fusion.2
  %convolution.3 = f32[16,64]{1,0:T(8,128)S(1)} convolution(%fusion.6, %fusion.8), dim_labels=bf_io-&gt;bf, metadata={op_name="jit(mini_attention)/jit(main)/matmul_1/dot_general" source_file="/home/ptoulme/tpu.py" source_line=35}
  %multiply.6 = f32[16,64]{1,0:T(8,128)} multiply(%convolution.3, %convolution.3), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/integer_pow" source_file="/home/ptoulme/tpu.py" source_line=41}
  %constant.14 = f32[]{:T(128)} constant(0)
  %reduce.0 = f32[16]{0:T(128)S(1)} reduce(%multiply.6, %constant.14), dimensions={1}, to_apply=%region_0.15, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=41}
  ROOT %tuple = (f32[16]{0:T(128)S(1)}, f32[16,64]{1,0:T(8,128)S(1)}) tuple(%reduce.0, %convolution.3)
}

%region_2.39 (Arg_0.36: f32[], Arg_1.37: f32[]) -&gt; f32[] {
  %Arg_1.37 = f32[]{:T(128)} parameter(1), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum"}
  %Arg_0.36 = f32[]{:T(128)} parameter(0), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum"}
  ROOT %add.38 = f32[]{:T(128)} add(%Arg_0.36, %Arg_1.37), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=50}
}

%fused_computation.4 (param_0.30: f32[16], param_1.33: f32[16,64], param_2.17: f32[16]) -&gt; f32[16] {
  %param_1.33 = f32[16,64]{1,0:T(8,128)S(1)} parameter(1)
  %param_2.17 = f32[16]{0:T(128)S(1)} parameter(2)
  %broadcast.32 = f32[16,64]{1,0:T(8,128)} broadcast(%param_2.17), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["8"]}}
  %divide.7 = f32[16,64]{1,0:T(8,128)} divide(%param_1.33, %broadcast.32), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %param_0.30 = f32[16]{0:T(128)S(1)} parameter(0)
  %broadcast.25 = f32[16,64]{1,0:T(8,128)} broadcast(%param_0.30), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["8"]}}
  %subtract.5 = f32[16,64]{1,0:T(8,128)} subtract(%divide.7, %broadcast.25), metadata={op_name="jit(mini_attention)/jit(main)/softmax/sub" source_file="/home/ptoulme/tpu.py" source_line=49}
  %exponential.5 = f32[16,64]{1,0:T(8,128)} exponential(%subtract.5), metadata={op_name="jit(mini_attention)/jit(main)/softmax/exp" source_file="/home/ptoulme/tpu.py" source_line=49}
  %constant.13 = f32[]{:T(128)} constant(0)
  ROOT %reduce.1 = f32[16]{0:T(128)S(1)} reduce(%exponential.5, %constant.13), dimensions={1}, to_apply=%region_2.39, metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=50}
}

%region_1.28 (Arg_0.25: f32[], Arg_1.26: f32[]) -&gt; f32[] {
  %Arg_1.26 = f32[]{:T(128)} parameter(1), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max"}
  %Arg_0.25 = f32[]{:T(128)} parameter(0), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max"}
  ROOT %maximum.27 = f32[]{:T(128)} maximum(%Arg_0.25, %Arg_1.26), metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max" source_file="/home/ptoulme/tpu.py" source_line=48}
}

%fused_computation.8 (param_0.32: f32[16,64], param_1.34: f32[16]) -&gt; f32[16] {
  %param_0.32 = f32[16,64]{1,0:T(8,128)S(1)} parameter(0)
  %param_1.34 = f32[16]{0:T(128)S(1)} parameter(1)
  %broadcast.35 = f32[16,64]{1,0:T(8,128)} broadcast(%param_1.34), dimensions={0}, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":["8"]}}
  %divide.9 = f32[16,64]{1,0:T(8,128)} divide(%param_0.32, %broadcast.35), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=42}
  %constant.15 = f32[]{:T(128)} constant(-inf)
  ROOT %reduce.2 = f32[16]{0:T(128)S(1)} reduce(%divide.9, %constant.15), dimensions={1}, to_apply=%region_1.28, metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max" source_file="/home/ptoulme/tpu.py" source_line=48}
}

%fused_computation.9 (param_0.29: f32[16]) -&gt; f32[16] {
  %param_0.29 = f32[16]{0:T(128)S(1)} parameter(0)
  %constant.12 = f32[]{:T(128)} constant(0.015625)
  %broadcast.37 = f32[16]{0:T(128)} broadcast(%constant.12), dimensions={}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":[]}}
  %multiply.7 = f32[16]{0:T(128)} multiply(%param_0.29, %broadcast.37), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/div" source_file="/home/ptoulme/tpu.py" source_line=41}
  %constant.16 = f32[]{:T(128)} constant(1e-06)
  %broadcast.36 = f32[16]{0:T(128)} broadcast(%constant.16), dimensions={}, backend_config={"flag_configs":[],"scoped_memory_configs":[],"used_scoped_memory_configs":[],"output_chunk_bound_config":{"output_chunk_bound":[]}}
  %add.5 = f32[16]{0:T(128)} add(%multiply.7, %broadcast.36), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/add" source_file="/home/ptoulme/tpu.py" source_line=41}
  ROOT %sqrt.5 = f32[16]{0:T(128)S(1)} sqrt(%add.5), metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/sqrt" source_file="/home/ptoulme/tpu.py" source_line=41}
}

ENTRY %main.47 (Arg_0.1: f32[16,64], Arg_1.2: f32[64,64], Arg_2.3: f32[64,32]) -&gt; f32[16,32] {
  %Arg_1.2 = f32[64,64]{1,0:T(8,128)} parameter(1), metadata={op_name="w1"}
  %copy-start = (f32[64,64]{1,0:T(8,128)S(1)}, f32[64,64]{1,0:T(8,128)}, u32[]{:S(2)}) copy-start(%Arg_1.2), cross_program_prefetch_index=0
  %Arg_2.3 = f32[64,32]{0,1:T(8,128)} parameter(2), metadata={op_name="w2"}
  %Arg_0.1 = f32[16,64]{1,0:T(8,128)} parameter(0), metadata={op_name="x"}
  %copy-done = f32[64,64]{1,0:T(8,128)S(1)} copy-done(%copy-start)
  %multiply_reduce_fusion = (f32[16]{0:T(128)S(1)}, f32[16,64]{1,0:T(8,128)S(1)}) fusion(%Arg_0.1, %copy-done), kind=kOutput, calls=%fused_computation.3, metadata={op_name="jit(mini_attention)/jit(main)/matmul_1/dot_general" source_file="/home/ptoulme/tpu.py" source_line=35}
  %copy-start.1 = (f32[64,32]{0,1:T(8,128)S(1)}, f32[64,32]{0,1:T(8,128)}, u32[]{:S(2)}) copy-start(%Arg_2.3)
  %get-tuple-element.1 = f32[16,64]{1,0:T(8,128)S(1)} get-tuple-element(%multiply_reduce_fusion), index=1, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=41}
  %get-tuple-element = f32[16]{0:T(128)S(1)} get-tuple-element(%multiply_reduce_fusion), index=0, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=41}
  %add_sqrt_fusion = f32[16]{0:T(128)S(1)} fusion(%get-tuple-element), kind=kLoop, calls=%fused_computation.9, metadata={op_name="jit(mini_attention)/jit(main)/rms_norm/sqrt" source_file="/home/ptoulme/tpu.py" source_line=41}
  %fusion.5 = f32[16]{0:T(128)S(1)} fusion(%get-tuple-element.1, %add_sqrt_fusion), kind=kLoop, calls=%fused_computation.8, metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_max" source_file="/home/ptoulme/tpu.py" source_line=48}
  %fusion.2 = f32[16]{0:T(128)S(1)} fusion(%fusion.5, %get-tuple-element.1, %add_sqrt_fusion), kind=kLoop, calls=%fused_computation.4, metadata={op_name="jit(mini_attention)/jit(main)/softmax/reduce_sum" source_file="/home/ptoulme/tpu.py" source_line=50}
  %copy-done.1 = f32[64,32]{0,1:T(8,128)S(1)} copy-done(%copy-start.1)
  ROOT %fusion = f32[16,32]{1,0:T(8,128)} fusion(%copy-done.1, %fusion.2, %fusion.5, %get-tuple-element.1, %add_sqrt_fusion), kind=kOutput, calls=%fused_computation, metadata={op_name="jit(mini_attention)/jit(main)/matmul_2/dot_general" source_file="/home/ptoulme/tpu.py" source_line=56}
}
</code></pre>
</article>
  </div>

  </div>
</div>

      </div>
      <div class="gist-meta">
        <a href="https://gist.github.com/patrick-toulme/23183f93a42054bd6fb5cdcd4efb72ab/raw/8e48a6298bdadbf8f7b50cc2b50e504a1bace21d/final_hlo.md" style="float:right" class="Link--inTextBlock">view raw</a>
        <a href="https://gist.github.com/patrick-toulme/23183f93a42054bd6fb5cdcd4efb72ab#file-final_hlo-md" class="Link--inTextBlock">
          final_hlo.md
        </a>
        hosted with &#10084; by <a class="Link--inTextBlock" href="https://github.com">GitHub</a>
      </div>
    </div>
</div>
</div><p>That covers the HLO optimization passes. At this point we have fused, scheduled operations with memory annotations &#8212; but it's still relatively hardware-agnostic. Next, the TPU backend translates each fusion into LLO (Low Level Operators), a representation that maps directly to physical hardware units: the MXU for matmuls, VPU for elementwise ops, XLU for transposes, and DMA engines for memory movement.</p><h2>LLO IR</h2><p>After HLO optimizations, XLA&#8217;s TPU backend translates each fusion into <strong>LLO (Low Level Operators)</strong>&#8212;a TPU-specific intermediate representation that maps directly to hardware units. Each fusion becomes a separate LLO program that goes through in my case 78 optimization passes before producing final VLIW bundles.</p><h3>TPU Hardware Units</h3><p>Before diving into LLO, here&#8217;s a quick orientation. Each TPU Trillium V6e TensorCore contains:</p><ul><li><p><strong>MXU (Matrix Unit)</strong>: Two 256&#215;256 systolic arrays for matrix multiplies. This is where the FLOPS come from.</p></li><li><p><strong>VPU (Vector Processing Unit)</strong>: Handles elementwise ops (add, mul, exp, etc.) across 8 sublanes &#215; 128 lanes.</p></li><li><p><strong>XLU (Transpose Unit)</strong>: Cross-lane shuffles, transposes, and permutations.</p></li><li><p><strong>Scalar Unit</strong>: Scalar operations, address calculation, and control flow.</p></li><li><p><strong>DMA engines</strong>: Async memory transfers between HBM and VMEM.</p></li></ul><p>The compiler&#8217;s job is to keep all of these busy simultaneously. A well-scheduled TPU program overlaps MXU matmuls with VPU elementwise ops with DMA transfers &#8212; all in the same VLIW bundle.</p><h3>Deep Dive: Compiling <code>multiply_reduce_fusion</code></h3><p>Let&#8217;s trace how <code>multiply_reduce_fusion</code>&#8212;which computes our first matrix multiply and the squared values for RMS normalization&#8212;transforms from initial LLO to final machine code.</p><p><strong>Initial LLO (Pass 02)</strong></p><p>The compiler first translates the fusion HLO region into a verbose, unscheduled LLO representation. Here&#8217;s how the matrix multiply setup begins:</p><pre><code><code>$region0: #{multiply_reduce_fusion}
  #allocation0 [shape = 'f32[1024]{0}', space=vmem, size = 0x1000, tag = 'scoped memory']
  #allocation1 [shape = 'f32[16]{0:T(1024)S(1)}', space=vmem, size = 0x1000, tag = 'reduce buffer']
  
  %s0 = inlined_call_operand.hbm [shape: f32[16,64], index: 0, kind: input]
  %s1 = inlined_call_operand.vmem [shape: f32[64,64], index: 1, kind: input]
  %s2 = inlined_call_operand.vmem [shape: f32[16], index: 2, kind: output]
  %s3 = inlined_call_operand.vmem [shape: f32[16,64], index: 3, kind: output]
</code></code></pre><p>The compiler allocates VMEM buffers and establishes that operand 0 comes from HBM (High Bandwidth Memory) while operand 1 is already in VMEM. This is a <strong>multi-output fusion</strong>&#8212;it produces both the matmul result (<code>%s3</code>) and the sum of squares (<code>%s2</code>) for RMS norm.</p><p>The initial MXU operations use <code>vmatpush</code> to load weight columns into the systolic array:</p><pre><code><code>%v63 = vld [vmem:[%s55] sm:$0xff]
%64 = vmatpush.bf16.msra.mxu0 %v63
...
%v93 = vld [vmem:[#allocation0] sm:$0xff]
%94 = vmatmul.bf16.gmra.mxu0 %v93
%v95 = vpop.f32.mrf.mxu0
</code></code></pre><p>This sequence streams weight tiles through the MXU in bf16 precision, then pops the accumulated f32 results.</p><p><strong>The Reduction Pattern</strong></p><p>After the matmul, we need to sum squared values across columns. The compiler generates a <strong>cross-lane reduction</strong> using the XLU (transpose unit):</p><pre><code><code>%141 = vxpose.xlu0.b32.start [1/2] (short) /*vx=*/%v138, /*width=*/128
%142 = vxpose.xlu0.b32.end [2/2] (short) /*vx=*/%v140, /*width=*/128
%v143 = vpop.trf.xlu0
%v144 = vpop.trf.xlu0
...
%v158 = vpop.trf.xlu0  // 16 pops total</code></code></pre><p>After transposing, a tree reduction sums the 16 lanes:</p><pre><code><code>%v161 = vadd.f32 0.0, %v143
%v165 = vadd.f32 %v161, %v144
%v169 = vadd.f32 %v165, %v145
...
%v221 = vadd.f32 %v217, %v158
</code></code></pre><p>Then a sublane rotation pattern completes the reduction:</p><pre><code><code>%v223 = vrot.slane %v221, 4   // rotate by 4
%v226 = vadd.f32 %v221, %v223
%v228 = vrot.slane %v226, 2   // rotate by 2
%v231 = vadd.f32 %v226, %v228
%v233 = vrot.slane %v231, 1   // rotate by 1
%v236 = vadd.f32 %v231, %v233
</code></code></pre><p>This classic parallel reduction pattern uses log&#8322;(n) steps&#8212;rotate by 4, 2, then 1&#8212;to sum 8 sublanes into a single scalar.</p><p><strong>Final VLIW Bundles (Pass 78)</strong></p><p>After all optimization passes, the compiler produces <strong>71 tightly-packed VLIW bundles</strong>. Each bundle groups independent operations that execute in parallel across hardware units:</p><p>This is VLIW (Very Long Instruction Word) execution &#8212; the compiler statically packs independent operations into fixed-width bundles, and the hardware executes everything in a bundle simultaneously without runtime dependency checking.</p><pre><code><code>0x9 : { %22 = dma.hbm_to_vmem [thread:$0] /*hbm=*/%s359_s0, ...
       %v28_v1 = vlaneseq  
       %v301_v2 = vmov 0.0
       %vm302_vm0 = vmmov 0
       %v247_v4 = vld [vmem:[%s360_s1 + $0x30] sm:$0xff]
       %v248_v5 = vld [vmem:[%s360_s1 + $0x38] sm:$0xff]
       %v249_v6 = vld [vmem:[%s360_s1 + $0x20] sm:$0xff] }</code></code></pre><p>This single bundle executes <strong>7 operations simultaneously</strong>: one DMA transfer from HBM, one lane sequence generation, two constant materializations, and three VMEM loads. The compiler has carefully scheduled these to avoid resource conflicts.</p><p>The MXU operations use the optimized <code>vmatpush3</code> instruction:</p><pre><code><code>0xc : { %261 = vmatpush3.bf16.msra.mxu0 %v61_v7
        %v89_v15 = vpack.c.bf16 %v253_v14, %v85_v13
        %297 = dma.done.wait [#allocation4], 256 }
</code></code></pre><p>Three operations in parallel: push a weight tile, pack bf16 values for the next push, and synchronize the DMA. The <code>bf16</code> suffix indicates this matmul uses bfloat16 precision on the MXU, which provides 2x throughput compared to f32.</p><p>The matmul result extraction uses a masked variant:</p><pre><code><code>0x14 : { %269 = vmatmul.mubr.msk.bf16.vlgmr.msra.gmra.mxu0 %vm257_vm2, %v258_v18 }
0x15 : { %v95_v19 = vpop.f32.mrf.mxu0 }
0x16 : { %v98_v20 = vmul.f32 %v95_v19, %v95_v19
         %105 = vst [vmem:[%s362_s3] sm:$0xff] /*vst_source=*/%v95_v19 }
</code></code></pre><p>The squaring happens immediately in bundle 0x16, overlapped with storing the matmul result&#8212;the multi-output fusion lets us reuse <code>%v95_v19</code> for both outputs.</p><div class="native-video-embed" data-component-name="VideoPlaceholder" data-attrs="{&quot;mediaUploadId&quot;:&quot;d9d4ccb7-bc2f-48c7-826b-98942dd8c27f&quot;,&quot;duration&quot;:null}"></div><p>In the above video, we can visualize this fusion on a TPU V6e Trillium. Note we are using one MXU in this visualization, while a Trillium has 2 256x256 MXUs.</p><div><hr></div><h3>LLO Compilation Overview: All Fusions</h3><p>Our <code>mini_attention</code> function compiles into <strong>five fusions</strong>, each with its own LLO compilation producing different bundle counts and hardware utilization patterns.</p><p><strong>multiply_reduce_fusion (71 bundles)</strong></p><p>This is our largest fusion, implementing both <code>matmul_1</code> (the first matrix multiply) and the sum-of-squares reduction for RMS normalization. The fusion demonstrates sophisticated hardware utilization:</p><ul><li><p>MXU0 handles the 16&#215;64 by 64&#215;64 bf16 matmul</p></li><li><p>VPU computes the squared values immediately after MXU results</p></li><li><p>XLU0 performs the cross-lane transpose for reduction</p></li><li><p>DMA streams the input activations from HBM</p></li></ul><p>The 71 bundles break down into: DMA setup, weight loading, matmul execution, squaring, transpose-based reduction, and sublane rotation reduction.</p><p><strong>add_sqrt_fusion (10 bundles)</strong></p><p>The smallest and simplest fusion computes <code>sqrt(mean + epsilon)</code> for RMS normalization:</p><pre><code><code>0x1 : { %v2_v0 = vld [vmem:[%s37_s0] sm:$0x1] }
0x2 : { %v5_v1 = vmul.f32 0.015625, %v2_v0 }
0x3 : { %v9_v2 = vadd.f32 1e-06, %v5_v1 }
0x4 : { %19 = vrsqrt.f32 %v9_v2
        %vm13_vm0 = vcmp.eq.f32.partialorder %v9_v2, inf
        %v16_v4 = vand.u32 2147483648, %v9_v2
        %vm15_vm1 = vcmp.eq.f32.partialorder %v9_v2, 0.0 }
0x5 : { %v20_v3 = vpop.eup %19 }
0x6 : { %v12_v5 = vmul.f32 %v20_v3, %v9_v2 }
0x7 : { %v14_v6 = vsel /*vm=*/%vm13_vm0, /*on_true_vy=*/%v9_v2, /*on_false_vx=*/%v12_v5 }
0x8 : { %v17_v7 = vsel /*vm=*/%vm15_vm1, /*on_true_vy=*/%v16_v4, /*on_false_vx=*/%v14_v6 }
0x9 : { %18 = vst [vmem:[%s38_s1] sm:$0x1] /*vst_source=*/%v17_v7 }
</code></code></pre><p>The constant <code>0.015625</code> is <code>1/64</code>&#8212;the mean computation. The <code>vrsqrt.f32</code> instruction computes reciprocal square root. Special case handling for infinity and zero ensures numerical correctness.</p><p><strong>fusion.5 (56 bundles)</strong></p><p>This fusion computes <code>reduce_max</code> across rows for the softmax numerically-stable computation. It follows a similar pattern to the sum reduction in multiply_reduce_fusion: XLU transpose followed by tree reduction, but using <code>max</code> operations instead of <code>add</code>.</p><p><strong>fusion.2 (65 bundles)</strong></p><p>Implements <code>exp(x - max) + reduce_sum</code>&#8212;the softmax numerator and denominator. The <code>vpow2</code> instruction computes the exponential followed by another transpose-and-reduce pattern for the sum.</p><p><strong>fusion (48 bundles)</strong></p><p>The final matmul (<code>matmul_2</code>) multiplying the softmax output by the value matrix. This fusion is smaller than multiply_reduce_fusion because it&#8217;s a pure matmul without the additional reduction operations.</p><h3>TLP: The Top Level Program</h3><p>The <strong>TLP (Top Level Program)</strong> orchestrates the entire mini_attention execution. Looking at the correct TLP for our function, we can see how it coordinates all five fusions and manages data movement.</p><p><strong>Program Structure Overview</strong></p><p>The mini_attention TLP is 174 bundles (0x00-0xad) and follows this execution flow:</p><pre><code><code>0x5f-0x65 : Program entry and mode checking
0x66-0x7d : DMA: Load weight matrix from HBM (copy-start)
0x7e-0x82 : Call multiply_reduce_fusion (matmul_1 + squares)
0x83-0x8d : DMA: Start loading value matrix (copy-start.1)
0x8e-0x91 : Call add_sqrt_fusion (sqrt for RMS norm)
0x92-0x96 : Call fusion.5 (reduce_max for softmax)
0x97-0x9b : Call fusion.2 (exp + reduce_sum for softmax)
0x9c-0xa0 : DMA: Wait for value matrix
0xa1-0xa6 : Call fusion (matmul_2)
0xa7-0xad : Program finalization
</code></code></pre><p><strong>DMA Orchestration and Overlapping</strong></p><p>The TLP carefully overlaps memory transfers with computation. First, it loads the weight matrix for the first matmul:</p><pre><code><code>0x77 : { %26 = dma.hbm_to_vmem [thread:$1] /*hbm=*/%s10_s14, 
             /*size_in_granules=*/%s22_s17, /*vmem=*/%s24_s19, 
             /*dst_syncflagno=*/[#allocation8] }
</code></code></pre><p>The DMA uses <code>thread:$1</code>&#8212;the TPU has multiple DMA engines, allowing concurrent transfers. After initiating this transfer, the TLP immediately waits for completion and then calls the first fusion:</p><pre><code><code>0x7a : { %293 = dma.done.wait [#allocation8], 1024 }
0x7b : { %294 = vsyncadd [#allocation8], 4294966272 }
0x7c : { %39 = vsyncpa [#allocation8], 1 }
0x81 : { %45 = inlined_call %s9_s0, %s299_s1, %s300_s2, %s301_s3 
              /* %multiply_reduce_fusion = fusion(%Arg_0.1, %copy-done) */ }
</code></code></pre><p><strong>Overlapping DMA with Compute</strong></p><p>While RMS normalization fusions run, the TLP starts loading the value matrix for the second matmul:</p><pre><code><code>0x8c : { %56 = dma.hbm_to_vmem [thread:$1] /*hbm=*/%s354_s25, 
             /*size_in_granules=*/512, /*vmem=*/%s54_s23, 
             /*dst_syncflagno=*/[#allocation13] }
</code></code></pre><p>This DMA runs in the background while <code>add_sqrt_fusion</code>, <code>fusion.5</code>, and <code>fusion.2</code> execute:</p><pre><code><code>0x90 : { %66 = inlined_call ... /* %add_sqrt_fusion */ }
0x95 : { %71 = inlined_call ... /* %fusion.5 */ }
0x9a : { %76 = inlined_call ... /* %fusion.2 */ }
</code></code></pre><p>Only before the final matmul does the TLP wait for the value matrix:</p><pre><code><code>0x9d : { %295 = dma.done.wait [#allocation13], 512 }
0xa5 : { %87 = inlined_call ... /* %fusion = final matmul */ }
</code></code></pre><p>This is <strong>double-buffering in action</strong>&#8212;the softmax computation hides the latency of loading the value matrix.</p><div class="native-video-embed" data-component-name="VideoPlaceholder" data-attrs="{&quot;mediaUploadId&quot;:&quot;7bbfeeac-68d7-42c9-a38b-02fb3d89a4a4&quot;,&quot;duration&quot;:null}"></div><p>The above video visualizes this DMA and compute overlap on TPU.</p><p><strong>Fusion Calling Convention</strong></p><p>Each fusion call passes VMEM buffer addresses via <code>inlined_call</code>. For example, the final matmul:</p><pre><code><code>0xa2 : { %s312_s0 = smov [#allocation12] /* materialized constant */
        %s355_s5 = sld [smem:[#allocation25_spill]]
        %s356_s5 = int_to_ptr.hbm [resolvable:$false] %s355_s5 }
0xa3 : { %s313_s1 = smov [#allocation16] /* materialized constant */
        %s314_s2 = smov [#allocation15] /* materialized constant */ }
0xa4 : { %s315_s3 = smov [#allocation11] /* materialized constant */
        %s316_s4 = smov [#allocation14] /* materialized constant */ }
0xa5 : { %87 = inlined_call %s312_s0, %s313_s1, %s314_s2, %s315_s3, %s316_s4, %s356_s5
              /* %fusion = fusion(%copy-done.1, %fusion.2, %fusion.5, 
                                   %get-tuple-element.1, %add_sqrt_fusion) */ }
</code></code></pre><p>The fusion receives six operands: the value matrix from HBM (<code>copy-done.1</code>), plus intermediate results from all previous fusions. Note the <code>#allocation25_spill</code>&#8212;this is a register spill to scalar memory, indicating the compiler ran out of scalar registers and had to temporarily store a value.</p><p><strong>Trace Points for Profiling</strong></p><p>Throughout execution, the TLP emits trace markers:</p><pre><code><code>0x78 : { %30 = vtrace 2415919104 }  // After copy-start
0x82 : { %49 = vtrace 2415919106 }  // After multiply_reduce_fusion
0x91 : { %69 = vtrace 2415919108 }  // After add_sqrt_fusion
0x96 : { %74 = vtrace 2415919109 }  // After fusion.5
0x9b : { %79 = vtrace 2415919110 }  // After fusion.2
0xa6 : { %90 = vtrace 2415919112 }  // After final fusion
</code></code></pre><p>These trace IDs let TPU profiling tools measure the time spent in each fusion, helping identify performance bottlenecks.</p><p><strong>Program Finalization</strong></p><p>The TLP concludes with synchronization:</p><pre><code><code>0xa9 : { %95 = vsettm %s317_s26 }        // Set timer mode
0xaa : { %96 = vdelay 1 }                 // Delay for pipeline drain
0xab : { %97 = sfence }                   // Memory fence
0xac : { %s318_s27 = smov 0 }
0xad : { %98 = sst [smem:[#allocation17]] %s318_s27 }  // Signal completion
</code></code></pre><p>The <code>sfence</code> ensures all memory operations have completed before the final store signals to the host that the TPU program has finished.</p><h2>Conclusion/Takeaways</h2><p>Eight lines of JAX code became 250 VLIW bundles across 5 fused kernels. Here&#8217;s what the compiler did to get there:</p><p><strong>The TPU compiler is a sophisticated codegen compiler.</strong> GPUs have codegen too (Triton, XLA/GPU, torch.compile), but the TPU compiler is doing something impressive: it takes your entire computation graph, fuses across operation boundaries, schedules VLIW bundles that saturate 5+ hardware units simultaneously, and orchestrates async DMA &#8212; all automatically. No manual tiling, no explicit shared memory management, no <code>@triton.autotune</code>. You write JAX and the compiler figures out the rest.</p><p><strong>It generalizes to novel workloads.</strong> This matters more than it might seem. On GPUs, peak performance often requires hand-tuned kernels &#8212; FlashAttention exists because the compiler couldn&#8217;t find that optimization automatically. The TPU compiler&#8217;s approach is different: rather than relying on a library of pre-optimized patterns, it reasons about your specific computation from first principles. A weird custom attention variant, a novel normalization scheme, some exotic activation function &#8212; the compiler will fuse them, schedule them, and generate reasonable code without anyone having to write a custom kernel. The performance ceiling might be lower than a hand-tuned GPU kernel, but the floor is much higher.</p><p><strong>Fusion is everything.</strong> The original ~25 HLO ops collapsed into 5 fusions. Each fusion keeps intermediates in VMEM, avoiding HBM round-trips. RMS norm and softmax never touch HBM &#8212; they execute entirely in fast on-chip memory.</p><p><strong>Async DMA hides memory latency.</strong> While the VPU computes softmax, the DMA engine loads the value matrix in the background. By the time softmax finishes, the weights are already in VMEM &#8212; zero stall before the final matmul.</p><p><strong>VLIW packs independent ops.</strong> A single bundle can execute a DMA transfer, three VMEM loads, and two vector ops simultaneously. The compiler statically schedules everything; the hardware just executes.</p><div><hr></div><h3>What does this mean for you?</h3><p>If you&#8217;re training or running inference on TPUs, you don&#8217;t need to understand any of this to get good performance &#8212; <strong>that&#8217;s the point</strong>. The compiler handles fusion, memory management, and scheduling automatically.</p><p>But if you&#8217;re debugging why a particular computation is slow, or curious whether a custom operation will perform well, <em><strong>you now know how to look under the hood</strong></em>. Dump the IR, find your fusion, count the bundles, check if DMA is overlapping with compute. The tools exist; they&#8217;re just undocumented.</p><p>And if you&#8217;re deciding between TPUs and GPUs for a new workload: <strong>TPUs reward experimentation</strong>. You can try unconventional architectures without writing custom kernels. The compiler will generate reasonable code for whatever you throw at it. That&#8217;s a different value proposition than &#8220;fastest possible matmul&#8221; &#8212; it&#8217;s &#8220;fast enough matmul for any shape, any fusion pattern, automatically.&#8221;</p><div><hr></div><p>If you want to explore this yourself, the dump flags are straightforward:</p><pre><code><code>XLA_FLAGS="--xla_dump_to=./hlo --xla_dump_hlo_pass_re=.*"
LIBTPU_INIT_ARGS="--xla_jf_dump_to=./llo --xla_jf_dump_llo_text=true"</code></code></pre><p>The HLO is readable. The LLO takes some squinting, but the patterns emerge &#8212; look for <code>vmatpush</code>/<code>vmatmul</code> pairs for matmuls, <code>vxpose</code>/<code>vpop.trf</code> for transposes, and <code>vrot.slane</code> for reductions.</p><p>Most of this compiler is closed-source, but the IRs tell the story.</p><div><hr></div><p>Questions? Message me on Linkedin: <a href="https://www.linkedin.com/in/patrick-toulme-150b041a5/">https://www.linkedin.com/in/patrick-toulme-150b041a5/</a></p><div class="subscription-widget-wrap-editor" data-attrs="{&quot;url&quot;:&quot;https://patricktoulme.substack.com/subscribe?&quot;,&quot;text&quot;:&quot;Subscribe&quot;,&quot;language&quot;:&quot;en&quot;}" data-component-name="SubscribeWidgetToDOM"><div class="subscription-widget show-subscribe"><div class="preamble"><p class="cta-caption">Thanks for reading Just a Byte - AI Compilers, Silicon, and Systems! Subscribe for free to receive new posts and support my work.</p></div><form class="subscription-widget-subscribe"><input type="email" class="email-input" name="email" placeholder="Type your email&#8230;" tabindex="-1"><input type="submit" class="button primary" value="Subscribe"><div class="fake-input-wrapper"><div class="fake-input"></div><div class="fake-button"></div></div></form></div></div><p></p>]]></content:encoded></item></channel></rss>