Spaces:
Sleeping
Sleeping
| """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 | |
| def load_retriever(): | |
| """Carrega o retriever (baixa o banco vetorial no 1º start). Compartilhado.""" | |
| from backend.retriever import Retriever | |
| return Retriever() | |