PowerMachine's picture
V2: 12.97M params, HW optimizer, parallel BBPE, DPO, reasoning, inference, 10 bugs fixed
8594de8 verified
Raw History Blame Contribute Delete
6.64 kB
"""generator.py — Inferência autoregressiva para BiGRU_T_version.
Implementa geração greedy + top-k sampling + streaming, reaproveitando a
filosofia do GRURingV139.generate() da fonte (xavante_work/flexnet/gru_ring_v13_9.py).
O UnifiedModel produz (batch, vocab_size) — 1 logit por forward. Para geração
autoregressiva, alimentamos a sequência crescente e usamos o último logit.
"""
from __future__ import annotations
import logging
import math
from typing import Iterator, List, Optional, Tuple
import torch
import torch.nn.functional as F
logger = logging.getLogger(__name__)
class BiGRUTGenerator:
"""Gerador autoregressivo para UnifiedModel.
Suporta:
- greedy decoding
- top-k sampling com temperatura
- repetition penalty
- max_new_tokens + EOS stop
- streaming (yield token a token)
Args:
model: UnifiedModel treinado
eos_token_id: ID do token de fim (para parar)
pad_token_id: ID do padding
max_seq_len: comprimento máximo de contexto (trunca à esquerda)
"""
def __init__(
self,
model,
eos_token_id: int = 2,
pad_token_id: int = 1,
max_seq_len: int = 64,
):
self.model = model
self.eos_token_id = eos_token_id
self.pad_token_id = pad_token_id
self.max_seq_len = max_seq_len
self.device = next(model.parameters()).device
@torch.no_grad()
def generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int = 32,
temperature: float = 1.0,
top_k: int = 0,
repetition_penalty: float = 1.0,
do_sample: bool = False,
) -> torch.Tensor:
"""Gera tokens autoregressivamente.
Args:
input_ids: (batch, T) tokens de prompt
max_new_tokens: máx. tokens a gerar
temperature: temperatura do sampling (1.0 = sem escala)
top_k: se > 0, amostra apenas dos top-k tokens
repetition_penalty: penaliza tokens já gerados (1.0 = sem pena)
do_sample: se False, greedy decoding
Returns:
generated_ids: (batch, T + max_new_tokens)
"""
self.model.eval()
batch_size = input_ids.size(0)
generated = input_ids.clone().to(self.device)
for _step in range(max_new_tokens):
# Trunca à esquerda se exceder max_seq_len
if generated.size(1) > self.max_seq_len:
context = generated[:, -self.max_seq_len:]
else:
context = generated
# Forward (sem hipótese — inferência usa só TrainT)
out = self.model(context, temperature=1.0, use_hypothesis=False)
y_hat = out[0] if isinstance(out, tuple) else out # (batch, vocab)
# Último logit
logits = y_hat # já é (batch, vocab) — 1 logit por forward
# Repetition penalty
if repetition_penalty != 1.0:
for b in range(batch_size):
for prev_token in generated[b].tolist():
if logits[b, prev_token] > 0:
logits[b, prev_token] /= repetition_penalty
else:
logits[b, prev_token] *= repetition_penalty
if do_sample and temperature > 0:
# Top-k sampling
if top_k > 0:
top_k = min(top_k, logits.size(-1))
values, _ = torch.topk(logits, top_k, dim=-1)
min_val = values[:, -1:].unsqueeze(-1)
logits = torch.where(
logits < min_val,
torch.full_like(logits, float("-inf")),
logits,
)
# Temperatura
logits = logits / max(temperature, 1e-8)
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
else:
# Greedy
next_token = logits.argmax(dim=-1, keepdim=True)
# Concatena
generated = torch.cat([generated, next_token], dim=1)
# Para se todos geraram EOS
if (next_token == self.eos_token_id).all():
break
return generated
@torch.no_grad()
def stream_generate(
self,
input_ids: torch.Tensor,
max_new_tokens: int = 32,
temperature: float = 1.0,
top_k: int = 0,
repetition_penalty: float = 1.0,
do_sample: bool = False,
) -> Iterator[torch.Tensor]:
"""Geração streaming — yields um token por vez.
Args: mesmos de generate()
Yields:
next_token: (batch, 1) tensor a cada iteração
"""
self.model.eval()
batch_size = input_ids.size(0)
generated = input_ids.clone().to(self.device)
for _step in range(max_new_tokens):
if generated.size(1) > self.max_seq_len:
context = generated[:, -self.max_seq_len:]
else:
context = generated
out = self.model(context, temperature=1.0, use_hypothesis=False)
y_hat = out[0] if isinstance(out, tuple) else out
logits = y_hat
if repetition_penalty != 1.0:
for b in range(batch_size):
for prev_token in generated[b].tolist():
if logits[b, prev_token] > 0:
logits[b, prev_token] /= repetition_penalty
else:
logits[b, prev_token] *= repetition_penalty
if do_sample and temperature > 0:
if top_k > 0:
top_k = min(top_k, logits.size(-1))
values, _ = torch.topk(logits, top_k, dim=-1)
min_val = values[:, -1:].unsqueeze(-1)
logits = torch.where(
logits < min_val,
torch.full_like(logits, float("-inf")),
logits,
)
logits = logits / max(temperature, 1e-8)
probs = F.softmax(logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
else:
next_token = logits.argmax(dim=-1, keepdim=True)
generated = torch.cat([generated, next_token], dim=1)
yield next_token
if (next_token == self.eos_token_id).all():
break
__all__ = ["BiGRUTGenerator"]