연구 코드가 잘 돌아도, 프로덕션 인프라가 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가 모두 지원되는지 미리 검증해.
변환된 모델이 원본과 비트 단위로 같은지 항상 검증해. 수치 차이가 미묘하게 다를 수 있어. 단위 테스트에서 난수 입력의 출력을 비교해.