File size: 2,991 Bytes
abdb45a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
"""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()