File size: 6,328 Bytes
961cf0c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
import torch
import torch.nn as nn
import timm
from torchvision import models
from transformers import AutoModel
from tab_transformer import TabTransformer


class loadModels:

    # ======================================================
    # Utilitário: controle de fine-tuning do backbone
    # ======================================================
    @staticmethod
    def set_backbone_train_mode(model, mode="frozen_weights", last_n_layers=1):
        if mode == "frozen_weights":
            for p in model.parameters():
                p.requires_grad = False

        elif mode == "unfrozen_weights":
            for p in model.parameters():
                p.requires_grad = True

        elif mode == "last_layer_unfrozen_weights":
            for p in model.parameters():
                p.requires_grad = False

            children = list(model.children())
            for layer in children[-last_n_layers:]:
                for p in layer.parameters():
                    p.requires_grad = True
        else:
            raise ValueError(f"Invalid backbone_train_mode: {mode}")

    # ======================================================
    # Image encoders
    # ======================================================
    @staticmethod
    def loadModelImageEncoder(
        cnn_model_name: str,
        backbone_train_mode: str = "frozen_weights"
    ):

        # ------------------------------
        # TorchVision CNNs
        # ------------------------------
        if cnn_model_name == "resnet-18":
            model = models.resnet18(pretrained=True)
            cnn_dim = 512
            model.fc = nn.Identity()

            loadModels.set_backbone_train_mode(
                model, backbone_train_mode, last_n_layers=1
            )

        elif cnn_model_name == "resnet-50":
            model = models.resnet50(pretrained=True)
            cnn_dim = 2048
            model.fc = nn.Identity()

            loadModels.set_backbone_train_mode(
                model, backbone_train_mode, last_n_layers=1
            )

        elif cnn_model_name == "densenet169":
            model = models.densenet169(pretrained=True)
            cnn_dim = 1664
            model.classifier = nn.Identity()

            if backbone_train_mode == "last_layer_unfrozen_weights":
                for p in model.parameters():
                    p.requires_grad = False
                for p in model.features.denseblock4.parameters():
                    p.requires_grad = True
            else:
                loadModels.set_backbone_train_mode(model, backbone_train_mode)

        elif cnn_model_name == "mobilenet-v2":
            model = models.mobilenet_v2(pretrained=True)
            cnn_dim = 1280
            model.classifier = nn.Identity()

            loadModels.set_backbone_train_mode(
                model, backbone_train_mode, last_n_layers=1
            )

        elif cnn_model_name == "efficientnet-b0":
            model = models.efficientnet_b0(pretrained=True)
            cnn_dim = 1280
            model.classifier = nn.Identity()

            loadModels.set_backbone_train_mode(
                model, backbone_train_mode, last_n_layers=1
        )

        elif cnn_model_name == "efficientnet-b4":
            model = models.efficientnet_b4(pretrained=True)
            cnn_dim = 1792
            model.classifier = nn.Identity()

            loadModels.set_backbone_train_mode(
                model, backbone_train_mode, last_n_layers=1
        )

        elif cnn_model_name == "efficientnet-b7":
            model = models.efficientnet_b7(pretrained=True)
            cnn_dim = 2560
            model.classifier = nn.Identity()

            loadModels.set_backbone_train_mode(
                model, backbone_train_mode, last_n_layers=1
            )

        # ------------------------------
        # timm models (ViT / Hybrid)
        # ------------------------------
        elif cnn_model_name in timm.list_models(pretrained=True):
            model = timm.create_model(cnn_model_name, pretrained=True)
            model.reset_classifier(0)
            cnn_dim = model.num_features

            if backbone_train_mode == "last_layer_unfrozen_weights":
                for p in model.parameters():
                    p.requires_grad = False

                # estratégia genérica: último estágio
                if hasattr(model, "stages"):
                    for p in model.stages[-1].parameters():
                        p.requires_grad = True
                elif hasattr(model, "blocks"):
                    for p in model.blocks[-1].parameters():
                        p.requires_grad = True
            else:
                loadModels.set_backbone_train_mode(model, backbone_train_mode)

        else:
            raise ValueError(f"Backbone '{cnn_model_name}' não implementado.")

        return model, cnn_dim

    # ======================================================
    # Text encoders
    # ======================================================
    @staticmethod
    def loadTextModelEncoder(
        text_model_encoder: str,
        train_mode: str = "frozen_weights"
    ):

        # ------------------------------
        # HuggingFace Transformers
        # ------------------------------
        if text_model_encoder in ["bert-base-uncased", "gpt2"]:
            model = AutoModel.from_pretrained(text_model_encoder)
            output_dim = model.config.hidden_size

            if train_mode == "unfrozen_weights":
                for p in model.parameters():
                    p.requires_grad = True
            else:
                for p in model.parameters():
                    p.requires_grad = False

            return model, output_dim, output_dim

        # ------------------------------
        # TabTransformer
        # ------------------------------
        elif text_model_encoder == "tab-transformer":
            categorical_indices = list(range(82))
            output_dim = 85

            model = TabTransformer(
                categorical_cardinalities=categorical_indices,
                num_continuous=4,
                output_dim=output_dim
            )

            return model, output_dim, output_dim

        else:
            raise ValueError(f"Text encoder '{text_model_encoder}' não suportado.")