Optimalizace buněčného automatu "Život" na GPU: PyTorch, CUDA a Triton
Buněčný automat Conwayho "Život" je ideální pro paralelní výpočty na GPU díky jednoduchým lokálním pravidlům. Každá buňka v mřížce N×N analyzuje 8 sousedů: živá buňka přežívá při 2–3 živých sousedech, mrtvá ožívá při přesně 3. Testování probíhá na Nvidia A40 s mřížkou 216×216 (4 GB v int8). Teoretický limit je 11,5 ms za iteraci, určený propustností paměti 696 GB/s.
Teoretická omezení a základní výpočty
Pro aktualizaci jedné buňky je potřeba načíst 9 bajtů a zapsat 1 bajt. Při 4 GB dat je minimální čas 4 GB × 2 / 696 GB/s = 11,5 ms. Výpočetní zátěž je minimální, bottleneck je paměť. Hranice mřížky ignorujeme pro zjednodušení.
PyTorch: od základní implementace k torch.compile
PyTorch používá box blur pro počítání sousedů místo standardních konvolucí float32.
def gol_torch_sum(x: torch.Tensor) -> torch.Tensor:
y = x[2:] + x[1:-1] + x[:-2]
z = y[:, 2:] + y[:, 1:-1] + y[:, :-2]
z = torch.nn.functional.pad(z, (1, 1, 1, 1), value=0)
return ((x == 1) & (z == 4)) | (z == 3).to(torch.int8)
Základní verze: 223 ms kvůli režii jednotlivých operací. torch.compile slučuje graf a optimalizuje: 38,1 ms (30 % od špičky). Automaticky generuje Triton-jádro.
CUDA: ruční správa vláken a cache
CUDA-jádro počítá jednu buňku na vlákno s nastavitelnými bloky (block_size_row × block_size_col).
__global__ void gol_kernel_i8(const int8_t* __restrict__ x_ptr,
int8_t* __restrict__ out_ptr,
int64_t rowstride, int64_t n) {
int64_t x = blockIdx.x * blockDim.x + threadIdx.x;
int64_t y = blockIdx.y * blockDim.y + threadIdx.y;
if (x >= n - 2 || y >= n - 2) return;
int8_t r00 = x_ptr[y * rowstride + x + 0 * rowstride + 0];
// ... (zbylých 8 sousedů)
int8_t sum = r00 + r01 + r02 + r10 + r12 + r20 + r21 + r22;
int8_t result = (r11 > 0) ? ((sum == 2) || (sum == 3) ? 1 : 0) : (sum == 3 ? 1 : 0);
out_ptr[(y + 1) * rowstride + (x + 1)] = result;
}
Cache L1 je kritická: bez ní >55 ms. Optimálně 1×128 (26 ms, 44 % špičky). Bloky ≤1024, násobky 32. Čtvercové bloky minimalizují obvod, obdélníkové využívají row-major.
Klíčové parametry bloků CUDA:
- Maximum 1024 vláken/blok
- Násobek 32 pro obsazenost
- Rovnováha registrů a shared memory
- Preferujte 1×128 pro úlohy vázané na paměť
Triton: vektorizace a automatická optimalizace
Triton zjednodušuje CUDA, přidává tensor-operace. Jádro načítá 3×3 bloky s maskami hranic.
@triton.jit
def gol_triton_2d_kernel(x_ptr, out_ptr, row_stride: tl.int64, N: tl.int64,
BLOCK_SIZE_ROW: tl.constexpr, BLOCK_SIZE_COL: tl.constexpr):
# offsets a masks pro 3x3
row00 = tl.load(x_ptr + row_offsets0 * row_stride + col_offsets0,
mask=row_mask0 & col_mask0, other=0)
# ... (9 načtení)
sum = row00 + row01 + row02 + row10 + row12 + row20 + row21 + row22
result = tl.where(row11 > 0, (sum == 2) | (sum == 3), sum == 3).to(tl.int8)
tl.store(out_ptr + row_offsets1 * row_stride + col_offsets1, result,
mask=row_mask1 & col_mask1)
Bloky po 1024 (8 buněk/vlákno, 128 vláken). Triton auto-vektorizuje a používá shared memory: 22,5 ms (51 % špičky).
Porovnání výkonnosti
| Framework | Čas (ms) | % od špičky | Poznámka |
|-----------|----------|-------------|----------|
| PyTorch | 223 | 5% | Overhead |
| torch.compile | 38,1 | 30% | Fusion |
| CUDA | 26 | 44% | 1×128 |
| Triton | 22,5 | 51% | Vectorized |
| Teorie | 11,5 | 100% | Memory |
Co je důležité
- Limit 11,5 ms dosažitelný při ideálním cachování
- Triton vede (22,5 ms) díky auto-vektorizaci
- Bloky 1×128 v CUDA optimální pro memory-bound
- torch.compile zrychluje PyTorch 5,8×
- Dále – grouped kernels CUDA s cyklem po skupinách buněk
— Editorial Team
Zatím žádné komentáře.