Repository navigation
Expand file tree
/
Copy pathpredict_and_visualize.py
More file actions
339 lines (276 loc) · 13.2 KB
/
Copy pathpredict_and_visualize.py
File metadata and controls
339 lines (276 loc) · 13.2 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
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
#!/usr/bin/env python3
"""
predict_and_visualize.py — I-JEPA Predictor Visualisation
Columns per row (one row = one image):
1. Original — unmodified image
2. Context — what the context encoder sees (targets + excluded greyed)
3. Target block 1 — context pixels + NN reconstruction for block 1 only
4. Target block 2 — context pixels + NN reconstruction for block 2 only
5. Target block 3 — context pixels + NN reconstruction for block 3 only
6. Target block 4 — context pixels + NN reconstruction for block 4 only
7. Full reconstruction— context pixels + NN reconstruction for all 4 blocks
NN reconstruction: for each target patch the gallery patch whose target-encoder
feature is cosine-nearest to the predictor's predicted feature is used.
Output: images/ijepa_predictions.png
Usage:
python predict_and_visualize.py
"""
import os
import sys
import numpy as np
import torch
import torch.nn.functional as F
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
from torchvision import transforms, datasets
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from models import IJEPA
# ─── Configuration ─────────────────────────────────────────────────────────────
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
IMG_SIZE = 224
PATCH_SIZE = 16
GRID_SIZE = IMG_SIZE // PATCH_SIZE # 14
NUM_PATCHES = GRID_SIZE * GRID_SIZE # 196
ENCODER_DIM = 768
ENCODER_DEPTH = 12
ENCODER_HEADS = 12
PRED_DIM = 384
PRED_DEPTH = 6
PRED_HEADS = 12
CHECKPOINT_PATH = 'checkpoints/ijepa_checkpoint_ep50.pth'
TARGET_ENC_PATH = 'checkpoints/ijepa_target_encoder_final.pth'
OUTPUT_PATH = 'images/ijepa_predictions.png'
GALLERY_SIZE = 200 # test images used to build the NN gallery
NUM_DEMO = 4 # rows in the output figure
NUM_TARGETS = 4 # target blocks per image
# ImageNet-style normalisation used during training
MEAN = torch.tensor([0.4467, 0.4398, 0.4066])
STD = torch.tensor([0.2603, 0.2566, 0.2713])
GREY = np.array([0.5, 0.5, 0.5], dtype=np.float32)
# RGB border colours for each target block (used in column titles)
BLOCK_COLORS = [
(0.95, 0.30, 0.30), # red
(0.30, 0.65, 0.95), # blue
(0.30, 0.85, 0.40), # green
(0.95, 0.85, 0.15), # yellow
]
# ─── Rendering helpers ─────────────────────────────────────────────────────────
def patch_rc(idx):
"""Flat patch index → (row, col) on the 14×14 grid."""
return divmod(int(idx), GRID_SIZE)
def grey_patches(base, patch_set):
"""Return a copy of base with every patch in patch_set replaced by grey."""
ps = PATCH_SIZE
out = base.copy()
for pidx in patch_set:
r, c = patch_rc(pidx)
y0, x0 = r * ps, c * ps
out[y0:y0+ps, x0:x0+ps] = GREY
return out
def place_nn_patches(base, tm_indices, pred_m, gallery_feats, gallery_patches):
"""
Replace patches in tm_indices with their nearest-neighbour gallery patches.
pred_m : (Nt, E) predictor output for this block, any device
gallery_feats : (G, E) L2-normalised float tensor on CPU
gallery_patches: list of (ps, ps, 3) float32 arrays, length G
"""
ps = PATCH_SIZE
out = base.copy()
pred_norm = F.normalize(pred_m.cpu().float(), dim=-1) # (Nt, E)
sims = pred_norm @ gallery_feats.T # (Nt, G)
nn_idx = sims.argmax(dim=-1) # (Nt,)
for local_i, pidx in enumerate(tm_indices.cpu().tolist()):
r, c = patch_rc(pidx)
y0, x0 = r * ps, c * ps
out[y0:y0+ps, x0:x0+ps] = gallery_patches[nn_idx[local_i].item()]
return out
# ─── Per-column panel builders ─────────────────────────────────────────────────
def panel_original(img_np):
"""Column 1: unmodified image."""
return img_np.copy()
def panel_context(img_np, ctx_mask, tgt_masks):
"""
Column 2: only context patches are shown.
Every patch the context encoder cannot see (targets + excluded) is grey.
"""
ctx_set = set(ctx_mask.cpu().tolist())
hidden = set(range(NUM_PATCHES)) - ctx_set
return grey_patches(img_np, hidden)
def panel_single_block(img_np, ctx_mask, tgt_masks, preds,
gallery_feats, gallery_patches, block_idx):
"""
Columns 3–6: context pixels + NN reconstruction for one target block.
All other patches (remaining target blocks + excluded) are grey.
"""
ctx_set = set(ctx_mask.cpu().tolist())
active_set = set(tgt_masks[block_idx].cpu().tolist())
# Start with grey everywhere
out = np.full_like(img_np, fill_value=0.5)
# Restore context patches
for pidx in ctx_set:
r, c = patch_rc(pidx)
y0, x0 = r * PATCH_SIZE, c * PATCH_SIZE
out[y0:y0+PATCH_SIZE, x0:x0+PATCH_SIZE] = img_np[y0:y0+PATCH_SIZE, x0:x0+PATCH_SIZE]
# Fill target block with NN patches
out = place_nn_patches(out, tgt_masks[block_idx], preds[block_idx],
gallery_feats, gallery_patches)
return out
def panel_full_reconstruction(img_np, ctx_mask, tgt_masks, preds,
gallery_feats, gallery_patches):
"""
Column 7: context pixels + NN reconstruction for all 4 target blocks.
Only excluded patches (if any) remain grey.
"""
ctx_set = set(ctx_mask.cpu().tolist())
tgt_all = set().union(*[set(tm.cpu().tolist()) for tm in tgt_masks])
# Grey for excluded patches only
excluded = set(range(NUM_PATCHES)) - ctx_set - tgt_all
out = grey_patches(img_np, excluded)
# Fill each target block with NN patches
for m_i in range(len(tgt_masks)):
out = place_nn_patches(out, tgt_masks[m_i], preds[m_i],
gallery_feats, gallery_patches)
return out
# ─── Main ──────────────────────────────────────────────────────────────────────
def main():
os.makedirs('images', exist_ok=True)
os.makedirs('data', exist_ok=True)
# 1. Load model ──────────────────────────────────────────────────────────────
print('Loading model …')
model = IJEPA(
img_size=IMG_SIZE, patch_size=PATCH_SIZE,
encoder_dim=ENCODER_DIM, encoder_depth=ENCODER_DEPTH,
encoder_heads=ENCODER_HEADS,
predictor_dim=PRED_DIM, predictor_depth=PRED_DEPTH,
predictor_heads=PRED_HEADS,
).to(DEVICE)
if os.path.exists(CHECKPOINT_PATH):
ck = torch.load(CHECKPOINT_PATH, map_location=DEVICE)
model.load_state_dict(ck['model_state_dict'])
print(f' Loaded {CHECKPOINT_PATH}')
elif os.path.exists(TARGET_ENC_PATH):
model.target_encoder.load_state_dict(
torch.load(TARGET_ENC_PATH, map_location=DEVICE)
)
print(f' Warning: full checkpoint not found; using {TARGET_ENC_PATH} (target encoder only)')
print(' Predictor output is random — NN panels will be uninformative.')
else:
sys.exit(
f'No checkpoint found at:\n {CHECKPOINT_PATH}\n {TARGET_ENC_PATH}\n'
'Run Self_Supervised_Learning.ipynb first.'
)
model.eval()
# 2. Dataset ─────────────────────────────────────────────────────────────────
print('Loading STL-10 test set …')
tf_norm = transforms.Compose([
transforms.Resize(IMG_SIZE),
transforms.CenterCrop(IMG_SIZE),
transforms.ToTensor(),
transforms.Normalize(MEAN.tolist(), STD.tolist()),
])
tf_raw = transforms.Compose([
transforms.Resize(IMG_SIZE),
transforms.CenterCrop(IMG_SIZE),
transforms.ToTensor(),
])
dset_norm = datasets.STL10(root='data', split='test', download=True, transform=tf_norm)
dset_raw = datasets.STL10(root='data', split='test', download=False, transform=tf_raw)
# 3. Build gallery feature bank ──────────────────────────────────────────────
print(f'Building gallery from {GALLERY_SIZE} images …')
gallery_feats = []
gallery_patches = []
loader = torch.utils.data.DataLoader(
dset_norm, batch_size=32, shuffle=False, num_workers=4
)
collected = 0
with torch.no_grad():
for batch_norm, _ in loader:
if collected >= GALLERY_SIZE:
break
n = min(batch_norm.shape[0], GALLERY_SIZE - collected)
batch_norm = batch_norm[:n].to(DEVICE)
feats = model.target_encoder(batch_norm) # (n, N, E)
feats = F.layer_norm(feats, (ENCODER_DIM,))
for i in range(n):
gallery_feats.append(feats[i].cpu())
raw_np = dset_raw[collected + i][0].permute(1, 2, 0).numpy()
for pidx in range(NUM_PATCHES):
r, c = patch_rc(pidx)
y0, x0 = r * PATCH_SIZE, c * PATCH_SIZE
gallery_patches.append(
raw_np[y0:y0+PATCH_SIZE, x0:x0+PATCH_SIZE].copy()
)
collected += n
gallery_feats = F.normalize(
torch.cat(gallery_feats, dim=0).float(), dim=-1
) # (G, E) pre-normalised for dot-product NN search
print(f' Gallery: {gallery_feats.shape[0]:,} patches from {collected} images')
# 4. Demo images ─────────────────────────────────────────────────────────────
demo_idx = [3, 47, 91, 143]
demo_norm = torch.stack([dset_norm[i][0] for i in demo_idx]).to(DEVICE)
demo_raw = [dset_raw[i][0].permute(1, 2, 0).numpy() for i in demo_idx]
# 5. Run predictor ────────────────────────────────────────────────────────────
print('Running forward pass …')
B = len(demo_idx)
with torch.no_grad():
ctx_masks, tgt_masks = model.masking(B, DEVICE)
ctx_out = model.context_encoder(demo_norm, ctx_masks)
predictions = model.predictor(ctx_out, tgt_masks, DEVICE) # [B][M] (Nt, E)
# 6. Render figure ────────────────────────────────────────────────────────────
print('Rendering figure …')
COLS = 2 + NUM_TARGETS + 1 # original + context + 4 blocks + full = 7
FIG_W = COLS * 2.5
FIG_H = NUM_DEMO * 2.5 + 0.7
fig, axes = plt.subplots(
NUM_DEMO, COLS,
figsize=(FIG_W, FIG_H),
gridspec_kw={'hspace': 0.04, 'wspace': 0.03},
)
fig.patch.set_facecolor('#0d0d0d')
col_labels = (
['Original', 'Context'] +
[f'Target {i+1}' for i in range(NUM_TARGETS)] +
['Reconstruction']
)
for col_i, label in enumerate(col_labels):
color = '#e0e0e0'
if 2 <= col_i < 2 + NUM_TARGETS:
r, g, b = BLOCK_COLORS[col_i - 2]
color = f'#{int(r*255):02x}{int(g*255):02x}{int(b*255):02x}'
axes[0, col_i].set_title(label, fontsize=9.5, color=color,
pad=6, fontweight='bold')
for row_i in range(NUM_DEMO):
img_np = demo_raw[row_i]
ctx_m = ctx_masks[row_i]
tgt_m = tgt_masks[row_i]
preds_i = predictions[row_i]
panels = (
[panel_original(img_np),
panel_context(img_np, ctx_m, tgt_m)] +
[panel_single_block(img_np, ctx_m, tgt_m, preds_i,
gallery_feats, gallery_patches, bi)
for bi in range(NUM_TARGETS)] +
[panel_full_reconstruction(img_np, ctx_m, tgt_m, preds_i,
gallery_feats, gallery_patches)]
)
for col_i, panel in enumerate(panels):
ax = axes[row_i, col_i]
ax.imshow(np.clip(panel, 0, 1))
ax.set_xticks([])
ax.set_yticks([])
# coloured border around each target-block column
if 2 <= col_i < 2 + NUM_TARGETS:
for spine in ax.spines.values():
spine.set_visible(True)
spine.set_edgecolor(BLOCK_COLORS[col_i - 2])
spine.set_linewidth(1.5)
else:
for spine in ax.spines.values():
spine.set_visible(False)
plt.savefig(OUTPUT_PATH, dpi=150, bbox_inches='tight',
facecolor='#0d0d0d', edgecolor='none')
print(f'\nSaved → {OUTPUT_PATH}')
if __name__ == '__main__':
main()