You’re at roughly 65 TFLOPS of useful compute (about 6 × params FLOPs per token, plus the recompute from gradient checkpointing), which leaves headroom: the LLMQ paper reports around 51% MFU on consumer GPUs with an optimized stack. In order of expected payoff:
torch.compile on the model. It fuses RMSNorm, RoPE and SwiGLU into fewer kernels and was worth about 1.5x in published single-GPU training benchmarks./mnt/c, pre-tokenize and pack sequences so there’s no padding, and use a fused optimizer.
Run the PyTorch profiler for a few steps first to see whether you’re bound by memory-bound ops, recompute or data , so you apply the right fix.
Also worth knowing: the big matmuls are compute-bound, but everything between them (norms, RoPE, SwiGLU, softmax, logits) is memory-bandwidth-bound, and on a 16 GB card with ~960 GB/s that’s where a lot of your step time goes. That’s why kernel fusion (compile, Liger) gives such large gains, and why FP8 helps beyond just faster matmuls.