test / skill_example /references /t4-optimization-guide.md
Jack-Khuu
Demo
88a1dd2
|
Raw History Blame Contribute Delete
21.8 kB
# T4 Turing GPU Optimization Guide
## Overview
The NVIDIA T4 GPU, based on the Turing architecture (sm_75), is widely deployed in cloud inference environments. While it lacks many features of newer GPUs (no BF16, limited shared memory), it remains important for cost-effective inference. This guide covers optimization strategies specific to T4 hardware.
## T4 Hardware Specifications
| Specification | Value |
|---|---|
| Architecture | Turing (sm_75) |
| Streaming Multiprocessors (SMs) | 40 |
| Memory Bandwidth | 320 GB/s (GDDR6) |
| Shared Memory per SM | 64 KB (configurable) |
| L2 Cache | 4 MB |
| FP32 CUDA Cores | 2560 |
| Tensor Cores (2nd gen) | 320 |
| Memory | 16 GB GDDR6 |
| Memory Bus Width | 256-bit |
| TDP | 70W |
| Max Threads per SM | 1024 |
| Max Threads per Block | 1024 |
| Warp Size | 32 |
| Max Warps per SM | 32 |
| Register File per SM | 64 KB |
| Max Registers per Thread | 255 |
## Key Limitations Compared to H100/A100
| Feature | T4 | A100 | H100 |
|---|---|---|---|
| BF16 Support | **No** | Yes | Yes |
| FP16 Tensor Cores | Yes (2nd gen) | Yes (3rd gen) | Yes (4th gen) |
| Memory Type | GDDR6 | HBM2e | HBM3 |
| Bandwidth | 320 GB/s | 2.0 TB/s | 3.35 TB/s |
| Shared Memory/SM | 64 KB | 164 KB | 192 KB |
| SMs | 40 | 108 | 132 |
| Memory | 16 GB | 40/80 GB | 80 GB |
| TDP | 70W | 400W | 700W |
| INT8 Tensor Cores | Yes | Yes | Yes |
| FP8 Support | No | No | Yes |
| TMA | No | No | Yes |
## No BF16 Support: FP16 Only
The T4 does **not** support BF16 (bfloat16). All kernels must use FP16 (float16) or FP32. This is the single most important consideration when porting kernels to T4.
### Converting BF16 Kernels to FP16
```cpp
// BEFORE (H100/A100 -- BF16):
#include <cuda_bf16.h>
__nv_bfloat16 val = __float2bfloat16(x);
float result = __bfloat162float(val);
__nv_bfloat162 pair = __halves2bfloat162(a, b);
// AFTER (T4 -- FP16):
#include <cuda_fp16.h>
half val = __float2half(x);
float result = __half2float(val);
half2 pair = __halves2half2(a, b);
```
### Vectorized FP16 Access on T4
```cpp
#include <cuda_fp16.h>
// Load 8 FP16 values (128 bits) via float4
__device__ __forceinline__ void load_fp16x8_t4(
const half* ptr,
int base_idx,
float vals[8]
) {
float4 packed = reinterpret_cast<const float4*>(ptr)[base_idx / 8];
half2* pairs = reinterpret_cast<half2*>(&packed);
#pragma unroll
for (int i = 0; i < 4; i++) {
vals[2 * i] = __half2float(__low2half(pairs[i]));
vals[2 * i + 1] = __half2float(__high2half(pairs[i]));
}
}
// Store 8 FP16 values (128 bits) via float4
__device__ __forceinline__ void store_fp16x8_t4(
half* ptr,
int base_idx,
const float vals[8]
) {
float4 packed;
half2* pairs = reinterpret_cast<half2*>(&packed);
#pragma unroll
for (int i = 0; i < 4; i++) {
pairs[i] = __halves2half2(
__float2half(vals[2 * i]),
__float2half(vals[2 * i + 1])
);
}
reinterpret_cast<float4*>(ptr)[base_idx / 8] = packed;
}
```
### FP16 Overflow Considerations
FP16 has a much narrower dynamic range than BF16:
| Property | FP16 | BF16 |
|---|---|---|
| Max value | 65,504 | ~3.4e38 |
| Min positive normal | ~6.1e-5 | ~1.2e-38 |
| Precision | ~3.3 decimal digits | ~3 decimal digits |
This means you must be more careful with:
```cpp
// DANGER: FP16 overflow in attention scores
// If seq_len is large and values are not scaled, scores can overflow 65504
float score = dot_product / sqrtf(head_dim); // Scale BEFORE exponentiation
// DANGER: FP16 underflow in normalization
// Very small variances can underflow
float variance = sum_sq / n;
float inv_std = rsqrtf(variance + epsilon); // epsilon prevents 1/0
// SAFE: Always compute reductions in FP32
float sum = 0.0f; // FP32 accumulator
for (int i = tid; i < n; i += blockDim.x) {
sum += __half2float(input[i]); // Convert to FP32 for accumulation
}
```
## Memory-Bound Optimization
The T4's 320 GB/s bandwidth is roughly 6x less than A100 and 10x less than H100. Most deep learning kernels on T4 will be memory-bound.
### Bandwidth Utilization Targets
| Achieved Bandwidth | Assessment | Action |
|---|---|---|
| > 250 GB/s (78%) | Excellent | Kernel is well-optimized |
| 200-250 GB/s (62-78%) | Good | Minor improvements possible |
| 100-200 GB/s (31-62%) | Fair | Check coalescing and vectorization |
| < 100 GB/s (31%) | Poor | Significant optimization needed |
### Maximizing Bandwidth on T4
```cpp
// Rule 1: Always use vectorized loads (float4 = 16 bytes per thread)
float4 data = reinterpret_cast<const float4*>(input)[tid];
// Rule 2: Ensure coalesced access
// Good: consecutive threads access consecutive addresses
int idx = blockIdx.x * blockDim.x + threadIdx.x;
float val = input[idx];
// Rule 3: Minimize global memory round-trips by fusing operations
// BAD: Two kernel launches = two memory round-trips
// rmsnorm(x, weight, temp);
// gelu(temp, output);
// GOOD: Single fused kernel
// fused_rmsnorm_gelu(x, weight, output);
// Rule 4: Use shared memory for data reuse
__shared__ float smem[TILE_SIZE];
// Load once from global memory, use multiple times from shared
```
### Arithmetic Intensity
To determine if a kernel is memory-bound or compute-bound on T4:
```
Arithmetic Intensity = FLOPs / Bytes Transferred
T4 compute-to-memory ratio:
FP32: 8.1 TFLOPS / 320 GB/s = 25.3 FLOP/byte
FP16: 65 TFLOPS / 320 GB/s = 203 FLOP/byte (with Tensor Cores)
If your kernel's arithmetic intensity < 25.3 (FP32) or < 203 (FP16):
=> Memory-bound: Focus on bandwidth optimization
Otherwise:
=> Compute-bound: Focus on Tensor Core utilization
```
Common operations and their arithmetic intensity:
| Operation | Approximate AI | Bound on T4 |
|---|---|---|
| RMSNorm | ~5 FLOP/byte | Memory-bound |
| GELU | ~3 FLOP/byte | Memory-bound |
| Softmax | ~5 FLOP/byte | Memory-bound |
| RoPE | ~4 FLOP/byte | Memory-bound |
| MatMul (large) | ~100+ FLOP/byte | Compute-bound |
| Attention (small seq) | ~10 FLOP/byte | Memory-bound |
## Smaller Tile Sizes
The T4's 64 KB shared memory per SM (compared to 164 KB on A100 or 192 KB on H100) requires smaller tile sizes.
### Tile Size Recommendations
| Operation | H100 Tile | A100 Tile | T4 Tile |
|---|---|---|---|
| MatMul | 128x128 | 128x128 | 64x64 |
| Attention | 128x64 | 128x64 | 64x32 |
| Reduction | N/A | N/A | N/A |
| Conv2D tile | 64x64 | 64x64 | 32x32 |
### Shared Memory Budget
```cpp
// T4: 64 KB shared memory max
// Common configurations:
// For reduction kernels (RMSNorm, LayerNorm):
// Only need warp reduction buffer -- minimal shared memory
__shared__ float warp_results[32]; // 128 bytes -- plenty of room
// For tiled matmul:
// Two tiles in shared memory
// 64x64 tile of FP16 = 64 * 64 * 2 = 8 KB per tile
// Two tiles = 16 KB -- fits easily
__shared__ half tile_a[64][64 + 1]; // +1 for bank conflict avoidance
__shared__ half tile_b[64][64 + 1];
// For attention with larger tiles:
// 64x32 Q tile + 32x64 K tile + 32x64 V tile
// = 64*32*2 + 32*64*2 + 32*64*2 = 4 + 4 + 4 = 12 KB
// Plus score buffer: 64*32*4 = 8 KB
// Total: ~20 KB -- fits in 64 KB
// WARNING: Do NOT use tile sizes designed for H100/A100
// __shared__ half tile[128][128]; // = 32 KB per tile -- too large for T4!
```
### Dynamic Shared Memory on T4
```cpp
// Query available shared memory
cudaDeviceProp prop;
cudaGetDeviceProperties(&prop, 0);
printf("Shared memory per block: %zu bytes\n", prop.sharedMemPerBlock);
printf("Shared memory per SM: %zu bytes\n", prop.sharedMemPerMultiprocessor);
// T4 output:
// Shared memory per block: 49152 bytes (48 KB default)
// Shared memory per SM: 65536 bytes (64 KB)
// To use more than 48 KB, request explicitly:
cudaFuncSetAttribute(
my_kernel,
cudaFuncAttributeMaxDynamicSharedMemorySize,
64 * 1024 // Request up to 64 KB
);
```
## INT8 Quantization
The T4 has strong INT8 Tensor Core support, making quantization an effective optimization strategy.
### INT8 Quantization for Inference
```python
import torch
from torch.quantization import quantize_dynamic
# Dynamic quantization (easiest)
model_int8 = quantize_dynamic(
model,
{torch.nn.Linear},
dtype=torch.qint8,
)
# Or using bitsandbytes for LLM quantization
import bitsandbytes as bnb
# 8-bit linear layer
linear_8bit = bnb.nn.Linear8bitLt(
input_features=4096,
output_features=4096,
bias=False,
has_fp16_weights=False,
)
```
### INT8 Kernel Example
```cpp
#include <cuda_runtime.h>
#include <cstdint>
// INT8 GEMV kernel for T4 (simplified)
__global__ void int8_gemv_kernel(
const int8_t* __restrict__ weight, // [out_features, in_features]
const int8_t* __restrict__ input, // [in_features]
float* __restrict__ output, // [out_features]
const float* __restrict__ scale_w, // [out_features] per-channel scale
const float scale_x, // input scale
const int in_features,
const int out_features
) {
const int row = blockIdx.x * blockDim.x + threadIdx.x;
if (row >= out_features) return;
const int8_t* w_row = weight + row * in_features;
// Accumulate dot product in INT32
int32_t acc = 0;
for (int i = 0; i < in_features; i += 4) {
// Load 4 INT8 values at once (32 bits)
int32_t w4 = reinterpret_cast<const int32_t*>(w_row)[i / 4];
int32_t x4 = reinterpret_cast<const int32_t*>(input)[i / 4];
int8_t* w_bytes = reinterpret_cast<int8_t*>(&w4);
int8_t* x_bytes = reinterpret_cast<int8_t*>(&x4);
acc += (int32_t)w_bytes[0] * (int32_t)x_bytes[0];
acc += (int32_t)w_bytes[1] * (int32_t)x_bytes[1];
acc += (int32_t)w_bytes[2] * (int32_t)x_bytes[2];
acc += (int32_t)w_bytes[3] * (int32_t)x_bytes[3];
}
// Dequantize to float
output[row] = (float)acc * scale_w[row] * scale_x;
}
```
## TensorRT Integration
TensorRT is highly optimized for T4 and often provides the best inference performance.
### Converting a Model to TensorRT
```python
import torch
import tensorrt as trt
# Export model to ONNX first
torch.onnx.export(
model,
dummy_input,
"model.onnx",
opset_version=17,
input_names=["input"],
output_names=["output"],
dynamic_axes={
"input": {0: "batch_size", 1: "seq_len"},
"output": {0: "batch_size", 1: "seq_len"},
},
)
# Build TensorRT engine
logger = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(logger)
network = builder.create_network(
1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
)
parser = trt.OnnxParser(network, logger)
with open("model.onnx", "rb") as f:
parser.parse(f.read())
config = builder.create_builder_config()
config.set_memory_pool_limit(trt.MemoryPoolType.WORKSPACE, 1 << 30) # 1 GB
# Enable FP16 (T4 supports this)
config.set_flag(trt.BuilderFlag.FP16)
# Enable INT8 for maximum performance
config.set_flag(trt.BuilderFlag.INT8)
engine = builder.build_serialized_network(network, config)
```
### TensorRT with Custom Plugins
If you have custom CUDA kernels, you can wrap them as TensorRT plugins:
```cpp
#include "NvInferPlugin.h"
class RMSNormPlugin : public nvinfer1::IPluginV2DynamicExt {
public:
// ... TensorRT plugin interface implementation
// This allows your custom RMSNorm to be used within TensorRT graphs
int enqueue(
const nvinfer1::PluginTensorDesc* inputDesc,
const nvinfer1::PluginTensorDesc* outputDesc,
const void* const* inputs,
void* const* outputs,
void* workspace,
cudaStream_t stream
) noexcept override {
// Launch your custom kernel here
const int hidden_size = inputDesc[0].dims.d[inputDesc[0].dims.nbDims - 1];
const int num_rows = /* compute total rows */;
rmsnorm_kernel<256><<<num_rows, 256, 0, stream>>>(
static_cast<const half*>(inputs[0]),
static_cast<const half*>(inputs[1]),
static_cast<half*>(outputs[0]),
hidden_size,
1e-6f
);
return 0;
}
};
```
## 16 GB Memory Strategies
The T4's 16 GB memory is a significant constraint for modern models. Here are strategies to work within this limit.
### Model Size Estimates
| Model | FP32 | FP16 | INT8 | INT4 |
|---|---|---|---|---|
| LLaMA-7B | 28 GB | 14 GB | 7 GB | 3.5 GB |
| LLaMA-13B | 52 GB | 26 GB | 13 GB | 6.5 GB |
| Mistral-7B | 28 GB | 14 GB | 7 GB | 3.5 GB |
| SD 1.5 UNet | 3.4 GB | 1.7 GB | 0.85 GB | 0.43 GB |
| SDXL UNet | 10 GB | 5 GB | 2.5 GB | 1.25 GB |
| SD3 Transformer | 8 GB | 4 GB | 2 GB | 1 GB |
### Memory Optimization Techniques
```python
import torch
# 1. Use FP16 (required on T4 -- no BF16!)
model = model.half() # Convert to FP16
# 2. Use quantization
from transformers import BitsAndBytesConfig
quantization_config = BitsAndBytesConfig(
load_in_8bit=True, # INT8 quantization
# Or for 4-bit:
# load_in_4bit=True,
# bnb_4bit_compute_dtype=torch.float16,
)
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-2-7b-hf",
quantization_config=quantization_config,
device_map="auto",
)
# 3. Use CPU offloading for diffusers
pipe = StableDiffusionPipeline.from_pretrained(
"runwayml/stable-diffusion-v1-5",
torch_dtype=torch.float16,
)
pipe.enable_model_cpu_offload() # Moves parts to CPU when not in use
# 4. Enable attention slicing to reduce peak memory
pipe.enable_attention_slicing(slice_size=1)
# 5. Use VAE tiling for high-resolution image generation
pipe.enable_vae_tiling()
# 6. Gradient checkpointing (for fine-tuning)
model.gradient_checkpointing_enable()
```
### Memory-Efficient Kernel Design
```cpp
// Design kernels that minimize temporary allocations
// BAD: Allocates a large temporary buffer
torch::Tensor bad_kernel(torch::Tensor input) {
auto temp = torch::empty_like(input); // Doubles memory usage!
// ... compute into temp ...
auto output = torch::empty_like(input);
// ... compute into output from temp ...
return output;
}
// GOOD: In-place computation or single output buffer
torch::Tensor good_kernel(torch::Tensor input) {
auto output = torch::empty_like(input);
// Compute directly into output, using registers and shared memory
// for intermediate values instead of global memory temp buffers
return output;
}
```
## Occupancy Tuning for 40 SMs
The T4 has only 40 SMs with a maximum of 1024 threads per SM (32 warps). This is significantly less than A100 (108 SMs, 2048 threads/SM) or H100 (132 SMs, 2048 threads/SM).
### Grid Sizing for T4
```cpp
// T4: 40 SMs, max 1024 threads/SM, max 32 warps/SM
int num_sms = 40;
// Minimum: 1 block per SM
int min_grid = 40;
// Recommended: 2-4 blocks per SM
int grid_2x = 80;
int grid_4x = 160;
// For memory-bound kernels: more blocks help hide latency
int grid_8x = 320;
```
### Block Size Recommendations for T4
| Kernel Type | Block Size | Blocks/SM | Warps/SM | Occupancy |
|---|---|---|---|---|
| Element-wise | 256 | 4 | 32 | 100% |
| Row reduction | 256 | 4 | 32 | 100% |
| Row reduction | 128 | 8 | 32 | 100% |
| Tiled matmul | 128 | 4 | 16 | 50% |
| Complex kernel | 64 | 8 | 16 | 50% |
### Register Pressure on T4
```cpp
// T4: 64 KB registers per SM = 16,384 32-bit registers
// Max 255 registers per thread
// With 1024 threads/SM: 16 registers per thread (if fully occupied)
// With 512 threads/SM: 32 registers per thread (50% occupancy)
// With 256 threads/SM: 64 registers per thread (25% occupancy)
// Use __launch_bounds__ to help the compiler
__global__ __launch_bounds__(256, 4) // 256 threads, aim for 4 blocks/SM
void t4_optimized_kernel(...) {
// With 4 blocks of 256 threads = 1024 threads/SM
// Budget: 16 registers per thread for 100% occupancy
// Realistically: 32-48 registers with some occupancy trade-off
}
```
## Cloud Instance Types
| Provider | Instance | GPUs | vCPUs | RAM | Price (approx) |
|---|---|---|---|---|---|
| AWS | g4dn.xlarge | 1x T4 | 4 | 16 GB | $0.526/hr |
| AWS | g4dn.2xlarge | 1x T4 | 8 | 32 GB | $0.752/hr |
| AWS | g4dn.12xlarge | 4x T4 | 48 | 192 GB | $3.912/hr |
| GCP | n1-standard-4 + T4 | 1x T4 | 4 | 15 GB | $0.35/hr |
| GCP | n1-standard-8 + T4 | 1x T4 | 8 | 30 GB | $0.45/hr |
| Azure | NC4as_T4_v3 | 1x T4 | 4 | 28 GB | $0.526/hr |
| Azure | NC8as_T4_v3 | 1x T4 | 8 | 56 GB | $0.752/hr |
### Choosing the Right Instance
- **g4dn.xlarge**: Best for single-model inference, batch size 1
- **g4dn.2xlarge**: Good for inference with larger batch sizes or CPU-heavy preprocessing
- **g4dn.12xlarge**: Multi-model serving or batch processing
## build.toml Configuration for T4
```toml
[build]
cuda-version = "12.4"
cuda-capabilities = ["7.5"] # T4 = sm_75
[kernel.rmsnorm]
src = ["src/rmsnorm_fp16.cu"] # FP16 version!
[kernel.gelu]
src = ["src/gelu_fp16.cu"]
[kernel.rope]
src = ["src/rope_fp16.cu"]
```
### Multi-Target Configuration
```toml
[build]
cuda-version = "12.4"
# Support T4, A100, and H100
cuda-capabilities = ["7.5", "8.0", "9.0"]
```
### Conditional Compilation for T4
```cpp
// Use preprocessor to handle architecture differences
#if __CUDA_ARCH__ >= 800
// A100/H100: Use BF16
#include <cuda_bf16.h>
using compute_t = __nv_bfloat16;
using compute_t2 = __nv_bfloat162;
#define FLOAT_TO_COMPUTE __float2bfloat16
#define COMPUTE_TO_FLOAT __bfloat162float
#else
// T4: Use FP16
#include <cuda_fp16.h>
using compute_t = half;
using compute_t2 = half2;
#define FLOAT_TO_COMPUTE __float2half
#define COMPUTE_TO_FLOAT __half2float
#endif
```
## T4 RMSNorm Example (FP16)
```cpp
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
template<int BLOCK_SIZE>
__global__ void rmsnorm_fp16_t4(
const half* __restrict__ input,
const half* __restrict__ weight,
half* __restrict__ output,
const int hidden_size,
const float epsilon
) {
const int row = blockIdx.x;
const int tid = threadIdx.x;
const half* x = input + row * hidden_size;
half* out = output + row * hidden_size;
// Compute sum of squares with FP32 accumulation
float sum_sq = 0.0f;
// Vectorized: process 8 FP16 values per iteration
const int num_vecs = hidden_size / 8;
for (int v = tid; v < num_vecs; v += BLOCK_SIZE) {
float4 packed = reinterpret_cast<const float4*>(x)[v];
half2* pairs = reinterpret_cast<half2*>(&packed);
#pragma unroll
for (int j = 0; j < 4; j++) {
float lo = __half2float(__low2half(pairs[j]));
float hi = __half2float(__high2half(pairs[j]));
sum_sq += lo * lo + hi * hi;
}
}
// Warp reduction
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
sum_sq += __shfl_xor_sync(0xffffffff, sum_sq, offset);
}
// Block reduction
__shared__ float warp_sums[BLOCK_SIZE / 32];
int warp_id = tid / 32;
int lane_id = tid % 32;
if (lane_id == 0) warp_sums[warp_id] = sum_sq;
__syncthreads();
__shared__ float rms_scale;
if (warp_id == 0) {
float val = (lane_id < BLOCK_SIZE / 32) ? warp_sums[lane_id] : 0.0f;
for (int offset = 16; offset > 0; offset >>= 1)
val += __shfl_xor_sync(0xffffffff, val, offset);
if (lane_id == 0)
rms_scale = rsqrtf(val / hidden_size + epsilon);
}
__syncthreads();
float scale = rms_scale;
// Vectorized normalization
for (int v = tid; v < num_vecs; v += BLOCK_SIZE) {
float4 in_packed = reinterpret_cast<const float4*>(x)[v];
float4 w_packed = reinterpret_cast<const float4*>(weight)[v];
half2* in_pairs = reinterpret_cast<half2*>(&in_packed);
half2* w_pairs = reinterpret_cast<half2*>(&w_packed);
float4 out_packed;
half2* out_pairs = reinterpret_cast<half2*>(&out_packed);
#pragma unroll
for (int j = 0; j < 4; j++) {
float x0 = __half2float(__low2half(in_pairs[j]));
float x1 = __half2float(__high2half(in_pairs[j]));
float w0 = __half2float(__low2half(w_pairs[j]));
float w1 = __half2float(__high2half(w_pairs[j]));
out_pairs[j] = __halves2half2(
__float2half(x0 * scale * w0),
__float2half(x1 * scale * w1)
);
}
reinterpret_cast<float4*>(out)[v] = out_packed;
}
}
torch::Tensor rmsnorm_forward(torch::Tensor input, torch::Tensor weight, double eps) {
TORCH_CHECK(input.dtype() == torch::kFloat16, "T4 requires FP16 input");
TORCH_CHECK(input.is_contiguous());
int hidden_size = input.size(-1);
auto input_2d = input.view({-1, hidden_size});
int num_rows = input_2d.size(0);
auto output = torch::empty_like(input_2d);
constexpr int BLOCK_SIZE = 256;
rmsnorm_fp16_t4<BLOCK_SIZE><<<num_rows, BLOCK_SIZE>>>(
reinterpret_cast<const half*>(input_2d.data_ptr<at::Half>()),
reinterpret_cast<const half*>(weight.data_ptr<at::Half>()),
reinterpret_cast<half*>(output.data_ptr<at::Half>()),
hidden_size,
static_cast<float>(eps)
);
return output.view_as(input);
}
```
## Summary
When targeting the T4:
1. **Use FP16 only** -- BF16 is not supported on Turing (sm_75)
2. **Optimize for 320 GB/s bandwidth** -- most kernels will be memory-bound
3. **Use smaller tile sizes** -- 64 KB shared memory limits tile dimensions
4. **Leverage INT8 Tensor Cores** for quantized inference
5. **Consider TensorRT** for production inference workloads
6. **Plan for 16 GB memory** -- use quantization, offloading, and attention slicing
7. **Tune for 40 SMs** -- grid sizes should be multiples of 40
8. **Guard against FP16 overflow** -- especially in attention score computation
9. **Use conditional compilation** (`__CUDA_ARCH__`) to support T4 alongside A100/H100