I'd be interested in contributing to the course, if you need anything specific.
Here's a tensorboard URL that will probably stop working within a few days. http://bulma.tensorfork.com:31337/#profile
You can view the memory profiler by using the dropdown menus on the left. Here's a particularly chonky CrossReplicaSum: https://i.imgur.com/CcdJzLj.png
The game here is to keep that number at the top -- peak memory usage -- below 15GB. In practice, TPUv3-8's run out of memory at around 14.5GB, which immediately crashes (and hence you can't profile it). So we're always trying to get as close to 15GB as possible.
The first thing you immediately notice is that real-life training runs are very spiky. Different parts of the pipeline end up allocating wildly different amounts of memory. There's almost no such thing as a constant memory usage pipeline (which I was dismayed to discover).
In this profiling run, you can see that there's a big ass-spike at ~4000 on the X axis. The green bar marks the lifetime of the operation causing the highest peak memory usage. Different operations depend on each other, forming a chain of allocations. a + b takes 'a' and 'b' as inputs, and any temporary tensor reachable by either 'a' or 'b' cannot be freed until a + b is finished executing. Ditto for all other operations.
So you see, it's easy to accidentally build a "tower" of allocations, rather than a flat line. Thus, your total model parameter count is severely limited compared to what it could be, since in this situation the only way to reduce memory usage (without rewriting the code) is to scale down the model params.
Hovering over the big ass-orange allocation, we see that the shape is -- gosh, tensorboard is infuriating sometimes. I tried to copy-paste the shape, but whenever I move the mouse off of the allocation, the info on the left vanishes. Anyway, the shape is F32[32,2048,1,12608][1,3,0,2]. It means the cross replica sum is happening across TPU cores 1, 3, 0, and 2; it's a float32 sum; the batch size is 32; the hidden dimension is 2048, and the vocab dimension is 12,608. Since it's across four cores, multiply that dim by 4, and the total vocab size is 50,432, which is exactly right for a GPT model (https://nv-adlr.github.io/MegatronLM has details).
So right away, we can see that (a) the non-peak memory usage is around 4GB or so, and (b) the peak mem usage of the spike is around 12GB. That means if we eliminate the spike, we can scale up our model by more than 3x, if usage scales linearly. (Sometimes you get lucky and it's linear, other times something is superlinear. It's more or less linear in my experience.)
So how do we eliminate the spike? Heck if I know how the Google pros do it, but my way of doing it is to unstack along the batch dimension and perform each operation sequentially.
In other words, the total memory usage here is O(32 * 2048 * 12,608) which is quite hefty. By unstacking along the batch dimension, you get 32 tensors, each of size 2048 by 12,608. Therefore, if you do each operation sequentially, the temporary buffer is now O(2048 * 12,608), giving us a 32x savings.
Is this slower? Surprisingly, more often than not, it's as fast or faster. The reason is subtle: slowdowns occur due to memory bandwidth and network bandwidth. As long as the unstack is strictly a memory bandwidth effect, then it's just as fast, because you're trading CPU cycles for memory -- and you have tons of CPU cycles here, since it's a TPU core. (The TPU core utilization in our experience is always around 30%, and we've never seen it higher than 65%.) So you should always, always make this trade whenever possible.
Network traffic is trickier. This is a cross-replica sum, which means it's sending the tensors across the network to each TPU core. The TPU cores are connected via a high speed interconnect nexus thingie, but like all bottlenecks, this one has a limit. It's a very high limit, but it's not endless.
The only solution I've found is to think of an idea and then test that idea. Reasoning from first principles almost never works for me. I've seen others solve problems by reasoning from first principles, so it's possible that I'm simply stupid. But I find it's much more effective to try as many ideas as possible, as quickly as possible. You often end up surprised.
I'll type more stuff later if I feel like it, or you can ask more followup questions. Feel free to DM me on twitter if you'd like to chat in realtime sometime.