Ray Serve 모델 멀티플렉싱
Ray Serve 모델 멀티플렉싱
비슷한 입력 형태를 가진 여러 모델을 하나의 레플리카 풀로 효율적으로 서빙하는 기법이 **모델 멀티플렉싱(model multiplexing)**이에요. 요청 헤더를 보고 해당 모델로 트래픽을 라우팅해서, 같은 형태지만 가중치만 다른 모델이 간헐적으로 호출되는 상황에서 비용을 줄이고 부하를 분산해요. 여기선 serve.multiplexed와 serve.get_multiplexed_model_id API로 멀티플렉싱 deployment를 작성하는 방법을 다뤄요.
왜 모델 멀티플렉싱인가
모델 멀티플렉싱은 비슷한 입력 유형을 가진 여러 모델을 레플리카 풀에서 효율적으로 서빙하는 기법이에요. 요청 헤더를 기준으로 해당 모델에 트래픽을 라우팅해요. 여러 모델을 레플리카 풀로 서빙하면 비용을 최적화하고 트래픽을 부하 분산할 수 있어요. 같은 형태지만 가중치만 다르고 드물게 호출되는 모델이 많은 경우에 특히 유용해요. deployment의 어떤 레플리카에 이미 모델이 로드돼 있으면, 그 모델을 향한 트래픽(요청 헤더 기준)은 불필요한 로드 시간 없이 자동으로 그 레플리카로 라우팅돼요.
멀티플렉싱 deployment 작성하기
멀티플렉싱 deployment를 작성하려면 serve.multiplexed와 serve.get_multiplexed_model_id API를 사용해요.
AWS S3 버킷에 여러 PyTorch 모델이 이런 구조로 있다고 해볼게요:
s3://my_bucket/1/model.pt
s3://my_bucket/2/model.pt
s3://my_bucket/3/model.pt
s3://my_bucket/4/model.pt
...
멀티플렉싱 deployment를 정의해요:
from ray import serve
import aioboto3
import torch
import starlette
@serve.deployment
class ModelInferencer:
def __init__(self):
self.bucket_name = "my_bucket"
@serve.multiplexed(max_num_models_per_replica=3)
async def get_model(self, model_id: str):
session = aioboto3.Session()
async with session.resource("s3") as s3:
obj = await s3.Bucket(self.bucket_name)
await obj.download_file(f"{model_id}/model.pt", f"model_{model_id}.pt")
return torch.load(f"model_{model_id}.pt", weights_only=False)
async def __call__(self, request: starlette.requests.Request):
model_id = serve.get_multiplexed_model_id()
model = await self.get_model(model_id)
return model.forward(torch.rand(64, 3, 512, 512))
entry = ModelInferencer.bind()
:::note 참고
serve.multiplexed API에는 max_num_models_per_replica 파라미터가 있어요. 이 값으로 단일 레플리카에 로드할 모델 수를 정해요. 모델 수가 max_num_models_per_replica보다 많아지면 Serve는 LRU 정책으로 가장 오래 사용되지 않은 모델을 퇴출해요.
:::
:::tip 팁
이 코드 예시는 PyTorch Model 객체를 사용해요. 자체 모델 클래스를 정의해서 쓸 수도 있어요. 모델이 퇴출될 때 리소스를 해제하려면 __del__ 메서드를 구현하세요. 모델이 퇴출될 때 Ray Serve가 내부적으로 __del__ 메서드를 호출해 리소스를 해제해요.
:::
serve.get_multiplexed_model_id는 요청 헤더에서 모델 ID를 가져와요. 이 ID는 get_model 함수에 전달돼요. 모델이 레플리카에 아직 캐시되지 않았다면 S3 버킷에서 로드하고, 이미 캐시됐다면 캐시된 모델을 반환해요.
:::note 참고 내부적으로 Serve 라우터는 요청 헤더의 모델 ID로 해당 레플리카에 트래픽을 라우팅해요. 모델을 가진 모든 레플리카가 과부하 상태면 Ray Serve는 새 레플리카로 요청을 보내고, 그 레플리카가 S3 버킷에서 모델을 로드·캐시해요. :::
특정 모델로 요청을 보내려면 요청 헤더에 serve_multiplexed_model_id 필드를 넣고, 보내려는 모델 ID를 값으로 설정해요.
import requests # noqa: E402
resp = requests.get(
"http://localhost:8000", headers={"serve_multiplexed_model_id": str("1")}
)
:::note 참고
serve_multiplexed_model_id는 요청 헤더에 필수로 있어야 하고, 값은 보내려는 모델 ID여야 해요. 요청 헤더에서 serve_multiplexed_model_id를 찾지 못하면 Serve는 이를 일반 요청으로 취급해 임의의 레플리카로 라우팅해요.
:::
위 코드를 실행하면 deployment 로그에 다음과 같은 줄이 보여요:
INFO 2023-05-24 01:19:03,853 default_Model default_Model#EjYmnQ CUpzhwUUNw / default replica.py:442 - Started executing request CUpzhwUUNw
INFO 2023-05-24 01:19:03,854 default_Model default_Model#EjYmnQ CUpzhwUUNw / default multiplex.py:131 - Loading model '1'.
INFO 2023-05-24 01:19:04,859 default_Model default_Model#EjYmnQ CUpzhwUUNw / default replica.py:542 - __CALL__ OK 1005.8ms
모델을 계속 로드해서 max_num_models_per_replica를 초과하면 가장 오래 사용되지 않은 모델이 퇴출되고 로그에 다음 줄이 보여요:
INFO 2023-05-24 01:19:15,988 default_Model default_Model#rimNjA WzjTbJvbPN / default replica.py:442 - Started executing request WzjTbJvbPN
INFO 2023-05-24 01:19:15,988 default_Model default_Model#rimNjA WzjTbJvbPN / default multiplex.py:145 - Unloading model '3'.
INFO 2023-05-24 01:19:15,988 default_Model default_Model#rimNjA WzjTbJvbPN / default multiplex.py:131 - Loading model '4'.
INFO 2023-05-24 01:19:16,993 default_Model default_Model#rimNjA WzjTbJvbPN / default replica.py:542 - __CALL__ OK 1005.7ms
handle의 options API로 특정 모델에 요청을 보낼 수도 있어요.
obj_ref = handle.options(multiplexed_model_id="1").remote("<your param>")
모델 구성(model composition)을 쓸 때는 Serve DeploymentHandle로 업스트림 deployment에서 멀티플렉싱 deployment로 요청을 보낼 수 있어요. options에 multiplexed_model_id를 설정해야 해요. 예를 들어:
from ray.serve.handle import DeploymentHandle # noqa: E402
@serve.deployment
class Downstream:
def __call__(self):
return serve.get_multiplexed_model_id()
@serve.deployment
class Upstream:
def __init__(self, downstream: DeploymentHandle):
self._h = downstream
async def __call__(self, request: starlette.requests.Request):
return await self._h.options(multiplexed_model_id="bar").remote()
serve.run(Upstream.bind(Downstream.bind()))
resp = requests.get("http://localhost:8000")
모델 ID 매칭 타임아웃 구성하기
serve_multiplexed_model_id가 있는 요청이 도착하면 Serve 라우터는 이미 모델을 로드한 레플리카와 매칭하려 해요. 타임아웃 안에 매칭되는 레플리카가 없으면 기본 라우팅 전략으로 폴백해서 모델을 즉시 로드하는 아무 레플리카로 보내요.
이 타임아웃은 RAY_SERVE_MULTIPLEXED_MODEL_ID_MATCHING_TIMEOUT_S 환경 변수로 설정할 수 있어요:
export RAY_SERVE_MULTIPLEXED_MODEL_ID_MATCHING_TIMEOUT_S=2.0
기본값: 1.0초. 같은 미로드 모델을 향한 요청이 동시에 많이 몰릴 때 썬더링 허드(thundering herd) 문제를 피하려고, 실제 타임아웃은 이 값과 value * 2 사이(예: 기본값 기준 1.0–2.0초)에서 무작위로 정해져요.
모델 로드에 오래 걸리고 이미 모델을 로드한 레플리카를 기다리는 편이 낫다면 타임아웃을 늘리세요. 아무 레플리카로 더 빨리 폴백하고 싶다면 줄이면 돼요.
배칭과 함께 쓰기
모델 멀티플렉싱은 @serve.batch 데코레이터와 결합해 효율적인 배치 추론을 할 수 있어요. 두 기능을 함께 쓰면 Ray Serve가 모델 ID별로 배치를 자동 분할해서, 각 배치가 같은 모델의 요청만 포함하도록 보장해요. 이렇게 하면 서로 다른 모델을 겨냥한 요청이 한 배치에 섞이는 문제를 막아줘요.
다음 예시는 멀티플렉싱과 배칭을 결합한 모습이에요:
from typing import List # noqa: E402
from starlette.requests import Request
@serve.deployment(max_ongoing_requests=15)
class BatchedMultiplexModel:
@serve.multiplexed(max_num_models_per_replica=3)
async def get_model(self, model_id: str):
# Load and return your model here
return model_id
@serve.batch(max_batch_size=10, batch_wait_timeout_s=0.1)
async def batched_predict(self, inputs: List[str]) -> List[str]:
# Get the model ID - this works correctly inside batched functions
# because all requests in the batch target the same model
model_id = serve.get_multiplexed_model_id()
model = await self.get_model(model_id)
# Process the batch with the loaded model
return [f"{model}:{inp}" for inp in inputs]
async def __call__(self, request: Request):
# Extract input from the request body
input_text = await request.body()
return await self.batched_predict(input_text.decode())
:::note 참고
serve.get_multiplexed_model_id()는 @serve.batch로 장식된 함수 안에서도 올바르게 동작해요. Ray Serve는 배치의 모든 요청이 같은 multiplexed_model_id를 갖도록 보장하므로, 이 값을 안전하게 사용해서 배치 전체에 적용할 적절한 모델을 로드할 수 있어요.
:::