"""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"]