본문으로 건너뛰기
김신건의 로그

TPU (Tensor Processing Unit): Google 의 ML ASIC

· 수정 · 📖 약 3분 · 1,115자/단어 #ml #hardware #tpu #google #asic #xla
TPU, Tensor Processing Unit, Google TPU, TPU pod, TPU v4, Trillium

정의

TPU (Tensor Processing Unit) 는 Google 이 딥러닝 워크로드 전용으로 설계한 ASIC (Application-Specific Integrated Circuit). 2015년 내부 투입, 2017년 ISCA 발표. GPU 처럼 범용 병렬 프로세서가 아닌 행렬 곱셈 특화 가속기.

핵심 컴퓨팅 단위는 Systolic Array 기반 MXU (Matrix Multiply Unit).

언제 쓰이나

  • 대규모 Transformer (LLM, Vision Transformer) 학습/추론
  • Google Gemini, PaLM, BERT 등 Google 모델 대부분 TPU 에서 학습
  • JAX/XLA 스택으로 분산 학습을 구현할 때
  • GPU 에 비해 비용 효율($/TFLOPS) 이 중요한 장기 학습

핵심 아키텍처

MXU (Matrix Multiply Unit)

Systolic Array 를 이용한 대규모 병렬 행렬곱. 가중치(weight)가 PE 격자에 정주(stationary)하고 활성화(activation)가 격자를 통과하며 MAC 수행.

flowchart TB
    HBM["HBM (고대역폭 메모리)"] --> UB["Unified Buffer\n(on-chip SRAM 32-128MB)"]
    UB --> WF["Weight FIFO"]
    WF --> MXU["MXU (Systolic Array)\nBF16 행렬곱"]
    MXU --> ACC["Accumulators\n(FP32 누적)"]
    ACC --> VU["Vector Unit\n(softmax, norm, ReLU 등)"]
    VU --> UB
    UB --> HBM

bfloat16

FP32 exponent(8bit) + 축소된 mantissa(7bit). GPU 의 FP16(5+10) 과 달리 FP32 와 exponent 범위 동일 → 오버플로/언더플로 없이 학습 안정성 유지.

형식exponentmantissa동적 범위
FP328bit23bitFP32
FP165bit10bit좁음 (오버플로 위험)
BF168bit7bitFP32 와 동일
FP8 E4M34bit3bit매우 좁음 (추론용)

세대별 발전

flowchart LR
    V1["TPU v1\n2015\n추론 전용\nINT8"] --> V2["TPU v2\n2017\n학습 지원\nBF16 + HBM"]
    V2 --> V3["TPU v3\n2018\n액체 냉각\n2x v2"]
    V3 --> V4["TPU v4\n2021\n4096 칩 pod\nOCS 상호연결"]
    V4 --> V5["TPU v5e/p\n2023\n비용/성능 두 트랙"]
    V5 --> V6["Trillium v6\n2024\n5x vs v5e"]

세대별 상세

세대출시MXU 크기피크 TFLOPS (BF16)HBM특징
v12015256×25692 TOPS (INT8)없음추론 전용, AlphaGo 사용
v22017128×128 × 245 TFLOPS16 GB최초 BF16 학습, HBM 도입
v32018128×128 × 2123 TFLOPS32 GB액체 냉각, v2 2배
v42021128×128 × 4275 TFLOPS32 GBOCS 광 상호연결, PaLM 학습
v5e2023-~197 TFLOPS16 GB비용 효율 트랙
v5p2023-~459 TFLOPS95 GB성능 트랙
Trillium (v6)2024-~918 TFLOPS32 GBv5e 대비 4.7x 컴퓨트

TPU Pod: 수천 칩 연결

개별 TPU 칩을 고속 인터커넥트로 묶은 단위. Pod 안에서 칩들이 dedicated ICI (Inter-Chip Interconnect) 로 연결.

flowchart TB
    subgraph "TPU v4 Pod (최대 4096 칩)"
        subgraph "Cube 0"
            T00["TPU 0"] <-->|"ICI"| T01["TPU 1"]
            T01 <-->|"ICI"| T02["TPU 2"]
            T02 <-->|"ICI"| T03["TPU 3"]
        end
        subgraph "Cube 1"
            T10["TPU 4"] <-->|"ICI"| T11["TPU 5"]
            T11 <-->|"ICI"| T12["TPU 6"]
        end
        Cube0["Cube 0"] <-->|"OCS (광 스위치)"| Cube1["Cube 1"]
    end

TPU v4 Pod 의 혁신: OCS (Optical Circuit Switch) 로 수백~수천 칩을 유연하게 연결. 임의의 topology 구성 가능.

Pod 크기칩 수피크 성능
TPU v4 슬라이스8~512수십 petaFLOPS
TPU v4 Pod4096~1 exaFLOPS

XLA: 컴파일러가 MXU 를 최대한 활용

TPU 는 XLA (Accelerated Linear Algebra) 컴파일러 없이는 제 성능이 나오지 않는다.

flowchart LR
    JAX["JAX / TF / PyTorch/XLA\n(Python 코드)"] --> HLO["HLO\n(High-Level Optimizer IR)"]
    HLO --> OPT["최적화 패스\n(fusion, tiling, layout)"]
    OPT --> LLO["LLO\n(Low-Level Optimizer)"]
    LLO --> TPU["TPU HW 커널\n(MXU + VU 명령)"]

XLA 가 자동으로 수행하는 최적화:

최적화효과
Op fusionsoftmax = exp + sum + div → 단일 커널
Tiling행렬을 MXU 크기에 맞게 분할
Layout optimization메모리 레이아웃을 MXU 친화적으로 전환
Rematerializationactivation checkpoint (메모리/컴퓨트 트레이드오프)
SPMD partitioning분산 학습 자동 샤딩

실전: JAX 로 TPU 활용

기본 설정

import jax
import jax.numpy as jnp

# 사용 가능한 TPU 장치 확인
devices = jax.devices('tpu')
print(f"TPU 장치 수: {len(devices)}")

# jit 컴파일: XLA 가 MXU 최적화 적용
@jax.jit
def linear(weights, x):
    return jnp.dot(x, weights)

# BF16 명시 사용
x = jnp.ones((1024, 512), dtype=jnp.bfloat16)
W = jnp.ones((512, 256), dtype=jnp.bfloat16)
y = linear(W, x)  # [1024, 256] BF16

분산 학습 (pmap / pjit)

from jax.experimental import mesh_utils
from jax.sharding import Mesh, PartitionSpec, NamedSharding

# TPU Pod: 64 칩을 8x8 mesh 로 배치
devices = mesh_utils.create_device_mesh((8, 8))
mesh = Mesh(devices, axis_names=('data', 'model'))

# 가중치 모델 병렬, 배치 데이터 병렬 샤딩
weight_sharding = NamedSharding(mesh, PartitionSpec('model', None))
data_sharding   = NamedSharding(mesh, PartitionSpec('data', None))

@jax.jit
def train_step(state, batch):
    # pjit 내부에서 자동 collective (all-reduce, all-gather)
    loss, grads = jax.value_and_grad(loss_fn)(state.params, batch)
    return state.apply_gradients(grads=grads), loss

Flax 모델 정의

from flax import linen as nn

class TransformerBlock(nn.Module):
    hidden: int
    heads: int

    @nn.compact
    def __call__(self, x):
        # 모든 dot/matmul 이 XLA 에 의해 MXU 최적화됨
        attn_out = nn.MultiHeadDotProductAttention(
            num_heads=self.heads
        )(x, x)
        x = nn.LayerNorm()(x + attn_out)
        mlp_out = nn.Dense(self.hidden * 4)(x)
        mlp_out = nn.gelu(mlp_out)
        mlp_out = nn.Dense(self.hidden)(mlp_out)
        return nn.LayerNorm()(x + mlp_out)

GPU 와 비교

flowchart LR
    subgraph "GPU (H100)"
        g1["HBM 80GB\n3.35 TB/s"] --> g2["L2 Cache 50MB"]
        g2 --> g3["132 SM\n각 128 CUDA + 4 Tensor Core"]
        g3 --> g4["CUDA / cuDNN\n광범위한 생태계"]
    end
    subgraph "TPU (v4)"
        t1["HBM 32GB\n1.2 TB/s"] --> t2["Unified Buffer 32MB"]
        t2 --> t3["MXU 4개\n(Systolic Array)"]
        t3 --> t4["XLA 컴파일\nJAX / TF / PyTorch/XLA"]
    end
항목GPU (H100)TPU v4
피크 TFLOPS (BF16)1,979275 (칩당)
MFU (실제 활용률)30-50%60-70%
메모리 대역폭3.35 TB/s1.2 TB/s
소프트웨어CUDA (광범위)XLA 중심
병렬 실행 모델SIMT (warp 32T)Systolic Array
지연시간낮음 (스트리밍)배치 지향
접근성AWS, Azure, GCP, on-premGCP TPU 만
유연성높음 (그래픽, HPC 포함)ML 특화

IMPORTANT

H100 의 피크 TFLOPS 가 훨씬 높지만, MFU 를 감안하면 실효 성능 차는 크게 줄어든다. TPU 는 행렬 연산에서 효율이 매우 높고, Pod 로 묶으면 H100 클러스터보다 통신 대역폭이 유리하다.

한계

한계상세
GCP 한정AWS, Azure, on-prem 불가. 벤더 종속
CUDA 생태계 없음PyTorch CUDA 확장 직접 사용 불가. XLA 기반 재작성 필요
동적 shape 비효율XLA 는 shape 별 재컴파일. 가변 길이 시퀀스 처리 까다로움
디버깅 어려움JIT 컴파일 후 실행 → 스택 트레이스 추적 복잡
소형 모델작은 모델/배치에서는 GPU 가 오히려 빠름

흔한 함정

WARNING

  1. jax.jit 없이 실행 = Python 레벨 eager 실행, MXU 최적화 전혀 없음. 항상 @jax.jit 필수.
  2. 행렬 차원이 128 배수 아님 = MXU padding 낭비. hidden_dim, intermediate 크기를 128 배수로.
  3. 동적 shape 남발 = 매번 재컴파일 발생. 고정 shape 또는 jax.vmap + padding.
  4. HBM 부족 무시 = TPU v4 HBM 32GB. 큰 모델은 분산 샤딩 (tensor/pipeline parallelism) 필수.
  5. PyTorch 습관 그대로 = TPU 에서 .cuda() 대신 jax.device_put(x, devices[0]) 사용.

관련 위키

이 글의 용어 (6개)
모델 양자화ml
정의 양자화 (Quantization) 는 신경망 모델의 가중치/활성화를 높은 정밀도 (FP32, FP16) 에서 낮은 정밀도 (INT8, INT4 등) 로 변환하여 메모리와 연…
분산 학습ml
정의 분산 학습 (Distributed Training) 은 한 또는 에 들어가지 않는 큰 모델을 여러 가속기에 나눠서 학습시키는 기법. 현대 LLM (Llama 405B, GP…
GPU: 그래픽/ML 병렬 프로세서ml
정의 GPU (Graphics Processing Unit) 는 원래 3D 그래픽 렌더링을 위해 설계됐지만, massively parallel 아키텍처가 딥러닝의 행렬 연산과 완…
HBMml
정의 HBM (High Bandwidth Memory)는 DRAM die 를 수직으로 적층해 매우 높은 대역폭을 제공하는 메모리 패키지. AI 시대 GPU/TPU 의 메모리 병목…
SPMDml
정의 SPMD (Single Program, Multiple Data)는 병렬 컴퓨팅의 대표 모델 중 하나. 모든 프로세스(또는 thread)가 같은 프로그램을 실행하되, 각자 …
Systolic Arrayml
정의 Systolic Array 는 격자 형태로 배치된 다수의 Processing Element (PE) 가, 입력 데이터가 박동(systolic)처럼 격자를 가로질러 흐르는 동…

💬 댓글

사이트 검색 / 명령어

검색

스크롤 = 확대/축소 · 드래그 = 이동 · 0 = 원래 크기 · ESC = 닫기