File size: 1,397 Bytes
6a04068
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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

import torch
import torch.nn as nn
import torchvision


def create_effnetb2_model(num_classes: int = 7,
                          seed: int=42):
  """Creates a PyTorch EfficientNetB2 feature extractor"""
  # Setup pretrained EffNetB2 weights
  weights = torchvision.models.EfficientNet_B2_Weights.DEFAULT
  # Get EffNetB2 transforms
  transforms = weights.transforms()
  # Setup pretrained model instance
  model = torchvision.models.efficientnet_b2(weights=weights)
  # Freeze the base layers in the model
  for param in model.parameters():
    param.requires_grad = False
  # Create classifier
  torch.manual_seed(seed)
  model.classifier = nn.Sequential(
      nn.Dropout(p=0.3, inplace=True),
      nn.Linear(in_features=1408, out_features=num_classes)
  )
  return model, transforms

def create_vit_model(num_classes:int=7,
                     seed:int=42):
  """Creates a PyTorch ViT pretrained feature extractor"""
  # Create Vit_B_16 pretrained weights, transforms and models
  weights = torchvision.models.ViT_B_16_Weights.DEFAULT
  transforms = weights.transforms()
  model = torchvision.models.vit_b_16(weights=weights)

  # Freeze all the base layers
  for param in model.parameters():
    param.requires_grad = False

  # Change classifier head
  model.heads = nn.Sequential(
      nn.Linear(in_features=768,
                out_features=num_classes)
  )
  return model, transforms