"""Carga do modelo + geração. O loop de decode manual (lens) chega no Bloco 2.""" import os import threading import torch from dotenv import load_dotenv from transformers import AutoTokenizer, TextIteratorStreamer from espelho import config def pick_device() -> str: # ZeroGPU: .to('cuda') no nível do módulo é o padrão exigido (emulação # CUDA fora de @spaces.GPU; GPU real dentro). if os.environ.get("SPACES_ZERO_GPU"): return "cuda" return "mps" if torch.backends.mps.is_available() else "cpu" def on_space() -> bool: return bool(os.environ.get("SPACE_ID")) def available_models() -> list[str]: """Candidatos com snapshot completo no cache local (roda offline). No ZeroGPU, modelos só podem ser carregados no boot (a emulação CUDA não intercepta .to('cuda') em runtime), então o Space carrega seu conjunto no startup (ver space/app.py) e o seletor troca entre os já carregados. """ from huggingface_hub import snapshot_download out = [] for repo in config.CANDIDATE_MODELS: try: snapshot_download( repo, local_files_only=True, ignore_patterns=["*.gguf", "original/*", "*.pth"], ) out.append(repo) except Exception: pass # incompleto ou ausente: fora do seletor return out def load_model(model_id: str | None = None): """Carrega tokenizer e modelo (text-only) em bf16 no MPS. Retorna (tokenizer, model). Para o Gemma 3 (checkpoint multimodal) usa a classe text-only Gemma3ForCausalLM; para o Gemma 2, Gemma2ForCausalLM com atenção eager (recomendação upstream por causa do softcapping). """ load_dotenv() model_id = model_id or config.ACTIVE_MODEL token = os.environ.get("HF_TOKEN") device = pick_device() tokenizer = AutoTokenizer.from_pretrained(model_id, token=token) # bf16 no MPS/GPU; em CPU (ex.: Space free tier) float32 é mais rápido. dtype = torch.bfloat16 if device != "cpu" else torch.float32 kwargs: dict = {"dtype": dtype, "token": token} if "gemma-3" in model_id: # Classe multimodal: par nativo do checkpoint 4b-it. A text-only # (Gemma3ForCausalLM) deixava as linhas de padding do vocab como # MISSING -> init aleatório -> ids órfãos amostrados em respostas # longas (ver DECISIONS.md). Usamos só o caminho de texto dela. from transformers import Gemma3ForConditionalGeneration as cls elif "gemma-2" in model_id: from transformers import Gemma2ForCausalLM as cls kwargs["attn_implementation"] = "eager" else: from transformers import AutoModelForCausalLM as cls model = cls.from_pretrained(model_id, **kwargs) model.to(device) model.eval() return tokenizer, model def text_submodel(model): """Submódulo decoder de texto (tem .norm e .layers), para carga text-only (model.model) ou multimodal (model.model.language_model).""" inner = getattr(model, "model", model) return getattr(inner, "language_model", inner) def text_config(model): """Config do caminho de texto (num_hidden_layers etc.), em ambas as cargas.""" cfg = model.config return getattr(cfg, "text_config", cfg) def build_inputs(tokenizer, messages: list[dict], device: str) -> torch.Tensor: """Aplica o chat template e retorna input_ids no device.""" # enable_thinking=False: desliga o bloco do Qwen3; variável extra # é ignorada por templates que não a usam (Gemma). out = tokenizer.apply_chat_template( messages, add_generation_prompt=True, return_tensors="pt", enable_thinking=False, ) # transformers 5 retorna BatchEncoding; versões antigas, o tensor direto. input_ids = out if isinstance(out, torch.Tensor) else out["input_ids"] return input_ids.to(device) def _sample(logits: torch.Tensor, temperature: float | None, top_k: int = 50, top_p: float = 1.0) -> int: """Próximo token a partir dos logits da última posição [1, V]. Aplica top-k/top-p como o generate oficial: sem truncar a cauda, a amostragem sorteia tokens raros de outros idiomas no meio do texto (vazamento multilíngue; ver DECISIONS.md). """ if temperature is None: return int(torch.argmax(logits, dim=-1)) logits = logits.float() / temperature k = min(top_k or logits.shape[-1], logits.shape[-1]) vals, idx = torch.topk(logits, k, dim=-1) # desc probs = torch.softmax(vals, dim=-1) if top_p is not None and top_p < 1.0: cum = probs.cumsum(dim=-1) # mantém o menor prefixo cuja massa acumulada atinge top_p probs = probs.masked_fill(cum - probs > top_p, 0.0) probs = probs / probs.sum(dim=-1, keepdim=True) choice = torch.multinomial(probs, num_samples=1) return int(idx.gather(-1, choice)) def generate_with_trace(tokenizer, model, messages: list[dict], max_new_tokens: int, temperature: float | None = None, seed: int | None = None, on_text=None): """Decode manual token a token com past_key_values, capturando os hidden states de todas as camadas a cada passo (prefill e decode). Retorna (texto, TurnTrace). `on_text(pedaço)` — se fornecido — recebe o texto incremental para streaming na UI. """ import time from espelho.lens import LogitLens, TurnTrace if seed is not None: torch.manual_seed(seed) lens = LogitLens(model, tokenizer) input_ids = build_inputs(tokenizer, messages, model.device.type) eos_ids = model.generation_config.eos_token_id if eos_ids is None: eos_ids = tokenizer.eos_token_id eos_ids = set(eos_ids) if isinstance(eos_ids, (list, tuple)) else {eos_ids} # Truncamento da amostragem igual ao generate oficial: usa o que o # checkpoint recomenda (4B: top_k=64, top_p=0.95) ou os defaults do # transformers (top_k=50, top_p=1.0) quando o config não define. top_k = model.generation_config.top_k or 50 top_p = model.generation_config.top_p or 1.0 # hidden por passo: {layer: vetor [d]} só das camadas da banda, # sempre na última posição da sequência daquele passo. steps: list[dict[int, torch.Tensor]] = [] def grab(hidden_states) -> None: steps.append({l: hidden_states[l][0, -1].detach() for l in lens.band}) t0 = time.perf_counter() with torch.no_grad(): out = model(input_ids=input_ids, use_cache=True, output_hidden_states=True) grab(out.hidden_states) t_prefill = time.perf_counter() - t0 generated: list[int] = [] decoded_prev = "" t1 = time.perf_counter() with torch.no_grad(): for _ in range(max_new_tokens): next_id = _sample(out.logits[:, -1], temperature, top_k, top_p) generated.append(next_id) if on_text is not None: decoded = tokenizer.decode(generated, skip_special_tokens=True) if len(decoded) > len(decoded_prev) and not decoded.endswith("�"): on_text(decoded[len(decoded_prev):]) decoded_prev = decoded if next_id in eos_ids: break out = model( input_ids=torch.tensor([[next_id]], device=model.device), past_key_values=out.past_key_values, use_cache=True, output_hidden_states=True, ) grab(out.hidden_states) t_decode = time.perf_counter() - t1 response = tokenizer.decode(generated, skip_special_tokens=True) if on_text is not None and len(response) > len(decoded_prev): on_text(response[len(decoded_prev):]) # Lens em lote: por camada da banda, empilha as posições e projeta. t2 = time.perf_counter() topk = { layer: lens.topk_batch(torch.stack([s[layer] for s in steps])) for layer in lens.band } t_lens = time.perf_counter() - t2 # Rótulo das posições: última do prompt + cada token gerado (exceto o # último, cujo hidden não é capturado quando é EOS/fim do loop). tokens = [""] + tokenizer.convert_ids_to_tokens(generated)[: len(steps) - 1] n = max(len(generated), 1) trace = TurnTrace( prompt=messages[-1]["content"], response=response, tokens=tokens, band_layers=list(lens.band), topk=topk, num_layers=lens.num_layers, timings={ "prefill_s": round(t_prefill, 3), "decode_s": round(t_decode, 3), "lens_s": round(t_lens, 3), "new_tokens": len(generated), "tok_s": round(n / t_decode, 2) if t_decode > 0 else None, }, ) return response, trace def stream_generate(tokenizer, model, messages: list[dict], max_new_tokens: int, temperature: float | None = None, seed: int | None = None): """Gera resposta em streaming; itera pedaços de texto.""" if seed is not None: torch.manual_seed(seed) input_ids = build_inputs(tokenizer, messages, model.device.type) streamer = TextIteratorStreamer( tokenizer, skip_prompt=True, skip_special_tokens=True ) gen_kwargs: dict = { "input_ids": input_ids, "max_new_tokens": max_new_tokens, "streamer": streamer, } if temperature is None: gen_kwargs["do_sample"] = False else: gen_kwargs["do_sample"] = True gen_kwargs["temperature"] = temperature thread = threading.Thread( target=model.generate, kwargs=gen_kwargs, daemon=True ) thread.start() yield from streamer thread.join()