V2: 12.97M params, HW optimizer, parallel BBPE, DPO, reasoning, inference, 10 bugs fixed
8594de8 verified Download src/bigru_t/inference/generator.py from PowerMachine/BiGRU_T_version: direct link, hf CLI and curl.
- Browser
- Download file 6.64 kB
-
https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/inference/generator.py
- Command line
-
hf download hf://PowerMachine/BiGRU_T_version/src/bigru_t/inference/generator.py
-
curl -L -o generator.py https://huggingface.co/PowerMachine/BiGRU_T_version/resolve/main/src/bigru_t/inference/generator.py
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 | |
| 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 | |
| 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"] | |