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: | |
| 1. 🔹 Simple tensor addition | |
| 2. 🔹 What actually runs on GPU | |
| 3. 🔹 How FlashAttention-3 differs | |
| 4. 🔹 A full stack diagram | |
| --- | |
| # 🧠 PART 1 — `c = a + b` (CUDA Tensor) | |
| Assume: | |
| ```python | |
| 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: | |
| ```python | |
| 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.