Update README: use snapshot_download for inference
Browse files
README.md
CHANGED
|
@@ -106,27 +106,29 @@ 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
|
| 110 |
```
|
| 111 |
|
| 112 |
### Inference
|
| 113 |
|
| 114 |
```python
|
|
|
|
| 115 |
import torch
|
|
|
|
| 116 |
from safetensors.torch import load_file
|
| 117 |
from torchvision.transforms import functional as TF
|
| 118 |
from PIL import Image
|
| 119 |
|
| 120 |
-
# 1.
|
| 121 |
-
|
| 122 |
-
|
| 123 |
|
| 124 |
-
# 2. Import the model
|
| 125 |
from models import create_model
|
| 126 |
-
|
| 127 |
-
# 3. Instantiate and load weights
|
| 128 |
model = create_model("unet3plus", backbone="efficientnet")
|
| 129 |
-
|
|
|
|
|
|
|
| 130 |
model.load_state_dict(state_dict)
|
| 131 |
model.eval()
|
| 132 |
|
|
|
|
| 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 |
|