Big Model Inference

Big Model Inference (대형 모델 추론)

가장 큰 발전 중 하나가 GPU 메모리에 완전히 들어가지 않는 모델로도 추론을 할 수 있게 해주는 Big Model Inference예요. 모델을 GPU·CPU·하드디스크에 조각내어 배치하고, 층을 통과할 때마다 필요 메모리만 옮겨 가며 추론하는 방식이에요.

출처: Big Model Inference (공식)

일반적인 모델 로드 방식

PyTorch 모델을 로드하는 일반적인 흐름은 이렇습니다. GPU 메모리를 넘어서는 모델이라면 여기서 메모리 에러가 나요.

import torch

my_model = ModelClass(...)
state_dict = torch.load(checkpoint_file)
my_model.load_state_dict(state_dict)

빈 모델 만들고 배치하기

Big Model Inference에서는 init_empty_weights 컨텍스트 매니저로 모델의 빈 뼈대를 먼저 만들어요. 파라미터가 없는 상태라 메모리를 거의 쓰지 않아요.

from accelerate import init_empty_weights

with init_empty_weights():
    my_model = ModelClass(...)

그 다음 load_checkpoint_and_dispatch()가 빈 모델에 체크포인트를 로드하고, 각 층의 가중치를 사용 가능한 디바이스에 배치(deploy)해요. 빠른 디바이스(GPU, MPS, XPU, NPU 등)부터 시작해 느린 디바이스(CPU, 하드디스크)로 옮겨 가요.

from accelerate import load_checkpoint_and_dispatch

model = load_checkpoint_and_dispatch(
    model, checkpoint=checkpoint_file, device_map="auto"
)

device_map="auto"로 설정하면 GPU의 가용 공간을 먼저 채우고, 이어서 CPU, 마지막으로 하드디스크까지 사용해요.

추론 실행하기

모델이 완전히 배치되면 추론을 할 수 있어요. 입력을 모델 파라미터와 같은 디바이스 유형으로 옮긴 뒤 실행하면 됩니다.

input = torch.randn(2, 3)
device_type = next(iter(model.parameters())).device.type
input = input.to(device_type)
output = model(input)

입력이 층을 통과할 때마다 CPU에서 GPU로(또는 디스크→CPU→GPU로) 전송되고, 출력이 계산된 뒤 그 층은 GPU에서 제거돼요. 오버헤드는 붙지만, 가장 큰 층 하나가 GPU에 들어가기만 하면 어떤 크기의 모델이든 실행할 수 있게 돼요.

여러 GPU를 써도 한 번에 하나의 GPU만 활성화돼요. 그래서 모델 병렬화에서 GPU는 이전 GPU가 출력을 보내줄 때까지 기다려야 해요. 이런 스크립트는 accelerate launchtorchrun 대신 일반 Python으로 실행하는 걸 권장해요.

Hugging Face 생태계에서 사용하기

Transformers나 Diffusers 같은 Hugging Face 라이브러리는 from_pretrained 생성자에서 Big Model Inference를 지원해요. device_map="auto"를 추가하기만 하면 돼요.

예를 들어 BigScience T0pp 110억 파라미터 모델을 이렇게 로드할 수 있어요.

from transformers import AutoModelForSeq2SeqLM

model = AutoModelForSeq2SeqLM.from_pretrained("bigscience/T0pp", device_map="auto")

torch_dtype을 지정하면 더 낮은 정밀도로 로드해 메모리를 아낄 수도 있어요.

from transformers import AutoModelForSeq2SeqLM

model = AutoModelForSeq2SeqLM.from_pretrained(
    "bigscience/T0pp", device_map="auto", torch_dtype=torch.float16
)

더 알아보기