andreribeiro87 commited on
Commit
e664e1c
·
verified ·
1 Parent(s): 5ce189d

Add AutoModel.from_pretrained support (unet3plus+efficientnet)

Browse files
README.md CHANGED
@@ -106,42 +106,33 @@ This model uses a custom PyTorch architecture. The model code is included in the
106
  ### Installation
107
 
108
  ```bash
109
- pip install torch torchvision timm safetensors huggingface_hub
110
  ```
111
 
112
  ### Inference
113
 
114
  ```python
115
- import sys
116
  import torch
117
- from huggingface_hub import snapshot_download
118
- from safetensors.torch import load_file
119
  from torchvision.transforms import functional as TF
120
  from PIL import Image
121
 
122
- # 1. Download the full repo (weights + model code) — handles LFS automatically
123
- repo_dir = snapshot_download("andreribeiro87/unet3plus-efficientnet-kvasir-seg")
124
- sys.path.insert(0, repo_dir)
125
-
126
- # 2. Import and instantiate the model
127
- from models import create_model
128
- model = create_model("unet3plus", backbone="efficientnet")
129
-
130
- # 3. Load weights
131
- state_dict = load_file(f"{repo_dir}/model.safetensors")
132
- model.load_state_dict(state_dict)
133
  model.eval()
134
 
135
- # 4. Preprocess an image
136
  image = Image.open("your_colonoscopy_image.jpg").convert("RGB")
137
  x = TF.to_tensor(TF.resize(image, [256, 256])).unsqueeze(0) # (1, 3, 256, 256)
138
 
139
- # 5. Predict
140
  with torch.no_grad():
141
- logit = model(x) # (1, 1, 256, 256)
142
- mask = (logit.sigmoid() > 0.5).squeeze() # bool tensor (256, 256)
143
 
144
- # 6. Convert to PIL
145
  pred_mask = TF.to_pil_image(mask.float())
146
  ```
147
 
 
106
  ### Installation
107
 
108
  ```bash
109
+ pip install torch torchvision timm transformers
110
  ```
111
 
112
  ### Inference
113
 
114
  ```python
 
115
  import torch
116
+ from transformers import AutoModel
 
117
  from torchvision.transforms import functional as TF
118
  from PIL import Image
119
 
120
+ # Load model — downloads weights + code automatically
121
+ model = AutoModel.from_pretrained(
122
+ "andreribeiro87/unet3plus-efficientnet-kvasir-seg",
123
+ trust_remote_code=True,
124
+ )
 
 
 
 
 
 
125
  model.eval()
126
 
127
+ # Preprocess
128
  image = Image.open("your_colonoscopy_image.jpg").convert("RGB")
129
  x = TF.to_tensor(TF.resize(image, [256, 256])).unsqueeze(0) # (1, 3, 256, 256)
130
 
131
+ # Predict
132
  with torch.no_grad():
133
+ outputs = model(pixel_values=x)
134
+ mask = (outputs["logits"].sigmoid() > 0.5).squeeze() # bool (256, 256)
135
 
 
136
  pred_mask = TF.to_pil_image(mask.float())
137
  ```
138
 
config.json CHANGED
@@ -1,10 +1,14 @@
1
  {
2
- "architecture": "unet3plus",
3
  "backbone": "efficientnet",
4
  "img_size": 256,
5
  "in_channels": 3,
6
  "out_channels": 1,
7
  "inter_ch": 64,
 
 
 
 
8
  "training": {
9
  "loss": "dice_focal",
10
  "focal_gamma": 1.1217038657488427,
 
1
  {
2
+ "model_type": "unet3plus",
3
  "backbone": "efficientnet",
4
  "img_size": 256,
5
  "in_channels": 3,
6
  "out_channels": 1,
7
  "inter_ch": 64,
8
+ "auto_map": {
9
+ "AutoConfig": "configuration_unet3plus.UNet3PlusConfig",
10
+ "AutoModel": "modeling_unet3plus.UNet3PlusForSegmentation"
11
+ },
12
  "training": {
13
  "loss": "dice_focal",
14
  "focal_gamma": 1.1217038657488427,
configuration_unet3plus.py ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from transformers import PretrainedConfig
2
+
3
+
4
+ class UNet3PlusConfig(PretrainedConfig):
5
+ model_type = "unet3plus"
6
+
7
+ def __init__(
8
+ self,
9
+ backbone: str = "efficientnet",
10
+ in_channels: int = 3,
11
+ out_channels: int = 1,
12
+ inter_ch: int = 64,
13
+ img_size: int = 256,
14
+ **kwargs,
15
+ ):
16
+ super().__init__(**kwargs)
17
+ self.backbone = backbone
18
+ self.in_channels = in_channels
19
+ self.out_channels = out_channels
20
+ self.inter_ch = inter_ch
21
+ self.img_size = img_size
model.safetensors CHANGED
@@ -1,3 +1,3 @@
1
  version https://git-lfs.github.com/spec/v1
2
- oid sha256:94528930d89b8b147886031f2dbec583407cc3662b59909454c7721e513da72a
3
- size 52240388
 
1
  version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4b20d1c13554c4da7ea0a35318b90fda300f220a023224abdac5d4e658d780a1
3
+ size 52243740
modeling_unet3plus.py ADDED
@@ -0,0 +1,48 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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_unet3plus import UNet3PlusConfig
13
+ from models import create_model
14
+
15
+
16
+ class UNet3PlusForSegmentation(PreTrainedModel):
17
+ """UNet 3+ segmentation model with EfficientNet-B0 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 = UNet3PlusConfig
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: UNet3PlusConfig):
29
+ super().__init__(config)
30
+ self.model = create_model(
31
+ architecture="unet3plus",
32
+ backbone=config.backbone,
33
+ in_channels=config.in_channels,
34
+ out_channels=config.out_channels,
35
+ inter_ch=config.inter_ch,
36
+ )
37
+
38
+ def forward(
39
+ self,
40
+ pixel_values: torch.Tensor,
41
+ labels: torch.Tensor | None = None,
42
+ **kwargs,
43
+ ):
44
+ logits = self.model(pixel_values)
45
+ loss = None
46
+ if labels is not None:
47
+ loss = F.binary_cross_entropy_with_logits(logits, labels.float())
48
+ return {"loss": loss, "logits": logits}