Repository navigation
Expand file tree
/
Copy pathplot_rule_examples.py
More file actions
232 lines (197 loc) · 9 KB
/
Copy pathplot_rule_examples.py
File metadata and controls
232 lines (197 loc) · 9 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
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
#!/usr/bin/env python3
"""Plot annotated example images for the COCOLogicV2 rules.
Produces two figures:
- ``rule_examples_grid.pdf`` — a 2x5 grid, one panel per rule (10 panels), with
the rule's friendly name printed below each image and bounding boxes drawn
for the COCO objects in that rule's whitelist (see ``RULE_CATEGORIES``).
- ``rule_examples_pair.pdf`` — a 2x1 vertical pair showing examples for
rule_1 ("Signal and Ride") and rule_5 ("Three of a Kind"), same annotation
style as the grid.
Image selection is per image_id, looked up in the two fewshot
("few-shot") COCOLogicV2 JSON files. Edit the ``RULE_IMAGE_IDS`` dict below
to choose different examples. Edit ``RULE_CATEGORIES`` to tune which COCO
categories are highlighted for each rule.
"""
import argparse
import json
import os
import sys
import warnings
import matplotlib.pyplot as plt
import seaborn as sns
from matplotlib.patches import Rectangle
from PIL import Image
BASE = os.path.dirname(os.path.abspath(__file__))
# ---------- USER-EDITABLE CONFIG ----------
# One fewshot COCO image_id per rule. Defaults are positive examples
# auto-selected from the fewshot train/test JSONs. Replace freely with any
# image_id that exists in the few-shot data.
RULE_IMAGE_IDS = {
1: "493806", # Signal and Ride
2: "441883", # Double Serving
3: "140500", # Herd Alone
4: "185360", # Either Dog or Car
5: "78858", # Three of a Kind
6: "399540", # Car Majority
7: "555048", # Empty Seat
8: "149623", # Single Mode Traffic
9: "518785", # Personal Transport
10: "560718", # Surf Trip
}
# COCO categories highlighted per rule. Names must match the standard COCO
# vocabulary (case-sensitive). Unknown names are silently dropped.
RULE_CATEGORIES = {
1: ["traffic light", "bicycle", "bus", "bus"], # Signal and Ride
2: ["bottle", "cup", "pizza"], # Double Serving
3: ["cow", "sheep", "elephant", "person"], # Herd Alone
4: ["dog", "car"], # Either Dog or Car
5: ["bottle", "cup"], # Three of a Kind
6: ["car", "truck"], # Car Majority
7: ["chair", "couch", "person"], # Empty Seat
8: ["car", "bus", "motorcycle", "bicycle"], # Single Mode Traffic
9: ["bicycle", "person", "car"], # Personal Transport
10: ["surfboard", "person"], # Surf Trip
}
# ---------- /CONFIG ----------
def load_image_index(data_dir):
"""Build {image_id_str: file_name} from both fewshot JSONs.
Also returns ``{rule_idx (1-indexed): rule_name}`` so the script doesn't
duplicate the rule names already in the dataset.
"""
files = [
os.path.join(data_dir, "cocologic_train_fewshot.json"),
os.path.join(data_dir, "cocologic_test_fewshot.json"),
]
image_index = {}
rule_names = {}
for path in files:
with open(path) as f:
data = json.load(f)
for img_id, info in data["images"].items():
image_index[str(img_id)] = info["file_name"]
for rule_idx in range(1, 11):
rule_key = f"rule_{rule_idx}"
rule_names[rule_idx] = data["rules"][rule_key]["name"]
return image_index, rule_names
def load_coco_annotations(coco_dir):
"""Load COCO instance annotations for train2017 and val2017.
Returns ``(by_image_id, cat_id_to_name, cat_name_to_id)`` where
``by_image_id[img_id]`` is the list of annotation dicts for that image.
"""
by_image_id = {}
cat_id_to_name = {}
cat_name_to_id = {}
for split in ("train2017", "val2017"):
ann_path = os.path.join(coco_dir, "annotations", f"instances_{split}.json")
if not os.path.exists(ann_path):
raise FileNotFoundError(
f"Missing {ann_path}. Download via "
"http://images.cocodataset.org/annotations/annotations_trainval2017.zip"
)
print(f"Loading {ann_path} ...", file=sys.stderr)
with open(ann_path) as f:
data = json.load(f)
for cat in data["categories"]:
cat_id_to_name[cat["id"]] = cat["name"]
cat_name_to_id[cat["name"]] = cat["id"]
for ann in data["annotations"]:
by_image_id.setdefault(ann["image_id"], []).append(ann)
return by_image_id, cat_id_to_name, cat_name_to_id
def build_palette(category_ids):
"""Stable color per category id, drawn from ``tab20``."""
palette = sns.color_palette("tab20", n_colors=20)
return {cid: palette[i % len(palette)] for i, cid in enumerate(sorted(category_ids))}
def draw_image_with_boxes(ax, image_path, annotations, allowed_cat_ids,
cat_id_to_name, palette):
"""Render the image and overlay bounding boxes + category-name labels."""
image = Image.open(image_path).convert("RGB")
ax.imshow(image)
for ann in annotations:
cat_id = ann["category_id"]
if cat_id not in allowed_cat_ids:
continue
x, y, w, h = ann["bbox"]
color = palette.get(cat_id, (1, 0, 0))
ax.add_patch(Rectangle(
(x, y), w, h, linewidth=2, edgecolor=color, facecolor="none",
))
ax.text(
x, max(y - 4, 4), cat_id_to_name[cat_id],
fontsize=12, color="white", va="bottom", ha="left",
bbox=dict(facecolor=color, pad=1.5, edgecolor="none"),
)
ax.set_xticks([])
ax.set_yticks([])
for spine in ax.spines.values():
spine.set_visible(False)
def plot_rule_grid(rule_indices, rows, cols, out_path,
image_index, rule_names, coco_annotations,
cat_id_to_name, cat_name_to_id, coco_dir):
"""Generic grid plotter. Walks ``rule_indices`` into a ``rows × cols`` figure."""
fig, axes = plt.subplots(rows, cols, figsize=(cols * 4, rows * 4.3))
# Flatten axes for row-major iteration; handles the 1-column case too.
axes_flat = [axes] if rows * cols == 1 else list(axes.flatten())
for ax, rule_idx in zip(axes_flat, rule_indices):
image_id = RULE_IMAGE_IDS.get(rule_idx)
file_name = image_index.get(str(image_id))
if file_name is None:
warnings.warn(
f"rule_{rule_idx}: image_id {image_id!r} not found in the "
"fewshot JSONs; leaving panel blank."
)
ax.text(0.5, 0.5, "(missing)", ha="center", va="center")
ax.set_xticks([]); ax.set_yticks([])
for spine in ax.spines.values():
spine.set_visible(False)
else:
allowed_names = RULE_CATEGORIES.get(rule_idx, [])
allowed_cat_ids = {
cat_name_to_id[n] for n in allowed_names if n in cat_name_to_id
}
anns = coco_annotations.get(int(image_id), [])
palette = build_palette(allowed_cat_ids)
image_path = os.path.join(coco_dir, "images", file_name)
draw_image_with_boxes(ax, image_path, anns, allowed_cat_ids,
cat_id_to_name, palette)
ax.set_xlabel(rule_names.get(rule_idx, f"rule_{rule_idx}"),
fontsize=15, labelpad=8)
fig.tight_layout()
# Tighten the vertical gap between rows (only applies when there is one).
if rows > 1:
fig.subplots_adjust(hspace=-0.3)
out_dir = os.path.dirname(out_path)
if out_dir:
os.makedirs(out_dir, exist_ok=True)
fig.savefig(out_path, bbox_inches="tight")
plt.close(fig)
print(f"Saved {out_path}", file=sys.stderr)
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--data_dir", default="cocologicv2_data",
help="Directory with the fewshot COCOLogicV2 JSON files.")
parser.add_argument("--coco_dir", default="../datasets/coco",
help="Root of the COCO dataset (contains images/ and annotations/).")
parser.add_argument("--output_dir", default=os.path.join(BASE, "plots"),
help="Directory to write rule_examples_grid.pdf and rule_examples_pair.pdf "
"(default: analysis/plots).")
args = parser.parse_args()
image_index, rule_names = load_image_index(args.data_dir)
coco_annotations, cat_id_to_name, cat_name_to_id = load_coco_annotations(args.coco_dir)
plot_rule_grid(
list(range(1, 11)), rows=2, cols=5,
out_path=os.path.join(args.output_dir, "rule_examples_grid.pdf"),
image_index=image_index, rule_names=rule_names,
coco_annotations=coco_annotations,
cat_id_to_name=cat_id_to_name, cat_name_to_id=cat_name_to_id,
coco_dir=args.coco_dir,
)
plot_rule_grid(
[1, 5], rows=2, cols=1,
out_path=os.path.join(args.output_dir, "rule_examples_pair.pdf"),
image_index=image_index, rule_names=rule_names,
coco_annotations=coco_annotations,
cat_id_to_name=cat_id_to_name, cat_name_to_id=cat_name_to_id,
coco_dir=args.coco_dir,
)
if __name__ == "__main__":
main()