Buckets:
what goes in when you do c= a+b
a = torch.randn(1024, 1024, device="cuda")
b = torch.randn(1024, 1024, device="cuda")
c = a + b
diagram
┌─────────────────────────────────────────┐
│ 1️⃣ Python Layer │
│ │
│ c = a + b │
│ ↓ │
│ calls: a.__add__(b) │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 2️⃣ PyTorch Python Binding │
│ │
│ torch.add(a, b) │
│ ↓ │
│ dispatch to ATen backend │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 3️⃣ ATen Dispatcher (C++) │
│ │
│ Checks device type │
│ - CPU? │
│ - CUDA? ← yes │
│ │
│ Select CUDA implementation │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 4️⃣ CUDA Runtime API │
│ │
│ Launch kernel: │
│ add_kernel<<<grid, block>>>() │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 5️⃣ GPU Execution │
│ │
│ Thousands of parallel threads run: │
│ │
│ idx = blockIdx.x * blockDim.x + │
│ threadIdx.x │
│ │
│ c[idx] = a[idx] + b[idx] │
│ │
│ Runs on CUDA cores │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 6️⃣ Result Tensor c │
│ │
│ Stored in GPU memory (VRAM) │
└─────────────────────────────────────────┘
full story
Perfect — let’s build a **clear end-to-end diagram** of what happens when:
```python
c = a + b
And then extend it to FlashAttention-3 (FA3).
I’ll show:
- 🔹 Simple tensor addition
- 🔹 What actually runs on GPU
- 🔹 How FlashAttention-3 differs
- 🔹 A full stack diagram
🧠 PART 1 — c = a + b (CUDA Tensor)
Assume:
a = torch.randn(1024, 1024, device="cuda")
b = torch.randn(1024, 1024, device="cuda")
c = a + b
🔹 FULL END-TO-END FLOW
┌─────────────────────────────────────────┐
│ 1️⃣ Python Layer │
│ │
│ c = a + b │
│ ↓ │
│ calls: a.__add__(b) │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 2️⃣ PyTorch Python Binding │
│ │
│ torch.add(a, b) │
│ ↓ │
│ dispatch to ATen backend │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 3️⃣ ATen Dispatcher (C++) │
│ │
│ Checks device type │
│ - CPU? │
│ - CUDA? ← yes │
│ │
│ Select CUDA implementation │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 4️⃣ CUDA Runtime API │
│ │
│ Launch kernel: │
│ add_kernel<<<grid, block>>>() │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 5️⃣ GPU Execution │
│ │
│ Thousands of parallel threads run: │
│ │
│ idx = blockIdx.x * blockDim.x + │
│ threadIdx.x │
│ │
│ c[idx] = a[idx] + b[idx] │
│ │
│ Runs on CUDA cores │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 6️⃣ Result Tensor c │
│ │
│ Stored in GPU memory (VRAM) │
└─────────────────────────────────────────┘
⚡ Important Insight
Python does zero math.
It only:
- Triggers operator
- Routes to C++
- Launches CUDA kernel
Actual addition happens in:
GPU assembly instructions running on CUDA cores.
🧠 PART 2 — What Actually Runs on GPU?
After compilation, the CUDA kernel becomes:
Python
↓
C++ CUDA code
↓
PTX (intermediate)
↓
SASS (real GPU assembly)
↓
CUDA cores execute instructions
Each GPU thread does:
LOAD a[i]
LOAD b[i]
ADD
STORE c[i]
Thousands of threads in parallel.
🚀 PART 3 — Now Add FlashAttention-3
Now let’s compare with:
output = flash_attn(q, k, v)
🔹 Naive Attention (Multiple Kernels)
Without FlashAttention:
QKᵀ → Kernel 1
Softmax → Kernel 2
Multiply V → Kernel 3
🚨 Problems:
- Writes big attention matrix to VRAM
- High memory bandwidth
- Multiple kernel launches
🔥 FlashAttention-3 Flow
FlashAttention-3 fuses everything into ONE kernel.
🔹 FA3 End-to-End Diagram
┌─────────────────────────────────────────┐
│ 1️⃣ Python Call │
│ │
│ flash_attn(q, k, v) │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 2️⃣ PyTorch C++ Extension Wrapper │
│ │
│ Calls custom CUDA op │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 3️⃣ Custom FA3 CUDA Kernel │
│ │
│ Single fused kernel: │
│ │
│ ┌──────────────────────────────┐ │
│ │ Load Q tile into shared mem │ │
│ │ Load K tile │ │
│ │ Compute partial QKᵀ │ │
│ │ Online softmax (stable) │ │
│ │ Multiply with V │ │
│ │ Write output │ │
│ └──────────────────────────────┘ │
│ │
│ Uses: │
│ - Shared memory │
│ - Warp shuffles │
│ - Tensor cores (HMMA instructions) │
└─────────────────────────────────────────┘
↓
┌─────────────────────────────────────────┐
│ 4️⃣ GPU Tensor Cores │
│ │
│ Matrix multiply instructions: │
│ HMMA.1688.F32 │
└─────────────────────────────────────────┘
🧠 Key Difference
a + b
- Simple elementwise kernel
- Memory bound
- Each thread independent
FlashAttention-3
- Fused mega-kernel
- Compute + memory optimized
- Tiled execution
- IO-aware
- Tensor core accelerated
- Warp-level parallelism
📊 Final Combined Stack Diagram
USER CODE
│
▼
Python Frontend
│
▼
PyTorch API Layer
│
▼
ATen C++ Dispatcher
│
┌──────────────┴──────────────┐
▼ ▼
Standard CUDA Op Custom CUDA Op
(torch.add) (FlashAttention-3)
│ │
▼ ▼
Simple CUDA Kernel Fused Attention Kernel
│ │
▼ ▼
CUDA Cores Tensor Cores + SM
│ │
▼ ▼
GPU Memory (VRAM) Shared Mem + Registers
🧠 Mental Model
Think of Python as:
🧑💼 Project manager
PyTorch C++ as:
📞 Operations manager
CUDA kernel as:
🏭 Factory machinery
GPU cores as:
⚙️ Workers on assembly lines
FlashAttention-3 is:
🚀 A hyper-optimized factory where all steps happen in one room without moving materials around.
Xet Storage Details
- Size:
- 12.7 kB
- Xet hash:
- 217101b4eae27a5d3b7a20ea1d5eac0493cafa09901ceef161e9f1df7ceed2f8
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.