I've been playing around with a transformer variant where, instead of one shared feed-forward layer, every single neuron gets its own tiny private network. It's a fun idea, but the first version I wrote ran painfully slow on my Mac, and this post is basically my notes on figuring out why, and what I tried to fix it.
The Setup
The model has 3 layers and around 2,574 "neurons" at full size (I mostly tested a shrunk-down version with way fewer, just to iterate faster). Each neuron has its own small 2-layer MLP with its own weights — nothing is shared between neurons at this stage. Each layer does roughly:
- Normal causal self-attention
- A regular shared MLP that mixes information across neurons
- The weird part: each of the ~2,574 neurons runs through its own tiny 2-layer network, completely independently
for each neuron n:
h = GELU( x[n] * in_w[n] + in_b[n] )
h = GELU( LN( h @ hw0[n] + hb0[n] ) * lw0[n] + lb0[n] ) + h
h = GELU( LN( h @ hw1[n] + hb1[n] ) * lw1[n] + lb1[n] ) + h
out[n] = h · out_w[n] + out_b[n]
Each of those tiny per-neuron networks works with a hidden size of 61. Multiply that out across 2,574 neurons and you get almost 2 million parameters just for this one part of the model, per layer.
Why It Was So Slow
My first working version used a grouped convolution (basically a trick to run 2,574 separate small networks without a Python for-loop). It gave the right answer, but it was really slow — like 4+ seconds per training step at full size, which is unusable.
My best guess is that the GPU wasn't the bottleneck — it was mostly sitting idle. What was probably happening is that this operation gets broken into thousands of tiny separate GPU calls instead of one big one, and each tiny call has some fixed overhead. So you end up paying that overhead thousands of times per step instead of once.
Basically: keep the math exactly the same, but find a way to do it in way fewer GPU calls.
What I Tried (Not All of It Worked)
I went through a handful of attempts. A couple of them had straight-up wrong gradients that I didn't catch until later, which was annoying but a good reminder to actually check the math instead of just checking that the code runs.
| Attempt | What I did | How it went |
|---|---|---|
| 1 | Batched the per-neuron math with matrix multiply, wrote my own backward pass | Gradients were wrong — I'd used the wrong formula for the GELU derivative |
| 2 | Fixed the GELU derivative | Forward pass looked right, but backward was still broken (a silly variable naming bug) |
| 3 | Just let torch.compile handle it | This actually worked well, and was the easiest option |
| 4 | Wrote a custom Metal shader by hand | Compiled, but was actually slower than doing nothing special, plus some accuracy issues |
| 5 | Went back and fixed the bugs from attempt 1-2 properly | Finally correct, but not especially fast |
| 6 | More Metal shader debugging | Hit some crashes and got weird all-zero outputs a couple times before sorting out the dispatch pattern |
| 7 | Used PyTorch's built-in helper for compiling Metal shaders instead of doing it by hand | Much less painful, and gave good results |
Attempt That Worked Best: Just Use torch.compile
The simplest fix, honestly, was realizing that a batched matrix multiply already does the whole per-neuron computation in one GPU call instead of thousands. Wrapping that in torch.compile let PyTorch fuse the surrounding steps (normalization, activation, etc.) together too, which cut things down even more.
@torch.compile(mode="max-autotune")
def internal_stage(x, in_w, in_b, hw0, hb0, lw0, lb0, hw1, hb1, lw1, lb1, out_w, out_b):
h0 = F.gelu(x.unsqueeze(-1) * in_w.unsqueeze(1) + in_b.unsqueeze(1))
pl0 = torch.bmm(h0, hw0) + hb0.unsqueeze(1)
var0, mean0 = torch.var_mean(pl0, dim=-1, keepdim=True)
n0 = (pl0 - mean0) * torch.rsqrt(var0 + 1e-5)
g0 = n0 * lw0.unsqueeze(1) + lb0.unsqueeze(1)
h1 = F.gelu(g0) + h0
# layer 2 is the same pattern
return (h2 * out_w.unsqueeze(1)).sum(dim=-1) + out_b.unsqueeze(1)
On the small test version this dropped from a noticeable delay to under a millisecond. On a single full-size run I measured around 220ms, down from over 4 seconds — a solid improvement, though I only ran it a few times and the numbers bounced around a fair bit run to run.
Attempt: Writing My Own Backward Pass
The compiled version above still had some jitter in timing, so I tried writing the backward pass by hand instead of letting PyTorch's autograd figure it out. The idea is that autograd normally has to track what happened separately for every one of the 2,574 neurons, which adds overhead. Doing the math myself skips that.
This meant getting the actual derivative of GELU right, not just an approximation:
def gelu_backward(x, grad_out):
cdf = 0.5 * (1.0 + torch.erf(x / 1.4142))
pdf = 0.3989 * torch.exp(-0.5 * x * x)
return grad_out * (cdf + x * pdf)
and the LayerNorm derivative:
def layernorm_backward(grad_output, x, mean, inv_std):
x_hat = (x - mean) * inv_std
g = grad_output
mean_g = g.mean(dim=-1, keepdim=True)
mean_gx = (g * x_hat).mean(dim=-1, keepdim=True)
return inv_std * (g - mean_g - x_hat * mean_gx)
One thing I got wrong initially: I was saving way too many intermediate tensors for backward — at full size that was something like 4-5GB, which just crashed on my 16GB machine. Saving only the essentials and recomputing the rest during backward fixed that.
Attempt: Writing a Metal Shader Myself
I also tried going lower-level and writing an actual Metal compute shader, thinking it might be faster than anything PyTorch could generate automatically.
The first version used a separate compiled library called from Python, which technically worked but was a pain — I ran into random crashes that seemed related to how memory buffers were being managed between Python and the native code, and never fully figured out the root cause. Debugging GPU code blind, without good error messages, is rough.
One annoying detail: Metal doesn't have a built-in erf() function, which GELU needs, so I had to use a polynomial approximation instead:
float gelu(float x) {
float a = fabs(x);
float t = 1.0 / (1.0 + 0.2316419 * a);
float b = ((((1.330274429*t - 1.821255978)*t + 1.781477937)*t
- 0.356563782)*t + 0.319381530)*t;
float cdf = 1.0 - 0.3989422804 * b * exp(-0.5 * a * a);
if (x < 0.0) cdf = 1.0 - cdf;
return x * cdf;
}
The approximation error is tiny (way below anything that would matter for training), but it does mean this version won't give bit-for-bit identical results to plain PyTorch, which is worth knowing if you care about exact reproducibility.
Eventually I found that PyTorch has a built-in way to compile Metal shaders directly, which handles all the memory management stuff automatically. Switching to that got rid of basically all the crashes and was just a lot less stressful to work with.
Where I Landed
The current version combines a custom Metal shader for the forward pass with the hand-written backward pass described above, plus the regular attention and shared MLP parts. It also has basic text generation (temperature + top-k sampling) so I can actually see the model produce something.
How Much I Actually Tested This
I only ran the full validation on the small test-size model, not the full 81-million-parameter one — I didn't want to risk breaking anything on the version that was actually training. So take the full-scale numbers with a grain of salt; they're extrapolated, not measured directly.
| Check | Result |
|---|---|
| Forward output matches the reference implementation | Yes, difference around 0.000006 |
| No broken (NaN) gradients | Yes, checked all 13 sets of gradients |
| Attention still works | Yes |
| Text generation runs without crashing | Yes |
| Optimizer updates every parameter | Yes |
| Speed (small model) | ~25ms per step |
Rough speed comparison
| Version | Small model | Full model (rough estimate) |
|---|---|---|
| Original (grouped conv) | — | ~4,300ms |
| torch.compile | 0.7ms | ~220ms |
| Hand-written backward | 40ms | not tested |
| Metal shader + hand-written backward | 25ms | not tested |
So the best confirmed full-scale improvement (roughly 19x faster) actually comes from the simplest fix — just letting torch.compile do its thing. The Metal shader version felt more stable and consistent in my small-scale tests, but I never got to properly benchmark it at full size, so I can't really say it's better overall yet.
What I'd Actually Take Away From This
Things that seemed to genuinely help
- The big one: batching the per-neuron math into one matrix multiply instead of thousands of tiny separate calls. Everything else is a smaller improvement on top of this.
- torch.compile is genuinely really good at fusing operations together automatically — got a big speedup for very little extra work.
- Writing the backward pass by hand can help, but only if you get the math exactly right, which took me a couple of tries.
- Recomputing values during backward instead of saving everything in memory is a good trick when you're memory constrained (16GB machines fill up fast).
- If you're going to write custom GPU shaders, use whatever built-in helper your framework provides instead of managing native memory by hand — it's just less error-prone.
Stuff I'm still unsure about, or didn't finish
- The backward pass still isn't a single fused Metal operation — it's still a bunch of regular PyTorch operations, just with the math figured out ahead of time instead of relying on autograd.
- The Metal version has a tiny numerical difference from regular PyTorch because of the GELU approximation. Small, but real.
- The shader assumes the per-neuron hidden size stays under 128 — bigger than that and I'd need to rewrite part of it.
- None of this was tested at full scale yet, only on a shrunk-down version. The full-scale numbers could be optimistic.
- This whole approach is specific to this particular "every neuron gets its own network" idea. For a normal transformer, this probably wouldn't help much, since normal transformers don't have this many-tiny-independent-operations problem in the first place.
- I caught more than one wrong-gradient bug that slipped through an earlier check — so honestly, I'd want to double check this backward pass again before trusting it for anything important.
Comments
Post a Comment