Brenow's picture
Upload folder using huggingface_hub
4397340 verified
Raw
History Blame Contribute Delete
9.72 kB
"""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()