gemma-search-doc / src /backend /model_loader.py
Felipe Albernaz
Assistente RAG do setor de energia (gemma-search + gemma-naive-rag)
abdb45a
Raw
History Blame Contribute Delete
2.99 kB
"""Carregamento do LLM (com troca/descarregamento) e do retriever (cacheado).
Só UM LLM fica carregado por vez (os dois Gemma não cabem juntos na memória).
Ao trocar de modelo na UI, o anterior é descarregado e o novo é carregado.
O retriever (embedding + banco) é compartilhado e cacheado.
"""
import gc
import logging
import streamlit as st
import torch
from transformers import AutoModelForCausalLM, AutoProcessor, AutoTokenizer
import config
logger = logging.getLogger(__name__)
def _load_tokenizer(model_id: str, revision: str):
"""Tokenizer do Gemma (com parse_response/apply_chat_template).
Tenta o AutoProcessor (multimodal); se faltar a dependência de visão, cai para
o AutoTokenizer — suficiente, pois o app é só texto.
"""
kwargs = dict(revision=revision, trust_remote_code=True, token=config.HF_TOKEN)
try:
return AutoProcessor.from_pretrained(model_id, **kwargs).tokenizer
except Exception as exc: # ex.: Gemma4Processor sem torchvision/pillow
logger.warning("AutoProcessor falhou (%s); usando AutoTokenizer.", exc)
return AutoTokenizer.from_pretrained(model_id, **kwargs)
def _load_llm(model_id: str, revision: str):
device = "cuda" if torch.cuda.is_available() else "cpu"
logger.info("Carregando LLM %s @ %s em %s", model_id, revision, device)
kwargs = dict(revision=revision, trust_remote_code=True, dtype=torch.bfloat16)
if config.HF_TOKEN:
kwargs["token"] = config.HF_TOKEN
if device == "cuda":
kwargs["device_map"] = "cuda"
else:
kwargs["device_map"] = "cpu"
kwargs["low_cpu_mem_usage"] = True
model = AutoModelForCausalLM.from_pretrained(model_id, **kwargs)
tokenizer = _load_tokenizer(model_id, revision)
return model, tokenizer
def _unload_current_llm():
"""Descarrega o LLM atual da sessão e libera memória."""
for key in ("llm_model", "llm_tokenizer"):
st.session_state.pop(key, None)
st.session_state.pop("llm_key", None)
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
def ensure_llm(model_key: str):
"""Garante que SÓ o LLM `model_key` está carregado (descarrega o outro se houver).
Mantém o modelo em st.session_state entre reruns; só recarrega quando troca.
"""
if st.session_state.get("llm_key") == model_key and st.session_state.get("llm_model") is not None:
return st.session_state.llm_model, st.session_state.llm_tokenizer
_unload_current_llm()
cfg = config.MODELS[model_key]
model, tokenizer = _load_llm(cfg["model_id"], cfg["revision"])
st.session_state.llm_model = model
st.session_state.llm_tokenizer = tokenizer
st.session_state.llm_key = model_key
return model, tokenizer
@st.cache_resource(show_spinner=False)
def load_retriever():
"""Carrega o retriever (baixa o banco vetorial no 1º start). Compartilhado."""
from backend.retriever import Retriever
return Retriever()