Download modeling_attention_unet.py from andreribeiro87/attention-unet-convnext-kvasir-seg: direct link, hf CLI and curl.
- Browser
- Download file 1.93 kB
-
https://huggingface.co/andreribeiro87/attention-unet-convnext-kvasir-seg/resolve/main/modeling_attention_unet.py
- Command line
-
hf download hf://andreribeiro87/attention-unet-convnext-kvasir-seg/modeling_attention_unet.py
-
curl -L -o modeling_attention_unet.py https://huggingface.co/andreribeiro87/attention-unet-convnext-kvasir-seg/resolve/main/modeling_attention_unet.py
1.93 kB
| import os | |
| import sys | |
| _here = os.path.dirname(os.path.abspath(__file__)) | |
| if _here not in sys.path: | |
| sys.path.insert(0, _here) | |
| import torch | |
| import torch.nn.functional as F | |
| from transformers import PreTrainedModel, PretrainedConfig | |
| from models import create_model | |
| class AttentionUNetConfig(PretrainedConfig): | |
| model_type = "attention_unet" | |
| def __init__( | |
| self, | |
| backbone: str = "convnext", | |
| in_channels: int = 3, | |
| out_channels: int = 1, | |
| img_size: int = 256, | |
| **kwargs, | |
| ): | |
| super().__init__(**kwargs) | |
| self.backbone = backbone | |
| self.in_channels = in_channels | |
| self.out_channels = out_channels | |
| self.img_size = img_size | |
| class AttentionUNetForSegmentation(PreTrainedModel): | |
| """Attention U-Net segmentation model with ConvNeXt-Tiny backbone. | |
| Returns a dict with key "logits" (raw sigmoid input, shape B×1×H×W). | |
| Pass pixel_values as a float32 tensor normalised to [0, 1], shape B×3×H×W. | |
| """ | |
| config_class = AttentionUNetConfig | |
| _tied_weights_keys = None | |
| # Class-level fallback so transformers v5 finalization never triggers | |
| # nn.Module.__getattr__ for this attribute (instance attr set in __init__ takes priority) | |
| all_tied_weights_keys = {} | |
| def __init__(self, config: AttentionUNetConfig): | |
| super().__init__(config) | |
| self.model = create_model( | |
| architecture="attention_unet", | |
| backbone=config.backbone, | |
| in_channels=config.in_channels, | |
| out_channels=config.out_channels, | |
| ) | |
| def forward( | |
| self, | |
| pixel_values: torch.Tensor, | |
| labels: torch.Tensor | None = None, | |
| **kwargs, | |
| ): | |
| logits = self.model(pixel_values) | |
| loss = None | |
| if labels is not None: | |
| loss = F.binary_cross_entropy_with_logits(logits, labels.float()) | |
| return {"loss": loss, "logits": logits} | |