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