본문 바로가기
C.W.K.
Stream
Lesson 01 of 05 · published

프로덕션 용 JAX 모델 내보내기

~8 min · ecosystem, jax, tutorial

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

연구 코드가 잘 돌아도, 프로덕션 인프라가 JAX가 아니면? 표준 답: jax2tf로 TensorFlow SavedModel 출력해.

import jax
import jax.numpy as jnp
from jax.experimental import jax2tf
import tensorflow as tf

# JAX model
@jax.jit
def my_model(params, x):
    return jax.nn.softmax(x @ params["W"] + params["b"])

# JAX → TF function
tf_fn = jax2tf.convert(my_model, with_gradient=False)

# wrapper Module
class MyTfModule(tf.Module):
    def __init__(self, params):
        self.params = tf.nest.map_structure(tf.Variable, params)

    @tf.function(input_signature=[tf.TensorSpec((None, 784), tf.float32)])
    def serve(self, x):
        return tf_fn(self.params, x)

# 저장
module = MyTfModule(params)
tf.saved_model.save(module, "my_saved_model/")

저장한 뒤 다른 환경에서 TensorFlow로 불러와:

loaded = tf.saved_model.load("my_saved_model/")
y = loaded.serve(x_input)   # JAX 가 없는 환경에서도 동작

TF Serving, 모바일용 TF Lite, 브라우저용 TFJS, Vertex AI 등과 호환돼.

ONNX 내보내기

JAX → ONNX 직접 변환은 여전히 (2026) 미숙. 일반적인 경로는 JAX → TF → ONNX야:

import tf2onnx

# 위에서 저장한 SavedModel 을 ONNX 로
spec = (tf.TensorSpec((None, 784), tf.float32, name="input"),)
model_proto, _ = tf2onnx.convert.from_keras(
    module,
    input_signature=spec,
    output_path="model.onnx",
)

ONNX 호환성은 점점 좋아지고 있어서 TensorRT, OpenVINO, Core ML 같은 다양한 추론 엔진으로 연결할 수 있어.

StableHLO, 새로운 IR

2024년부터, Google이 StableHLO를 cross-프레임워크 IR로 밀고 있어. JAX, TensorFlow, PyTorch 모두 내보낼 수 있어. XLA 기반 추론 인프라가 직접 받을 수 있어:

from jax.experimental import export

# JAX 함수 → StableHLO
exported = export.export(my_model)(params, jax.ShapeDtypeStruct((None, 784), jnp.float32))
serialized = exported.serialize()

# 다른 곳에서 로드
loaded = export.deserialize(serialized)
y = loaded.call(params, x_input)

장기적으로 StableHLO의 역할이 커질 수 있지만, 현재는 TensorFlow SavedModel이 가장 검증된 경로야.

ONNX 우회, JAX 직접 배포

인프라가 JAX를 지원하면 JAX로 직접 배포할 수도 있어:

# FastAPI server
from fastapi import FastAPI
import jax
import pickle

app = FastAPI()

# load params at startup
with open("params.pkl", "rb") as f:
    params = pickle.load(f)

# pre-compile
dummy_input = jnp.zeros((1, 784))
jit_predict = jax.jit(my_model)
_ = jit_predict(params, dummy_input)   # warm-up

@app.post("/predict")
def predict(data: dict):
    x = jnp.array(data["input"])
    y = jit_predict(params, x)
    return {"output": y.tolist()}

Google 내부에서는 JAX를 프로덕션에서 직접 실행하는 사례가 많지만, 외부에서는 기존 인프라와의 호환성 때문에 TensorFlow나 ONNX로 변환하는 경우가 더 흔해.

🚀 배포 전략 결정

(1) 인프라가 TF/Vertex AI 위주, jax2tf → SavedModel. (2) 모바일 / edge, jax2tf → TF Lite. (3) 브라우저, jax2tf → TFJS. (4) JAX를 직접 실행할 수 있음 (FastAPI, custom server), 그냥 JAX. (5) ONNX가 필요, JAX → TF → ONNX. 변환마다 약간의 지원 손실 가능, 중요한 op가 모두 지원되는지 미리 검증해.

변환된 모델이 원본과 비트 단위로 같은지 항상 검증해. 수치 차이가 미묘하게 다를 수 있어. 단위 테스트에서 난수 입력의 출력을 비교해.

Code

import jax
from jax import export
import jax.numpy as jnp
import numpy as np

# 1. Create a JIT-transformed function
@jax.jit
def predict(params, x):
    h = jax.nn.relu(x @ params['w1'] + params['b1'])
    return h @ params['w2'] + params['b2']

# 2. Define input shapes for export
params_shapes = {
    'w1': jax.ShapeDtypeStruct((784, 256), jnp.float32),
    'b1': jax.ShapeDtypeStruct((256,), jnp.float32),
    'w2': jax.ShapeDtypeStruct((256, 10), jnp.float32),
    'b2': jax.ShapeDtypeStruct((10,), jnp.float32),
}
x_shape = jax.ShapeDtypeStruct((1, 784), jnp.float32)

# 3. Export to StableHLO
exported = export.export(predict)(params_shapes, x_shape)

# 4. Get the StableHLO module (MLIR text)
stablehlo_module = exported.mlir_module()

# 5. Serialize for later use
serialized = export.export(predict)(params_shapes, x_shape).serialize()
# Can be saved to disk and loaded in a different process/language
# Export with dynamic batch dimension
scope = export.SymbolicScope()
dynamic_x = jax.ShapeDtypeStruct(
    export.symbolic_shape("batch, 784", scope=scope),
    jnp.float32,
)

exported_dynamic = export.export(predict)(params_shapes, dynamic_x)
# Now the exported model accepts any batch size
# You can pack StableHLO into a TensorFlow SavedModel
# for serving with TensorFlow Serving
from jax.experimental.jax2tf import convert as jax2tf_convert

# Note: the recommended path is now:
# 1. Export to StableHLO via jax.export
# 2. Load StableHLO into TF SavedModel if needed for TF Serving
# 3. Or use StableHLO directly with XLA-compatible runtimes

External links

Exercise

학습된 모델을 jax2tf로 변환해 TensorFlow SavedModel로 내보내. 일반 TensorFlow에서 다시 불러와 예측이 같은지 검증해. 이 내보내기 경로가 JAX를 연구 전용 도구로 남길지 프로덕션까지 가져갈지에 어떤 영향을 주는지 적어.

Progress

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

댓글 0

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

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