Views
No views yet
1from torchvision.models.segmentation import deeplabv3_resnet50, DeepLabV3_ResNet50_Weights
2
3def create_model(num_classes) -> torch.nn.Module:
4 model = deeplabv3_resnet50(weights=DeepLabV3_ResNet50_Weights.DEFAULT)
5
6 old_layer: nn.Conv2d = model.classifier[-1]
7 in_ch = old_layer.in_channels
8 model.classifier[-1] = nn.Conv2d(in_ch, num_classes, kernel_size=1)
9
10 # for p in model.parameters():
11 # p.requires_grad = False
12
13 return model