"""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()