Spaces:
Sleeping
Sleeping
| """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 <think> 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 = ["<fim do prompt>"] + 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() | |