Buckets:
| """ | |
| Claim 1 (b, c) + real-feature (a) — exercise all three DSP modules on the | |
| *real* pretrained backbones (DINOv2 ViT-L/14 and CLIP ViT-B/16), using the | |
| official iCVTEAM/DSP code. | |
| (b) Semantic Anchoring — FrozenDinoV2Encoder over N exemplars of a category; | |
| aggregate DINOv2 features into a per-category anchor | |
| and show cross-exemplar stability. | |
| (a) Primitive Imbuing — repo's exact closed-form ridge step (Eq. 7) applied | |
| to the *real* DINOv2 patch tokens; recon MSE < K-Means. | |
| (c) Conceptual Steering — repo's CAMGenerator (text-driven GradCAM on CLIP) | |
| on a real image; save heatmap + the L1 CAM-diff | |
| weight map that drives the saliency-aware MSE loss. | |
| Weights are downloaded at runtime. Results + figures go to ./outputs. | |
| """ | |
| import os, sys, json, urllib.request | |
| import numpy as np | |
| import torch | |
| import torch.nn.functional as F | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "DSP_src")) | |
| os.makedirs("outputs", exist_ok=True) | |
| DEV = "cuda" if torch.cuda.is_available() else "cpu" | |
| OUT = {"device": DEV, "torch": torch.__version__} | |
| torch.manual_seed(0); np.random.seed(0) | |
| DINO_URL = "https://dl.fbaipublicfiles.com/dinov2/dinov2_vitl14/dinov2_vitl14_pretrain.pth" | |
| CLIP_URL = "https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt" | |
| def dl(url, path): | |
| if not os.path.exists(path): | |
| print(f"downloading {url}", flush=True) | |
| urllib.request.urlretrieve(url, path) | |
| print(f"{path}: {os.path.getsize(path)/1e6:.1f} MB", flush=True) | |
| return path | |
| # ---- synthetic but structured "exemplars": textured crops of C categories ----- | |
| # Each category = a distinct low-freq colour/texture pattern; 5 exemplars each | |
| # (mirrors the 5-shot regime). DINOv2 should map same-category crops close. | |
| def make_exemplars(n_cat=3, k_shot=5, size=224): | |
| ys, xs = torch.meshgrid(torch.linspace(0, 1, size), torch.linspace(0, 1, size), indexing="ij") | |
| imgs, labels = [], [] | |
| for c in range(n_cat): | |
| for k in range(k_shot): | |
| fr = 4 + 3 * c | |
| phase = 0.3 * k | |
| r = 0.5 + 0.5 * torch.sin(2 * np.pi * fr * xs + phase) | |
| g = 0.5 + 0.5 * torch.sin(2 * np.pi * fr * ys + phase + c) | |
| b = 0.5 + 0.5 * torch.cos(2 * np.pi * fr * (xs + ys) + phase) | |
| img = torch.stack([r, g, b]) + 0.05 * torch.randn(3, size, size) | |
| imgs.append(img.clamp(0, 1)); labels.append(c) | |
| return torch.stack(imgs), torch.tensor(labels) | |
| # ============================ (b) Semantic Anchoring ============================ | |
| print("== Semantic Anchoring (real DINOv2 ViT-L/14) ==", flush=True) | |
| from models.dsp.modules import FrozenDinoV2Encoder | |
| dino_w = dl(DINO_URL, "outputs/dinov2_vitl14.pth") | |
| dino = FrozenDinoV2Encoder(dino_w, device=DEV).to(DEV).eval() | |
| imgs, labels = make_exemplars() | |
| imgs = F.interpolate(imgs, size=(224, 224)) | |
| pt_list, cls_list = [], [] | |
| with torch.no_grad(): | |
| for i in range(0, imgs.shape[0], 3): # small batches for 6GB-class GPUs | |
| b = imgs[i:i + 3].to(DEV) | |
| pt_list.append(dino(b, mode="x_norm_patchtokens").cpu()) | |
| cls_list.append(dino(b, mode="x_norm_clstoken").cpu()) | |
| patch_tokens = torch.cat(pt_list) # [N, 256, 1024] | |
| cls = F.normalize(torch.cat(cls_list), dim=-1) # [N, 1024] | |
| # per-category anchor = mean CLS over the 5 exemplars (aggregated categorical semantics) | |
| anchors = torch.stack([F.normalize(cls[labels == c].mean(0), dim=-1) for c in labels.unique()]) | |
| # stability: within-category cosine sim to anchor vs cross-category | |
| within, cross = [], [] | |
| for i in range(len(cls)): | |
| ci = labels[i].item() | |
| within.append(F.cosine_similarity(cls[i], anchors[ci], dim=0).item()) | |
| for cj in range(len(anchors)): | |
| if cj != ci: | |
| cross.append(F.cosine_similarity(cls[i], anchors[cj], dim=0).item()) | |
| OUT["anchoring_within_cat_cos_mean"] = float(np.mean(within)) | |
| OUT["anchoring_cross_cat_cos_mean"] = float(np.mean(cross)) | |
| OUT["anchoring_separation"] = float(np.mean(within) - np.mean(cross)) | |
| OUT["dino_patch_tokens_shape"] = list(patch_tokens.shape) | |
| print(json.dumps({k: OUT[k] for k in list(OUT)[-4:]}, indent=2), flush=True) | |
| # ===================== (a) Primitive Imbuing on real features =================== | |
| print("== Primitive Imbuing: closed-form ridge on real DINOv2 tokens ==", flush=True) | |
| def repo_closed_form_W(T, Pn, lam, K): | |
| inv = torch.inverse(Pn @ Pn.t() + lam * torch.eye(K, device=T.device, dtype=T.dtype)) | |
| return (T @ Pn.t()) @ inv | |
| def recon_mse(T, P, lam, K): | |
| Pn = F.normalize(P, p=2, dim=-1) | |
| return F.mse_loss(repo_closed_form_W(T, Pn, lam, K) @ Pn, T) | |
| # flatten category-0 real patch tokens -> target matrix T (fp64 for stability) | |
| T = F.normalize(patch_tokens[labels == 0].reshape(-1, 1024).double().cpu(), dim=-1).detach() | |
| K, lam = 64, 0.1 | |
| # K-Means init (cosine) using the repo's kmeans if available, else simple | |
| from models.dsp.kmeans_pytorch import kmeans | |
| _, C = kmeans(X=T, num_clusters=K, distance="cosine", device=torch.device("cpu"), tqdm_flag=False) | |
| P0 = F.normalize(C.double(), dim=-1) | |
| base = recon_mse(T, P0, lam, K).item() | |
| P = torch.nn.Parameter(P0.clone()); opt = torch.optim.Adam([P], lr=0.05) | |
| for _ in range(50): | |
| opt.zero_grad(); l = recon_mse(T, P, lam, K); l.backward(); opt.step() | |
| with torch.no_grad(): P.data.copy_(F.normalize(P.data, dim=-1)) | |
| fin = recon_mse(T, P.detach(), lam, K).item() | |
| OUT["imbuing_real_kmeans_mse"] = base | |
| OUT["imbuing_real_final_mse"] = fin | |
| OUT["imbuing_real_gain_pct"] = (base - fin) / base * 100 | |
| print(json.dumps({k: OUT[k] for k in list(OUT)[-3:]}, indent=2), flush=True) | |
| # ========================= (c) Conceptual Steering ============================= | |
| print("== Conceptual Steering: text-driven GradCAM (real CLIP ViT-B/16) ==", flush=True) | |
| try: | |
| from models.dsp.CAM.cam_generator import CAMGenerator | |
| clip_w = dl(CLIP_URL, "outputs/ViT-B-16.pt") | |
| cats = ["stripes", "waves", "grid"] | |
| cam_gen = CAMGenerator(categories=cats, clip_path=clip_w) | |
| cam_gen.to(DEV, torch.float32) | |
| # one real exemplar image of category 0, full-image bbox | |
| img0 = imgs[0:1].to(DEV) | |
| captions = [["", "stripes"]] # captions[0][1:] are the fg labels | |
| bboxes = [[[0.1, 0.1, 0.9, 0.9]]] # xyxy normalised | |
| refined_cams, keys = cam_gen(img0, captions, bboxes, gt_bboxes_only=False) | |
| cam = refined_cams[0].detach().cpu().numpy() | |
| OUT["cam_shape"] = list(cam.shape) | |
| OUT["cam_min"] = float(cam.min()); OUT["cam_max"] = float(cam.max()) | |
| OUT["cam_foreground_frac_gt_0p5"] = float((cam > 0.5).mean()) | |
| # the saliency-aware weight in _mse_loss = L1(CAM_gt, CAM_gen); here compute | |
| # a self-consistency weight L1(CAM(img), CAM(img+noise)) as a mechanism demo | |
| img_noisy = (img0 + 0.15 * torch.randn_like(img0)).clamp(0, 1) | |
| cams_n, _ = cam_gen(img_noisy, captions, bboxes, gt_bboxes_only=False) | |
| w = F.l1_loss(refined_cams.detach(), cams_n.detach(), reduction="none") | |
| OUT["steering_cam_diff_weight_mean"] = float(w.mean().item()) | |
| # save heatmap overlay | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| fig, ax = plt.subplots(1, 3, figsize=(11, 4)) | |
| ax[0].imshow(img0[0].permute(1, 2, 0).cpu().numpy()); ax[0].set_title("exemplar (cat 'stripes')") | |
| ax[1].imshow(cam, cmap="jet"); ax[1].set_title("text-driven GradCAM") | |
| ax[2].imshow(img0[0].permute(1, 2, 0).cpu().numpy()) | |
| import cv2 | |
| hm = cv2.resize((cam * 255).astype("uint8"), (224, 224)) | |
| ax[2].imshow(hm, cmap="jet", alpha=0.5); ax[2].set_title("overlay") | |
| for a in ax: a.axis("off") | |
| plt.tight_layout(); plt.savefig("outputs/conceptual_steering_gradcam.png", dpi=110) | |
| print(json.dumps({k: OUT[k] for k in list(OUT)[-5:]}, indent=2), flush=True) | |
| except Exception as e: | |
| import traceback; traceback.print_exc() | |
| OUT["steering_error"] = repr(e) | |
| with open("outputs/modules_verify.json", "w") as f: | |
| json.dump(OUT, f, indent=2) | |
| print("\nRESULTS_JSON " + json.dumps(OUT)) | |
Xet Storage Details
- Size:
- 8.11 kB
- Xet hash:
- 68572559895343ec215624d13824c7a30a57dfe3412b006b95613f39a527c79f
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.