File size: 1,363 Bytes
3275441 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 | """
Xavante - text_encoder.py
Responsabilidade: Encoder de texto com embedding reconfiguravel + GRU hierarquica
+ multi-head attention (Teoremas 7.1, 11.1, 12.1).
"""
from __future__ import annotations
import logging
import torch
import torch.nn as nn
from ..model.attention_multimodal import MultiHeadAttention
from ..model.embedding_reconfig import ReconfigurableEmbedding
from ..model.gru_hierarchy import GRUHierarchy
logger = logging.getLogger(__name__)
class TextEncoder(nn.Module):
def __init__(
self,
vocab_size: int = 32128,
d_model: int = 512,
n_heads: int = 8,
n_gru_levels: int = 3,
max_seq_len: int = 2048,
):
super().__init__()
self.embedding = ReconfigurableEmbedding(vocab_size, d_model, padding_idx=0)
self.pos_emb = nn.Embedding(max_seq_len, d_model)
self.gru = GRUHierarchy(d_model, d_model, n_levels=n_gru_levels)
self.attn = MultiHeadAttention(d_model, n_heads)
self.ln = nn.LayerNorm(d_model)
def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
B, L = input_ids.shape
pos = torch.arange(L, device=input_ids.device).unsqueeze(0).expand(B, L)
x = self.embedding(input_ids) + self.pos_emb(pos)
x = self.gru(x)
x = self.attn(x)
return self.ln(x)
__all__ = ["TextEncoder"]
|