Gemma-3 12B 모델과 VRAM 단편화 문제
Gemma-3 12B와 같은 대규모 멀티모달 AI 모델을 운용할 때, 초기 추론 속도는 빠르지만 대화가 길어지거나 이미지를 여러 장 처리하다 보면 시스템이 지연되거나 CUDA out of memory (OOM) 오류가 발생하는 경우가 많습니다. 이는 하드웨어의 성능 부족보다는 VRAM(비디오 RAM) 관리 및 메모리 풀링 메커니즘에서 기인하는 문제입니다.
BF16 정밀도 기준으로 Gemma-3 12B 모델의 가중치만 로드해도 약 24GB의 VRAM이 필요합니다. 여기에 이미지 전처리 텐서, 다중 턴 대화에 따른 Key-Value(KV) 캐시, 그리고 추론 중 생성되는 임시 활성화 값들이 누적되면 메모리 공간은 급격히 고갈됩니다. PyTorch의 torch.cuda.empty_cache() 함수를 활용하여 이러한 단편화를 동적으로 해결하는 전략을 살펴봅니다.
VRAM 소비의 3가지 계층 구조
GPU 메모리 사용량은 크게 세 가지 계층으로 나눌 수 있습니다.
- 모델 가중치 (Model Weights): 모델 로드 시 할당되는 고정 비용입니다. Gemma-3 12B의 경우 약 24GB를 점유하며, 프로세스가 종료될 때까지 유지됩니다.
- 추론 캐시 (Inference Cache): 텍스트 생성 속도를 높이기 위한 KV 캐시입니다. 대화 턴이 늘어날수록 선형적으로 증가하며, Flash Attention 사용 시 추가적인 버퍼가 필요합니다.
- 임시 텐서 및 단편화 (Temporary Tensors & Fragmentation): 이미지 임베딩, 중간 히든 스테이트, 그래디언트 버퍼 등이 포함됩니다. PyTorch는 메모리 할당 속도를 높이기 위해 메모리 풀(Pool)을 사용하는데, 변수를 삭제(
del)하더라도 OS로 반환되지 않고 풀에 남아 단편화를 유발합니다.
다음 코드는 다중 턴 추론 환경에서 VRAM 단편화가 어떻게 발생하는지 시뮬레이션합니다.
import torch
import gc
def analyze_vram_fragmentation():
"""Gemma-3 추론 과정에서의 VRAM 단편화 시뮬레이션"""
# 1. 모델 가중치 로드 (약 24GB 가정)
base_weights = torch.empty(6_000_000_000, device='cuda', dtype=torch.bfloat16)
print(f"가중치 로드 후 할당된 VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f} GB")
# 2. 다중 턴 대화 및 이미지 처리 시뮬레이션
for step in range(10):
# 이미지 임베딩 텐서
vision_embeds = torch.randn(1, 3, 512, 512, device='cuda')
# 어텐션 히든 스테이트
hidden_activations = torch.randn(1, 1024, 4096, device='cuda')
# 누적되는 KV 캐시
kv_states = torch.randn(32, step + 1, 16, 128, device='cuda')
# 연산 수행
_ = vision_embeds.sum()
_ = hidden_activations.mean()
print(f"턴 {step+1} 종료 후 할당된 VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f} GB")
# 3. 참조 해제 및 가비지 컬렉션
del vision_embeds, hidden_activations, kv_states
gc.collect()
print(f"변수 삭제 후 할당된 VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f} GB")
# 4. 캐시 비우기
torch.cuda.empty_cache()
print(f"캐시 초기화 후 할당된 VRAM: {torch.cuda.memory_allocated() / 1024**3:.2f} GB")
empty_cache()의 전략적 호출 시점
torch.cuda.empty_cache()는 PyTorch 메모리 풀에 남아 있는 미사용 캐시를 OS로 반환합니다. 하지만 이 함수는 CUDA 동기화를 유발하고 오버헤드가 있기 때문에, 매 추론 스텝마다 호출하는 것은 성능을 크게 저하시킵니다. 따라서 자연스러운 중단점(Breakpoint)에서 호출하는 것이 핵심입니다.
세션 및 컨텍스트 초기화
사용자가 대화 기록을 지우거나 새로운 이미지를 업로드할 때는 이전 컨텍스트가 더 이상 필요 없는 시점이므로 캐시를 반환하기 가장 좋습니다.
import streamlit as st
import torch
import gc
class SessionController:
@staticmethod
def clear_history():
"""대화 기록 초기화 및 VRAM 캐시 반환"""
if 'chat_logs' in st.session_state:
st.session_state.chat_logs = []
if 'vision_data' in st.session_state:
del st.session_state.vision_data
gc.collect()
torch.cuda.empty_cache()
st.toast("대화 및 VRAM 캐시가 초기화되었습니다.")
@staticmethod
def update_vision_input(file_obj):
"""새로운 이미지 업로드 시 이전 텐서 메모리 해제"""
if 'vision_data' in st.session_state:
del st.session_state.vision_data
torch.cuda.empty_cache()
# 이미지 전처리 로직 (생략)
# processed_tensor = preprocess(file_obj)
# st.session_state.vision_data = processed_tensor
스마트 VRAM 최적화 매니저
고정된 턴 수마다 캐시를 비우는 대신, 실제 메모리 단편화 정도를 모니터링하여 임계치를 초과할 때만 동적으로 해제하는 방식이 훨씬 효율적입니다.
class VRAMOptimizer:
def __init__(self, frag_threshold_gb=1.5, interval=5):
self.threshold_bytes = frag_threshold_gb * (1024 ** 3)
self.interval = interval
self.step_counter = 0
def monitor_and_optimize(self):
self.step_counter += 1
if self.step_counter % self.interval == 0:
allocated = torch.cuda.memory_allocated()
reserved = torch.cuda.memory_reserved()
fragmentation = reserved - allocated
if fragmentation > self.threshold_bytes:
print(f"단편화 감지: {fragmentation / 1024**3:.2f} GB. 캐시 최적화 실행.")
torch.cuda.empty_cache()
return True
return False
Streamlit 기반 멀티모달 애플리케이션 통합
앞서 구현한 최적화 로직을 Gemma-3 모델을 사용하는 Streamlit 웹 애플리케이션에 통합해 보겠습니다. 추론 엔진과 세션 관리를 분리하여 메모리 누수를 방지합니다.
import streamlit as st
import torch
from transformers import AutoModelForCausalLM, AutoProcessor
class MultimodalEngine:
def __init__(self):
self.model = None
self.processor = None
self.optimizer = VRAMOptimizer(frag_threshold_gb=2.0, interval=3)
self.turns = 0
def load_assets(self):
torch.cuda.empty_cache()
self.model = AutoModelForCausalLM.from_pretrained(
"google/gemma-3-12b-it",
torch_dtype=torch.bfloat16,
device_map="auto"
)
self.processor = AutoProcessor.from_pretrained("google/gemma-3-12b-it")
def infer(self, text_prompt, image_tensor=None):
self.turns += 1
if self.optimizer.monitor_and_optimize():
st.toast("VRAM 단편화 최적화 완료", icon="🛠️")
inputs = self.processor(
text=text_prompt,
images=image_tensor,
return_tensors="pt"
).to(self.model.device)
with torch.no_grad():
generated_ids = self.model.generate(**inputs, max_new_tokens=256)
decoded_text = self.processor.decode(generated_ids[0], skip_special_tokens=True)
# 추론 후 임시 변수 즉시 해제
del inputs, generated_ids
return decoded_text
def run_app():
st.set_page_config(page_title="Gemma-3 Studio", layout="wide")
if 'engine' not in st.session_state:
st.session_state.engine = MultimodalEngine()
st.session_state.engine.load_assets()
engine = st.session_state.engine
prompt = st.chat_input("메시지를 입력하세요...")
if prompt:
response = engine.infer(prompt, st.session_state.get('vision_data'))
st.chat_message("assistant").write(response)
실시간 VRAM 텔레메트리 대시보드
메모리 상태를 시각화하면 단편화 시점을 정확히 파악하고 디버깅하는 데 큰 도움이 됩니다. Altair를 사용하여 실시간으로 할당된 메모리와 단편화된 메모리를 추적하는 사이드바 대시보드를 구현합니다.
import pandas as pd
import altair as alt
import time
class TelemetryPanel:
def __init__(self):
self.logs = []
def record(self):
alloc = torch.cuda.memory_allocated() / 1024**3
res = torch.cuda.memory_reserved() / 1024**3
self.logs.append({
"time": time.time(),
"allocated": alloc,
"fragmented": res - alloc
})
if len(self.logs) > 50:
self.logs = self.logs[-50:]
def render(self):
st.sidebar.subheader("📊 VRAM 텔레메트리")
if not self.logs:
return
df = pd.DataFrame(self.logs)
chart = alt.Chart(df).transform_fold(
['allocated', 'fragmented'],
as_=['metric', 'value']
).mark_line().encode(
x='time:T',
y='value:Q',
color='metric:N'
).properties(height=200)
st.sidebar.altair_chart(chart, use_container_width=True)
# 단편화가 심할 경우 수동 제어권 제공
if df['fragmented'].iloc[-1] > 2.0:
if st.sidebar.button("수동 캐시 비우기"):
torch.cuda.empty_cache()
st.sidebar.success("캐시가 성공적으로 반환되었습니다.")
주의할 점은 empty_cache()를 호출한 직후에 거대한 텐서를 다시 할당해야 하는 상황이라면, 메모리 풀을 거치지 않고 CUDA 드라이버 레벨에서 새로 할당을 요청해야 하므로 오버헤드가 발생할 수 있다는 것입니다. 따라서 캐시 초기화는 사용자가 타이핑을 하거나 이미지를 선택하는 등 GPU가 유휴 상태인 구간에서 수행되도록 설계하는 것이 안정적인 멀티모달 서비스 구축의 핵심입니다.