본문 바로가기
C.W.K.
Stream
Lesson 03 of 06 · published

JAX vs PyTorch vs TensorFlow

~12 min · origins, jax, tutorial

Level 0호기심
0 XP0/73 lessons0/17 achievements
0/100 XP to next level100 XP to go0% complete

"PyTorch로 이미 다 할 수 있는데 왜 JAX를 써야 하지?" 합리적인 질문이야. 두 도구는 같은 문제를 서로 다른 관점에서 풀며, 강점을 보이는 영역도 달라.

PyTorch는 동적 실행과 eager 방식, 객체지향 API를 중심으로 해. nn.Module을 상속하고 순전파 메서드를 작성한 뒤 .backward()를 호출하면 그래디언트가 매개변수에 저장돼. 처음 배우는 사람에게 친숙한 대신 계산 그래프의 구조가 코드에 완전히 드러나지 않아 컴파일러 최적화에 제약이 생길 수 있어.

TensorFlow는 처음에는 계산 그래프를 먼저 만드는 방식이었지만 2.x부터 eager mode가 기본이 됐어. TFX, TF Serving, TF Lite, TFJS처럼 산업용 인프라가 강력한 반면 연구 현장에서는 사용 비중이 줄어드는 흐름이 있어.

JAX는 함수와 변환을 중심에 둬. 코어에는 nn.Module 같은 추상화가 없고 Flax나 Equinox가 그 역할을 따로 제공해. jit, grad, vmap, pmap이 일급 변환이어서 자유롭게 합성할 수 있고, 컴파일을 전제로 한 구조라 XLA가 깊이 최적화할 수 있어. 반면 순수 함수, 함수형 사고, PRNG 키 같은 개념을 먼저 익혀야 하는 진입 장벽이 있어.

         | PyTorch          | TensorFlow       | JAX
---------|------------------|------------------|------------------
스타일   | OOP, eager       | OOP/graph hybrid | functional
미분     | tensor.backward()| tape 기반        | jax.grad (함수)
batch    | 손으로 처리      | 손으로 처리      | jax.vmap
multi-GPU| DDP / FSDP       | tf.distribute    | pmap / sharding
연구     | 1 위 (분야 다수) | 감소             | 상승 (DeepMind 등)

🧭 어느 걸 골라야 하나

처음 머신러닝 프레임워크를 배우는 사람에게는 PyTorch가 접근하기 쉬워. 큰 기업의 기존 프로덕션 인프라에서는 TensorFlow도 여전히 중요한 선택지야. 함수 변환을 많이 합성하거나 TPU를 활용하는 연구라면 JAX가 잘 맞을 수 있어. 결국 둘 이상을 알면 문제에 맞춰 고를 수 있고, 이 퀘스트의 목적은 그중 JAX를 선택할 이유를 이해하는 데 있어.

JAX는 PyTorch를 없애기 위해 등장한 도구가 아니야. 같은 문제에 다른 답을 제시하며, 두 접근이 나란히 발전하는 시대에 우리는 살고 있어.

Code

# PyTorch style (for comparison)
import torch
import torch.nn as nn

class Model(nn.Module):
    def __init__(self):
        super().__init__()
        self.linear = nn.Linear(2, 1)

    def forward(self, x):
        return self.linear(x)

model = Model()
x = torch.tensor([[1.0, 2.0]])
y = model(x)          # Eager execution, tape is recording
loss = y.sum()
loss.backward()        # Replay tape to get gradients
print(model.linear.weight.grad)
# JAX style
import jax
import jax.numpy as jnp

def predict(params, x):
    w, b = params
    return jnp.dot(x, w) + b

def loss_fn(params, x, y):
    pred = predict(params, x)
    return jnp.mean((pred - y) ** 2)

# Parameters are just arrays in a tuple — no special Variable type
params = (jnp.array([0.5, 0.3]), jnp.array(0.1))
x = jnp.array([[1.0, 2.0]])
y = jnp.array([1.5])

# Gradient is a function, not a method on a loss object
grads = jax.grad(loss_fn)(params, x, y)
print(grads)  # Tuple of gradient arrays matching params structure

External links

Exercise

같은 단일 계층 회귀 모델을 PyTorch와 JAX로 각각 구현해. 합성 데이터로 100단계 학습하고 실행 시간을 측정해 봐. 그래디언트 상태의 소유자, 장치 관리, 루프의 가독성이라는 관점에서 API 차이를 적어. 우열을 정하려 하지 말고 JAX가 더 단순한 점 세 가지와 PyTorch가 더 단순한 점 세 가지를 찾아.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.