# Contenido de: VisionEnsembleModel.py import torch import torch.nn as nn import timm class VisionEnsembleModel(nn.Module): """ La misma clase de modelo que definiste en Colab. """ def __init__(self, num_classes, cnn_model_name='efficientnet_b2', vit_model_name='vit_small_patch16_224'): super().__init__() # Usamos pretrained=False aquí porque cargaremos nuestros propios pesos. # Timm cargará los pesos preentrenados si no encuentra un state_dict local, # pero es más limpio ser explícito. Al final, los sobrescribiremos. 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