BiGRU_T_version / src /bigru_t /multimodal /text_encoder.py
PowerMachine's picture
Upload folder using huggingface_hub
3275441 verified
Raw History Blame Contribute Delete
1.36 kB
"""
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"]