Skip to main content

Trying to Speed Up a Weird Neural Net Idea on My Mac's GP

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:

  1. Normal causal self-attention
  2. A regular shared MLP that mixes information across neurons
  3. 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.

AttemptWhat I didHow it went
1Batched the per-neuron math with matrix multiply, wrote my own backward passGradients were wrong — I'd used the wrong formula for the GELU derivative
2Fixed the GELU derivativeForward pass looked right, but backward was still broken (a silly variable naming bug)
3Just let torch.compile handle itThis actually worked well, and was the easiest option
4Wrote a custom Metal shader by handCompiled, but was actually slower than doing nothing special, plus some accuracy issues
5Went back and fixed the bugs from attempt 1-2 properlyFinally correct, but not especially fast
6More Metal shader debuggingHit some crashes and got weird all-zero outputs a couple times before sorting out the dispatch pattern
7Used PyTorch's built-in helper for compiling Metal shaders instead of doing it by handMuch 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.

CheckResult
Forward output matches the reference implementationYes, difference around 0.000006
No broken (NaN) gradientsYes, checked all 13 sets of gradients
Attention still worksYes
Text generation runs without crashingYes
Optimizer updates every parameterYes
Speed (small model)~25ms per step

Rough speed comparison

VersionSmall modelFull model (rough estimate)
Original (grouped conv)—~4,300ms
torch.compile0.7ms~220ms
Hand-written backward40msnot tested
Metal shader + hand-written backward25msnot 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

Popular posts from this blog

Outside the X bubble of software developers and tokenmaxers, how are non-technical professionals actually using AI in production to run real, revenue-generating businesses?

If you scroll tech Twitter, the discourse is dominated by "tokenmaxing" — people flexing multi-million token usage, chaining complex autonomous agent loops, and obsessing over API token throughput. But far away from that hype, how does AI function in a real enterprise setting where revenue, accuracy, and client relationships are on the line? I recently had a deep conversation with a medical professional who offers a grounded perspective missing from mainstream tech discussions. He is a co-founder of a medical communications agency based in Jeddah, working face-to-face with major pharmaceutical companies and healthcare professionals. His agency bridges the gap between medical science and advertising — building research-backed, highly regulated campaigns. Having worked in the industry for over a decade — spanning 9-to-5 corporate roles, freelancing, and now running an agency — his experience provides a clear benchmark for where AI provides genuine leverage, where it expos...

Shovels, Chatbots, and the AI Bubble Nobody Else Is In

I talked to a pharma executive last week. Big company, UK based, he runs things from what I think is the Middle East control room in Jeddah. Senior guy, more than a decade in the industry. I asked him the obvious question, how is he using AI day to day. He does not, not really. He uses Copilot to proofread emails and juggle the occasional idea, because Copilot is the only thing his work laptop will let him touch. He told me he thought about running a second laptop off the company firewall so he could actually explore what is out there, but keeping two work laptops was too much of a hassle. So he stayed on one laptop, inside the firewall, and stayed a chatbot user. That is not a story about one lazy executive. That is a story about an entire class of sectors, heavy IP, heavy security, heavy regulation, where the friction of adopting anything beyond a sanctioned chatbot is high enough that the frontier simply does not reach them yet. Pharma, medicine, anywhere production touches actu...