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

Update README: use snapshot_download for inference

Browse files
Files changed (1) hide show
  1. README.md +10 -8
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. Clone or download the repo files
121
- # git clone https://huggingface.co/andreribeiro87/unet3plus-efficientnet-kvasir-seg
122
- # cd unet3plus-efficientnet-kvasir-seg
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
- state_dict = load_file("model.safetensors")
 
 
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