Gemma-3 멀티모달 모델의 GPU 메모리 최적화: PyTorch 캐시 동적 할당 및 해제 전략

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 메모리 사용량은 크게 세 가지 계층으로 나눌 수 있습니다.

  1. 모델 가중치 (Model Weights): 모델 로드 시 할당되는 고정 비용입니다. Gemma-3 12B의 경우 약 24GB를 점유하며, 프로세스가 종료될 때까지 유지됩니다.
  2. 추론 캐시 (Inference Cache): 텍스트 생성 속도를 높이기 위한 KV 캐시입니다. 대화 턴이 늘어날수록 선형적으로 증가하며, Flash Attention 사용 시 추가적인 버퍼가 필요합니다.
  3. 임시 텐서 및 단편화 (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가 유휴 상태인 구간에서 수행되도록 설계하는 것이 안정적인 멀티모달 서비스 구축의 핵심입니다.

태그: PyTorch Gemma-3 VRAM CUDA Streamlit

8월 3일 13:06에 게시됨