File size: 41,698 Bytes
5dc80b3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065 1066 1067 1068 1069 1070 1071 1072 1073 1074 1075 1076 1077 1078 1079 1080 1081 1082 1083 1084 1085 1086 1087 1088 1089 1090 1091 1092 1093 1094 1095 1096 1097 1098 1099 1100 1101 1102 1103 1104 1105 1106 1107 1108 1109 1110 1111 1112 1113 1114 1115 1116 1117 1118 1119 1120 1121 1122 1123 1124 1125 1126 1127 1128 1129 1130 1131 1132 1133 1134 1135 1136 1137 1138 1139 1140 1141 1142 1143 1144 1145 1146 1147 1148 1149 1150 1151 1152 1153 1154 1155 1156 1157 1158 1159 1160 1161 1162 1163 1164 1165 1166 1167 1168 1169 1170 1171 1172 1173 1174 1175 1176 1177 1178 1179 1180 1181 1182 1183 1184 1185 1186 1187 1188 1189 1190 1191 1192 1193 1194 | # End-to-End Explanation: PyTorch & Triton Implementation
## For someone with basic Python knowledge
---
## PART 1: What Are the Available Benchmarks?
Before diving into code, let's understand what benchmarks are available to measure performance.
### Overview of Benchmarking Suite
The benchmarking system measures how fast the model runs and how efficiently it uses memory. It compares two setups:
- **Tiered Model**: Uses the SRAM/DRAM memory hierarchy (what we want to study)
- **Baseline Model**: Standard model without memory tiering
### Files Involved: `benchmark.py` and `run_benchmark.py`
#### A. `run_benchmark.py` - The Command-Line Interface
This file lets you **run benchmarks from the terminal**. Think of it as a control panel.
**Key command-line options:**
```bash
# See available benchmarks and options
python run_benchmark.py --help
# Compare tiered vs baseline models
python run_benchmark.py --mode compare --batch-sizes 1,8,32 --seq-lens 64,128
# Benchmark only the tiered (memory-aware) model
python run_benchmark.py --mode tiered --warmup 5 --iterations 50 --output results.json
# Quick test (for learning)
python run_benchmark.py --mode tiered --warmup 1 --iterations 3 --batch-sizes 2 --seq-lens 16
```
**What each option means:**
| Option | Meaning |
|--------|---------|
| `--mode` | What to benchmark: `tiered` (new model), `baseline` (standard), or `compare` (both) |
| `--batch-sizes` | How many samples to process at once (comma-separated: `1,8,32` means test with 1, 8, and 32 samples) |
| `--seq-lens` | Length of input sequences: `64,128` means test with sequences of 64 and 128 tokens |
| `--hidden-size` | Internal dimension of the model (default: 512) |
| `--H-cycles` | Number of times H-level (slow) processes data (default: 2) |
| `--L-cycles` | Number of times L-level (fast) processes data (default: 2) |
| `--H-layers` | Number of H-level transformer layers (default: 4) |
| `--L-layers` | Number of L-level transformer layers (default: 4) |
| `--warmup` | How many runs to discard before measuring (to let GPU settle) |
| `--iterations` | How many actual measurements to take |
| `--output` | Save results to JSON file |
---
### B. What Does `benchmark.py` Actually Measure?
Inside `benchmark.py`, the `BenchmarkResult` class records these metrics:
#### **Latency Metrics** (How fast things run, in microseconds `ΞΌs`)
- **L-level latency**: Time for the "fast" tier to process data
- **H-level latency**: Time for the "slow" tier to process data
- **Total inference latency**: End-to-end prediction time (in milliseconds `ms`)
- **h_over_l_latency_ratio**: How many times slower H-level is than L-level
**Why this matters:** If H-level takes 10Γ longer than L-level, we know the memory hierarchy is working β slow operations vs. fast ones.
#### **Memory Metrics** (How much GPU memory used)
- **sram_peak_mb**: Peak memory in the "fast" tier
- **dram_peak_mb**: Peak memory in the "slow" tier
- **total_gpu_memory_mb**: Total GPU memory used
#### **Transfer Metrics** (Cost of moving data between tiers, in microseconds)
- **h_l_transfer_mean_us**: Average time to copy data from HβL
- **l_h_transfer_mean_us**: Average time to copy data from LβH
**Why this matters:** If transfers are expensive, the tiering strategy might not be worth it.
#### **Hit Rate & Efficiency**
- **sram_hit_rate**: How often data we tried to put in "fast" memory actually fit (0.0 = never, 1.0 = always)
- **memory_efficiency**: Useful compute time / total time (higher is better)
**Why this matters:** If hit_rate is low, data keeps spilling to slow memory and we're not getting the optimization benefit.
#### **Triton Kernel Probes** (Direct measurement of memory latency)
- **triton_sram_probe_latency_us**: Measured latency of "fast" memory via Triton kernels
- **triton_dram_probe_latency_us**: Measured latency of "slow" memory via Triton kernels
**Why this matters:** These are the "ground truth" measurements showing the actual speed difference between SRAM and DRAM.
---
### C. How Benchmarks Are Run
**Flow of a benchmark:**
1. **Create dummy data** (synthetic input tensors)
2. **Warmup phase** (run N times to let GPU caches warm up)
- These runs are NOT counted in final results
3. **Timing phase** (run N times with GPU timers active)
- Start GPU timer
- Run model inference
- Stop GPU timer
- Record elapsed time
4. **Collect statistics**: min, max, mean, std dev from all the timed runs
5. **Measure memory** (peak usage during runs)
6. **Calculate derived metrics** (ratios, efficiency scores)
7. **Output results** (to console and/or JSON file)
---
## PART 2: PyTorch Implementation β Line-by-Line
### File 1: `memory_tier.py` β The Memory Manager
This is the **PyTorch foundation** for the two-tier memory system. It mimics how GPU memory is organized (SRAM = fast cache vs. DRAM = main memory).
#### **Header & Imports (Lines 1-22)**
```python
"""
Memory Tier Manager for HRM SRAM/DRAM implementation.
Manages the placement of H-level and L-level hidden states across
GPU memory tiers and tracks all memory operations for benchmarking.
- SRAM tier: Uses CUDA pinned memory + explicit prefetching.
L-level states are kept GPU-resident with minimal transfers.
- DRAM tier: Standard GPU global memory with transfer tracking.
H-level states go through normal allocation paths.
"""
import time
from typing import Dict, List, Optional, Tuple
from dataclasses import dataclass, field
from contextlib import contextmanager
import torch
```
**Explanation:**
- The docstring explains the purpose: manage two memory tiers
- `dataclass` and `field` are used to create structured data containers
- `contextmanager` creates a context (like `with` statements in Python)
- `torch` is imported to use PyTorch tensor operations
#### **Section 1: MemoryEvent β Tracking Individual Operations (Lines 26-31)**
```python
@dataclass
class MemoryEvent:
"""A single tracked memory operation."""
tier: str # 'sram' or 'dram'
operation: str # 'alloc', 'load', 'store', 'transfer'
bytes: int
duration_us: float # microseconds
timestamp: float
```
**Explanation:**
- A `@dataclass` is like a **lightweight container** for data
- Each `MemoryEvent` records ONE memory operation (e.g., "allocated 1024 bytes in SRAM in 50 microseconds")
- Fields:
- `tier`: Which tier (fast or slow)?
- `operation`: What happened? (allocate new memory, load, store, transfer between tiers)
- `bytes`: How much data?
- `duration_us`: How long did it take? (in microseconds, where 1000 ΞΌs = 1 ms)
- `timestamp`: When did it happen?
**Real example:**
```python
event = MemoryEvent(
tier='sram',
operation='alloc',
bytes=65536,
duration_us=123.45,
timestamp=1704067200.123
)
# Created a record: "SRAM allocation of 65536 bytes took 123.45 microseconds"
```
#### **Section 2: TierStats β Accumulated Statistics (Lines 35-56)**
```python
@dataclass
class TierStats:
"""Accumulated statistics for one memory tier."""
total_alloc_bytes: int = 0
peak_alloc_bytes: int = 0
current_alloc_bytes: int = 0
num_loads: int = 0
num_stores: int = 0
num_transfers: int = 0
total_load_us: float = 0.0
total_store_us: float = 0.0
total_transfer_us: float = 0.0
hit_count: int = 0
miss_count: int = 0
@property
def hit_rate(self) -> float:
total = self.hit_count + self.miss_count
return self.hit_count / total if total > 0 else 0.0
@property
def avg_load_us(self) -> float:
return self.total_load_us / self.num_loads if self.num_loads > 0 else 0.0
@property
def avg_store_us(self) -> float:
return self.total_store_us / self.num_stores if self.num_stores > 0 else 0.0
```
**Explanation:**
- This accumulates **total statistics** for a tier (all operations combined)
- `total_alloc_bytes`: Total memory ever allocated
- `peak_alloc_bytes`: Maximum memory used at any point
- `current_alloc_bytes`: Memory currently in use
- `num_loads`, `num_stores`: Counters for how many times data was read/written
- `hit_count` / `miss_count`: Success/failure for fitting data in SRAM
- **Properties** (`@property`): Computed on-the-fly from raw counts
- `hit_rate = hit_count / (hit_count + miss_count)` β percentage of successful SRAM fits
- `avg_load_us = total_load_us / num_loads` β average speed per load operation
- `avg_store_us = total_store_us / num_stores` β average speed per store operation
**Why properties are useful:**
Instead of storing `hit_rate` separately and having to update it constantly, we compute it whenever asked: `stats.hit_rate` automatically calculates the current rate.
---
#### **Section 3: MemoryTierManager β The Main Manager Class (Lines 60-90)**
```python
class MemoryTierManager:
"""Coordinates SRAM/DRAM memory placement and tracking for HRM.
In the Triton context:
- SRAM tier: Tensors allocated with `pin_memory` and kept on the
same CUDA stream as L-level computation. Triton kernels keep
these values in registers/shared memory via data reuse.
- DRAM tier: Standard `torch.cuda` tensors. Triton kernels load
these from global memory each time.
"""
def __init__(
self,
device: torch.device,
enable_tracking: bool = True,
sram_capacity_mb: float = 48.0, # Typical L2 cache size
):
self.device = device
self.enable_tracking = enable_tracking
self.sram_capacity_bytes = int(sram_capacity_mb * 1024 * 1024)
# State registries
self._sram_tensors: Dict[str, torch.Tensor] = {}
self._dram_tensors: Dict[str, torch.Tensor] = {}
# Event log
self._events: List[MemoryEvent] = []
self._sram_stats = TierStats()
self._dram_stats = TierStats()
# CUDA events for GPU timing
self._use_cuda = device.type == 'cuda'
if self._use_cuda:
self._sram_stream = torch.cuda.Stream(device=device)
self._dram_stream = torch.cuda.Stream(device=device)
else:
self._sram_stream = None
self._dram_stream = None
```
**Explanation:**
- **`__init__` method**: Initializes the manager when created
- `device`: Where to allocate (CPU or GPU?)
- `enable_tracking`: Should we record all operations?
- `sram_capacity_mb`: How much "fast" memory is available? (Typical GPU L2 cache = 48 MB)
- **State registries:**
- `_sram_tensors`: Dictionary storing all tensors in SRAM (key = name, value = tensor)
- `_dram_tensors`: Dictionary storing all tensors in DRAM
- Example: `_sram_tensors['layer1_hidden'] = torch.tensor(...)`
- **Event log:**
- `_events`: A list of `MemoryEvent` objects (every operation is recorded)
- `_sram_stats`, `_dram_stats`: Running statistics for each tier
- **CUDA streams:**
- A stream is like a **"lane"** for GPU operations (operations in same lane execute sequentially, different lanes can run in parallel)
- Two separate streams (`_sram_stream`, `_dram_stream`) let us potentially run fast and slow operations concurrently
---
#### **Section 4A: SRAM Allocation (Lines 95-123)**
```python
def alloc_sram(self, name: str, shape: Tuple, dtype: torch.dtype) -> torch.Tensor:
"""Allocate a tensor in the SRAM tier (GPU-resident, pinned)."""
# Calculate size in bytes
nbytes = torch.tensor([], dtype=dtype).element_size()
for s in shape:
nbytes *= s
# Check capacity: does it fit?
if self._sram_stats.current_alloc_bytes + nbytes > self.sram_capacity_bytes:
# Doesn't fit β spill to DRAM and record a "miss"
self._sram_stats.miss_count += 1
return self.alloc_dram(name, shape, dtype)
# It fits! Record a "hit"
self._sram_stats.hit_count += 1
# Time the allocation
t0 = self._timer_start()
tensor = torch.zeros(shape, dtype=dtype, device=self.device)
# Hint: keep GPU-resident (don't swap to CPU)
if self._use_cuda:
with torch.cuda.stream(self._sram_stream):
tensor = tensor.contiguous()
self._sram_tensors[name] = tensor
dur = self._timer_end(t0)
# Update statistics
self._sram_stats.total_alloc_bytes += nbytes
self._sram_stats.current_alloc_bytes += nbytes
self._sram_stats.peak_alloc_bytes = max(
self._sram_stats.peak_alloc_bytes,
self._sram_stats.current_alloc_bytes,
)
# Record the operation
self._record_event('sram', 'alloc', nbytes, dur)
return tensor
```
**Explanation (line by line):**
1. **Calculate size:** Convert shape into bytes
- Example: shape `(128, 512)` with `float32` (4 bytes) = 128 Γ 512 Γ 4 = 262,144 bytes
2. **Capacity check:** Does this allocation fit in SRAM?
- If `current_alloc_bytes + nbytes > sram_capacity_bytes`, it doesn't fit
- **Fall back to DRAM** and **count as a miss** (failed to use SRAM)
- Otherwise, **count as a hit** (successfully allocated in SRAM)
3. **Timing:** Record when allocation starts (`t0`)
4. **Create tensor:** `torch.zeros(shape, ...)` creates a tensor of zeros
- `shape`: dimensions
- `dtype`: data type (e.g., float32)
- `device`: 'cuda' or 'cpu'
5. **GPU optimization:** `tensor.contiguous()` makes the data **contiguous in memory** (important for GPU performance)
- Done on the SRAM stream to suggest GPU placement
6. **Store in registry:** Save for later: `_sram_tensors[name] = tensor`
7. **Update stats:**
- Add `nbytes` to total and current
- Update peak if we're now using more than before
8. **Record event:** Add this operation to the event log
---
#### **Section 4B: DRAM Allocation (Lines 126-142)**
```python
def alloc_dram(self, name: str, shape: Tuple, dtype: torch.dtype) -> torch.Tensor:
"""Allocate a tensor in the DRAM tier (standard GPU memory)."""
nbytes = torch.tensor([], dtype=dtype).element_size()
for s in shape:
nbytes *= s
t0 = self._timer_start()
tensor = torch.zeros(shape, dtype=dtype, device=self.device)
self._dram_tensors[name] = tensor
dur = self._timer_end(t0)
self._dram_stats.total_alloc_bytes += nbytes
self._dram_stats.current_alloc_bytes += nbytes
self._dram_stats.peak_alloc_bytes = max(
self._dram_stats.peak_alloc_bytes,
self._dram_stats.current_alloc_bytes,
)
self._record_event('dram', 'alloc', nbytes, dur)
return tensor
```
**Explanation:**
- **Almost identical to `alloc_sram`**, but:
- **No capacity check** (unrestricted DRAM)
- **No stream optimization** (standard allocation)
- **Stores in `_dram_tensors`** instead of `_sram_tensors`
---
#### **Section 5: Cross-Tier Transfers (Lines 147-177)**
```python
def transfer_sram_to_dram(self, name: str) -> torch.Tensor:
"""Copy a tensor from SRAM tier to DRAM tier."""
src = self._sram_tensors[name]
t0 = self._timer_start()
dst = src.clone()
if self._use_cuda:
torch.cuda.synchronize(self.device)
dur = self._timer_end(t0)
self._dram_tensors[name + '_from_sram'] = dst
self._sram_stats.num_transfers += 1
self._sram_stats.total_transfer_us += dur
self._record_event('sram', 'transfer', src.nelement() * src.element_size(), dur)
return dst
def transfer_dram_to_sram(self, name: str) -> torch.Tensor:
"""Copy a tensor from DRAM tier to SRAM tier."""
src = self._dram_tensors[name]
t0 = self._timer_start()
dst = src.clone()
if self._use_cuda:
torch.cuda.synchronize(self.device)
dur = self._timer_end(t0)
self._sram_tensors[name + '_from_dram'] = dst
self._dram_stats.num_transfers += 1
self._dram_stats.total_transfer_us += dur
self._record_event('dram', 'transfer', src.nelement() * src.element_size(), dur)
return dst
```
**Explanation:**
- **HβL transfer (SRAM to DRAM):**
1. Get source tensor from SRAM registry
2. `clone()` = make a complete copy
3. `torch.cuda.synchronize()` = wait for GPU to finish (so we time actual operation)
4. Calculate duration
5. Store in DRAM registry with new name
6. Update transfer counter and total transfer time
7. Record the event
- **LβH transfer (DRAM to SRAM):** Same logic in reverse
**Why clone?** We don't want to move the original; we want a copy at the destination.
**Why synchronize?** GPU operations often run asynchronously (GPU schedules them but CPU moves on). Synchronize makes CPU wait, so we measure actual GPU time.
---
#### **Section 6: Context Managers for Streams (Lines 182-198)**
```python
@contextmanager
def sram_context(self):
"""Context manager that runs operations on the SRAM stream."""
if self._use_cuda and self._sram_stream is not None:
with torch.cuda.stream(self._sram_stream):
yield self._sram_stream
else:
yield None
@contextmanager
def dram_context(self):
"""Context manager that runs operations on the DRAM stream."""
if self._use_cuda and self._dram_stream is not None:
with torch.cuda.stream(self._dram_stream):
yield self._dram_stream
else:
yield None
```
**Explanation:**
A `@contextmanager` lets you use `with` statements:
```python
# Usage:
with memory_manager.sram_context() as stream:
# Operations here run on the SRAM stream
x = torch.matmul(a, b) # GPU runs this on _sram_stream
y = x + 1
# When we exit the block, stream is restored
```
**Why useful?** Different operations can run on different streams in parallel, potentially overlapping computation and memory transfer.
---
#### **Section 7: Statistics & Reporting (Lines 203-243)**
```python
def get_stats(self) -> Dict:
"""Return all memory tier statistics."""
return {
'sram': {
'peak_mb': self._sram_stats.peak_alloc_bytes / (1024 * 1024),
'current_mb': self._sram_stats.current_alloc_bytes / (1024 * 1024),
'hit_rate': self._sram_stats.hit_rate,
'num_loads': self._sram_stats.num_loads,
'num_stores': self._sram_stats.num_stores,
'num_transfers': self._sram_stats.num_transfers,
'avg_load_us': self._sram_stats.avg_load_us,
'avg_store_us': self._sram_stats.avg_store_us,
'total_transfer_us': self._sram_stats.total_transfer_us,
},
'dram': {
# Similar for DRAM...
},
'num_events': len(self._events),
}
def get_events(self) -> List[MemoryEvent]:
"""Return raw event log."""
return list(self._events)
def reset_stats(self):
"""Clear all statistics and event log."""
self._events.clear()
self._sram_stats = TierStats()
self._dram_stats = TierStats()
def free_all(self):
"""Release all managed tensors."""
self._sram_tensors.clear()
self._dram_tensors.clear()
self._sram_stats.current_alloc_bytes = 0
self._dram_stats.current_alloc_bytes = 0
```
**Explanation:**
- **`get_stats()`**: Returns a dictionary with all collected metrics (ready to print or save)
- Converts bytes to MB (divide by 1024Β²)
- Includes hit rate, counts, timing
- **`get_events()`**: Returns the raw event log (for detailed analysis)
- **`reset_stats()`**: Clear counters before a new benchmark run (don't contaminate results with previous runs)
- **`free_all()`**: Release memory and reset current allocations (cleanup)
---
#### **Section 8: Internal Timing Helpers (Lines 248-257)**
```python
def _timer_start(self) -> float:
if self._use_cuda:
torch.cuda.synchronize(self.device)
return time.perf_counter()
def _timer_end(self, t0: float) -> float:
if self._use_cuda:
torch.cuda.synchronize(self.device)
return (time.perf_counter() - t0) * 1e6 # β microseconds
def _record_event(self, tier: str, op: str, nbytes: int, dur_us: float):
if self.enable_tracking:
self._events.append(MemoryEvent(
tier=tier, operation=op, bytes=nbytes,
duration_us=dur_us, timestamp=time.time(),
))
```
**Explanation:**
- **`_timer_start()`:**
- Synchronize GPU first (wait for all pending operations)
- Record CPU time
- Returns the start time
- **`_timer_end(t0)`:**
- Synchronize GPU (wait for all operations since start)
- Calculate elapsed time
- Convert to microseconds (Γ 1,000,000)
- **`_record_event()`:**
- If tracking is enabled, create a `MemoryEvent` and add to log
- Otherwise, skip (for performance when we don't need detailed logs)
---
### File 2: `softmax.py` β Simple PyTorch Reference
```python
import torch
import torch.nn.functional as F
sample = torch.tensor([[1,2,3,4,5], [5,4,3,2,1]], dtype=torch.float32, device='cuda')
ref = F.softmax(sample, dim=1)
print(f"Softmax result: {ref=}")
```
**Explanation:**
This is a **minimal reference implementation** showing:
1. Create a sample tensor: 2 rows Γ 5 columns of numbers
2. Apply softmax along dimension 1 (across columns)
3. Print the result
**What softmax does:** Converts scores to probabilities (sum to 1). Example:
- Input `[1, 2, 3, 4, 5]`
- Output ~`[0.67, 0.18, 0.05, 0.01, 0.006]` (largest inputs become largest probabilities)
This is a building block used in attention mechanisms in transformers.
---
## PART 3: Triton Implementation β Line-by-Line
### File: `triton_kernels.py` β GPU Kernels for Memory-Aware Compute
Triton is a language for writing GPU kernels (low-level GPU programs). Unlike PyTorch which relies on pre-built library functions, Triton lets us control **exactly** how data flows through GPU memory.
#### **Header & Imports (Lines 1-23)**
```python
"""
Triton kernels for HRM SRAM/DRAM memory-tiered operations.
Key idea: In Triton, SRAM = registers + shared memory (managed by compiler
within a block). DRAM = global memory (GPU HBM). By structuring kernels to keep
L-level state tile-resident (loaded once, reused many times within a block),
we ensure L-level stays in SRAM. H-level state is loaded from global memory
(DRAM) each cycle, paying the full memory bandwidth cost.
This gives us real, measurable latency differences that map to the HRM's
hierarchical update frequencies.
"""
import torch
import triton
import triton.language as tl
import math
```
**Explanation:**
- Key insight: **Data locality matters enormously** on GPUs
- SRAM (L1/L2 cache, registers, shared memory): ~1-4 cycles latency, 100s of GB/s bandwidth
- DRAM (global memory, HBM): ~200-400 cycles latency, 100s of GB/s like SRAM but way higher latency!
**The trick:** Keep "important" data (L-level) in SRAM by reusing it within a kernel block. Load "less critical" data (H-level) fresh each time from DRAM.
---
#### **Kernel 1: SRAM-Resident RMS-Norm + Residual (Lines 28-64)**
```python
@triton.jit
def _rms_norm_residual_fused_kernel(
X_ptr, # Input tensor (residual branch)
Residual_ptr, # Residual connection input
Out_ptr, # Output tensor
N: tl.constexpr, # Hidden dimension (constexpr β compiler tiles in SRAM)
eps: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""Fused RMS-norm + residual add.
By making N and BLOCK_N constexpr, the compiler keeps the entire hidden
vector in registers/shared-memory (SRAM) across the norm computation.
This is the kernel used for L-level (fast path).
"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK_N)
mask = cols < N
# ---- Load both inputs into SRAM (registers) in one shot ----
x = tl.load(X_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32)
r = tl.load(Residual_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32)
# Residual add β stays in registers
h = x + r
# RMS norm β entirely in registers, no global memory round-trip
variance = tl.sum(h * h, axis=0) / N
h_norm = h * tl.math.rsqrt(variance + eps)
# ---- Store back to global memory ----
tl.store(Out_ptr + row * N + cols, h_norm.to(tl.bfloat16), mask=mask)
```
**Explanation (line by line):**
1. **Function signature:**
- `@triton.jit`: This is GPU code (JIT-compiled, not Python)
- `X_ptr`, `Residual_ptr`, `Out_ptr`: **Pointers** to tensors in GPU memory
- `N: tl.constexpr`: "constexpr" = compiler treats as a constant (compile-time value, not runtime)
- `BLOCK_N: tl.constexpr`: Block size (constexpr tells compiler to optimize for fixed size)
2. **Get thread ID:**
- `row = tl.program_id(0)`: Which "row" is this GPU thread processing?
- On GPU, thousands of threads run in parallel; each processes one row
3. **Create index array:**
- `cols = tl.arange(0, BLOCK_N)`: Array `[0, 1, 2, ..., BLOCK_N-1]`
- `mask = cols < N`: Boolean mask to handle if N < BLOCK_N (some threads might not have data)
4. **Load data into registers:**
- `tl.load(X_ptr + row * N + cols, ...)`: Load elements from memory
- `row * N`: Start at row-th row
- `+ cols`: Load all columns for this row
- `mask=mask`: Only load valid columns
- `other=0.0`: If invalid column, use 0.0
- `.to(tl.float32)`: Convert to 32-bit floats for math
- **KEY**: All data now in registers/shared memory (SRAM)!
5. **Compute residual:**
- `h = x + r`: Add the two inputs (element-wise, all in registers)
6. **Compute RMS norm:**
- RMS = sqrt(mean(xΒ²))
- `variance = tl.sum(h * h, axis=0) / N`: Compute mean of squares
- `h_norm = h * tl.math.rsqrt(variance + eps)`: Normalize by reciprocal square root
- `eps`: Small number to avoid division by zero
- **KEY**: All arithmetic in registers!
7. **Store result:**
- `tl.store(Out_ptr + ..., h_norm.to(tl.bfloat16), mask=mask)`
- Write normalized result back to global memory
- `.to(tl.bfloat16)`: Convert to lower precision for storage (saves memory bandwidth)
**Why this is fast:**
- Load data once β keep in SRAM β do many operations β store once
- If we did this in PyTorch with separate ops (add, then norm), we'd load/store twice
---
#### **Kernel 2: DRAM-Sourced RMS-Norm + Residual (Lines 68-95)**
```python
@triton.jit
def _rms_norm_residual_dram_kernel(
X_ptr,
Residual_ptr,
Out_ptr,
N: tl.constexpr,
eps: tl.constexpr,
BLOCK_N: tl.constexpr,
):
"""RMS-norm + residual for H-level.
Structurally identical but designed to be called with larger strides
and without re-use inside a meta-kernel. Each call does a full
DRAM round-trip, modeling the slower H-level memory access pattern.
"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK_N)
mask = cols < N
# Global memory load (DRAM)
x = tl.load(X_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32)
r = tl.load(Residual_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32)
h = x + r
variance = tl.sum(h * h, axis=0) / N
h_norm = h * tl.math.rsqrt(variance + eps)
tl.store(Out_ptr + row * N + cols, h_norm.to(tl.bfloat16), mask=mask)
```
**Explanation:**
**Identical kernel code**, but:
- **Intent is different**: This is called less frequently (H-level runs slower)
- **Memory behavior differs** via **calling context**:
- SRAM kernel: Data reused many times within tight loops β stays in cache
- DRAM kernel: Data used once then discarded β evicted from cache
**In code**, they're identical because the difference is **how often and how much** they're called, not the kernel itself.
---
#### **Kernel 3: SwiGLU Activation (Lines 100-126)**
```python
@triton.jit
def _swiglu_fused_sram_kernel(
GateUp_ptr, # [rows, 2 * inter] β gate and up projections concatenated
Out_ptr, # [rows, inter]
inter: tl.constexpr,
BLOCK_INTER: tl.constexpr,
):
"""Fused SiLU(gate) * up in a single kernel pass.
For L-level: The gate and up vectors are loaded once into registers
and the activation is computed without spilling to DRAM.
"""
row = tl.program_id(0)
cols = tl.arange(0, BLOCK_INTER)
mask = cols < inter
# Load gate and up from contiguous memory β both go into SRAM
gate = tl.load(GateUp_ptr + row * 2 * inter + cols, mask=mask, other=0.0).to(tl.float32)
up = tl.load(GateUp_ptr + row * 2 * inter + inter + cols, mask=mask, other=0.0).to(tl.float32)
# SiLU(gate) * up β entirely in registers
silu_gate = gate * tl.sigmoid(gate)
result = silu_gate * up
tl.store(Out_ptr + row * inter + cols, result.to(tl.bfloat16), mask=mask)
```
**Explanation:**
**SwiGLU** = A gating mechanism in transformers. Formula: `output = sigmoid(gate) * up`
- **Input layout:** `GateUp_ptr` contains concatenated `[gate | up]` (e.g., first half is gate, second half is up)
- **Load both halves:**
- Gate: `row * 2 * inter + cols` (first half)
- Up: `row * 2 * inter + inter + cols` (second half)
- **Compute SiLU:**
- `tl.sigmoid(gate)`: Apply sigmoid (s-shaped function) to gate
- `silu_gate = gate * sigmoid(gate)`: SiLU = gate * sigmoid(gate)
- `result = silu_gate * up`: Gated output
- **Store:** Result back to global memory
**Why fused?** If done separately:
1. Load gate
2. Compute sigmoid
3. Write intermediate
4. Load intermediate
5. Compute SiLU
6. Write intermediate
7. Load up
8. Compute product
9. Write result
Fused does it in one kernel β one load, one store.
---
#### **Kernel 4: State Transfer (Lines 131-148)**
```python
@triton.jit
def _state_transfer_kernel(
Src_ptr,
Dst_ptr,
numel: tl.constexpr,
BLOCK: tl.constexpr,
):
"""Explicit memory copy kernel for cross-tier state transfer.
Used when H-level needs to read L-level output (or vice versa).
Triton compiles this into optimized async memcpy instructions.
"""
pid = tl.program_id(0)
offsets = pid * BLOCK + tl.arange(0, BLOCK)
mask = offsets < numel
data = tl.load(Src_ptr + offsets, mask=mask, other=0.0)
tl.store(Dst_ptr + offsets, data, mask=mask)
```
**Explanation:**
Simple memcpy kernel (copy data from source to destination):
- `pid = tl.program_id(0)`: Program/thread block ID
- `offsets = pid * BLOCK + tl.arange(0, BLOCK)`: Calculate which elements this block handles
- If BLOCK=1024, pid=0 handles offsets 0-1023, pid=1 handles 1024-2047, etc.
- `mask = offsets < numel`: Don't copy past the end
- Load and store in parallel across many threads
**Why separate kernel?** GPU memcopy can be optimized by the compiler (uses memory controllers, not just compute cores).
---
#### **Kernel 5: Memory Latency Probe (Lines 153-176)**
```python
@triton.jit
def _memory_latency_probe_kernel(
Data_ptr,
Out_ptr,
N: tl.constexpr,
BLOCK_N: tl.constexpr,
NUM_ITERS: tl.constexpr,
):
"""Probe kernel to measure effective memory latency.
Performs NUM_ITERS dependent loads to measure true SRAM vs DRAM latency.
The data dependency chain prevents compiler reordering.
"""
pid = tl.program_id(0)
cols = tl.arange(0, BLOCK_N)
mask = cols < N
# Initial load from global memory
acc = tl.load(Data_ptr + pid * N + cols, mask=mask, other=0.0)
# Dependent iteration chain β forces sequential memory access
for _ in range(NUM_ITERS):
# This stays in SRAM (registers) because acc is reused
acc = acc * 1.00001 + 0.00001
tl.store(Out_ptr + pid * N + cols, acc, mask=mask)
```
**Explanation:**
This kernel **measures latency** by creating a **dependency chain** that can't be optimized away:
1. Load initial data
2. For N iterations:
- Multiply by 1.00001 + add 0.00001 (cheap operations)
- Result depends on previous iteration (creates dependency)
3. Store result
**Why not just load/store?** The compiler could optimize away separate load-stores, but a **dependency chain** forces real latency measurement.
**For SRAM data:**
- Data stays in registers
- All iterations hit registers (super fast)
- Total time β NUM_ITERS Γ 1 cycle β very fast
**For DRAM data:**
- Data in global memory
- Each iteration reloads from DRAM
- Total time β NUM_ITERS Γ 200-400 cycles β slow!
**Result:** By comparing SRAM vs DRAM probe times, we measure the latency difference.
---
#### **Python Wrappers (Lines 184-265)**
```python
def _next_power_of_2(n: int) -> int:
return 1 << (n - 1).bit_length()
def triton_rms_norm_residual_sram(
x: torch.Tensor,
residual: torch.Tensor,
eps: float = 1e-5,
) -> torch.Tensor:
"""SRAM-optimized fused RMS-norm + residual for L-level."""
assert x.shape == residual.shape
assert x.is_contiguous() and residual.is_contiguous()
rows, N = x.shape[0] * (x.shape[1] if x.ndim == 3 else 1), x.shape[-1]
flat_x = x.reshape(rows, N)
flat_r = residual.reshape(rows, N)
out = torch.empty_like(flat_x)
BLOCK_N = _next_power_of_2(N)
_rms_norm_residual_fused_kernel[(rows,)](
flat_x, flat_r, out,
N=N, eps=eps, BLOCK_N=BLOCK_N,
)
return out.reshape(x.shape)
```
**Explanation:**
- **`_next_power_of_2(n)`**: Find smallest power of 2 β₯ n
- Example: `_next_power_of_2(512)` β 512, `_next_power_of_2(513)` β 1024
- (Bit manipulation: `(n-1).bit_length()` gives number of bits, `1 << x` is 2^x)
- GPUs work best with powers of 2 (thread block sizes)
- **`triton_rms_norm_residual_sram(...)`**: Python wrapper to call the Triton kernel
- Checks inputs are same shape and contiguous
- Flattens to 2D (rows Γ hidden_size)
- Creates output tensor
- Picks block size (nearest power of 2)
- Launches kernel with `[(rows,)]` β one thread block per row
- Reshapes output back to original
**Why a wrapper?** Triton kernels are GPU code; we need Python to:
1. Prepare data (reshape, allocate output)
2. Launch the kernel (call the JIT-compiled function)
3. Return result to CPU
---
## PART 4: How PyTorch & Triton Work Together
### The Full Pipeline
```
βββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
β Input Data (e.g., transformer input) β
ββββββββββββββββββββββββββββββββββ¬βββββββββββββββββββββββββββββββββ
β
ββββββββββββββΌβββββββββββββ
β MemoryTierManager β
β allocates L-level βββββ β
β allocates H-level βββ β β
ββββββββββββββββββ¬βββββ β β
β β β
ββββββββββββββββββββββΌβββ β β
β L-level (Fast Path) β β β
β Triton kernels: β β β
β - SRAM RMS+Residual β β β
β - SRAM SwiGLU β β β
β (Data in regs/cache) β β β
ββββββββββββββββββ¬βββββββ β β
β β β
ββββββββββββββββββββββΌβββββββββ β β
β H-level (Slow Path) β β β
β Triton kernels: β β β
β - DRAM RMS+Residual β β β
β (Data in global memory) β β β
ββββββββββββββββββ¬ββββββββββββ β β
β β β
βββββββββββββββββββββΌββββββββββββββ β
β Triton State Transfer: ββ β
β Move LβH and HβL states ββββββββΌβ β
β (what MemoryTierManager times) β β
βββββββββββββββββββββ¬βββββββββββββ β
β β
ββββββββββββββΌβββββββββββββ β
β Output β β
β (result tensor) β β
ββββββββββββββββββββββββββ β
β
βββββββββββββββββββββΌ
β Statistics collected:
β - L/H latency
β - Transfer time
β - Memory usage
β - Hit rate
```
### Real Example: Forward Pass
```python
# 1. Create memory manager
mem_mgr = MemoryTierManager(device='cuda')
# 2. Allocate L-level (fast) hidden state
L_hidden = mem_mgr.alloc_sram('L_state', shape=(batch, hidden_dim), dtype=torch.float32)
# β Records allocation in SRAM if it fits, otherwise spills to DRAM
# 3. Allocate H-level (slow) hidden state
H_hidden = mem_mgr.alloc_dram('H_state', shape=(batch, hidden_dim), dtype=torch.float32)
# β Records allocation in DRAM
# 4. L-level computes using SRAM kernel (fast)
with mem_mgr.sram_context():
L_out = triton_rms_norm_residual_sram(L_hidden, residual, eps=1e-5)
L_out = triton_swiglu_sram(gate_up, inter_dim)
# All data in registers/shared memory
# Time recorded: ~10-100 microseconds per operation
# 5. H-level computes using DRAM kernel (slower)
with mem_mgr.dram_context():
H_out = triton_rms_norm_residual_dram(H_hidden, residual, eps=1e-5)
# Data loaded from global memory each time
# Time recorded: ~100-1000 microseconds per operation
# 6. Transfer L-level output to H-level
L_to_H_output = mem_mgr.transfer_sram_to_dram('L_out')
# β Records transfer size and time
# 7. Collect metrics
stats = mem_mgr.get_stats()
print(f"L/H latency ratio: {stats['H_latency'] / stats['L_latency']}")
# β Usually 10-100x difference
```
---
## PART 5: Bringing It All Together
### What Happens When You Run a Benchmark
```
python run_benchmark.py --mode tiered --iterations 10
```
1. **Setup:**
- Create MemoryTierManager
- Create model (HRM_Tiered)
- Create dummy batch of data
2. **Warmup phase (first 5 runs):**
- Forward pass (not timed)
- GPU caches warm up, compilers warm up
- Discarded from results
3. **Benchmark phase (10 timed runs):**
- For each iteration:
- Start GPU timer
- Forward pass:
- L-level uses SRAM kernels (fast)
- H-level uses DRAM kernels (slow)
- Record latencies in mem_mgr
- Stop GPU timer
- Record elapsed time
4. **Statistics:**
- Calculate mean, min, max, std dev of times
- Get memory stats (peak, current usage)
- Get transfer stats (HβL copy times)
- Calculate derived metrics (ratios, efficiency)
- Get Triton probe latencies (direct SRAM vs DRAM measurement)
5. **Output:**
- Print table of results
- Save to JSON
- Optionally generate plots
---
## Key Insights
### Why Two Implementations?
| Aspect | PyTorch | Triton |
|--------|---------|--------|
| **What it expresses** | Algorithmic intent | Micro-architecture intent |
| **Memory control** | Coarse (allocate tensor) | Fine (exact register usage) |
| **Performance** | Relies on library kernels | Direct GPU control |
| **Latency** | Can hide memory issues | Exposes memory latency differences |
| **Usability** | Easy to write | Low-level, harder to write |
### Why SRAM vs DRAM?
**The 200Γ Latency Difference:**
- SRAM (L2 cache): Load time ~4 cycles = ~2 nanoseconds = **very fast**
- DRAM (HBM): Load time ~400 cycles = **200 nanoseconds = 200x slower!**
By **exploiting this difference**, we can:
- Keep "important" (L-level) data in fast SRAM
- Allow "less important" (H-level) data to use slow DRAM
- Model hierarchical computation: fast frequent + slow infrequent
### Why Fuse Operations?
**Without fusion:**
```
Load β Compute β Store β Load β Compute β Store
```
Multiple memory round-trips!
**With fusion:**
```
Load β Compute β Compute β Compute β Store
```
One load, many operations, one store. **Massive speedup!**
---
## Testing It Yourself
### Quick test (2 minutes):
```bash
cd test-env/HRM_optimised
python run_benchmark.py --mode tiered --warmup 1 --iterations 3 --batch-sizes 2 --seq-lens 16
```
### Full benchmark (10 minutes):
```bash
python run_benchmark.py --mode compare --warmup 5 --iterations 20 --batch-sizes 1,8,32 --seq-lens 64,128
```
### Analyze results:
```python
import json
with open('benchmark_results/results.json') as f:
results = json.load(f)
print(f"L/H latency ratio: {results[0]['h_over_l_latency_ratio']}")
print(f"SRAM hit rate: {results[0]['sram_hit_rate']}")
print(f"Memory efficiency: {results[0]['memory_efficiency']}")
```
---
## Summary
You now understand:
β
**Benchmarks** β What metrics are collected and why
β
**PyTorch layer** β How `MemoryTierManager` tracks two memory tiers
β
**Triton layer** β How kernels control GPU memory access patterns
β
**Integration** β How they work together to model hierarchical computation
β
**Performance** β Why SRAM vs DRAM and data fusion matter for speed
The key idea: **Different parts of the model operate at different timescales (L-level fast, H-level slow), and we can model this using GPU memory hierarchy (SRAM fast, DRAM slow).**
|