File size: 1,979 Bytes
7b4e05b
143710c
 
 
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
 
 
7b4e05b
143710c
 
 
 
 
7b4e05b
143710c
 
 
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
 
 
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
143710c
7b4e05b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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

"""Encoder implementations.

Note: ResNet50 itself is NOT run here during training -- features are

precomputed/cached by scripts/extract_features.py (see src/features/extractor.py).

This module only holds the small trainable projection layer that maps cached

CNN features (e.g. 2048-d) down to the shared embedding space (e.g. 256-d).

"""

from __future__ import annotations

import torch

import torch.nn as nn

from src.models.base import BaseEncoder

class ResNet50Encoder(BaseEncoder):

    """Linear projection of cached, frozen ResNet50 features -> embed_dim."""

    def __init__(self, feature_dim: int = 2048, embed_dim: int = 256, freeze: bool = True, **kwargs):

        super().__init__()

        self.freeze = freeze

        self.linear = nn.Linear(feature_dim, embed_dim)

        self.relu = nn.ReLU()

        self.dropout = nn.Dropout(0.5)

    def forward(self, image_features: torch.Tensor) -> torch.Tensor:

        x = self.linear(image_features)

        x = self.relu(x)

        x = self.dropout(x)

        return x

class ResNet50AttentionEncoder(BaseEncoder):

    """Projects a SPATIAL GRID of cached, frozen ResNet50 features -> embed_dim.

    Unlike ResNet50Encoder (which projects one pooled global vector), this

    projects each of the 49 spatial grid positions independently, preserving

    per-region information for an attention-based decoder to attend over.

    """

    def __init__(self, feature_dim: int = 2048, embed_dim: int = 256, freeze: bool = True, **kwargs):

        super().__init__()

        self.freeze = freeze

        self.linear = nn.Linear(feature_dim, embed_dim)

        self.relu = nn.ReLU()

        self.dropout = nn.Dropout(0.5)

    def forward(self, image_features: torch.Tensor) -> torch.Tensor:

        # image_features: (batch, num_pixels, feature_dim) -- num_pixels=49 for a 7x7 grid

        x = self.linear(image_features)

        x = self.relu(x)

        x = self.dropout(x)

        return x