Repository navigation
Expand file tree
/
Copy pathgetting_started.py
More file actions
210 lines (160 loc) · 8.19 KB
/
Copy pathgetting_started.py
File metadata and controls
210 lines (160 loc) · 8.19 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
import os
import torch
import numpy as np
import SimpleITK as sitk
from glob import glob
from PIL import Image
from huggingface_hub import snapshot_download # Install huggingface_hub if not already installed
import dotenv
from timeit import default_timer as timer
dotenv.load_dotenv() # Load environment variables from .env file
# --- Download Trained Model Weights (~400MB) ---
# License reminder: The official nnInteractive checkpoint is licensed under
# Creative Commons Attribution Non Commercial Share Alike 4.0 (CC BY-NC-SA 4.0).
# See the License section of this readme!.
HF_TOKEN =os.getenv("HF_TOKEN") # Set Hugging Face cache directory
DOWNLOAD_DIR = os.getenv("MODEL_DIR")# Specify the download directory
DATA_DIR = os.getenv("DATA_DIR") # Specify the data directory (not used in this script but may be useful for your projects)
REPO_ID = "nnInteractive/nnInteractive"
MODEL_NAME = "nnInteractive_v1.0" # Updated models may be available in the future
download_path = snapshot_download(
repo_id=REPO_ID,
allow_patterns=[f"{MODEL_NAME}/*"],
local_dir=DOWNLOAD_DIR
)
# The model is now stored in DOWNLOAD_DIR/MODEL_NAME.
# --- Initialize Inference Session ---
from nnInteractive import inference
from nnInteractive.inference.inference_session import nnInteractiveInferenceSession
session = nnInteractiveInferenceSession(
device=torch.device("cuda:0"), # Set inference device
use_torch_compile=False, # Experimental: Not tested yet
verbose=False,
torch_n_threads=os.cpu_count(), # Use available CPU cores
do_autozoom=True, # Enables AutoZoom for better patching
use_pinned_memory=True, # Optimizes GPU memory transfers
)
# Load the trained model
model_path = os.path.join(DOWNLOAD_DIR, MODEL_NAME)
session.initialize_from_trained_model_folder(model_path)
# --- Load Video Frames as Volume ---
# DO NOT preprocess the image in any way. Give it to nnInteractive as it is! DO NOT apply level window, DO NOT normalize
# intensities and never ever convert an image with higher precision (float32, uint16, etc) to uint8!
# The ONLY instance where some preprocesing makes sense is if your original image is too large to be reasonably used.
# This may be the case, for example, for some microCT images. In this case you can consider downsampling.
# Load images from a video as a greyscale volume (e.g. shape (1, x, y, z) where z is the number of frames)
# Load all video frames
images_dir = os.path.join(DATA_DIR, "Experiment_Set", "JPEGImages")
annotations_dir = os.path.join(DATA_DIR, "Experiment_Set", "Annotations")
# Get all image files sorted by name
image_files = sorted(glob(os.path.join(images_dir, "*.jpg")))
if not image_files:
raise ValueError(f"No images found in {images_dir}")
print(f"Loading {len(image_files)} frames...")
# Load frames and convert to greyscale
n_max_frames = 1 # Set this to your frame count you want to process at once.
frame_count = 0
frames = []
for img_path in image_files:
if frame_count >= n_max_frames:
break
frame_sitk = sitk.ReadImage(img_path)
frame = sitk.GetArrayFromImage(frame_sitk) # Shape: (H, W) or (H, W, C)
# Convert to greyscale if RGB
if frame.ndim == 3 and frame.shape[-1] == 3:
frame = frame.mean(axis=-1) # Average RGB channels
frames.append(frame)
frame_count += 1
# Stack frames into volume: (1, H, W, num_frames)
img = np.stack(frames, axis=-1)[None, ...] # Shape: (1, H, W, num_frames)
# Load first frame mask as prompt
first_frame_name = os.path.basename(image_files[0]).replace('.jpg', '.png')
first_mask_path = os.path.join(annotations_dir, first_frame_name)
if not os.path.exists(first_mask_path):
raise ValueError(f"First frame mask not found: {first_mask_path}")
# Use PIL to load the paletted PNG properly (without colormap applied)
input_mask_pil = Image.open(first_mask_path)
input_mask = np.array(input_mask_pil) # Shape: (H, W) for paletted images
# If mask is still RGB (colormap was applied), take first channel
if input_mask.ndim == 3:
input_mask = input_mask[:, :, 0]
# Extract only label 1 from the mask (masks contain labels 1-17)
LABEL_ID = 1 # Change this to segment different labels
input_mask_binary = (input_mask == LABEL_ID).astype(np.uint8)
print(f"Volume shape: {img.shape}")
print(f"First frame mask shape: {input_mask_binary.shape}")
print(f"Label {LABEL_ID} pixels found: {input_mask_binary.sum()}")
# Validate input dimensions
if img.ndim != 4:
raise ValueError("Input image must be 4D with shape (1, x, y, z)")
session.set_image(img)
# --- Define Output Buffer ---
target_tensor = torch.zeros(img.shape[1:], dtype=torch.uint8) # Must be 3D (x, y, z)
session.set_target_buffer(target_tensor)
# --- Interacting with the Model ---
# Use the first frame mask as the initial prompt for segmentation
# Prepare the mask volume: same shape as img (without the channel dimension)
mask_volume = np.zeros(img.shape[1:], dtype=np.uint8) # Shape: (H, W, num_frames)
mask_volume[:, :, 0] = input_mask_binary # Place mask in first frame (z=0)
# Add the first frame mask as a lasso (closed contour) interaction
print("Adding first frame mask as segmentation prompt...")
# Start the timer here to measure inference time
start = timer()
#session.add_lasso_interaction(mask_volume, include_interaction=True)
session.add_initial_seg_interaction(mask_volume, run_prediction=True)
# The model will now propagate this segmentation across all frames
# You can add additional interactions to refine the segmentation if needed:
# Example: Add a **positive** point interaction on a specific frame
# POINT_COORDINATES = (x, y, frame_index) # Example: (50, 60, 5) for frame 5
# session.add_point_interaction(POINT_COORDINATES, include_interaction=True)
# Example: Add a **negative** point interaction
# session.add_point_interaction(POINT_COORDINATES, include_interaction=False)
# Example: Add a 2D bounding box on a specific frame
# BBOX_COORDINATES = [[x1, x2], [y1, y2], [frame, frame+1]] # Last dim must be [d, d+1]
# session.add_bbox_interaction(BBOX_COORDINATES, include_interaction=True)
# You can combine any number of interactions as needed.
# The model refines the segmentation result incrementally with each new interaction.
# --- Retrieve Results ---
# The target buffer holds the segmentation result for all frames
results = session.target_buffer.clone() # Shape: (H, W, num_frames)
# End the timer after retrieving results
end = timer()
print(f"Inference time: {end - start}")
# OR (equivalent)
# results = target_tensor.clone()
print(f"Segmentation complete! Result shape: {results.shape}")
print(results)
# Optional: Save results for each frame
output_dir = os.path.join(DATA_DIR, "Experiment_Set", "Predictions")
os.makedirs(output_dir, exist_ok=True)
# Extract palette from input mask for colormap visualization
palette = input_mask_pil.palette
if palette is None:
# If no palette, create a default one
palette = Image.new('P', (256, 1))
palette.putdata(range(256))
for frame_idx, img_file in enumerate(image_files):
if frame_idx >= n_max_frames:
break
frame_name = os.path.basename(img_file).replace('.jpg', '.png')
output_path = os.path.join(output_dir, frame_name)
# Extract frame segmentation
frame_seg = results[:, :, frame_idx].numpy().astype(np.uint8)
# Create paletted image with the same colormap as input mask
frame_seg_pil = Image.fromarray(frame_seg, mode='P')
if palette is not None and hasattr(input_mask_pil, 'palette'):
frame_seg_pil.palette = input_mask_pil.palette
# Save as PNG with palette
frame_seg_pil.save(output_path)
print(f"Saved {min(n_max_frames, len(image_files))} segmentation masks to {output_dir}")
# Cloning is required because the buffer will be **reused** for the next object.
# Alternatively, set a new target buffer for each object:
# session.set_target_buffer(torch.zeros(img.shape[1:], dtype=torch.uint8))
# --- Start a New Object Segmentation ---
session.reset_interactions() # Clears the target buffer and resets interactions
# Now you can start segmenting the next object in the image.
# --- Set a New Image ---
# Setting a new image also requires setting a new matching target buffer
# session.set_image(NEW_IMAGE)
# session.set_target_buffer(torch.zeros(NEW_IMAGE.shape[1:], dtype=torch.uint8))
# Enjoy!