attention-unet-convnext-kvasir-seg / modeling_attention_unet.py
andreribeiro87's picture
Add AutoModel.from_pretrained support (attention_unet+convnext)
60ad774 verified
Raw History Blame Contribute Delete
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}