File size: 910 Bytes
3f6e26d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch.nn as nn

class ConnectorTransformer(nn.Module):
    def __init__(self, input_dim, output_dim, depth=2, num_heads=4, mlp_ratio=4.0):
        super().__init__()
        if input_dim != output_dim:
            self.input_proj = nn.Linear(input_dim, output_dim)
        else:
            self.input_proj = nn.Identity()
            
        hidden_size = output_dim
        
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=hidden_size, 
            nhead=num_heads, 
            dim_feedforward=int(hidden_size * mlp_ratio), 
            activation="gelu",
            batch_first=True,
            norm_first=True
        )
        self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=depth)
        self.norm = nn.LayerNorm(hidden_size)

    def forward(self, x):
        x = self.input_proj(x)
        x = self.encoder(x)
        x = self.norm(x)
        return x