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