File size: 1,858 Bytes
143710c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7b4e05b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
143710c
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
"""CaptionModel: wires encoder + decoder together.

This class never needs to change when you swap architectures -- it only
relies on the BaseEncoder/BaseDecoder contract.
"""
from __future__ import annotations

import torch
import torch.nn as nn

from src.models.base import BaseDecoder, BaseEncoder
from src.models.registry import build_decoder, build_encoder


class CaptionModel(nn.Module):
    def __init__(self, encoder: BaseEncoder, decoder: BaseDecoder):
        super().__init__()
        self.encoder = encoder
        self.decoder = decoder

    @classmethod
    def from_config(cls, config: dict, vocab_size: int) -> "CaptionModel":
        encoder = build_encoder(config)
        decoder = build_decoder(config, vocab_size)
        return cls(encoder, decoder)

    def forward(self, image_features: torch.Tensor, input_seq: torch.Tensor) -> torch.Tensor:
        image_embed = self.encoder(image_features)
        return self.decoder.forward(image_embed, input_seq)

    @torch.no_grad()
    def generate(
        self,
        image_features: torch.Tensor,
        start_idx: int,
        end_idx: int,
        max_len: int,
    ) -> list[int]:
        self.eval()
        image_embed = self.encoder(image_features)
        return self.decoder.generate(image_embed, start_idx, end_idx, max_len)
    
    #for beam search
    @torch.no_grad()
    def generate(
        self,
        image_features: torch.Tensor,
        start_idx: int,
        end_idx: int,
        max_len: int,
        decoding: str = "greedy",
        beam_width: int = 3,
    ) -> list[int]:
        self.eval()
        image_embed = self.encoder(image_features)
        if decoding == "beam":
            return self.decoder.generate_beam(image_embed, start_idx, end_idx, max_len, beam_width)
        return self.decoder.generate(image_embed, start_idx, end_idx, max_len)