Skip to content

Commit 1f5b0a5

Browse files
authored
Merge pull request #400 from JdeRobot/dph/image_detection
Improve image resize management for object detection pipeline
2 parents e8e74ee + 17ced2a commit 1f5b0a5

15 files changed

Lines changed: 654 additions & 203 deletions

.streamlit/config.toml

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,2 @@
1+
[server]
2+
maxUploadSize = 1000 # Sets the limit to 1GB

app.py

Lines changed: 111 additions & 86 deletions
Original file line numberDiff line numberDiff line change
@@ -1,69 +1,8 @@
11
import streamlit as st
2-
import sys
3-
import subprocess
42
from tabs.dataset_viewer import dataset_viewer_tab
53
from tabs.inference import inference_tab
64
from tabs.evaluator import evaluator_tab
7-
8-
9-
def browse_folder():
10-
"""
11-
Opens a native folder selection dialog and returns the selected folder path.
12-
Works on Windows, macOS, and Linux (with zenity or kdialog).
13-
Returns None if cancelled or error.
14-
"""
15-
try:
16-
if sys.platform.startswith("win"):
17-
script = (
18-
"Add-Type -AssemblyName System.windows.forms;"
19-
"$f=New-Object System.Windows.Forms.FolderBrowserDialog;"
20-
'if($f.ShowDialog() -eq "OK"){Write-Output $f.SelectedPath}'
21-
)
22-
result = subprocess.run(
23-
["powershell", "-NoProfile", "-Command", script],
24-
capture_output=True,
25-
text=True,
26-
timeout=30,
27-
)
28-
folder = result.stdout.strip()
29-
return folder if folder else None
30-
elif sys.platform == "darwin":
31-
script = (
32-
'POSIX path of (choose folder with prompt "Select dataset folder:")'
33-
)
34-
result = subprocess.run(
35-
["osascript", "-e", script], capture_output=True, text=True, timeout=30
36-
)
37-
folder = result.stdout.strip()
38-
return folder if folder else None
39-
else:
40-
# Linux: try zenity, then kdialog
41-
for cmd in [
42-
[
43-
"zenity",
44-
"--file-selection",
45-
"--directory",
46-
"--title=Select dataset folder",
47-
],
48-
[
49-
"kdialog",
50-
"--getexistingdirectory",
51-
"--title",
52-
"Select dataset folder",
53-
],
54-
]:
55-
try:
56-
result = subprocess.run(
57-
cmd, capture_output=True, text=True, timeout=30
58-
)
59-
folder = result.stdout.strip()
60-
if folder:
61-
return folder
62-
except (subprocess.TimeoutExpired, FileNotFoundError, Exception):
63-
continue
64-
return None
65-
except Exception:
66-
return None
5+
from perceptionmetrics.utils.gui import browse_folder
676

687

698
def browse_dataset_path():
@@ -80,13 +19,13 @@ def browse_dataset_path():
8019

8120
# Initialize commonly used session state keys
8221
st.session_state.setdefault("dataset_path", "")
83-
st.session_state.setdefault("dataset_type", "COCO")
84-
st.session_state.setdefault("split", "val")
22+
st.session_state.setdefault("dataset_type", "YOLO")
23+
st.session_state.setdefault("split", "test")
8524
st.session_state.setdefault("config_option", "Manual Configuration")
8625
st.session_state.setdefault("confidence_threshold", 0.5)
8726
st.session_state.setdefault("nms_threshold", 0.5)
8827
st.session_state.setdefault("max_detections", 100)
89-
st.session_state.setdefault("device", "cpu")
28+
st.session_state.setdefault("device", "cuda")
9029
st.session_state.setdefault("batch_size", 1)
9130
st.session_state.setdefault("evaluation_step", 5)
9231
st.session_state.setdefault("detection_model", None)
@@ -185,15 +124,6 @@ def browse_dataset_path():
185124
step=1,
186125
key="max_detections",
187126
)
188-
st.number_input(
189-
"Image Resize Height",
190-
min_value=1,
191-
max_value=4096,
192-
value=640,
193-
step=1,
194-
key="resize_height",
195-
help="Height to resize images for inference",
196-
)
197127
with col2:
198128
st.selectbox(
199129
"Device",
@@ -226,16 +156,83 @@ def browse_dataset_path():
226156
key="evaluation_step",
227157
help="Update UI with intermediate metrics every N images (0 = disable intermediate updates)",
228158
)
229-
st.number_input(
230-
"Image Resize Width",
231-
min_value=1,
232-
max_value=4096,
233-
value=640,
234-
step=1,
235-
key="resize_width",
236-
help="Width to resize images for inference",
159+
160+
st.write("---")
161+
st.write("**Image Size Configuration**")
162+
163+
# Resize Logic
164+
enable_resize = st.checkbox(
165+
"Enable Resize", value=True, key="enable_resize"
166+
)
167+
168+
if enable_resize:
169+
resize_strategy = st.radio(
170+
"Resize Strategy",
171+
["Fixed Dimensions", "Min Side"],
172+
key="resize_strategy",
173+
horizontal=True,
174+
label_visibility="collapsed",
237175
)
238176

177+
if resize_strategy == "Fixed Dimensions":
178+
c1, c2 = st.columns(2)
179+
with c1:
180+
st.number_input(
181+
"Image Resize Height",
182+
min_value=1,
183+
max_value=4096,
184+
value=640,
185+
step=1,
186+
key="resize_height",
187+
help="Height to resize images for inference",
188+
)
189+
with c2:
190+
st.number_input(
191+
"Image Resize Width",
192+
min_value=1,
193+
max_value=4096,
194+
value=640,
195+
step=1,
196+
key="resize_width",
197+
help="Width to resize images for inference",
198+
)
199+
else:
200+
st.number_input(
201+
"Min Side",
202+
min_value=1,
203+
max_value=4096,
204+
value=640,
205+
step=1,
206+
key="min_side",
207+
help="Minimum size of the shorter side of the image",
208+
)
209+
210+
# Crop Logic
211+
enable_crop = st.checkbox("Enable Center Crop", key="enable_crop")
212+
213+
if enable_crop:
214+
c1, c2 = st.columns(2)
215+
with c1:
216+
st.number_input(
217+
"Crop Height",
218+
min_value=1,
219+
max_value=4096,
220+
value=640,
221+
step=1,
222+
key="crop_height",
223+
help="Center crop height",
224+
)
225+
with c2:
226+
st.number_input(
227+
"Crop Width",
228+
min_value=1,
229+
max_value=4096,
230+
value=640,
231+
step=1,
232+
key="crop_width",
233+
help="Center crop width",
234+
)
235+
239236
# Load model action in sidebar
240237
from perceptionmetrics.models.torch_detection import TorchImageDetectionModel
241238
import json, tempfile
@@ -283,20 +280,48 @@ def browse_dataset_path():
283280
device = st.session_state.get("device", "cpu")
284281
batch_size = int(st.session_state.get("batch_size", 1))
285282
evaluation_step = int(st.session_state.get("evaluation_step", 5))
286-
resize_height = int(st.session_state.get("resize_height", 640))
287-
resize_width = int(st.session_state.get("resize_width", 640))
288283
model_format = st.session_state.get("model_format", "torchvision")
284+
285+
# Resize Logic extraction
286+
enable_resize = st.session_state.get("enable_resize", True)
287+
resize_cfg = None
288+
if enable_resize:
289+
resize_strategy = st.session_state.get(
290+
"resize_strategy", "Fixed Dimensions"
291+
)
292+
if resize_strategy == "Fixed Dimensions":
293+
resize_height = int(
294+
st.session_state.get("resize_height", 640)
295+
)
296+
resize_width = int(
297+
st.session_state.get("resize_width", 640)
298+
)
299+
resize_cfg = {
300+
"height": resize_height,
301+
"width": resize_width,
302+
}
303+
else:
304+
min_side = int(st.session_state.get("min_side", 640))
305+
resize_cfg = {"min_side": min_side}
306+
289307
config_data = {
290308
"confidence_threshold": confidence_threshold,
291309
"nms_threshold": nms_threshold,
292310
"max_detections_per_image": max_detections,
293311
"device": device,
294312
"batch_size": batch_size,
295313
"evaluation_step": evaluation_step,
296-
"resize_height": resize_height,
297-
"resize_width": resize_width,
298314
"model_format": model_format.lower(),
299315
}
316+
if resize_cfg is not None:
317+
config_data["resize"] = resize_cfg
318+
319+
if enable_crop:
320+
crop_height = int(st.session_state.get("crop_height", 640))
321+
crop_width = int(st.session_state.get("crop_width", 640))
322+
crop_cfg = {"height": crop_height, "width": crop_width}
323+
config_data["crop"] = crop_cfg
324+
300325
with tempfile.NamedTemporaryFile(
301326
delete=False, suffix=".json", mode="w"
302327
) as tmp_cfg:
Lines changed: 111 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,111 @@
1+
import argparse
2+
import os
3+
4+
from perceptionmetrics.datasets.yolo import YOLODataset
5+
from perceptionmetrics.models.torch_detection import TorchImageDetectionModel
6+
7+
8+
def parse_args() -> argparse.Namespace:
9+
"""Parse user input arguments
10+
11+
:return: parsed arguments
12+
:rtype: argparse.Namespace
13+
"""
14+
parser = argparse.ArgumentParser(
15+
description="Evaluate and Visualize YOLO detection models."
16+
)
17+
18+
parser.add_argument(
19+
"--model", type=str, required=True, help="Scripted pytorch model"
20+
)
21+
parser.add_argument(
22+
"--ontology",
23+
type=str,
24+
required=True,
25+
help="JSON file containing model output ontology",
26+
)
27+
parser.add_argument(
28+
"--model_cfg",
29+
type=str,
30+
required=True,
31+
help="JSON file with model configuration",
32+
)
33+
parser.add_argument(
34+
"--dataset_fname", type=str, required=True, help="YOLO-like YAML dataset file"
35+
)
36+
parser.add_argument(
37+
"--dataset_dir",
38+
type=str,
39+
required=True,
40+
help="Directory containing the dataset images",
41+
)
42+
parser.add_argument(
43+
"--split",
44+
type=str,
45+
required=True,
46+
help="Name of the split to be evaluated",
47+
)
48+
parser.add_argument(
49+
"--metrics_fname",
50+
type=str,
51+
required=True,
52+
help="CSV file where the evaluation results will be stored",
53+
)
54+
parser.add_argument(
55+
"--outdir",
56+
type=str,
57+
required=False,
58+
help="Directory where the evaluation results and visualizations will be stored. Only used if 'save_visualizations' or 'results_per_sample' are True",
59+
)
60+
parser.add_argument(
61+
"--results_per_sample",
62+
action="store_true",
63+
help="Save results per sample",
64+
)
65+
parser.add_argument(
66+
"--save_visualizations",
67+
action="store_true",
68+
help="Save visualizations for each sample",
69+
)
70+
71+
args = parser.parse_args()
72+
73+
if args.save_visualizations or args.results_per_sample:
74+
if args.outdir is None:
75+
raise ValueError(
76+
"'outdir' must be specified when 'save_visualizations' or 'results_per_sample' are True"
77+
)
78+
else:
79+
args.outdir = None
80+
81+
return args
82+
83+
84+
def main():
85+
"""Main function"""
86+
args = parse_args()
87+
88+
# Load image detection model
89+
model = TorchImageDetectionModel(args.model, args.model_cfg, args.ontology)
90+
91+
# Load dataset for evaluation
92+
dataset = YOLODataset(args.dataset_fname, args.dataset_dir)
93+
94+
# Evaluation loop
95+
results = model.eval(
96+
dataset,
97+
split=args.split,
98+
predictions_outdir=args.outdir,
99+
results_per_sample=args.results_per_sample,
100+
save_visualizations=args.save_visualizations,
101+
)
102+
103+
# Store evaluation results
104+
metrics_df = results["metrics_df"]
105+
os.makedirs(os.path.dirname(args.metrics_fname), exist_ok=True)
106+
metrics_df.to_csv(args.metrics_fname)
107+
print(f"Evaluation results saved to {args.metrics_fname}")
108+
109+
110+
if __name__ == "__main__":
111+
main()

examples/tasm_image.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -18,11 +18,12 @@ def parse_args() -> argparse.Namespace:
1818
"""
1919
parser = argparse.ArgumentParser()
2020
parser.add_argument(
21-
"--model_weights", type=str, required=True, help="Tensorflow model weights in HDF5 format"
22-
)
23-
parser.add_argument(
24-
"--model_name", type=str, required=True, help="TASM model name"
21+
"--model_weights",
22+
type=str,
23+
required=True,
24+
help="Tensorflow model weights in HDF5 format",
2525
)
26+
parser.add_argument("--model_name", type=str, required=True, help="TASM model name")
2627
parser.add_argument(
2728
"--ontology",
2829
type=str,

examples/tutorial_image_detection.ipynb

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -250,7 +250,7 @@
250250
" image = Image.open(image_path).convert('RGB')\n",
251251
" print(f\" Testing inference on: {image_path}\")\n",
252252
" # Run inference\n",
253-
" predictions = detection_model.inference(image)\n",
253+
" predictions = detection_model.predict(image)\n",
254254
"\n",
255255
" print(f\" Found {len(predictions['boxes'])} detections\")\n",
256256
" # Get ground truth for comparison\n",

0 commit comments

Comments
 (0)