CoolFace
Modelpublic

TriEightz/PneumoniaChestXRay-ConvNextBase

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes20downloads
ConvNextBase.py43 linesDownload Raw Back to root
1from transformers import PretrainedConfig, PreTrainedModel2from typing import List3from torchvision.models import convnext_base, ConvNeXt_Base_Weights4 5 6class ConvNextBaseConfig(PretrainedConfig):7    model_type = "ConvNext"8 9    def __init__(10        self,11        **kwargs,12    ):13        super().__init__(**kwargs)14 15 16class ConvNextBaseModel(PreTrainedModel):17    config_class = ConvNextBaseConfig18 19    def __init__(self, config):20        super().__init__(config)21        self.model = convnext_base()22 23    def forward(self, tensor):24        return self.model(tensor)25 26 27class ConvNextBaseModelForImageClassification(PreTrainedModel):28    config_class = ConvNextBaseConfig29 30    def __init__(self, config):31        super().__init__(config)32        self.model = convnext_base()33 34    def forward(self, tensor, labels=None):35        logits = self.model(tensor)36        if labels is not None:37            loss = torch.nn.cross_entropy(logits, labels)38            return {"loss": loss, "logits": logits}39        return {"logits": logits}40 41ConvNextBaseConfig.register_for_auto_class()42ConvNextBaseModel.register_for_auto_class("AutoModel")43ConvNextBaseModelForImageClassification.register_for_auto_class("AutoModelForImageClassification")