andreribeiro87 commited on
Commit
6996ae6
·
verified ·
1 Parent(s): c3d1bc2

Add AutoModel.from_pretrained support (attention_unet+convnext)

Browse files
README.md CHANGED
@@ -90,42 +90,33 @@ This model uses a custom PyTorch architecture. The model code is included in the
90
  ### Installation
91
 
92
  ```bash
93
- pip install torch torchvision timm safetensors huggingface_hub
94
  ```
95
 
96
  ### Inference
97
 
98
  ```python
99
- import sys
100
  import torch
101
- from huggingface_hub import snapshot_download
102
- from safetensors.torch import load_file
103
  from torchvision.transforms import functional as TF
104
  from PIL import Image
105
 
106
- # 1. Download the full repo (weights + model code) — handles LFS automatically
107
- repo_dir = snapshot_download("andreribeiro87/attention-unet-convnext-kvasir-seg")
108
- sys.path.insert(0, repo_dir)
109
-
110
- # 2. Import and instantiate the model
111
- from models import create_model
112
- model = create_model("attention_unet", backbone="convnext")
113
-
114
- # 3. Load weights
115
- state_dict = load_file(f"{repo_dir}/model.safetensors")
116
- model.load_state_dict(state_dict)
117
  model.eval()
118
 
119
- # 4. Preprocess an image
120
  image = Image.open("your_colonoscopy_image.jpg").convert("RGB")
121
  x = TF.to_tensor(TF.resize(image, [256, 256])).unsqueeze(0) # (1, 3, 256, 256)
122
 
123
- # 5. Predict
124
  with torch.no_grad():
125
- logit = model(x) # (1, 1, 256, 256)
126
- mask = (logit.sigmoid() > 0.5).squeeze() # bool tensor (256, 256)
127
 
128
- # 6. Convert to PIL
129
  pred_mask = TF.to_pil_image(mask.float())
130
  ```
131
 
 
90
  ### Installation
91
 
92
  ```bash
93
+ pip install torch torchvision timm transformers
94
  ```
95
 
96
  ### Inference
97
 
98
  ```python
 
99
  import torch
100
+ from transformers import AutoModel
 
101
  from torchvision.transforms import functional as TF
102
  from PIL import Image
103
 
104
+ # Load model — downloads weights + code automatically
105
+ model = AutoModel.from_pretrained(
106
+ "andreribeiro87/attention-unet-convnext-kvasir-seg",
107
+ trust_remote_code=True,
108
+ )
 
 
 
 
 
 
109
  model.eval()
110
 
111
+ # Preprocess
112
  image = Image.open("your_colonoscopy_image.jpg").convert("RGB")
113
  x = TF.to_tensor(TF.resize(image, [256, 256])).unsqueeze(0) # (1, 3, 256, 256)
114
 
115
+ # Predict
116
  with torch.no_grad():
117
+ outputs = model(pixel_values=x)
118
+ mask = (outputs["logits"].sigmoid() > 0.5).squeeze() # bool (256, 256)
119
 
 
120
  pred_mask = TF.to_pil_image(mask.float())
121
  ```
122
 
config.json CHANGED
@@ -1,9 +1,13 @@
1
  {
2
- "architecture": "attention_unet",
3
  "backbone": "convnext",
4
  "img_size": 256,
5
  "in_channels": 3,
6
  "out_channels": 1,
 
 
 
 
7
  "training": {
8
  "loss": "bce_dice",
9
  "lr": 0.001,
 
1
  {
2
+ "model_type": "attention_unet",
3
  "backbone": "convnext",
4
  "img_size": 256,
5
  "in_channels": 3,
6
  "out_channels": 1,
7
+ "auto_map": {
8
+ "AutoConfig": "configuration_attention_unet.AttentionUNetConfig",
9
+ "AutoModel": "modeling_attention_unet.AttentionUNetForSegmentation"
10
+ },
11
  "training": {
12
  "loss": "bce_dice",
13
  "lr": 0.001,
configuration_attention_unet.py ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import PretrainedConfig
2
+
3
+
4
+ class AttentionUNetConfig(PretrainedConfig):
5
+ model_type = "attention_unet"
6
+
7
+ def __init__(
8
+ self,
9
+ backbone: str = "convnext",
10
+ in_channels: int = 3,
11
+ out_channels: int = 1,
12
+ img_size: int = 256,
13
+ **kwargs,
14
+ ):
15
+ super().__init__(**kwargs)
16
+ self.backbone = backbone
17
+ self.in_channels = in_channels
18
+ self.out_channels = out_channels
19
+ self.img_size = img_size
model.safetensors CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:b0c29289995adfc04f79dc1c5a0f1c0e0f85b3b19f11f0fab4ee0dd9699a2a5d
3
- size 139209476
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:29d062c1486094cb2acbd057b6fa40f8883c48fe60db26f6db5332c2bd3a4ed4
3
+ size 139211100
modeling_attention_unet.py ADDED
@@ -0,0 +1,47 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import sys
3
+
4
+ _here = os.path.dirname(os.path.abspath(__file__))
5
+ if _here not in sys.path:
6
+ sys.path.insert(0, _here)
7
+
8
+ import torch
9
+ import torch.nn.functional as F
10
+ from transformers import PreTrainedModel
11
+
12
+ from configuration_attention_unet import AttentionUNetConfig
13
+ from models import create_model
14
+
15
+
16
+ class AttentionUNetForSegmentation(PreTrainedModel):
17
+ """Attention U-Net segmentation model with ConvNeXt-Tiny backbone.
18
+
19
+ Returns a dict with key "logits" (raw sigmoid input, shape B×1×H×W).
20
+ Pass pixel_values as a float32 tensor normalised to [0, 1], shape B×3×H×W.
21
+ """
22
+ config_class = AttentionUNetConfig
23
+ _tied_weights_keys = None
24
+ # Class-level fallback so transformers v5 finalization never triggers
25
+ # nn.Module.__getattr__ for this attribute (instance attr set in __init__ takes priority)
26
+ all_tied_weights_keys = {}
27
+
28
+ def __init__(self, config: AttentionUNetConfig):
29
+ super().__init__(config)
30
+ self.model = create_model(
31
+ architecture="attention_unet",
32
+ backbone=config.backbone,
33
+ in_channels=config.in_channels,
34
+ out_channels=config.out_channels,
35
+ )
36
+
37
+ def forward(
38
+ self,
39
+ pixel_values: torch.Tensor,
40
+ labels: torch.Tensor | None = None,
41
+ **kwargs,
42
+ ):
43
+ logits = self.model(pixel_values)
44
+ loss = None
45
+ if labels is not None:
46
+ loss = F.binary_cross_entropy_with_logits(logits, labels.float())
47
+ return {"loss": loss, "logits": logits}