GPU에서 Conway의 생명 게임 최적화: PyTorch, CUDA, Triton 활용
Conway의 생명 게임 셀룰러 오토마톤은 간단한 지역 규칙 덕분에 병렬 GPU 컴퓨팅에 완벽하게 적합합니다. N×N 그리드의 각 셀은 8개의 이웃을 분석합니다: 살아있는 셀은 2~3개의 살아있는 이웃이 있을 때 생존하고, 죽은 셀은 정확히 3개의 살아있는 이웃이 있을 때 살아납니다. 테스트는 Nvidia A40에서 216×216 그리드(int8 기준 4 GB)로 수행되었습니다. 이론적 한계는 메모리 대역폭 696 GB/s로 결정된 반복당 11.5ms입니다.
이론적 한계와 기본 계산
단일 셀 업데이트에는 9바이트 로드와 1바이트 쓰기가 필요합니다. 4 GB 데이터의 경우 최소 시간은 4 GB × 2 / 696 GB/s = 11.5ms입니다. 계산 부하는 최소이며 메모리가 병목 현상입니다. 단순화를 위해 그리드 경계는 무시됩니다.
PyTorch: 기본 구현부터 torch.compile까지
PyTorch는 표준 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)
기본 버전: 개별 연산 오버헤드로 인해 223ms. torch.compile은 그래프를 병합하고 최적화합니다: 38.1ms(피크의 30%). 자동으로 Triton 커널을 생성합니다.
CUDA: 수동 스레드 및 캐시 관리
CUDA 커널은 구성 가능한 블록(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];
// ... (나머지 8개 이웃)
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;
}
L1 캐시가 중요합니다: 없으면 >55ms. 최적은 1×128(26ms, 피크의 44%). 블록 ≤1024, 32의 배수. 정사각형 블록은 둘레를 최소화하고, 직사각형 블록은 행 우선을 활용합니다.
주요 CUDA 블록 매개변수:
- 최대 1024 스레드/블록
- 점유율을 위한 32의 배수
- 레지스터와 공유 메모리 균형
- 메모리 제한 작업에 1×128 선호
Triton: 벡터화 및 자동 최적화
Triton은 텐서 연산을 추가하여 CUDA를 단순화합니다. 커널은 경계 마스크로 3×3 블록을 로드합니다.
@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):
# 3x3을 위한 오프셋과 마스크
row00 = tl.load(x_ptr + row_offsets0 * row_stride + col_offsets0,
mask=row_mask0 & col_mask0, other=0)
# ... (9개 로드)
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)
1024 블록(스레드당 8 셀, 128 스레드). Triton은 자동 벡터화 및 공유 메모리 사용: 22.5ms(피크의 51%).
성능 비교
| 프레임워크 | 시간(ms) | 피크 대비 % | 참고 |
|-----------|------------|-----------|------------|
| PyTorch | 223 | 5% | 오버헤드 |
| torch.compile | 38.1 | 30% | 퓨전 |
| CUDA | 26 | 44% | 1×128 |
| Triton | 22.5 | 51% | 벡터화됨 |
| 이론 | 11.5 | 100% | 메모리 |
핵심 요점
- 11.5ms 한계 완벽한 캐싱으로 달성 가능
- Triton 선두 (22.5ms) 자동 벡터화 덕분
- CUDA의 1×128 블록 메모리 제한 작업에 최적
- torch.compile PyTorch를 5.8배 가속
- 다음 단계: 셀 그룹 루프를 통한 그룹화 CUDA 커널
— Editorial Team
아직 댓글이 없습니다.