Spaces:
Sleeping
Sleeping
| # Contenido de: VisionEnsembleModel.py | |
| import torch | |
| import torch.nn as nn | |
| import timm | |
| class VisionEnsembleModel(nn.Module): | |
| def __init__(self, num_classes, cnn_model_name='efficientnet_b2', vit_model_name='vit_small_patch16_224'): | |
| super().__init__() | |
| self.cnn = timm.create_model(cnn_model_name, pretrained=False, num_classes=num_classes) | |
| cnn_features = self.cnn.get_classifier().in_features | |
| self.cnn.reset_classifier(0) | |
| self.vit = timm.create_model(vit_model_name, pretrained=False, num_classes=num_classes) | |
| vit_features = self.vit.head.in_features | |
| self.vit.head = nn.Identity() | |
| self.classifier = nn.Sequential( | |
| nn.BatchNorm1d(cnn_features + vit_features), | |
| nn.Linear(cnn_features + vit_features, 512), | |
| nn.ReLU(), | |
| nn.Dropout(0.5), | |
| nn.Linear(512, num_classes) | |
| ) | |
| def forward(self, image): | |
| cnn_feat = self.cnn(image) | |
| vit_feat = self.vit(image) | |
| combined = torch.cat([cnn_feat, vit_feat], dim=1) | |
| output = self.classifier(combined) | |
| return output |