File size: 5,194 Bytes
3346833
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
from .attention import *
from .embeddings import *



def create_mask_decoder(target,padding_idx=0):
    seq_len = target.size(1)
    pad_mask = (target != padding_idx).unsqueeze(1).unsqueeze(2)
    seq_mask = torch.tril(torch.ones(seq_len, seq_len)).to(target.device)
    trg_mask = pad_mask * seq_mask
    return trg_mask


class feedforward(nn.Module):
    def __init__(self,embed_dim=512,ff_hidden_dim=1024,dropout_rate = 0.1):
        super().__init__()
        self.mlp = nn.Sequential(
            nn.Linear(embed_dim, ff_hidden_dim),
            nn.GELU(),
            nn.Dropout(p=dropout_rate),
            nn.Linear(ff_hidden_dim, embed_dim),
            nn.GELU(),
            nn.Dropout(p=dropout_rate),
        )
    def forward(self, x):
        outputs = self.mlp(x)
        return outputs


class Encoder(nn.Module):
    def __init__(self, embed_dim=512, depth=4, heads=2, ff_hidden_dim=1024, dropout_rate = 0.1):
        super().__init__()
        self.layers = nn.ModuleList([])
        for _ in range(depth):
            self.layers.append(nn.ModuleList([
                MultiHeadSelfAttention(embed_dim,heads),
                nn.LayerNorm(embed_dim),
                feedforward(embed_dim,ff_hidden_dim, dropout_rate),
                nn.LayerNorm(embed_dim),
            ]))
    def forward(self, source,target=None,mask=None):
        x = source.clone()
        for attn,norm, ff,norm2 in self.layers:
            x = attn(Q=x,K=x,V=x,attn_mask = mask) + x
            x = norm(x)
            x = ff(x) + x
            x = norm(x)
        return x

class Decoder(nn.Module):
    def __init__(self, embed_dim=512, depth=4, heads=2, ff_hidden_dim=1024, dropout_rate = 0.1):
        super().__init__()
        self.layers = nn.ModuleList([])
        for _ in range(depth):
            self.layers.append(nn.ModuleList([
                MultiHeadSelfAttention(embed_dim,heads), # masked self attention
                nn.LayerNorm(embed_dim),
                MultiHeadSelfAttention(embed_dim,heads), # cross atention
                nn.LayerNorm(embed_dim),
                feedforward(embed_dim,ff_hidden_dim, dropout_rate), # feed forward
                nn.LayerNorm(embed_dim)
            ]))
    def forward(self,encoder_out,decoder_in,mask=None):
        x = decoder_in.clone()
        for self_attn,norm1,cross_attn,norm2, ff,norm3 in self.layers:
            x = self_attn(Q=x,K=x,V=x,attn_mask = mask) + x
            x = norm1(x)
            x = cross_attn(Q=x,V=encoder_out,K=encoder_out) + x
            x = norm2(x)
            x = ff(x) + x
            x = norm3(x)
        return x

class Transformer(nn.Module):
    def __init__(self,height=224,
                width=224,
                n_channels=3,
                patch_size=16,
                dim=512,
                encoder_head=2,
                encoder_feed_forward=1024,
                encoder_depth=2,
                decoder_head=2,
                decoder_feed_forward=1024,
                decoder_depth=1,
                targ_len = 10,
                vocab_size = 100,
                padding_idx = 0):

        super().__init__()
        self.image_embedding = patch_embedding(height,width,n_channels,patch_size,dim)
        self.decoder_embedding = nn.Embedding(vocab_size,dim, padding_idx=padding_idx)
        self.decoder_positional = PositionalEncoding(dim, targ_len)
        self.n_patchs = height*width//(patch_size**2)
        # Create a diagonal attention mask
        self.diag_attn_mask = ~torch.eye(self.n_patchs, dtype=torch.bool)
        self.encoder = Encoder(dim,encoder_depth, encoder_head,encoder_feed_forward)
        self.decoder = Decoder(dim,decoder_depth, decoder_head,decoder_feed_forward)
        self.padding_idx = padding_idx
        self.fc = nn.Linear(dim,vocab_size)

    def encode(self,image):
        device_ = image.device
        x = self.image_embedding(image)
        out = self.encoder(x,x,mask=self.diag_attn_mask.to(device_))
        return out


    def decode(self,encoder_out,text):
        device_ = encoder_out.device
        decoder_mask = create_mask_decoder(text,padding_idx=self.padding_idx)
        y = self.decoder_positional(self.decoder_embedding(text))
        out = self.decoder(encoder_out,y,mask=decoder_mask.to(device_))

        return self.fc(out)
        
    def forward(self, image, text):
        encoder_out = self.encode(image)
        outputs = self.decode(encoder_out,text)
        return outputs


@torch.no_grad()
def greedy_decoding(model, image, max_len, start_idx,end_idx):
    model.eval()
    decoder_input = torch.ones(1, 1, dtype=torch.long) * start_idx
    decoded_outputs = decoder_input
    decoded_outputs = decoded_outputs.to(image.device)
    encoder_out = model.encode(image)

    for i in range(max_len):
        decoder_output = model.decode(encoder_out,decoded_outputs)  # Decode next token
        next_word = decoder_output[:, -1].argmax(dim=-1).unsqueeze(-1)  # Get most probable word index
        if next_word==end_idx:
            break
        decoded_outputs = torch.cat([decoded_outputs, next_word], dim=-1)  # Append next word to output sequence
        
    return decoded_outputs[:, 1:]