junwatu commited on
Commit
c881b77
·
verified ·
1 Parent(s): 71334dd

Upload folder using huggingface_hub

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +4 -0
  2. LICENSE +21 -0
  3. README.md +252 -0
  4. claim1.py +59 -0
  5. claim2.py +73 -0
  6. claim5.py +107 -0
  7. claim5_official.py +79 -0
  8. claim6.py +34 -0
  9. configs/dsp-dior.yaml +99 -0
  10. configs/dsp-exdark.yaml +91 -0
  11. configs/dsp-ruod.yaml +91 -0
  12. crossover.py +42 -0
  13. databuilders/__init__.py +1 -0
  14. databuilders/dior.py +113 -0
  15. datamodules/__init__.py +2 -0
  16. datamodules/loader.py +109 -0
  17. datamodules/ref_table.py +76 -0
  18. datamodules/transforms/__init__.py +1 -0
  19. datamodules/transforms/layout_transform.py +132 -0
  20. figs/Main.png +3 -0
  21. fonts/GILI____.TTF +0 -0
  22. fonts/Rainbow-Party-2.ttf +3 -0
  23. gicdm_core.py +216 -0
  24. infer.sh +174 -0
  25. main.py +12 -0
  26. models/__init__.py +1 -0
  27. models/dsp/CAM/__init__.py +1 -0
  28. models/dsp/CAM/cam_generator.py +133 -0
  29. models/dsp/CAM/clip/__init__.py +1 -0
  30. models/dsp/CAM/clip/bpe_simple_vocab_16e6.txt.gz +3 -0
  31. models/dsp/CAM/clip/clip.py +245 -0
  32. models/dsp/CAM/clip/model.py +517 -0
  33. models/dsp/CAM/clip/simple_tokenizer.py +132 -0
  34. models/dsp/CAM/pytorch_grad_cam/__init__.py +14 -0
  35. models/dsp/CAM/pytorch_grad_cam/ablation_cam.py +134 -0
  36. models/dsp/CAM/pytorch_grad_cam/ablation_cam_multilayer.py +136 -0
  37. models/dsp/CAM/pytorch_grad_cam/ablation_layer.py +131 -0
  38. models/dsp/CAM/pytorch_grad_cam/activations_and_gradients.py +55 -0
  39. models/dsp/CAM/pytorch_grad_cam/base_cam.py +227 -0
  40. models/dsp/CAM/pytorch_grad_cam/eigen_cam.py +23 -0
  41. models/dsp/CAM/pytorch_grad_cam/eigen_grad_cam.py +21 -0
  42. models/dsp/CAM/pytorch_grad_cam/fullgrad_cam.py +95 -0
  43. models/dsp/CAM/pytorch_grad_cam/grad_cam.py +24 -0
  44. models/dsp/CAM/pytorch_grad_cam/grad_cam_plusplus.py +32 -0
  45. models/dsp/CAM/pytorch_grad_cam/guided_backprop.py +100 -0
  46. models/dsp/CAM/pytorch_grad_cam/layer_cam.py +36 -0
  47. models/dsp/CAM/pytorch_grad_cam/score_cam.py +63 -0
  48. models/dsp/CAM/pytorch_grad_cam/utils/__init__.py +4 -0
  49. models/dsp/CAM/pytorch_grad_cam/utils/find_layers.py +30 -0
  50. models/dsp/CAM/pytorch_grad_cam/utils/image.py +90 -0
.gitattributes CHANGED
@@ -33,3 +33,7 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ figs/Main.png filter=lfs diff=lfs merge=lfs -text
37
+ fonts/Rainbow-Party-2.ttf filter=lfs diff=lfs merge=lfs -text
38
+ scripts/data_process/ruod/008431.jpg filter=lfs diff=lfs merge=lfs -text
39
+ scripts/evaluation/FasterRCNN_score-mmdet/configs/reppoints/reppoints.png filter=lfs diff=lfs merge=lfs -text
LICENSE ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ MIT License
2
+
3
+ Copyright (c) 2026 CVTEAM
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
README.md ADDED
@@ -0,0 +1,252 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Envisioning Beyond the Few: Disentangled Semantics and Primitives for Few-Shot Atypical Layout-to-Image Generation
2
+
3
+ **ICML 2026**
4
+
5
+ **Authors:** Nan Bao, Yifan Zhao, Wenzhuang Wang, Jia Li
6
+
7
+ ![Main](figs/Main.png)
8
+
9
+ ## Environment Setup
10
+
11
+ We use two separate environments:
12
+
13
+ 1. **Main environment** for core training and inference.
14
+
15
+ ```bash
16
+ conda create -n dsp python=3.10.20
17
+ conda activate dsp
18
+ pip install torch==2.6.0 torchvision==0.21.0 torchaudio==2.6.0 --index-url https://download.pytorch.org/whl/cu126
19
+ pip install datasets==4.8.5 pillow==12.2.0 accelerate==1.13.0 transformers==5.8.1 diffusers==0.38.0 safetensors==0.8.0rc0 tensorboard==2.20.0 opencv-python==4.13.0.92 einops==0.8.2 imagesize==2.0.0 peft==0.19.1 ttach==0.0.3 ftfy==6.3.1 albumentations==2.0.8
20
+ ```
21
+
22
+ 2. **Evaluation environment** for MMDetection/MMEngine compatibility. It is used for evaluation with MMDetection/MMEngine due to strict version constraints, and also supports YOLO-based evaluation.
23
+
24
+ ```bash
25
+ conda create -n dsp-eval python=3.10.20
26
+ conda activate dsp-eval
27
+ conda install mkl==2023.1.0 numpy==1.26.4
28
+ conda install pytorch==2.1.2 torchvision==0.16.2 torchaudio==2.1.2 pytorch-cuda=12.1 -c pytorch -c nvidia
29
+ pip install mmengine==0.10.7 tqdm==4.67.3 shapely==2.1.2 scipy==1.15.3 terminaltables==3.1.10 ultralytics==8.4.50 pycocotools==2.0.11 https://download.openmmlab.com/mmcv/dist/cu121/torch2.1.0/mmcv-2.1.0-cp310-cp310-manylinux1_x86_64.whl "numpy<2.0.0" "setuptools<70.0.0"
30
+ ```
31
+
32
+ ## Set Environment Variables
33
+
34
+ Set the root path of this project:
35
+
36
+ ```bash
37
+ export DSP_PROJECT_DIR=/path/to/DSP # replace with the actual path
38
+ ```
39
+
40
+ It is recommended to add this line to `~/.bashrc` or `~/.zshrc` for persistence.
41
+
42
+ ## Pretrained Models Preparation
43
+
44
+ 1. We use several pretrained models as external dependencies. Please download them manually from the following sources:
45
+ - [stable-diffusion-v1-5](https://huggingface.co/stable-diffusion-v1-5/stable-diffusion-v1-5)
46
+ - [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14)
47
+ - [dinov2_vitl14_pretrain.pth](https://dl.fbaipublicfiles.com/dinov2/dinov2_vitl14/dinov2_vitl14_pretrain.pth)
48
+ - [ViT-B-16.pt](https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt)
49
+
50
+ 2. After downloading, organize the pretrained weights under `./pretrained` as follows:
51
+
52
+ ```bash
53
+ pretrained
54
+ ├── stable-diffusion-v1-5
55
+ │ └── ...
56
+ ├── clip-vit-large-patch14
57
+ │ └── ...
58
+ ├── dinov2_vitl14_pretrain.pth
59
+ └── ViT-B-16.pt
60
+ ```
61
+
62
+ You may either copy or symlink the files. We recommend using symbolic links:
63
+
64
+ ```bash
65
+ ln -s /path/to/stable-diffusion-v1-5 ./pretrained/stable-diffusion-v1-5
66
+ ln -s /path/to/clip-vit-large-patch14 ./pretrained/clip-vit-large-patch14
67
+ ln -s /path/to/dinov2_vitl14_pretrain.pth ./pretrained/dinov2_vitl14_pretrain.pth
68
+ ln -s /path/to/ViT-B-16.pt ./pretrained/ViT-B-16.pt
69
+ ```
70
+
71
+ ## Data Preparation
72
+
73
+ 1. We use several public datasets. Please download them manually from the following sources:
74
+
75
+ - [DIOR](https://gcheng-nwpu.github.io/#Datasets)
76
+ - [RUOD](https://github.com/xiaoDetection/RUOD)
77
+ - [ExDark](https://github.com/cs-chan/Exclusively-Dark-Image-Dataset/tree/master/Dataset)
78
+
79
+ 2. Unzip the downloaded datasets and organize the external dataset directories as follows:
80
+
81
+ ```bash
82
+ DIOR-VOC
83
+ ├── Annotations
84
+ │ ├── Horizontal_Bounding_Boxes
85
+ │ └── Oriented_Bounding_Boxes
86
+ └── VOC2007
87
+ ├── ImageSets
88
+ │ ├── Layout
89
+ │ ├── Main
90
+ │ └── Segmentation
91
+ └── JPEGImages
92
+ ```
93
+
94
+ ```bash
95
+ RUOD
96
+ ├── Environment_pic
97
+ │ ├── blur
98
+ │ ├── color
99
+ │ └── light
100
+ ├── Environmet_ANN
101
+ ├── RUOD_ANN
102
+ └── RUOD_pic
103
+ ├── test
104
+ └── train
105
+ ```
106
+
107
+ ```bash
108
+ ExDark
109
+ ├── annos
110
+ ├── imageclasslist.txt
111
+ └── images
112
+ ```
113
+
114
+ 3. Run data preprocessing scripts located in `./scripts/data_process`, after updating all hard-coded paths (e.g., `/path/to/DIOR_VOC`, `/path/to/RUOD`, `/path/to/ExDark`) in the scripts to match the local setup. Execute them in order.
115
+
116
+ The preprocessing outputs will be generated under `./data` with the following structure:
117
+
118
+ ```bash
119
+ data
120
+ ├── DIOR
121
+ │ ├── dior_emb.pt
122
+ │ ├── images -> /path/to/DIOR-VOC/VOC2007/JPEGImages
123
+ │ ├── metadatas
124
+ │ └── patches
125
+ ├── EXDARK
126
+ │ ├── exdark_emb.pt
127
+ │ ├── images
128
+ │ ├── metadatas
129
+ │ └── patches
130
+ └── RUOD
131
+ ├── images -> /path/to/RUOD/RUOD_pic
132
+ ├── metadatas
133
+ ├── patches
134
+ └── ruod_emb.pt
135
+ ```
136
+
137
+ ## Training and Inference
138
+
139
+ We provide three example configurations in `./configs`: `dsp-dior.yaml`, `dsp-ruod.yaml`, and `dsp-exdark.yaml`.
140
+
141
+ > **Argument Description:**
142
+ > - **config:** configuration file for model and dataset setup.
143
+ > - **metaseed:** seed generator identifier for deterministic sampling.
144
+ > - **num_seed:** number of sampling seeds for few-shot evaluation.
145
+ > - **k_shot:** number of samples per category in few-shot setting.
146
+ > - **run_id:** identifier for different runs.
147
+ > - **gpu_ids:** GPU device indices for execution.
148
+ > - **iter:** number of bootstrap iterations for FID.
149
+
150
+ ### Base Phase Training
151
+
152
+ ```bash
153
+ bash train_base.sh --config "dsp-dior"
154
+ bash train_base.sh --config "dsp-ruod"
155
+ bash train_base.sh --config "dsp-exdark"
156
+ ```
157
+
158
+ ### Novel Phase Training
159
+
160
+ ```bash
161
+ bash train_novel.sh --config "dsp-dior" --metaseed "aaa" --num_seed 50 --k_shot "5" --run_id "1" --gpu_ids "0,1,2,3"
162
+ bash train_novel.sh --config "dsp-ruod" --metaseed "aaa" --num_seed 50 --k_shot "5" --run_id "1" --gpu_ids "0,1,2,3"
163
+ bash train_novel.sh --config "dsp-exdark" --metaseed "aaa" --num_seed 50 --k_shot "5" --run_id "1" --gpu_ids "0,1,2,3"
164
+ ```
165
+
166
+ ### Inference
167
+
168
+ ```bash
169
+ bash infer.sh --config "dsp-dior" --metaseed "aaa" --num_seed 50 --k_shot "5" --run_id "1" --ckpt "100" --gpu_ids "0,1,2,3" --max_infer_size 50
170
+ bash infer.sh --config "dsp-ruod" --metaseed "aaa" --num_seed 50 --k_shot "5" --run_id "1" --ckpt "100" --gpu_ids "0,1,2,3" --max_infer_size 50
171
+ bash infer.sh --config "dsp-exdark" --metaseed "aaa" --num_seed 50 --k_shot "5" --run_id "1" --ckpt "100" --gpu_ids "0,1,2,3" --max_infer_size 50
172
+ ```
173
+
174
+ ## Evaluation
175
+
176
+ ### Preparation
177
+
178
+ Download the YOLO and Faster R-CNN weights from [this link](https://drive.google.com/drive/folders/1FWN02KEuGPdEkXv38MmT8-D4-uQcAn4_?usp=sharing). Place them under `./pretrained`. The expected directory structure is as follows:
179
+
180
+ ```bash
181
+ pretrained
182
+ ├── evaluation
183
+ │ ├── mmdet
184
+ │ │ ├── faster_rcnn_r50_fpn_1x-dior
185
+ │ │ │ └── epoch_12.pth
186
+ │ │ ├── faster_rcnn_r50_fpn_1x-exdark
187
+ │ │ │ └── epoch_12.pth
188
+ │ │ └── faster_rcnn_r50_fpn_1x-ruod
189
+ │ │ └── epoch_12.pth
190
+ │ └── yolo
191
+ │ └── best.pt
192
+ └── ... (pretrained models for training)
193
+ ```
194
+
195
+ ### YOLO (mAP / AP50 / AP75)
196
+
197
+ > **Note:** In yolo-wrapper-dior.sh, the `--xml_folder` path should be set to the DIOR annotation directory (`/path/to/DIOR-VOC/Annotations/Horizontal_Bounding_Boxes`).
198
+
199
+ ```bash
200
+ cd $DSP_PROJECT_DIR/scripts/evaluation/yoloscore-dior
201
+ bash yolo-wrapper-dior.sh --config "dsp-dior" --metaseed "aaa" --num_seed 50 --ckpt "100" --k_shot "5" --run_id "1" --gpu_ids 0
202
+ ```
203
+
204
+ ### Faster R-CNN (mAP / AP50 / AP75)
205
+
206
+ ```bash
207
+ cd $DSP_PROJECT_DIR/scripts/evaluation/FasterRCNN_score-mmdet
208
+ bash test-wrapper-dior.sh --config "dsp-dior" --metaseed "aaa" --num_seed 50 --ckpt "100" --k_shot "5" --run_id "1" --gpu_ids 0
209
+ bash test-wrapper-ruod.sh --config "dsp-ruod" --metaseed "aaa" --num_seed 50 --ckpt "100" --k_shot "5" --run_id "1" --gpu_ids 0
210
+ bash test-wrapper-exdark.sh --config "dsp-exdark" --metaseed "aaa" --num_seed 50 --ckpt "100" --k_shot "5" --run_id "1" --gpu_ids 0
211
+ ```
212
+
213
+ ### Bootstrap FID
214
+
215
+ ```bash
216
+ cd $DSP_PROJECT_DIR/scripts/evaluation/bootstrap_fid
217
+ python boot_fid-dior.py --config dsp-dior -run_id 1 -num_seeds 50 --iter 50 --k_shot 5
218
+ python boot_fid-ruod.py --config dsp-ruod -run_id 1 -num_seeds 50 --iter 50 --k_shot 5
219
+ python boot_fid-exdark.py --config dsp-exdark -run_id 1 -num_seeds 50 --iter 50 --k_shot 5
220
+ ```
221
+
222
+ Bootstrap FID results will be saved under `./metrics/BootstrapFID`.
223
+
224
+ ### Detection Metric Summarization
225
+
226
+ ```bash
227
+ cd $DSP_PROJECT_DIR/scripts/evaluation/summarize
228
+ bash summarize-wrapper.sh --config "dsp-dior" --k_shot "5" --run_id "1" --ckpt "100" --metaseed "aaa" --num_seed 50
229
+ bash summarize-wrapper.sh --config "dsp-ruod" --k_shot "5" --run_id "1" --ckpt "100" --metaseed "aaa" --num_seed 50
230
+ bash summarize-wrapper.sh --config "dsp-exdark" --k_shot "5" --run_id "1" --ckpt "100" --metaseed "aaa" --num_seed 50
231
+ ```
232
+
233
+ Detection evaluation results (mAP / AP50 / AP75, YOLO and Faster R-CNN) will be summarized in `./metrics`.
234
+
235
+ ## Acknowledgement
236
+
237
+ Our work is based on [stable diffusion](https://github.com/compvis/stable-diffusion), [diffusers](https://github.com/huggingface/diffusers), [CLIP](https://github.com/openai/CLIP), [DINOv2](https://github.com/facebookresearch/dinov2), [CC-Diff](https://github.com/AZZMM/CC-Diff), [MIGC](https://github.com/limuloo/MIGC), [GradCAM](https://github.com/linyq2117/CLIP-ES), and [kmeans_pytorch](https://github.com/subhadarship/kmeans_pytorch). Thanks for these great projects!
238
+
239
+ ## Citation
240
+
241
+ If you find our work useful for your research, please cite the following paper.
242
+
243
+ ```bib
244
+ @inproceedings{
245
+ bao2026envisioning,
246
+ title={Envisioning Beyond the Few: Disentangled Semantics and Primitives for Few-Shot Atypical Layout-to-Image Generation},
247
+ author={Bao, Nan and Zhao, Yifan and Wang, Wenzhuang and Li, Jia},
248
+ booktitle={Forty-third International Conference on Machine Learning},
249
+ year={2026},
250
+ url={https://openreview.net/forum?id=Jva4wVEySO}
251
+ }
252
+ ```
claim1.py ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Claim 1 verification: GICDM out-of-sample generated-point scaling (Eq. 1, Algorithm 1).
2
+
3
+ Test scenario (Figure 1 of paper): real samples drawn from a 60/40 mixture of two
4
+ hyperspheres with radii (r1, r2); generated samples from a mixture with swapped
5
+ radii and proportions. The two sets are disjoint, so all fidelity/coverage metrics
6
+ should score 0 in the ideal case. Standard metrics fail in high dimension due to
7
+ hubness; GICDM-corrected metrics should remain ~0.
8
+ """
9
+ import numpy as np
10
+ import sys, os
11
+ sys.path.insert(0, os.path.dirname(__file__))
12
+ from gicdm_core import (pairwise_sq_dists, icdm_scaling, gicdm,
13
+ clipped_density, clipped_coverage, hubness_stats)
14
+
15
+
16
+ def sample_mixture_spheres(d, n, r1, r2, p1):
17
+ n1 = int(round(p1 * n))
18
+ n2 = n - n1
19
+ def sphere(r, m):
20
+ v = np.random.normal(size=(m, d))
21
+ v /= np.linalg.norm(v, axis=1, keepdims=True)
22
+ return v * r
23
+ return np.vstack([sphere(r1, n1), sphere(r2, n2)])
24
+
25
+
26
+ def run(d, n, seed=0):
27
+ np.random.seed(seed)
28
+ Xr = sample_mixture_spheres(d, n, 3.0, 5.0, 0.6)
29
+ # swapped radii & proportions
30
+ Xg = sample_mixture_spheres(d, n, 5.0, 3.0, 0.4)
31
+ K1, K2 = 5, 50 # K2 = 10*K1
32
+
33
+ # raw metric
34
+ cd_raw = clipped_density(Xr, Xg, k=5)
35
+ cc_raw = clipped_coverage(Xr, Xg, k=5)
36
+
37
+ # GICDM-corrected
38
+ D_final, keep, Drr_gicdm = gicdm(Xr, Xg, K1, K2)
39
+ dissim = (D_final, Drr_gicdm) # (generated-to-real, real-to-real) GICDM dissimilarities
40
+ cd_gicdm = clipped_density(Xr, Xg, k=5, dissim=dissim, keep=keep)
41
+ cc_gicdm = clipped_coverage(Xr, Xg, k=5, dissim=dissim, keep=keep)
42
+
43
+ # Hubness in raw real-real space
44
+ h5_raw, A5_raw = hubness_stats(pairwise_sq_dists(Xr), k=5)
45
+
46
+ return dict(d=d, n=n, cd_raw=cd_raw, cc_raw=cc_raw,
47
+ cd_gicdm=cd_gicdm, cc_gicdm=cc_gicdm,
48
+ kept=float(keep.mean()), h5_raw=h5_raw, A5_raw=A5_raw)
49
+
50
+
51
+ if __name__ == "__main__":
52
+ import json
53
+ out = []
54
+ for d in [10, 50, 100, 500, 1000]:
55
+ r = run(d, 1000, seed=42)
56
+ out.append(r)
57
+ print(f"d={d:5d} CD raw={r['cd_raw']:.3f} GICDM={r['cd_gicdm']:.3f} | "
58
+ f"CC raw={r['cc_raw']:.3f} GICDM={r['cc_gicdm']:.3f} | kept={r['kept']:.2f} h5={r['h5_raw']:.2f}")
59
+ json.dump(out, open("results/claim1_hypersphere.json", "w"))
claim2.py ADDED
@@ -0,0 +1,73 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Claim 2 verification: Proposition 5.1 & Corollary 5.2.
2
+
3
+ Prop 5.1: p_hat_{mu,K}(x_i) = 1/(N V_d mu_i^d) * (1/K Sum_k k^{1/d})^d is a local
4
+ density estimator (asymptotically p(x_i)).
5
+ Cor 5.2: at ICDM convergence, mu_i ~ mu_bar for all i => p_hat equal everywhere
6
+ (density uniformized).
7
+
8
+ We verify numerically:
9
+ (a) p_hat_{mu,K} correlates with the true density in a mixture of Gaussians.
10
+ (b) after ICDM, the spread of mu_i collapses (std -> ~0) while raw std is large,
11
+ confirming density gradient removal.
12
+ """
13
+ import numpy as np
14
+ import sys, os
15
+ sys.path.insert(0, os.path.dirname(__file__))
16
+ from gicdm_core import pairwise_sq_dists, icdm_scaling, knn_distances
17
+
18
+
19
+ def unit_ball_volume(d):
20
+ from math import gamma, pi
21
+ return (pi ** (d / 2)) / gamma(d / 2 + 1)
22
+
23
+
24
+ def p_hat_mu_k(X, K):
25
+ D = pairwise_sq_dists(X)
26
+ mu = knn_distances(D, K, self_included=True).mean(axis=1)
27
+ N, d = X.shape
28
+ Vd = unit_ball_volume(d)
29
+ w = (np.sum([k ** (1.0 / d) for k in range(1, K + 1)]) / K) ** d
30
+ return 1.0 / (N * Vd * mu ** d) * w, mu
31
+
32
+
33
+ def run():
34
+ rng = np.random.default_rng(0)
35
+ # Mixture of two Gaussians with different variances -> different densities
36
+ d = 20
37
+ n1, n2 = 800, 800
38
+ X = np.vstack([
39
+ rng.normal(0, 0.5, size=(n1, d)),
40
+ rng.normal(3, 2.0, size=(n2, d)),
41
+ ])
42
+ true_dens = np.concatenate([
43
+ (2 * np.pi * 0.25) ** (-d / 2) * np.ones(n1),
44
+ (2 * np.pi * 4.0) ** (-d / 2) * np.ones(n2),
45
+ ])
46
+ K = 20
47
+ phat, mu = p_hat_mu_k(X, K)
48
+ # correlation between estimator and true density (log scale, robust)
49
+ log_corr = np.corrcoef(np.log(phat), np.log(true_dens))[0, 1]
50
+
51
+ # ICDM convergence: spread of mu before/after
52
+ mu_before = mu.copy()
53
+ delta, mu_after = icdm_scaling(pairwise_sq_dists(X), K, n_iter=10, return_mu=True)
54
+ res = dict(
55
+ d=d, N=len(X), K=K,
56
+ log_corr_density=float(log_corr),
57
+ mu_before_std=float(mu_before.std()),
58
+ mu_before_cv=float(mu_before.std() / mu_before.mean()),
59
+ mu_after_std=float(mu_after.std()),
60
+ mu_after_cv=float(mu_after.std() / mu_after.mean()),
61
+ )
62
+ print("Prop 5.1 log-correlation p_hat vs true density:", round(res['log_corr_density'], 3))
63
+ print("mu_i std before ICDM:", round(res['mu_before_std'], 4),
64
+ "| after ICDM:", round(res['mu_after_std'], 5))
65
+ print("mu_i CV before ICDM:", round(res['mu_before_cv'], 4),
66
+ "| after ICDM:", round(res['mu_after_cv'], 5))
67
+ return res
68
+
69
+
70
+ if __name__ == "__main__":
71
+ import json
72
+ res = run()
73
+ json.dump(res, open("results/claim2_density.json", "w"))
claim5.py ADDED
@@ -0,0 +1,107 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Claim 5 verification (scaled reproduction): Raisa et al. (2025) synthetic benchmark.
2
+
3
+ We implement a subset of the Raisa et al. test scenarios that exercise the two
4
+ metrics the paper reports gains on (Clipped Density & Clipped Coverage):
5
+ - GAUSSIAN MEAN DIFFERENCE (Purpose + Bounds: score should be 1 at equality)
6
+ - GAUSSIAN STD DEVIATION DIFFERENCE (Purpose: fails without GICDM)
7
+ - HYPERSPHERE SURFACE (Purpose: real vs generated on sphere shell)
8
+ - MODE COLLAPSE (Purpose / Bounds)
9
+ - SPHERE VS TORUS
10
+
11
+ For each scenario we evaluate Clipped Density and Clipped Coverage with and without
12
+ GICDM and check whether the metric behaves as the test expects. We tally pass rates
13
+ across the implemented scenarios and compare the *direction* of the paper's reported
14
+ gains (8/14->10/14 Purpose for Clipped Density; 8/13->11/13 Bounds).
15
+
16
+ NOTE: this is a scaled reproduction of the statistical mechanism (full 14 Purpose /
17
+ 13 Bounds tests require the official benchmark suite). It validates that GICDM
18
+ removes hubness-induced failures on representative scenarios.
19
+ """
20
+ import numpy as np
21
+ import sys, os
22
+ sys.path.insert(0, os.path.dirname(__file__))
23
+ from gicdm_core import pairwise_sq_dists, gicdm, clipped_density, clipped_coverage
24
+
25
+
26
+ def run_gicdm_metrics(Xr, Xg, k=5, K1=5, K2=50):
27
+ Df, keep, Drr_g = gicdm(Xr, Xg, K1, K2)
28
+ cd = clipped_density(Xr, Xg, k=k, dissim=(Df, Drr_g), keep=keep)
29
+ cc = clipped_coverage(Xr, Xg, k=k, dissim=(Df, Drr_g), keep=keep)
30
+ cd0 = clipped_density(Xr, Xg, k=k)
31
+ cc0 = clipped_coverage(Xr, Xg, k=k)
32
+ return dict(cd0=cd0, cc0=cc0, cd=cd, cc=cc)
33
+
34
+
35
+ def scenario_gauss_mean(d=100, n=1500, shift=0.0, seed=0):
36
+ rng = np.random.default_rng(seed)
37
+ Xr = rng.normal(0, 1, size=(n, d))
38
+ Xg = rng.normal(shift, 1, size=(n, d))
39
+ return Xr, Xg
40
+
41
+
42
+ def scenario_gauss_std(d=100, n=1500, s=1.0, seed=0):
43
+ rng = np.random.default_rng(seed)
44
+ Xr = rng.normal(0, 1, size=(n, d))
45
+ Xg = rng.normal(0, s, size=(n, d))
46
+ return Xr, Xg
47
+
48
+
49
+ def scenario_hypersphere(d=100, n=1500, r_r=1.0, r_g=1.0, seed=0):
50
+ rng = np.random.default_rng(seed)
51
+ def shell(r, m):
52
+ v = rng.normal(size=(m, d)); v /= np.linalg.norm(v, axis=1, keepdims=True)
53
+ return v * r
54
+ return shell(r_r, n), shell(r_g, n)
55
+
56
+
57
+ def scenario_mode_collapse(d=100, n=1500, n_modes=5, collapse=False, seed=0):
58
+ rng = np.random.default_rng(seed)
59
+ centers = rng.normal(0, 5, size=(n_modes, d))
60
+ if collapse:
61
+ # generated collapses to a single mode
62
+ c = centers[0]
63
+ Xg = c[None, :] + rng.normal(0, 0.3, size=(n, d))
64
+ else:
65
+ sel = rng.integers(0, n_modes, size=n)
66
+ Xg = centers[sel] + rng.normal(0, 0.3, size=(n, d))
67
+ selr = rng.integers(0, n_modes, size=n)
68
+ Xr = centers[selr] + rng.normal(0, 0.3, size=(n, d))
69
+ return Xr, Xg
70
+
71
+
72
+ def main():
73
+ results = {}
74
+ # GAUSSIAN MEAN DIFFERENCE: at shift=0 generated==real -> CD/CC should be ~1 (ideal)
75
+ Xr, Xg = scenario_gauss_mean(shift=0.0)
76
+ m = run_gicdm_metrics(Xr, Xg)
77
+ results['gauss_mean_ideal'] = m # expect cd0,cd ~1 (pass)
78
+
79
+ # GAUSSIAN STD DEVIATION DIFFERENCE: mismatch in scale -> raw metric distorted in HD
80
+ Xr, Xg = scenario_gauss_std(s=1.5)
81
+ m = run_gicdm_metrics(Xr, Xg)
82
+ results['gauss_std'] = m
83
+
84
+ # HYPERSPHERE SURFACE: real on r=1, generated on r=1.0 -> equal -> CD ~1
85
+ Xr, Xg = scenario_hypersphere(r_r=1.0, r_g=1.0)
86
+ m = run_gicdm_metrics(Xr, Xg)
87
+ results['hypersphere_equal'] = m
88
+
89
+ # MODE COLLAPSE: generated collapses to 1 mode -> coverage should drop
90
+ Xr, Xg = scenario_mode_collapse(collapse=True)
91
+ m = run_gicdm_metrics(Xr, Xg)
92
+ results['mode_collapse'] = m
93
+
94
+ # MODE collapse absent -> coverage ~1
95
+ Xr, Xg = scenario_mode_collapse(collapse=False)
96
+ m = run_gicdm_metrics(Xr, Xg)
97
+ results['mode_full'] = m
98
+
99
+ for k_, v in results.items():
100
+ print(f"{k_:22s} CD raw={v['cd0']:.3f} GICDM={v['cd']:.3f} | "
101
+ f"CC raw={v['cc0']:.3f} GICDM={v['cc']:.3f}")
102
+ import json
103
+ json.dump(results, open("results/claim5_benchmark.json", "w"))
104
+
105
+
106
+ if __name__ == "__main__":
107
+ main()
claim5_official.py ADDED
@@ -0,0 +1,79 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Claim 5 (faithful, scaled): Raisa et al. synthetic benchmark mechanism.
2
+
3
+ Uses the OFFICIAL GICDM repo (github.com/nicolassalvy/GICDM) Clipped Density /
4
+ Clipped Coverage with the standard vs GICDM DataProcessor. We run representative
5
+ Raisa scenarios on modest synthetic data (scaled down N,d) and tally whether each
6
+ metric behaves correctly with/without GICDM. We report the *direction* of change
7
+ and note the paper's full pass counts (Clipped Density Purpose 8/14 -> 10/14,
8
+ Bounds 8/13 -> 11/13; Clipped Coverage 8/14->10/14, 9/13->11/13).
9
+ """
10
+ import numpy as np
11
+ import sys, os
12
+ sys.path.insert(0, "/Users/equan_p/Developer/playground/ICML-2/official/GICDM")
13
+ from metrics.hubness_processor.standard import DataProcessorStandard
14
+ from metrics.hubness_processor.gicdm import GICDM
15
+ from metrics.metrics.clipped_density_coverage import ClippedDensityCoverage
16
+
17
+
18
+ def make_processor(Xr, use_gicdm, K=5):
19
+ if use_gicdm:
20
+ return GICDM(Xr, K=K, n_jobs=4, scale_factor=10)
21
+ return DataProcessorStandard(Xr, K=K, n_jobs=4)
22
+
23
+
24
+ def cd_cc(Xr, Xg, use_gicdm, K=5):
25
+ dp = make_processor(Xr, use_gicdm, K)
26
+ cdc = ClippedDensityCoverage(dp)
27
+ cd = cdc.clipped_density(Xg)
28
+ cc = cdc.clipped_coverage(Xg)
29
+ return float(cd), float(cc)
30
+
31
+
32
+ def sphere(d, n, r, rng):
33
+ v = rng.normal(size=(n, d)); v /= np.linalg.norm(v, axis=1, keepdims=True)
34
+ return v * r
35
+
36
+
37
+ def main():
38
+ rng = np.random.default_rng(0)
39
+ d, n = 100, 1500
40
+ rows = []
41
+ scenarios = {}
42
+
43
+ # 1. GAUSSIAN MEAN DIFFERENCE: shift=0 -> identical -> CD,CC should be ~1
44
+ Xr = rng.normal(0, 1, size=(n, d)); Xg0 = rng.normal(0, 1, size=(n, d))
45
+ scenarios['gauss_mean_equal'] = (Xr.copy(), Xg0.copy())
46
+
47
+ # 2. GAUSSIAN STD DEVIATION DIFFERENCE: scale mismatch
48
+ Xg_std = rng.normal(0, 1.5, size=(n, d))
49
+ scenarios['gauss_std_1p5'] = (Xr.copy(), Xg_std.copy())
50
+
51
+ # 3. HYPERSPHERE SURFACE equal radius (should be ~1)
52
+ sr = sphere(d, n, 1.0, rng); sg = sphere(d, n, 1.0, rng)
53
+ scenarios['hypersphere_equal'] = (sr, sg)
54
+
55
+ # 4. MODE COLLAPSE: gen collapses to one of 5 modes
56
+ centers = rng.normal(0, 5, size=(5, d))
57
+ Xr_m = centers[rng.integers(0, 5, n)] + rng.normal(0, 0.3, (n, d))
58
+ Xg_mc = centers[0] + rng.normal(0, 0.3, (n, d))
59
+ scenarios['mode_collapse'] = (Xr_m.copy(), Xg_mc.copy())
60
+ Xg_mf = centers[rng.integers(0, 5, n)] + rng.normal(0, 0.3, (n, d))
61
+ scenarios['mode_full'] = (Xr_m.copy(), Xg_mf.copy())
62
+
63
+ # 5. SPHERE vs TORUS-like (sphere vs scaled sphere offset)
64
+ st = sphere(d, n, 1.0, rng) + 3.0 # offset sphere -> out of manifold
65
+ scenarios['sphere_offset'] = (sr.copy(), st.copy())
66
+
67
+ results = {}
68
+ for name, (Xr_, Xg_) in scenarios.items():
69
+ cd0, cc0 = cd_cc(Xr_, Xg_, False)
70
+ cdg, ccg = cd_cc(Xr_, Xg_, True)
71
+ results[name] = dict(cd_raw=cd0, cc_raw=cc0, cd_gicdm=cdg, cc_gicdm=ccg)
72
+ print(f"{name:18s} CD raw={cd0:.3f} GICDM={cdg:.3f} | CC raw={cc0:.3f} GICDM={ccg:.3f}")
73
+
74
+ import json
75
+ json.dump(results, open("/Users/equan_p/Developer/playground/ICML-2/repro_gicdm/results/claim5_official.json", "w"))
76
+
77
+
78
+ if __name__ == "__main__":
79
+ main()
claim6.py ADDED
@@ -0,0 +1,34 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Claim 6 verification: hubness statistics h5_1(1%) and A5 (Table 4).
2
+
3
+ Reproduce Figure 7's Gaussian hubness-evolution: as dimension d increases, h5_1(1%)
4
+ and A5 rise (hubness appears). Confirm the landmark reference values:
5
+ - no-hubness baseline (low d Gaussian): h5_1(1%) slightly above 2, A5 < 0.01
6
+ - high-d Gaussian: h5_1(1%) several, A5 ~0.1+ (matches ImageNet/DINOv2 order)
7
+
8
+ We also reproduce Table 4's cross-embedding ordering qualitatively by noting the
9
+ paper's empirical numbers; full reproduction of the exact 16 embeddings requires the
10
+ official datasets/checkpoints (see logbook: github.com/nicolassalvy/GICDM). Here we
11
+ validate the *statistical methodology* and the baseline/no-hubness reference.
12
+ """
13
+ import numpy as np
14
+ import sys, os
15
+ sys.path.insert(0, os.path.dirname(__file__))
16
+ from gicdm_core import pairwise_sq_dists, hubness_stats
17
+
18
+
19
+ def gaussian_hubness(N=20000, dims=(10, 20, 50, 100, 200, 500, 1000, 2000), seed=0):
20
+ rng = np.random.default_rng(seed)
21
+ rows = []
22
+ for d in dims:
23
+ X = rng.normal(0, 1, size=(N, d))
24
+ D = pairwise_sq_dists(X)
25
+ h5, A5 = hubness_stats(D, k=5)
26
+ rows.append(dict(d=int(d), N=N, h5=float(h5), A5=float(A5)))
27
+ print(f"d={d:5d} h5_1(1%)={h5:.2f} A5={A5:.3f}")
28
+ return rows
29
+
30
+
31
+ if __name__ == "__main__":
32
+ import json
33
+ rows = gaussian_hubness()
34
+ json.dump(rows, open("results/claim6_gaussian_hubness.json", "w"))
configs/dsp-dior.yaml ADDED
@@ -0,0 +1,99 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ seed: 42
2
+ task_name: dsp-dior
3
+ ckpt_dir: ./ckpt
4
+ accelerator:
5
+ gradient_accumulation_steps: 80
6
+ mixed_precision: 'no'
7
+ report_to: tensorboard
8
+ model:
9
+ name: dsp
10
+ sd15_weight_path: ./pretrained/stable-diffusion-v1-5
11
+ clip_weight_path: ./pretrained/clip-vit-large-patch14
12
+ dinov2_vitl14_path: ./pretrained/dinov2_vitl14_pretrain.pth
13
+ clip_vit_b16_path: ./pretrained/ViT-B-16.pt
14
+ exemplar_pool:
15
+ data_embeds_dict_path: ./data/DIOR/dior_emb.pt
16
+ exemplar_pool_path: ./data/DIOR/images
17
+ image_proj_model:
18
+ dim: 1280
19
+ depth: 4
20
+ dim_head: 64
21
+ num_queries: [16, 8, 8]
22
+ ff_mult: 4
23
+ dataset:
24
+ name: dior
25
+ config: default
26
+ resolution: 512
27
+ ref_resolution: 224
28
+ categories:
29
+ all: [
30
+ vehicle, baseballfield, groundtrackfield, windmill, bridge,
31
+ overpass, ship, airplane, tenniscourt, airport,
32
+ expressway-service-area, basketballcourt, stadium, storagetank, chimney,
33
+ dam, expressway-toll-station, golffield, trainstation, harbor
34
+ ]
35
+ base: [
36
+ vehicle, baseballfield, groundtrackfield, bridge, overpass,
37
+ ship, airplane, tenniscourt, expressway-service-area, basketballcourt,
38
+ stadium, storagetank, expressway-toll-station, golffield, harbor
39
+ ]
40
+ novel: [windmill, airport, chimney, dam, trainstation]
41
+ data_files:
42
+ train:
43
+ base:
44
+ default: ./data/DIOR/metadatas/data_setting1/train_base.jsonl
45
+ novel:
46
+ airport: ./data/DIOR/metadatas/data_setting1/train_novel_airport.jsonl
47
+ chimney: ./data/DIOR/metadatas/data_setting1/train_novel_chimney.jsonl
48
+ dam: ./data/DIOR/metadatas/data_setting1/train_novel_dam.jsonl
49
+ trainstation: ./data/DIOR/metadatas/data_setting1/train_novel_trainstation.jsonl
50
+ windmill: ./data/DIOR/metadatas/data_setting1/train_novel_windmill.jsonl
51
+ infer:
52
+ base:
53
+ default: ./data/DIOR/metadatas/data_setting1/test_base.jsonl
54
+ novel:
55
+ airport: ./data/DIOR/metadatas/data_setting1/test_novel_airport.jsonl
56
+ chimney: ./data/DIOR/metadatas/data_setting1/test_novel_chimney.jsonl
57
+ dam: ./data/DIOR/metadatas/data_setting1/test_novel_dam.jsonl
58
+ trainstation: ./data/DIOR/metadatas/data_setting1/test_novel_trainstation.jsonl
59
+ windmill: ./data/DIOR/metadatas/data_setting1/test_novel_windmill.jsonl
60
+ column_names: [image, captions, bndboxes, obboxes]
61
+ novel_settings:
62
+ k_shot: 5
63
+ shuffle_seed: 42
64
+ image_patch_path: ./data/DIOR/patches
65
+ transform: LayoutTransform
66
+ training:
67
+ entry: train_dsp
68
+ optimizer:
69
+ name: AdamWGating
70
+ learning_rate: 1.e-4
71
+ weight_decay: 1.e-2
72
+ adam_beta: [0.9, 0.999]
73
+ adam_epsilon: 1.e-08
74
+ scheduler:
75
+ lr_scheduler: constant
76
+ lr_warmup_steps: 0
77
+ batch_size: 1
78
+ num_workers: 0
79
+ noise_offset: 0
80
+ input_perturbation: 0
81
+ prediction_type: null
82
+ max_grad_norm: 1.
83
+ base:
84
+ num_train_epochs: 100
85
+ max_train_steps: null
86
+ ckpt_interval_steps: 1200
87
+ novel:
88
+ num_train_epochs: null
89
+ max_train_steps: 100
90
+ ckpt_interval_steps: 100
91
+ base_ckpt_steps: 1200
92
+ inference:
93
+ entry: infer_dsp
94
+ output_dir: ./outputs
95
+ seed: 42
96
+ base:
97
+ ckpt_steps: 1200
98
+ novel:
99
+ ckpt_steps: 100
configs/dsp-exdark.yaml ADDED
@@ -0,0 +1,91 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ seed: 42
2
+ task_name: dsp-exdark
3
+ ckpt_dir: ./ckpt
4
+ accelerator:
5
+ gradient_accumulation_steps: 80
6
+ mixed_precision: 'no'
7
+ report_to: tensorboard
8
+ model:
9
+ name: dsp
10
+ sd15_weight_path: ./pretrained/stable-diffusion-v1-5
11
+ clip_weight_path: ./pretrained/clip-vit-large-patch14
12
+ dinov2_vitl14_path: ./pretrained/dinov2_vitl14_pretrain.pth
13
+ clip_vit_b16_path: ./pretrained/ViT-B-16.pt
14
+ exemplar_pool:
15
+ data_embeds_dict_path: ./data/EXDARK/exdark_emb.pt
16
+ exemplar_pool_path: ./data/EXDARK/images/train
17
+ image_proj_model:
18
+ dim: 1280
19
+ depth: 4
20
+ dim_head: 64
21
+ num_queries: [16, 8, 8]
22
+ ff_mult: 4
23
+ dataset:
24
+ name: dior
25
+ config: default
26
+ resolution: 512
27
+ ref_resolution: 224
28
+ categories:
29
+ all: [
30
+ bicycle, boat, bottle, bus, car, cat,
31
+ chair, cup, dog, motorbike, people, table
32
+ ]
33
+ base: [bicycle, boat, bottle, car, cat, chair, cup, people]
34
+ novel: [bus, dog, motorbike, table]
35
+ data_files:
36
+ train:
37
+ base:
38
+ default: ./data/EXDARK/metadatas/data_setting1/train_base.jsonl
39
+ novel:
40
+ bus: ./data/EXDARK/metadatas/data_setting1/train_novel_bus.jsonl
41
+ dog: ./data/EXDARK/metadatas/data_setting1/train_novel_dog.jsonl
42
+ motorbike: ./data/EXDARK/metadatas/data_setting1/train_novel_motorbike.jsonl
43
+ table: ./data/EXDARK/metadatas/data_setting1/train_novel_table.jsonl
44
+ infer:
45
+ base:
46
+ default: [./data/EXDARK/metadatas/data_setting1/test_base.jsonl, ./data/EXDARK/metadatas/data_setting1/val_base.jsonl]
47
+ novel:
48
+ bus: [./data/EXDARK/metadatas/data_setting1/val_novel_bus.jsonl, ./data/EXDARK/metadatas/data_setting1/test_novel_bus.jsonl]
49
+ dog: [./data/EXDARK/metadatas/data_setting1/val_novel_dog.jsonl, ./data/EXDARK/metadatas/data_setting1/test_novel_dog.jsonl]
50
+ motorbike: [./data/EXDARK/metadatas/data_setting1/val_novel_motorbike.jsonl, ./data/EXDARK/metadatas/data_setting1/test_novel_motorbike.jsonl]
51
+ table: [./data/EXDARK/metadatas/data_setting1/val_novel_table.jsonl, ./data/EXDARK/metadatas/data_setting1/test_novel_table.jsonl]
52
+ column_names: [image, captions, bndboxes, obboxes]
53
+ novel_settings:
54
+ k_shot: 5
55
+ shuffle_seed: 42
56
+ image_patch_path: ./data/EXDARK/patches
57
+ transform: LayoutTransform
58
+ training:
59
+ entry: train_dsp
60
+ optimizer:
61
+ name: AdamWGating
62
+ learning_rate: 1.e-4
63
+ weight_decay: 1.e-2
64
+ adam_beta: [0.9, 0.999]
65
+ adam_epsilon: 1.e-08
66
+ scheduler:
67
+ lr_scheduler: constant
68
+ lr_warmup_steps: 0
69
+ batch_size: 1
70
+ num_workers: 0
71
+ noise_offset: 0
72
+ input_perturbation: 0
73
+ prediction_type: null
74
+ max_grad_norm: 1.
75
+ base:
76
+ num_train_epochs: 100
77
+ max_train_steps: null
78
+ ckpt_interval_steps: 600
79
+ novel:
80
+ num_train_epochs: null
81
+ max_train_steps: 100
82
+ ckpt_interval_steps: 100
83
+ base_ckpt_steps: 600
84
+ inference:
85
+ entry: infer_dsp
86
+ output_dir: ./outputs
87
+ seed: 42
88
+ base:
89
+ ckpt_steps: 600
90
+ novel:
91
+ ckpt_steps: 100
configs/dsp-ruod.yaml ADDED
@@ -0,0 +1,91 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ seed: 42
2
+ task_name: dsp-ruod
3
+ ckpt_dir: ./ckpt
4
+ accelerator:
5
+ gradient_accumulation_steps: 80
6
+ mixed_precision: 'no'
7
+ report_to: tensorboard
8
+ model:
9
+ name: dsp
10
+ sd15_weight_path: ./pretrained/stable-diffusion-v1-5
11
+ clip_weight_path: ./pretrained/clip-vit-large-patch14
12
+ dinov2_vitl14_path: ./pretrained/dinov2_vitl14_pretrain.pth
13
+ clip_vit_b16_path: ./pretrained/ViT-B-16.pt
14
+ exemplar_pool:
15
+ data_embeds_dict_path: ./data/RUOD/ruod_emb.pt
16
+ exemplar_pool_path: ./data/RUOD/images/train
17
+ image_proj_model:
18
+ dim: 1280
19
+ depth: 4
20
+ dim_head: 64
21
+ num_queries: [16, 8, 8]
22
+ ff_mult: 4
23
+ dataset:
24
+ name: dior
25
+ config: default
26
+ resolution: 512
27
+ ref_resolution: 224
28
+ categories:
29
+ all: [
30
+ holothurian, echinus, scallop, starfish, fish,
31
+ corals, diver, cuttlefish, turtle, jellyfish,
32
+ ]
33
+ base: [holothurian, echinus, scallop, starfish, fish, diver]
34
+ novel: [corals, cuttlefish, turtle, jellyfish]
35
+ data_files:
36
+ train:
37
+ base:
38
+ default: ./data/RUOD/metadatas/data_setting1/train_base.jsonl
39
+ novel:
40
+ corals: ./data/RUOD/metadatas/data_setting1/train_novel_corals.jsonl
41
+ cuttlefish: ./data/RUOD/metadatas/data_setting1/train_novel_cuttlefish.jsonl
42
+ turtle: ./data/RUOD/metadatas/data_setting1/train_novel_turtle.jsonl
43
+ jellyfish: ./data/RUOD/metadatas/data_setting1/train_novel_jellyfish.jsonl
44
+ infer:
45
+ base:
46
+ default: ./data/RUOD/metadatas/data_setting1/test_base.jsonl
47
+ novel:
48
+ corals: ./data/RUOD/metadatas/data_setting1/test_novel_corals.jsonl
49
+ cuttlefish: ./data/RUOD/metadatas/data_setting1/test_novel_cuttlefish.jsonl
50
+ turtle: ./data/RUOD/metadatas/data_setting1/test_novel_turtle.jsonl
51
+ jellyfish: ./data/RUOD/metadatas/data_setting1/test_novel_jellyfish.jsonl
52
+ column_names: [image, captions, bndboxes, obboxes]
53
+ novel_settings:
54
+ k_shot: 5
55
+ shuffle_seed: 42
56
+ image_patch_path: ./data/RUOD/patches
57
+ transform: LayoutTransform
58
+ training:
59
+ entry: train_dsp
60
+ optimizer:
61
+ name: AdamWGating
62
+ learning_rate: 1.e-4
63
+ weight_decay: 1.e-2
64
+ adam_beta: [0.9, 0.999]
65
+ adam_epsilon: 1.e-08
66
+ scheduler:
67
+ lr_scheduler: constant
68
+ lr_warmup_steps: 0
69
+ batch_size: 1
70
+ num_workers: 0
71
+ noise_offset: 0
72
+ input_perturbation: 0
73
+ prediction_type: null
74
+ max_grad_norm: 1.
75
+ base:
76
+ num_train_epochs: 100
77
+ max_train_steps: null
78
+ ckpt_interval_steps: 1200
79
+ novel:
80
+ num_train_epochs: null
81
+ max_train_steps: 100
82
+ ckpt_interval_steps: 100
83
+ base_ckpt_steps: 1200
84
+ inference:
85
+ entry: infer_dsp
86
+ output_dir: ./outputs
87
+ seed: 42
88
+ base:
89
+ ckpt_steps: 1200
90
+ novel:
91
+ ckpt_steps: 100
crossover.py ADDED
@@ -0,0 +1,42 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Proposition 5.3: crossover dimension d* for standard Gaussian N(0, I_d).
2
+
3
+ d* solves P( median NND_k^2 <= E||X||^2 ) = 1/2, i.e. the integral
4
+ int_0^infty P( Bin(N-1, F_{chi2_d}(lambda=r)) >= k ) f_{chi2_d}(r) dr = 1/2
5
+ where F_{chi2_d}(lambda) is the noncentral chi2 CDF with noncentrality lambda,
6
+ f_{chi2_d} is the chi2 density (central, lambda=0), and r = ||X||^2 ~ chi2_d.
7
+ """
8
+ import numpy as np
9
+ from scipy import integrate, stats
10
+
11
+
12
+ def crossover_dimension(N, k, d_grid=None):
13
+ """Numerically solve Proposition 5.3 for d* given N and k."""
14
+ if d_grid is None:
15
+ d_grid = np.arange(2, 4001) # up to 4000 dims
16
+
17
+ def median_prob(d):
18
+ # P( Bin(N-1, F_{chi2_d}(lambda=r)) >= k ) averaged over r ~ chi2_d
19
+ def integrand(r):
20
+ # noncentral chi2 cdf with noncentrality lambda = r
21
+ # F_{chi2_d}(lambda=r)(d) = P( chi2_d(r) <= d )
22
+ p = stats.ncx2.cdf(d, d, r)
23
+ return stats.binom.cdf(k - 1, N - 1, p, loc=1) # P(Bin >= k) = 1 - P(Bin <= k-1)
24
+ # integrate r over chi2_d density
25
+ val, _ = integrate.quad(lambda r: stats.chi2.pdf(r, d) * (1 - stats.binom.cdf(k - 1, N - 1, stats.ncx2.cdf(d, d, r))),
26
+ 0, 200 + 20 * d, limit=200)
27
+ return val
28
+
29
+ probs = []
30
+ for d in d_grid:
31
+ probs.append(median_prob(d))
32
+ probs = np.array(probs)
33
+ # find d* where prob crosses 1/2
34
+ idx = np.argmin(np.abs(probs - 0.5))
35
+ return d_grid[idx], probs
36
+
37
+
38
+ if __name__ == "__main__":
39
+ for N in [1000, 10000, 50000]:
40
+ for k in [1, 5, 10, 20]:
41
+ dstar, _ = crossover_dimension(N, k)
42
+ print(f"N={N:6d} k={k:3d} -> d*={dstar}")
databuilders/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .dior import Dior as dior
databuilders/dior.py ADDED
@@ -0,0 +1,113 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import json
3
+ import datasets
4
+ from PIL import Image
5
+ from dataclasses import dataclass
6
+
7
+ @dataclass
8
+ class DiorConfig(datasets.BuilderConfig):
9
+ """BuilderConfig for Dior dataset."""
10
+ pass
11
+
12
+
13
+ class Dior(datasets.GeneratorBasedBuilder):
14
+ """DIOR Dataset."""
15
+
16
+ VERSION = datasets.Version("1.0.0")
17
+
18
+ BUILDER_CONFIG_CLASS = DiorConfig
19
+
20
+ BUILDER_CONFIGS = [
21
+ DiorConfig(
22
+ name="default",
23
+ description="Default configuration for the DIOR dataset.",
24
+ ),
25
+ ]
26
+
27
+ DEFAULT_CONFIG_NAME = "default"
28
+
29
+ def _info(self):
30
+
31
+ return datasets.DatasetInfo(
32
+ description="The DIOR Dataset.",
33
+ features=datasets.Features({
34
+ "image": datasets.Image(),
35
+ "captions": datasets.Sequence(datasets.Value("string")),
36
+ "bndboxes": datasets.Array2D(shape=(None, 4), dtype="float32"),
37
+ "obboxes": datasets.Array2D(shape=(None, 8), dtype="float32"),
38
+ "dataid": datasets.Value("string")
39
+ }),
40
+ homepage="http://www.example.com/",
41
+ citation="",
42
+ )
43
+
44
+ def _split_generators(self, dl_manager):
45
+
46
+ data_files = dl_manager.download_and_extract(self.config.data_files)
47
+
48
+ if not data_files or not isinstance(data_files, dict):
49
+ raise ValueError(
50
+ "This builder requires you to pass data_files as a dictionary."
51
+ "for example: data_files={'train': 'path/to/train.jsonl', 'test': 'path/to/test.jsonl'}"
52
+ )
53
+ split_generators = []
54
+
55
+ for split_name, metadata_list in data_files.items():
56
+ split_generators.append(
57
+ datasets.SplitGenerator(
58
+ name=split_name,
59
+ gen_kwargs={"metadata_list": metadata_list},
60
+ )
61
+ )
62
+
63
+ return split_generators
64
+
65
+ def _generate_examples(self, metadata_list):
66
+
67
+ idx = 0
68
+
69
+ for metadata_path in metadata_list:
70
+ # print(f"--> [Generator] Processing file: {metadata_path}")
71
+ base_dir = os.path.dirname(os.path.abspath(metadata_path))
72
+
73
+ with open(metadata_path, 'r', encoding='utf-8') as f:
74
+ for i, line in enumerate(f):
75
+ try:
76
+ data = json.loads(line)
77
+
78
+ absolute_image_path = os.path.join(base_dir, data["file_name"])
79
+ # image = Image.open(absolute_image_path).convert("RGB")
80
+
81
+ example = {
82
+ "image": absolute_image_path,
83
+ "captions": data.get("captions", []),
84
+ "bndboxes": data.get("bndboxes", []),
85
+ "obboxes": data.get("obboxes", []),
86
+ "dataid": os.path.splitext(os.path.basename(data["file_name"]))[0]
87
+ }
88
+
89
+ yield idx, example
90
+
91
+ idx += 1
92
+
93
+ except Exception as e:
94
+ print(f" - Skip Invalid Data ({metadata_path} Line {i}): {e}")
95
+
96
+
97
+ if __name__ == "__main__":
98
+ data_dir = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), "data", "DIOR")
99
+ # data_files = {
100
+ # "train": os.path.join(data_dir, "train_meta.jsonl"),
101
+ # "test": os.path.join(data_dir, "test_meta.jsonl"),
102
+ # }
103
+ data_files = {
104
+ "train": [os.path.join(data_dir, "train_meta_sample1.jsonl"),
105
+ os.path.join(data_dir, "train_meta_sample2.jsonl")],
106
+ "test": os.path.join(data_dir, "infer_meta_sample1.jsonl"),
107
+ }
108
+
109
+ builder = Dior(data_files=data_files, config_name="default")
110
+ builder.download_and_prepare()
111
+ dataset = builder.as_dataset()
112
+
113
+ print(dataset)
datamodules/__init__.py ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ from .loader import Loader
2
+ from .ref_table import RefTable
datamodules/loader.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from . import transforms
2
+ from .ref_table import RefTable
3
+ from utils import get_ckpt_path
4
+ import databuilders
5
+ import datasets
6
+ import torch
7
+ import json
8
+ import os
9
+
10
+
11
+ class Loader:
12
+ def __init__(self, config, image_processor, split='train', logger=None):
13
+ self.logger = logger
14
+ self.split, self.phase = split, config.phase
15
+ self.data_name = config.dataset.name
16
+ # self.data_files = config.dataset.data_files
17
+ self.data_files = config.dataset.data_files[self.split][self.phase]
18
+ self.categories = config.dataset.categories[self.phase]
19
+ self.image_column, self.caption_column, self.bbox_column, self.obbox_column = config.dataset.column_names
20
+ self.batch_size = config.training.batch_size
21
+ self.num_workers = config.training.num_workers
22
+ self.max_inference_size = config.inference.get('max_inference_size', None)
23
+ self.novel_sample_dict = None
24
+
25
+ self.set_dataset()
26
+ if self.phase == 'novel':
27
+ self.k_shot = config.dataset.novel_settings.k_shot
28
+ self.shuffle_seed = config.dataset.novel_settings.get('shuffle_seed', 42)
29
+ self.dump_file = os.path.join(get_ckpt_path(config), 'novel_sample_dict.json')
30
+ if self.split == 'train':
31
+ self.sample_dataset()
32
+ elif self.split == 'infer':
33
+ self.load_novel_sample_dict()
34
+ self.sample_dataset_infer_phase()
35
+
36
+ self.dataset = datasets.concatenate_datasets(self.dataset.values())
37
+ self.ref_table = RefTable(config, filter_dict=self.novel_sample_dict)
38
+ self.transform = getattr(transforms, config.dataset.get('transform', 'DefaultTransform'))(config, image_processor, self.split, ref_table=self.ref_table())
39
+ self.dataset = self.dataset.with_transform(self.transform)
40
+
41
+ def set_dataset(self):
42
+ builder = getattr(databuilders, self.data_name.lower(), None)
43
+ if builder is None:
44
+ raise ValueError(f"Unknown dataset: {self.data_name}")
45
+ builder = builder(data_files=self.data_files)
46
+ builder.download_and_prepare()
47
+ self.dataset = builder.as_dataset()
48
+
49
+ def sample_dataset(self):
50
+ ''' The data sample logics for few-shot learning. '''
51
+ self.novel_sample_dict = {}
52
+ for category in self.categories:
53
+ self.dataset[category] = self.dataset[category].shuffle(self.shuffle_seed).select(range(min(self.k_shot, len(self.dataset[category]))))
54
+ self.novel_sample_dict[category] = list(self.dataset[category]['dataid'])
55
+
56
+ def sample_dataset_infer_phase(self):
57
+ if self.max_inference_size is not None:
58
+ # self.dataset['default'] = self.dataset['default'].shuffle(self.shuffle_seed).select(range(self.max_inference_size))
59
+ for category in self.categories:
60
+ self.dataset[category] = self.dataset[category].shuffle(self.shuffle_seed).select(range(min(self.max_inference_size, len(self.dataset[category]))))
61
+
62
+ def collate_fn(self, examples):
63
+ images = torch.stack([example[self.image_column] for example in examples])
64
+ images = images.to(memory_format=torch.contiguous_format).float()
65
+ captions = [example[self.caption_column] for example in examples]
66
+ bboxes = [example[self.bbox_column] for example in examples]
67
+ obboxes = [example[self.obbox_column] for example in examples]
68
+ if isinstance(examples[0]["instances"], list):
69
+ instances = [example["instances"] for example in examples]
70
+ else:
71
+ instances = torch.stack([example["instances"] for example in examples])
72
+ dataid = [example["dataid"] for example in examples]
73
+ # Custom Keys
74
+ if 'masks' in examples[0].keys():
75
+ masks = torch.stack([example["masks"] for example in examples])
76
+ return {self.image_column: images, self.caption_column: captions, self.bbox_column: bboxes, self.obbox_column: obboxes, "instances": instances, "dataid": dataid, "masks": masks}
77
+ if 'parallels' in examples[0].keys():
78
+ # parallels = torch.cat([example["parallels"] for example in examples])
79
+ # images = torch.cat([images, parallels], dim=0)
80
+ parallels = [example["parallels"] for example in examples]
81
+ return {self.image_column: images, self.caption_column: captions, self.bbox_column: bboxes, self.obbox_column: obboxes, "instances": instances, "dataid": dataid, "parallels": parallels}
82
+ return {self.image_column: images, self.caption_column: captions, self.bbox_column: bboxes, self.obbox_column: obboxes, "instances": instances, "dataid": dataid}
83
+
84
+ def __call__(self):
85
+ # for i in range(6):
86
+ # _ = self.dataset[i]
87
+ dataloader = torch.utils.data.DataLoader(
88
+ self.dataset,
89
+ shuffle=True,
90
+ collate_fn=self.collate_fn,
91
+ batch_size=self.batch_size,
92
+ num_workers=self.num_workers,
93
+ )
94
+ return dataloader
95
+
96
+ def dump_novel_sample_dict(self):
97
+ if self.phase == 'novel' and self.logger is not None:
98
+ self.logger.info(f'Novel Sample Dict: \n{json.dumps(self.novel_sample_dict, indent=2)}', main_process_only=False)
99
+ self.logger.info(f'Novel Ref Table: \n{json.dumps(self.ref_table("novel"), indent=2)}', main_process_only=False)
100
+ self.logger.info(f'Dump novel_sample_dict at {self.dump_file}')
101
+ with open(self.dump_file, 'w') as f:
102
+ json.dump(self.novel_sample_dict, f, indent=2)
103
+
104
+ def load_novel_sample_dict(self):
105
+ if self.phase == 'novel' and self.logger is not None:
106
+ self.logger.info(f'Load novel_sample_dict at {self.dump_file}')
107
+ with open(self.dump_file, 'r') as f:
108
+ self.novel_sample_dict = json.load(f)
109
+ self.logger.info(f'Novel Sample Dict: \n{json.dumps(self.novel_sample_dict, indent=2)}', main_process_only=False)
datamodules/ref_table.py ADDED
@@ -0,0 +1,76 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import functools
2
+ import imagesize
3
+ import torch
4
+ import os
5
+ import json
6
+
7
+
8
+ def singleton(cls):
9
+ instances = {}
10
+ def get_instance(*args, **kwargs):
11
+ if cls not in instances:
12
+ instances[cls] = cls(*args, **kwargs)
13
+ return instances[cls]
14
+ return get_instance
15
+
16
+
17
+ @singleton
18
+ class RefTable:
19
+ def __init__(self, config, filter_dict=None):
20
+ self.phase = config.phase
21
+ self.categories_base, self.categories_novel = config.dataset.categories.base, config.dataset.categories.novel
22
+ self.categories = config.dataset.categories.get('all', None) or (self.categories_novel + self.categories_base)
23
+ self.image_patch_path = config.dataset.image_patch_path
24
+ self.filter_dict = filter_dict
25
+ self.augment = config.dataset.get('ref_augment', False)
26
+
27
+ cache_path = os.path.join(self.image_patch_path, "image_sizes.json")
28
+ if not os.path.exists(cache_path):
29
+ raise FileNotFoundError(f"File Not Found: {cache_path}")
30
+ with open(cache_path, 'r') as f:
31
+ self.img_size_cache = json.load(f)
32
+ self.ref_table, self.base_ref_table, self.novel_ref_table = self.build_ref_table()
33
+
34
+ def build_ref_table(self):
35
+ ref_table, base_ref_table, novel_ref_table = {}, {}, {}
36
+ for category in self.categories:
37
+ category_dir = os.path.join(self.image_patch_path, category)
38
+ patch_list = os.listdir(category_dir)
39
+ patch_list = list(filter(lambda s: s.endswith('.jpg'), patch_list))
40
+
41
+ # Filter the patch list to avoid data leakage of few-shot learning
42
+ if self.phase == 'novel' and category in self.categories_novel:
43
+ assert self.filter_dict is not None
44
+ patch_list = list(filter(lambda patch_name: patch_name.rsplit('_', 1)[0] in self.filter_dict[category], patch_list))
45
+ if self.augment:
46
+ aug_category_dir = os.path.join(self.image_patch_path, category, 'augmented')
47
+ aug_patch_list = os.listdir(aug_category_dir)
48
+ aug_patch_list = list(filter(lambda patch_name: patch_name.startswith(tuple(self.filter_dict[category])), aug_patch_list))
49
+ aug_patch_list = list(map(lambda patch_name: f"augmented/{patch_name}", aug_patch_list))
50
+ patch_list += aug_patch_list
51
+
52
+ def get_size(img_name):
53
+ return self.img_size_cache.get(img_name, [0, 1])
54
+
55
+ patch_list = sorted(patch_list, key = lambda img: get_size(img)[0] * get_size(img)[1], reverse=True)
56
+
57
+ # Only build novel ref table at novel phase
58
+ if self.phase == 'novel' and category in self.categories_novel:
59
+ novel_ref_table[category] = {
60
+ img: get_size(img)[0] / get_size(img)[1] for img in patch_list[:200]
61
+ }
62
+ elif category in self.categories_base:
63
+ base_ref_table[category] = {
64
+ img: get_size(img)[0] / get_size(img)[1] for img in patch_list[:200]
65
+ }
66
+
67
+ ref_table = novel_ref_table | base_ref_table
68
+ return ref_table, base_ref_table, novel_ref_table
69
+
70
+ def __call__(self, phase=None):
71
+ if phase is None:
72
+ return self.ref_table
73
+ if phase == 'base':
74
+ return self.base_ref_table
75
+ if phase == 'novel':
76
+ return self.novel_ref_table
datamodules/transforms/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .layout_transform import LayoutTransform
datamodules/transforms/layout_transform.py ADDED
@@ -0,0 +1,132 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import albumentations as A
2
+ from torchvision import transforms
3
+ from PIL import Image
4
+ import numpy as np
5
+ import functools
6
+ import imagesize
7
+ import torch
8
+ import cv2
9
+ import os
10
+
11
+
12
+ class LayoutTransform:
13
+ def __init__(self, config, image_processor, split, ref_table=None, filter_dict=None):
14
+ self.split, self.phase = split, config.phase
15
+ if split == 'train' and self.phase == 'novel':
16
+ self.image_transforms = A.Compose([
17
+ A.OneOf([
18
+ A.RandomSizedBBoxSafeCrop(height=config.dataset.resolution, width=config.dataset.resolution, erosion_rate=0.0, interpolation=cv2.INTER_CUBIC, p=0.3),
19
+ A.Resize(height=config.dataset.resolution, width=config.dataset.resolution, interpolation=cv2.INTER_CUBIC, p=0.7),
20
+ ], p=1.0),
21
+ A.Normalize(mean=[0.5], std=[0.5]),
22
+ A.pytorch.ToTensorV2(),
23
+ ], bbox_params=A.BboxParams(format='albumentations', label_fields=['labels'], min_area=0, min_visibility=0.0))
24
+ elif split == 'infer' or self.phase == 'base':
25
+ self.image_transforms = A.Compose([
26
+ A.Resize(config.dataset.resolution, config.dataset.resolution),
27
+ A.Normalize(mean=[0], std=[1]),
28
+ A.pytorch.ToTensorV2(),
29
+ ])
30
+ else:
31
+ raise ValueError("Invalid mode for Transform.")
32
+ self.image_patch_path = config.dataset.image_patch_path
33
+ self.ref_resolution = config.dataset.ref_resolution
34
+ self.image_column, self.caption_column, self.bbox_column, self.obbox_column = config.dataset.column_names
35
+ self.image_processor = image_processor
36
+ if ref_table is not None:
37
+ self.ref_table = ref_table
38
+ else:
39
+ self.categories = config.dataset.categories[self.phase]
40
+ self.filter_dict = filter_dict
41
+ self.ref_table = self.build_ref_table()
42
+ self.k_shot = config.dataset.novel_settings.k_shot
43
+ self.top_k = 1
44
+
45
+ def build_ref_table(self):
46
+ ref_table = {}
47
+ for category in self.categories:
48
+ category_dir = os.path.join(self.image_patch_path, category)
49
+ patch_list = os.listdir(category_dir)
50
+ # Filter the patch list to avoid data leakage of few-shot learning
51
+ if self.phase == 'novel':
52
+ assert self.filter_dict is not None
53
+ patch_list = list(filter(lambda patch_name: patch_name.split('_')[0] in self.filter_dict[category], patch_list))
54
+ patch_list = sorted(patch_list, key = lambda img: functools.reduce(lambda x, y: x*y, imagesize.get(os.path.join(category_dir, img))), reverse=True)
55
+ ref_table[category] = {img: functools.reduce(lambda x, y: x/y, imagesize.get(os.path.join(category_dir, img))) for img in patch_list[:200]}
56
+ return ref_table
57
+
58
+ @staticmethod
59
+ def find_nearest(array, value):
60
+ array = np.asarray(array)
61
+ idx = (np.abs(array/value - 1)).argmin()
62
+ return idx
63
+
64
+ @staticmethod
65
+ def find_k_nearest(array, value, k=1):
66
+ array = np.asarray(array)
67
+ dist = np.abs(array / value - 1)
68
+ idxs = np.argsort(dist)[:k]
69
+ return idxs
70
+
71
+ def get_instances(self, examples):
72
+ instances = []
73
+ for index, (caption, bboxes) in enumerate(zip(examples[self.caption_column], examples[self.bbox_column])):
74
+ categories = caption[1:]
75
+ instances_per_example = []
76
+ for name, bbox in zip(categories, bboxes):
77
+ if name == '':
78
+ instances_per_example.append(torch.zeros([self.top_k, 3, self.ref_resolution, self.ref_resolution]))
79
+ else:
80
+ value = (bbox[2] - bbox[0]) / max(bbox[3] - bbox[1], 1e-8)
81
+ chosen_idxs = self.find_k_nearest(list(self.ref_table[name].values()), value, k=self.top_k)
82
+ instances_per_bbox = []
83
+ for idx in chosen_idxs:
84
+ chosen_file = list(self.ref_table[name].keys())[idx]
85
+ img = Image.open(os.path.join(self.image_patch_path, name, chosen_file)).convert('RGB')
86
+ img = self.image_processor(images=img, return_tensors="pt")['pixel_values'].squeeze(0)
87
+ instances_per_bbox.append(img)
88
+ instances_per_example.append(torch.stack(instances_per_bbox))
89
+ instances.append(torch.stack(instances_per_example))
90
+ return instances
91
+
92
+ def train_transform(self, examples):
93
+ images, bboxes, obboxes = examples[self.image_column], examples[self.bbox_column], examples[self.obbox_column]
94
+ global_prompt = [caption[:1] for caption in examples[self.caption_column]]
95
+ captions = [caption[1:] for caption in examples[self.caption_column]]
96
+
97
+ new_images, new_bboxes, new_obboxes, new_captions = [], [], [], []
98
+ for i in range(len(images)):
99
+ num_instances = sum(bool(s) for s in captions[i])
100
+ bboxes_i = bboxes[i][:num_instances]
101
+ captions_i = captions[i][:num_instances]
102
+ transformed = self.image_transforms(image=images[i], bboxes=bboxes_i, labels=captions_i)
103
+ transformed["obboxes"] = []
104
+ for xmin, ymin, xmax, ymax in transformed["bboxes"]:
105
+ transformed["obboxes"].append([xmin, ymin, xmax, ymin, xmax, ymax, xmin, ymax])
106
+ for _ in range(len(captions[i]) - num_instances):
107
+ transformed["bboxes"].append([0,0,0,0])
108
+ transformed["obboxes"].append([0,0,0,0,0,0,0,0])
109
+ transformed["labels"].append("")
110
+ new_images.append(transformed["image"])
111
+ new_bboxes.append(transformed["bboxes"])
112
+ new_obboxes.append(transformed["obboxes"])
113
+ new_captions.append(global_prompt[i] + transformed["labels"])
114
+ examples[self.image_column] = new_images
115
+ examples[self.bbox_column] = new_bboxes
116
+ examples[self.obbox_column] = new_obboxes
117
+ examples[self.caption_column] = new_captions
118
+ return examples
119
+
120
+ def infer_transform(self, examples):
121
+ examples[self.image_column] = [self.image_transforms(image=image)['image'] for image in examples[self.image_column]]
122
+ return examples
123
+
124
+ def __call__(self, examples):
125
+ # print("Received keys:", examples.keys())
126
+ examples[self.image_column] = [np.array(image.convert("RGB")) for image in examples[self.image_column]]
127
+ if self.split == "train":
128
+ examples = self.train_transform(examples)
129
+ elif self.split == "infer":
130
+ examples = self.infer_transform(examples)
131
+ examples["instances"] = self.get_instances(examples)
132
+ return examples
figs/Main.png ADDED

Git LFS Details

  • SHA256: f8e407fae7cd699632315e76dfc45a7181d5e9933639833c7dc95254d399d2ba
  • Pointer size: 131 Bytes
  • Size of remote file: 442 kB
fonts/GILI____.TTF ADDED
Binary file (69.4 kB). View file
 
fonts/Rainbow-Party-2.ttf ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:8210f1b01b549890cc25f2028c980f87b4d2a45dbc3c64be8b60e8ec9f4d9f84
3
+ size 114632
gicdm_core.py ADDED
@@ -0,0 +1,216 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Core GICDM implementation reproducing the paper arXiv:2602.16449.
2
+
3
+ Implements:
4
+ - ICDM iterative hubness reduction (Section 3 / Algorithm in paper lines 485-498)
5
+ - GICDM out-of-sample generated-point scaling (Eq. 1, Algorithm 1)
6
+ - Hubness statistics h5_1(1%) and A5 (Table 4)
7
+ - Clipped Density / Clipped Coverage fidelity & coverage metrics
8
+ - Crossover dimension d* (Proposition 5.3)
9
+ - Raisa et al. (2025) synthetic benchmark scenarios
10
+ """
11
+ import numpy as np
12
+ from scipy import stats
13
+ from scipy.special import ncfdtr, ncfdtri
14
+
15
+
16
+ # ----------------------------------------------------------------------------
17
+ # Distances & k-nearest-neighbour helpers
18
+ # ----------------------------------------------------------------------------
19
+ def pairwise_sq_dists(X, Y=None):
20
+ if Y is None:
21
+ Y = X
22
+ Xn = np.sum(X ** 2, axis=1, keepdims=True)
23
+ Yn = np.sum(Y ** 2, axis=1, keepdims=True)
24
+ D2 = Xn + Yn.T - 2.0 * (X @ Y.T)
25
+ np.maximum(D2, 0, out=D2)
26
+ return np.sqrt(D2)
27
+
28
+
29
+ def knn_distances(D, k, self_included=True):
30
+ """Return (N, k) sorted distances to k nearest neighbours.
31
+
32
+ If self_included, the 0-th distance (0) is the point itself.
33
+ """
34
+ N = D.shape[0]
35
+ if self_included:
36
+ idx = np.argpartition(D, k, axis=1)[:, :k]
37
+ else:
38
+ # exclude self (diagonal)
39
+ Dc = D + np.eye(N) * 1e18
40
+ idx = np.argpartition(Dc, k, axis=1)[:, :k]
41
+ idx = idx[np.arange(N)[:, None], np.argsort(D[np.arange(N)[:, None], idx], axis=1)]
42
+ return D[np.arange(N)[:, None], idx]
43
+
44
+
45
+ # ----------------------------------------------------------------------------
46
+ # ICDM (Iterative Contextual Dissimilarity Measure)
47
+ # Matches official implementation: metrics/hubness_processor/hubness_reduction_methods.py
48
+ # ----------------------------------------------------------------------------
49
+ def _average_k_dist(D, k):
50
+ """Average of the k nearest distances (including self at distance 0),
51
+ i.e. sum of (k+1) smallest distances / k."""
52
+ k_dists = np.partition(D, k, axis=1)[:, : k + 1]
53
+ return k_dists.sum(axis=1) / k
54
+
55
+
56
+ def icdm_scaling(D, K, n_iter=10, return_mu=False):
57
+ """ICDM scaling factors delta_i, matching official `icdm_delta_low_memory`.
58
+
59
+ Iteratively applies NICDM:
60
+ d_{ij} <- d_{ij} / (sqrt(r_i) sqrt(r_j)) with r_i = average of k nearest dists (k=K)
61
+ delta_i <- delta_i / sqrt(r_i)
62
+ Final: secondary dissimilarity d^T_{ij} = d_{ij} delta_i delta_j.
63
+ Returns delta_i (length N) and optionally final average-neighbour distances.
64
+ """
65
+ N = D.shape[0]
66
+ d = D.copy()
67
+ deltas = np.ones(N)
68
+ for _ in range(n_iter):
69
+ r = _average_k_dist(d, K)
70
+ sqrt_r = np.maximum(np.sqrt(r), 1e-12)
71
+ d = d / (sqrt_r[:, None] * sqrt_r[None, :])
72
+ deltas = deltas / sqrt_r
73
+ if return_mu:
74
+ mu_final = _average_k_dist(d, K)
75
+ return deltas, mu_final
76
+ return deltas
77
+
78
+
79
+ # ----------------------------------------------------------------------------
80
+ # GICDM (Algorithm 1)
81
+ # ----------------------------------------------------------------------------
82
+ def gicdm(Xr, Xg, K1, K2, q=0.95, n_iter=10, return_info=False):
83
+ """Generative ICDM (Algorithm 1), matching the official implementation.
84
+
85
+ Xr : (N, d) real points
86
+ Xg : (M, d) generated points
87
+ K1, K2 : two filter scales (K2 = 10*K1 in the paper). For each scale k,
88
+ ICDM is applied with neighbourhood size 2k (paper: GICDM K = 2k).
89
+ Returns:
90
+ D_final : (M, N) GICDM real-to-generated dissimilarity matrix
91
+ keep : (M,) boolean mask of generated points passing multi-scale filter
92
+ Drr_gicdm : (N, N) GICDM real-to-real dissimilarity matrix
93
+ """
94
+ Drr = pairwise_sq_dists(Xr)
95
+ Drg = pairwise_sq_dists(Xg, Xr) # (M, N): generated-to-real
96
+ N, M = Xr.shape[0], Xg.shape[0]
97
+
98
+ keep = np.ones(M, dtype=bool)
99
+ info = {}
100
+
101
+ for k in (K1, K2):
102
+ # ICDM with neighbourhood 2k
103
+ delta_r = icdm_scaling(Drr, K=2 * k, n_iter=n_iter)
104
+
105
+ # --- real filter (Algorithm 1 lines 3-6) ---
106
+ # k nearest real neighbours (exclude self)
107
+ nn_r = np.argpartition(Drr, k, axis=1)[:, : k + 1]
108
+ is_self = nn_r == np.arange(N)[:, None]
109
+ no_self = ~is_self.any(axis=1)
110
+ is_self[no_self, -1] = True
111
+ nn_r = nn_r[~is_self].reshape(N, k)
112
+ real_avg_delta = delta_r[nn_r].mean(axis=1)
113
+ r_ri = np.abs(real_avg_delta - delta_r) / real_avg_delta
114
+ T_k = np.quantile(r_ri, q)
115
+
116
+ # --- generated filter (Algorithm 1 lines 9-14) ---
117
+ # k+1 nearest real neighbours of each generated point
118
+ nn_g = np.argpartition(Drg.T, k, axis=1)[:, : k + 1]
119
+ # r_synthetic = average of (delta_real_nn * d_orig) over k+1 neighbours
120
+ d_g_nn = Drg[np.arange(M)[:, None], nn_g]
121
+ delta_r_g_nn = delta_r[nn_g]
122
+ r_synthetic = (delta_r_g_nn * d_g_nn).sum(axis=1) / (k + 1)
123
+ delta_g = 1.0 / r_synthetic # Eq. (1) with mu_bar absorbed
124
+ synth_avg_delta = delta_r_g_nn.mean(axis=1)
125
+ r_gj = np.abs(synth_avg_delta - delta_g) / synth_avg_delta
126
+ keep &= (r_gj <= T_k)
127
+
128
+ info[k] = dict(T_k=float(T_k), delta_g=delta_g)
129
+
130
+ # final GICDM dissimilarities (Algorithm 1 line 17)
131
+ delta_g_K1 = info[K1]['delta_g']
132
+ delta_r_K1 = icdm_scaling(Drr, K=2 * K1, n_iter=n_iter)
133
+ D_final = Drg * (delta_r_K1[None, :] * delta_g_K1[None, :])
134
+ Drr_gicdm = Drr * np.outer(delta_r_K1, delta_r_K1)
135
+
136
+ if return_info:
137
+ return D_final, keep, info, (delta_r_K1, delta_g_K1)
138
+ return D_final, keep, Drr_gicdm
139
+
140
+
141
+ # ----------------------------------------------------------------------------
142
+ # Hubness statistics
143
+ # ----------------------------------------------------------------------------
144
+ def k_occurrence(D, k=5):
145
+ """O_k(x_i) = number of points for which x_i is among their k NN (excl self)."""
146
+ N = D.shape[0]
147
+ Dc = D + np.eye(N) * 1e18
148
+ knn = np.argsort(Dc, axis=1)[:, :k]
149
+ occ = np.zeros(N, dtype=int)
150
+ for j in range(k):
151
+ occ[knn[:, j]] += 1
152
+ return occ
153
+
154
+
155
+ def hubness_stats(D, k=5, q=0.01):
156
+ """Return h5_1(1%) and A5 (proportion of antihubs)."""
157
+ occ = k_occurrence(D, k)
158
+ mean_occ = occ.mean()
159
+ n = len(occ)
160
+ topq = max(1, int(np.floor(q * n)))
161
+ top_vals = np.sort(occ)[::-1][:topq]
162
+ h5 = top_vals.mean() / mean_occ if mean_occ > 0 else np.nan
163
+ A5 = float(np.mean(occ == 0))
164
+ return float(h5), A5
165
+
166
+
167
+ # ----------------------------------------------------------------------------
168
+ # Clipped Density / Clipped Coverage (Salvy et al. 2026)
169
+ # ----------------------------------------------------------------------------
170
+ def clipped_density(Xr, Xg, k=5, dissim=None, keep=None):
171
+ """Clipped Density fidelity metric.
172
+
173
+ For each generated point, distance to its k-th real NN; threshold = distance
174
+ from each real point to its k-th real NN (clip). Score averages clip term.
175
+ If dissim ('gicdm') is provided, use GICDM dissimilarities instead of raw dist.
176
+ keep: boolean mask of generated points to include (filtered-out points get 0).
177
+ """
178
+ Drg = pairwise_sq_dists(Xg, Xr) # (M,N)
179
+ Drr = pairwise_sq_dists(Xr)
180
+ if dissim is None:
181
+ d_rg = Drg
182
+ d_rr = Drr
183
+ else:
184
+ d_rg, d_rr = dissim
185
+ # k-th NN distance for each real point (in its own set)
186
+ kth_real = np.sort(d_rr + np.eye(len(Xr)) * 1e18, axis=1)[:, k - 1]
187
+ kth_gen = np.sort(d_rg, axis=1)[:, k - 1]
188
+ # clip each generated point's distance at its matched real threshold
189
+ thresh = kth_real[np.argmin(d_rg, axis=1)]
190
+ clip = np.clip(kth_gen / thresh, 0, 1)
191
+ if keep is not None:
192
+ clip = clip * keep.astype(float) # filtered points -> fidelity 0
193
+ return float(np.mean(clip))
194
+
195
+
196
+ def clipped_coverage(Xr, Xg, k=5, dissim=None, keep=None):
197
+ """Clipped Coverage: for each real point, does a kept generated point fall
198
+ within its k-th NN threshold?"""
199
+ Drg = pairwise_sq_dists(Xg, Xr)
200
+ Drr = pairwise_sq_dists(Xr)
201
+ if dissim is None:
202
+ d_rg = Drg
203
+ d_rr = Drr
204
+ else:
205
+ d_rg, d_rr = dissim
206
+ kth_real = np.sort(d_rr + np.eye(len(Xr)) * 1e18, axis=1)[:, k - 1]
207
+ min_d = d_rg.min(axis=0)
208
+ covered = (min_d <= kth_real).astype(float)
209
+ if keep is not None:
210
+ # a generated point only contributes if it is kept
211
+ d_rg_k = d_rg[keep, :]
212
+ if d_rg_k.shape[0] == 0:
213
+ return 0.0
214
+ min_d2 = d_rg_k.min(axis=0)
215
+ covered = (min_d2 <= kth_real).astype(float)
216
+ return float(np.mean(covered))
infer.sh ADDED
@@ -0,0 +1,174 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ set -e
3
+
4
+ CONFIGS=()
5
+ SEEDS=()
6
+ CKPTS=()
7
+ RUN_IDS=()
8
+ GPU_IDS=""
9
+ METASEED=""
10
+ NUM_SEED=-1
11
+ META_START=0
12
+ MAX_INFER_SIZE=""
13
+ K_SHOTS=()
14
+
15
+ while [[ $# -gt 0 ]]; do
16
+ case $1 in
17
+ --config)
18
+ shift
19
+ CONFIGS=($1)
20
+ ;;
21
+ --seed)
22
+ shift
23
+ SEEDS=($1)
24
+ ;;
25
+ --metaseed)
26
+ shift
27
+ METASEED=$1
28
+ ;;
29
+ --num_seed)
30
+ shift
31
+ NUM_SEED=$1
32
+ ;;
33
+ --meta_start)
34
+ shift
35
+ META_START=$1
36
+ ;;
37
+ --ckpt)
38
+ shift
39
+ CKPTS=($1)
40
+ ;;
41
+ --run_id)
42
+ shift
43
+ RUN_IDS=($1)
44
+ ;;
45
+ --k_shot)
46
+ shift
47
+ K_SHOTS=($1)
48
+ ;;
49
+ --gpu_ids)
50
+ shift
51
+ GPU_IDS=$1
52
+ ;;
53
+ --max_infer_size)
54
+ shift
55
+ MAX_INFER_SIZE=$1
56
+ ;;
57
+ *)
58
+ echo "Unknown argument: $1"
59
+ exit 1
60
+ ;;
61
+ esac
62
+ shift
63
+ done
64
+
65
+ if [[ ${#CONFIGS[@]} -eq 0 ]] \
66
+ || [[ ${#CKPTS[@]} -eq 0 ]] \
67
+ || [[ -z "$GPU_IDS" ]] \
68
+ || [[ ${#K_SHOTS[@]} -eq 0 ]] \
69
+ || [[ ${#RUN_IDS[@]} -eq 0 ]]; then
70
+ echo "Error: missing required arguments."
71
+ echo "You must provide: --config, --gpu_ids, --run_id, --k_shot"
72
+ echo "And one of: --seed OR --metaseed"
73
+ exit 1
74
+ fi
75
+
76
+ if [[ -z "$METASEED" && ${#SEEDS[@]} -eq 0 ]]; then
77
+ echo "Error: either --seed or --metaseed must be provided."
78
+ exit 1
79
+ fi
80
+
81
+ IFS=',' read -r -a GPU_ARRAY <<< "$GPU_IDS"
82
+ NUM_PROCESSES=${#GPU_ARRAY[@]}
83
+
84
+ if [[ -n "$METASEED" ]]; then
85
+ if [[ $NUM_SEED -lt 0 ]]; then
86
+ echo "Error: you must set --num_seed when using --metaseed."
87
+ exit 1
88
+ fi
89
+
90
+ echo "[INFO] Using metaseed=$METASEED"
91
+ echo "[INFO] Will use seeds index range [$META_START, $NUM_SEED)"
92
+
93
+ mapfile -t ALL_SEEDS < <(
94
+ shuf -i 0-9999 \
95
+ --random-source=<(awk -v s="$METASEED" 'BEGIN { while (1) printf "%s", s }') \
96
+ | head -n $NUM_SEED
97
+ )
98
+
99
+ # SEEDS=("${ALL_SEEDS[@]:META_START:NUM_SEED-META_START}")
100
+ SEEDS=("${ALL_SEEDS[@]:$META_START:$((NUM_SEED - META_START))}")
101
+
102
+ echo "[INFO] Total Seeds Generated: ${#ALL_SEEDS[@]}"
103
+ echo "[INFO] Using Seeds: ${SEEDS[@]}"
104
+ fi
105
+
106
+ echo "CONFIGS: ${CONFIGS[@]}"
107
+ echo "SEEDS: ${SEEDS[@]}"
108
+ echo "RUN_IDS: ${RUN_IDS[@]}"
109
+ echo "K_SHOTS: ${K_SHOTS[@]}"
110
+ echo "CKPTS: ${CKPTS[@]}"
111
+ echo "GPU_IDS: ${GPU_IDS}"
112
+ echo "NUM_PROCESSES: $NUM_PROCESSES"
113
+ echo "META_START: $META_START"
114
+ echo "NUM_SEED: $NUM_SEED"
115
+ echo "Actual inference episodes: ${#SEEDS[@]}"
116
+ echo "MAX_INFER_SIZE: ${MAX_INFER_SIZE:-"(unset)"}"
117
+
118
+ EXTRA_ARG=""
119
+ if [[ -n "$MAX_INFER_SIZE" ]]; then
120
+ EXTRA_ARG="-M $MAX_INFER_SIZE"
121
+ fi
122
+
123
+ run_with_retry(){
124
+ local cmd=("$@")
125
+ local max=3
126
+ local attempt=1
127
+
128
+ while (( attempt <= max )); do
129
+ echo "Attempt $attempt/$max: ${cmd[*]}"
130
+
131
+ set +e
132
+ "${cmd[@]}"
133
+ status=$?
134
+ set -e
135
+
136
+ if [[ $status -eq 0 ]]; then
137
+ return 0
138
+ fi
139
+
140
+ ((attempt++))
141
+ sleep 2
142
+ done
143
+
144
+ return 1
145
+ }
146
+
147
+ for k_shot in "${K_SHOTS[@]}"; do
148
+ for run_id in "${RUN_IDS[@]}"; do
149
+ for config in "${CONFIGS[@]}"; do
150
+ for seed in "${SEEDS[@]}"; do
151
+ for ckpt in "${CKPTS[@]}"; do
152
+ echo "config: $config, seed: $seed, ckpt: $ckpt, run: $run_id"
153
+
154
+ max_wait=12
155
+ ckpt_path="./ckpt/$config/novel/run-$run_id/$k_shot-shot/shuffle_seed-$seed/checkpoint-$ckpt"
156
+ while [[ ! -s "$ckpt_path" ]]; do
157
+ echo "waiting for ckpt: $ckpt_path"
158
+ sleep 300
159
+ ((max_wait--))
160
+ done
161
+
162
+ if [[ ! -s "$ckpt_path" ]]; then
163
+ echo "[ERROR] ckpt not found after waiting, skip."
164
+ continue
165
+ fi
166
+
167
+ run_with_retry \
168
+ accelerate launch --multi_gpu --gpu_ids $GPU_IDS --num_processes $NUM_PROCESSES main.py \
169
+ --config ./configs/$config.yaml -m infer -p novel -s $seed -k $k_shot -r $run_id -c $ckpt $EXTRA_ARG
170
+ done
171
+ done
172
+ done
173
+ done
174
+ done
main.py ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from utils import load_config
2
+ import variants
3
+
4
+ if __name__ == '__main__':
5
+ config = load_config()
6
+
7
+ if config.mode == 'train':
8
+ entry = getattr(variants.train, config.training.get('entry', 'train_dsp'))
9
+ elif config.mode == 'infer':
10
+ entry = getattr(variants.infer, config.inference.get('entry', 'infer_dsp'))
11
+
12
+ entry(config)
models/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from . import dsp
models/dsp/CAM/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .cam_generator import CAMGenerator
models/dsp/CAM/cam_generator.py ADDED
@@ -0,0 +1,133 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .pytorch_grad_cam import GradCAM
2
+ from .pytorch_grad_cam.utils.image import scale_cam_image
3
+ from . import clip
4
+ from .utils import reshape_transform, zeroshot_classifier, ClipOutputTarget, scoremap2bbox
5
+
6
+ from torchvision import transforms
7
+ import torch
8
+ import numpy as np
9
+ import cv2
10
+
11
+ class CAMGenerator:
12
+ def __init__(self, categories, clip_path):
13
+ # self.device = "cuda" if torch.cuda.is_available() else "cpu"
14
+
15
+ self.clip_path = clip_path
16
+ self.clip_model, _ = clip.load(self.clip_path, device="cpu")
17
+ self.target_layers = [self.clip_model.visual.transformer.resblocks[-1].ln_1]
18
+ self.cam = GradCAM(model=self.clip_model, target_layers=self.target_layers, reshape_transform=reshape_transform, use_cuda=True)
19
+
20
+ self.categories = categories
21
+ # if categories is None:
22
+ # self.categories = ['vehicle', 'baseballfield', 'groundtrackfield', 'windmill', 'bridge', 'overpass', 'ship', 'airplane', 'tenniscourt', 'airport',
23
+ # 'expressway-service-area', 'basketballcourt', 'stadium', 'storagetank', 'chimney', 'dam', 'expressway-toll-station', 'golffield', 'trainstation', 'harbor']
24
+ self.background_categories = ['ground','land','grass','tree','building','wall','sky','lake','water','river','sea','railway','railroad','keyboard','helmet',
25
+ 'cloud','house','mountain','ocean','road','rock','street','valley','bridge','sign',]
26
+
27
+ self.normalize = transforms.Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711))
28
+
29
+ def _prepare(self):
30
+ self.bg_text_features = zeroshot_classifier(self.background_categories, ['a clean origami {}.'], self.clip_model, self.device)
31
+ self.fg_text_features = zeroshot_classifier(self.categories, ['a clean origami {}.'], self.clip_model, self.device)
32
+
33
+ def to(self, device, dtype):
34
+ self.device = device
35
+ self.clip_model.to(device)
36
+ self.cam.set_device(device)
37
+ self._prepare()
38
+
39
+ def re_normalize(self, image):
40
+ # image = (image / 2) + 0.5
41
+ image = self.normalize(image)
42
+ return image
43
+
44
+ def get_label_list(self, captions):
45
+ label_list = []
46
+ label_id_list = []
47
+ for caption in captions:
48
+ if caption in self.categories and caption not in label_list:
49
+ label_list.append(caption)
50
+ label_id_list.append(self.categories.index(caption))
51
+ return label_list, label_id_list
52
+
53
+ def get_label_list_with_bboxes(self, captions, bboxes):
54
+ label_list, label_id_list, bboxes_list = [], [], []
55
+ for caption, bbox in zip(captions, bboxes):
56
+ if caption not in self.categories:
57
+ continue
58
+ if caption not in label_list:
59
+ label_list.append(caption)
60
+ label_id_list.append(self.categories.index(caption))
61
+ bboxes_list.append([bbox])
62
+ else:
63
+ bboxes_list[label_list.index(caption)].append(bbox)
64
+ return label_list, label_id_list, bboxes_list
65
+
66
+ def __call__(self, image, captions, bboxes, gt_bboxes_only=False):
67
+ image = self.re_normalize(image)
68
+ # label_list, label_id_list = self.get_label_list(captions[0][1:])
69
+ label_list, label_id_list, bboxes_list = self.get_label_list_with_bboxes(captions[0][1:], bboxes[0])
70
+ h, w = image.shape[-2], image.shape[-1]
71
+ image_features, attn_weight_list = self.clip_model.encode_image(image, h, w)
72
+
73
+ bg_features_temp = self.bg_text_features
74
+ fg_features_temp = self.fg_text_features[label_id_list]
75
+ text_features_temp = torch.cat([fg_features_temp, bg_features_temp], dim=0)
76
+ input_tensor = [image_features, text_features_temp, h, w]
77
+
78
+ keys, refined_cam_list = [], []
79
+ for idx, (label, bbox) in enumerate(zip(label_list, bboxes_list)):
80
+ keys.append(self.categories.index(label))
81
+ targets = [ClipOutputTarget(label_list.index(label))]
82
+ grayscale_cam, logits_per_image, attn_weight_last = self.cam(input_tensor=input_tensor, targets=targets, target_size=None)
83
+ grayscale_cam = grayscale_cam[0, :] # [32, 32]
84
+ # grayscale_cam_highres = cv2.resize(grayscale_cam, (ori_width, ori_height))
85
+
86
+ if idx == 0:
87
+ attn_weight_list.append(attn_weight_last)
88
+ attn_weight = [aw[:, 1:, 1:] for aw in attn_weight_list] # (b, hxw, hxw)
89
+ attn_weight = torch.stack(attn_weight, dim=0)[-8:] # [8, 1, 1024, 1024]
90
+ attn_weight = torch.mean(attn_weight, dim=0) # [1, 1024, 1024]
91
+ # attn_weight = attn_weight[0].detach() # [1024, 1024] # original detach
92
+ attn_weight = attn_weight[0] #.detach() # [1024, 1024]
93
+ attn_weight = attn_weight.float()
94
+
95
+ gt_box, gt_cnt = (np.array(bbox) * [grayscale_cam.shape[1], grayscale_cam.shape[0], grayscale_cam.shape[1], grayscale_cam.shape[0]]).astype(int), len(bbox)
96
+ if gt_bboxes_only:
97
+ box, cnt = gt_box, gt_cnt
98
+ else:
99
+ box, cnt = scoremap2bbox(scoremap=grayscale_cam.cpu().data.numpy(), threshold=0.4, multi_contour_eval=True)
100
+ box, cnt = np.concatenate([box, gt_box], axis=0), cnt + gt_cnt
101
+ aff_mask = torch.zeros_like(grayscale_cam)
102
+ for i_ in range(cnt):
103
+ x0_, y0_, x1_, y1_ = box[i_]
104
+ aff_mask[y0_:y1_, x0_:x1_] = 1
105
+ aff_mask = aff_mask.view(1, grayscale_cam.shape[0] * grayscale_cam.shape[1])
106
+
107
+ aff_mat = attn_weight
108
+ trans_mat = aff_mat / torch.sum(aff_mat, dim=0, keepdim=True)
109
+ trans_mat = trans_mat / torch.sum(trans_mat, dim=1, keepdim=True)
110
+ for _ in range(2):
111
+ trans_mat = trans_mat / torch.sum(trans_mat, dim=0, keepdim=True)
112
+ trans_mat = trans_mat / torch.sum(trans_mat, dim=1, keepdim=True)
113
+ trans_mat = (trans_mat + trans_mat.transpose(1, 0)) / 2
114
+ for _ in range(1):
115
+ trans_mat = torch.matmul(trans_mat, trans_mat)
116
+
117
+ trans_mat = trans_mat * aff_mask
118
+
119
+ cam_to_refine = grayscale_cam.view(-1, 1)
120
+ cam_refined = torch.matmul(trans_mat, cam_to_refine).reshape(h //16, w // 16)
121
+ cam_refined = cam_refined - cam_refined.min()
122
+ cam_refined = cam_refined / (cam_refined.max() + 1e-7)
123
+ refined_cam_list.append(cam_refined)
124
+
125
+ keys = torch.tensor(keys)
126
+ refined_cams = torch.stack(refined_cam_list, dim=0)
127
+
128
+ return refined_cams, keys
129
+
130
+
131
+ if __name__ == '__main__':
132
+ cam = CAMGenerator()
133
+
models/dsp/CAM/clip/__init__.py ADDED
@@ -0,0 +1 @@
 
 
1
+ from .clip import *
models/dsp/CAM/clip/bpe_simple_vocab_16e6.txt.gz ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:924691ac288e54409236115652ad4aa250f48203de50a9e4722a6ecd48d6804a
3
+ size 1356917
models/dsp/CAM/clip/clip.py ADDED
@@ -0,0 +1,245 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import hashlib
2
+ import os
3
+ import urllib
4
+ import warnings
5
+ from typing import Any, Union, List
6
+ from packaging import version
7
+
8
+ import torch
9
+ from PIL import Image
10
+ from torchvision.transforms import Compose, Resize, CenterCrop, ToTensor, Normalize
11
+ from tqdm import tqdm
12
+
13
+ from .model import build_model
14
+ from .simple_tokenizer import SimpleTokenizer as _Tokenizer
15
+ from collections import OrderedDict
16
+
17
+ try:
18
+ from torchvision.transforms import InterpolationMode
19
+ BICUBIC = InterpolationMode.BICUBIC
20
+ except ImportError:
21
+ BICUBIC = Image.BICUBIC
22
+
23
+
24
+ if version.parse(torch.__version__) < version.parse("1.7.1"):
25
+ warnings.warn("PyTorch version 1.7.1 or higher is recommended")
26
+
27
+
28
+ __all__ = ["available_models", "load", "tokenize"]
29
+ _tokenizer = _Tokenizer()
30
+
31
+ _MODELS = {
32
+ "RN50": "https://openaipublic.azureedge.net/clip/models/afeb0e10f9e5a86da6080e35cf09123aca3b358a0c3e3b6c78a7b63bc04b6762/RN50.pt",
33
+ "RN101": "https://openaipublic.azureedge.net/clip/models/8fa8567bab74a42d41c5915025a8e4538c3bdbe8804a470a72f30b0d94fab599/RN101.pt",
34
+ "RN50x4": "https://openaipublic.azureedge.net/clip/models/7e526bd135e493cef0776de27d5f42653e6b4c8bf9e0f653bb11773263205fdd/RN50x4.pt",
35
+ "RN50x16": "https://openaipublic.azureedge.net/clip/models/52378b407f34354e150460fe41077663dd5b39c54cd0bfd2b27167a4a06ec9aa/RN50x16.pt",
36
+ "RN50x64": "https://openaipublic.azureedge.net/clip/models/be1cfb55d75a9666199fb2206c106743da0f6468c9d327f3e0d0a543a9919d9c/RN50x64.pt",
37
+ "ViT-B/32": "https://openaipublic.azureedge.net/clip/models/40d365715913c9da98579312b702a82c18be219cc2a73407c4526f58eba950af/ViT-B-32.pt",
38
+ "ViT-B/16": "https://openaipublic.azureedge.net/clip/models/5806e77cd80f8b59890b7e101eabd078d9fb84e6937f9e85e4ecb61988df416f/ViT-B-16.pt",
39
+ "ViT-L/14": "https://openaipublic.azureedge.net/clip/models/b8cca3fd41ae0c99ba7e8951adf17d267cdb84cd88be6f7c2e0eca1737a03836/ViT-L-14.pt",
40
+ "ViT-L/14@336px": "https://openaipublic.azureedge.net/clip/models/3035c92b350959924f9f00213499208652fc7ea050643e8b385c2dac08641f02/ViT-L-14-336px.pt",
41
+ }
42
+
43
+
44
+ def _download(url: str, root: str):
45
+ os.makedirs(root, exist_ok=True)
46
+ filename = os.path.basename(url)
47
+
48
+ expected_sha256 = url.split("/")[-2]
49
+ download_target = os.path.join(root, filename)
50
+
51
+ if os.path.exists(download_target) and not os.path.isfile(download_target):
52
+ raise RuntimeError(f"{download_target} exists and is not a regular file")
53
+
54
+ if os.path.isfile(download_target):
55
+ if hashlib.sha256(open(download_target, "rb").read()).hexdigest() == expected_sha256:
56
+ return download_target
57
+ else:
58
+ warnings.warn(f"{download_target} exists, but the SHA256 checksum does not match; re-downloading the file")
59
+
60
+ with urllib.request.urlopen(url) as source, open(download_target, "wb") as output:
61
+ with tqdm(total=int(source.info().get("Content-Length")), ncols=80, unit='iB', unit_scale=True, unit_divisor=1024) as loop:
62
+ while True:
63
+ buffer = source.read(8192)
64
+ if not buffer:
65
+ break
66
+
67
+ output.write(buffer)
68
+ loop.update(len(buffer))
69
+
70
+ if hashlib.sha256(open(download_target, "rb").read()).hexdigest() != expected_sha256:
71
+ raise RuntimeError(f"Model has been downloaded but the SHA256 checksum does not not match")
72
+
73
+ return download_target
74
+
75
+
76
+ def _convert_image_to_rgb(image):
77
+ return image.convert("RGB")
78
+
79
+
80
+ def _transform(n_px):
81
+ return Compose([
82
+ Resize(n_px, interpolation=BICUBIC),
83
+ CenterCrop(n_px),
84
+ _convert_image_to_rgb,
85
+ ToTensor(),
86
+ Normalize((0.48145466, 0.4578275, 0.40821073), (0.26862954, 0.26130258, 0.27577711)),
87
+ ])
88
+
89
+
90
+ def available_models() -> List[str]:
91
+ """Returns the names of available CLIP models"""
92
+ return list(_MODELS.keys())
93
+
94
+
95
+ def load(name: str, device: Union[str, torch.device] = "cuda" if torch.cuda.is_available() else "cpu", jit: bool = False, download_root: str = None):
96
+ """Load a CLIP model
97
+
98
+ Parameters
99
+ ----------
100
+ name : str
101
+ A model name listed by `clip.available_models()`, or the path to a model checkpoint containing the state_dict
102
+
103
+ device : Union[str, torch.device]
104
+ The device to put the loaded model
105
+
106
+ jit : bool
107
+ Whether to load the optimized JIT model or more hackable non-JIT model (default).
108
+
109
+ download_root: str
110
+ path to download the model files; by default, it uses "~/.cache/clip"
111
+
112
+ Returns
113
+ -------
114
+ model : torch.nn.Module
115
+ The CLIP model
116
+
117
+ preprocess : Callable[[PIL.Image], torch.Tensor]
118
+ A torchvision transform that converts a PIL image into a tensor that the returned model can take as its input
119
+ """
120
+ if name in _MODELS:
121
+ model_path = _download(_MODELS[name], download_root or os.path.expanduser("~/.cache/clip"))
122
+ elif os.path.isfile(name):
123
+ model_path = name
124
+ else:
125
+ raise RuntimeError(f"Model {name} not found; available models = {available_models()}")
126
+
127
+ with open(model_path, 'rb') as opened_file:
128
+ try:
129
+ # loading JIT archive
130
+ model = torch.jit.load(opened_file, map_location=device if jit else "cpu").eval()
131
+ state_dict = None
132
+ except RuntimeError:
133
+ # loading saved state dict
134
+ if jit:
135
+ warnings.warn(f"File {model_path} is not a JIT archive. Loading as a state dict instead")
136
+ jit = False
137
+ if 'RN50' in model_path:
138
+ state_dict = torch.load(opened_file, map_location="cpu")
139
+ else:
140
+ state_dict0 = torch.load(model_path, map_location="cpu")
141
+ state_dict = OrderedDict()
142
+ for k in state_dict0.keys():
143
+ state_dict[k.replace('module.', '')] = state_dict0[k]
144
+
145
+
146
+ if not jit:
147
+ model = build_model(state_dict or model.state_dict()).to(device)
148
+ if str(device) == "cpu":
149
+ model.float()
150
+ return model, _transform(model.visual.input_resolution)
151
+
152
+ # patch the device names
153
+ device_holder = torch.jit.trace(lambda: torch.ones([]).to(torch.device(device)), example_inputs=[])
154
+ device_node = [n for n in device_holder.graph.findAllNodes("prim::Constant") if "Device" in repr(n)][-1]
155
+
156
+ def patch_device(module):
157
+ try:
158
+ graphs = [module.graph] if hasattr(module, "graph") else []
159
+ except RuntimeError:
160
+ graphs = []
161
+
162
+ if hasattr(module, "forward1"):
163
+ graphs.append(module.forward1.graph)
164
+
165
+ for graph in graphs:
166
+ for node in graph.findAllNodes("prim::Constant"):
167
+ if "value" in node.attributeNames() and str(node["value"]).startswith("cuda"):
168
+ node.copyAttributes(device_node)
169
+
170
+ model.apply(patch_device)
171
+ patch_device(model.encode_image)
172
+ patch_device(model.encode_text)
173
+
174
+ # patch dtype to float32 on CPU
175
+ if str(device) == "cpu":
176
+ float_holder = torch.jit.trace(lambda: torch.ones([]).float(), example_inputs=[])
177
+ float_input = list(float_holder.graph.findNode("aten::to").inputs())[1]
178
+ float_node = float_input.node()
179
+
180
+ def patch_float(module):
181
+ try:
182
+ graphs = [module.graph] if hasattr(module, "graph") else []
183
+ except RuntimeError:
184
+ graphs = []
185
+
186
+ if hasattr(module, "forward1"):
187
+ graphs.append(module.forward1.graph)
188
+
189
+ for graph in graphs:
190
+ for node in graph.findAllNodes("aten::to"):
191
+ inputs = list(node.inputs())
192
+ for i in [1, 2]: # dtype can be the second or third argument to aten::to()
193
+ if inputs[i].node()["value"] == 5:
194
+ inputs[i].node().copyAttributes(float_node)
195
+
196
+ model.apply(patch_float)
197
+ patch_float(model.encode_image)
198
+ patch_float(model.encode_text)
199
+
200
+ model.float()
201
+
202
+ return model, _transform(model.input_resolution.item())
203
+
204
+
205
+ def tokenize(texts: Union[str, List[str]], context_length: int = 77, truncate: bool = False) -> Union[torch.IntTensor, torch.LongTensor]:
206
+ """
207
+ Returns the tokenized representation of given input string(s)
208
+
209
+ Parameters
210
+ ----------
211
+ texts : Union[str, List[str]]
212
+ An input string or a list of input strings to tokenize
213
+
214
+ context_length : int
215
+ The context length to use; all CLIP models use 77 as the context length
216
+
217
+ truncate: bool
218
+ Whether to truncate the text in case its encoding is longer than the context length
219
+
220
+ Returns
221
+ -------
222
+ A two-dimensional tensor containing the resulting tokens, shape = [number of input strings, context_length].
223
+ We return LongTensor when torch version is <1.8.0, since older index_select requires indices to be long.
224
+ """
225
+ if isinstance(texts, str):
226
+ texts = [texts]
227
+
228
+ sot_token = _tokenizer.encoder["<|startoftext|>"]
229
+ eot_token = _tokenizer.encoder["<|endoftext|>"]
230
+ all_tokens = [[sot_token] + _tokenizer.encode(text) + [eot_token] for text in texts]
231
+ if version.parse(torch.__version__) < version.parse("1.8.0"):
232
+ result = torch.zeros(len(all_tokens), context_length, dtype=torch.long)
233
+ else:
234
+ result = torch.zeros(len(all_tokens), context_length, dtype=torch.int)
235
+
236
+ for i, tokens in enumerate(all_tokens):
237
+ if len(tokens) > context_length:
238
+ if truncate:
239
+ tokens = tokens[:context_length]
240
+ tokens[-1] = eot_token
241
+ else:
242
+ raise RuntimeError(f"Input {texts[i]} is too long for context length {context_length}")
243
+ result[i, :len(tokens)] = torch.tensor(tokens)
244
+
245
+ return result
models/dsp/CAM/clip/model.py ADDED
@@ -0,0 +1,517 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from collections import OrderedDict
2
+ from typing import Tuple, Union
3
+
4
+ import numpy as np
5
+ import torch
6
+ import torch.nn.functional as F
7
+ from torch import nn
8
+
9
+ def upsample_pos_emb(emb, new_size):
10
+ # upsample the pretrained embedding for higher resolution
11
+ # emb size NxD
12
+ first = emb[:1, :]
13
+ emb = emb[1:, :]
14
+ N, D = emb.size(0), emb.size(1)
15
+ size = int(np.sqrt(N))
16
+ assert size * size == N
17
+ #new_size = size * self.upsample
18
+ emb = emb.permute(1, 0)
19
+ emb = emb.view(1, D, size, size).contiguous()
20
+ # emb = F.upsample(emb, size=new_size, mode='bilinear',)
21
+ emb = F.interpolate(emb, size=new_size, mode='bilinear', align_corners=False)
22
+ emb = emb.view(D, -1).contiguous()
23
+ emb = emb.permute(1, 0)
24
+ emb = torch.cat([first, emb], 0)
25
+ emb = nn.parameter.Parameter(emb.half())
26
+ return emb
27
+
28
+ class Bottleneck(nn.Module):
29
+ expansion = 4
30
+
31
+ def __init__(self, inplanes, planes, stride=1):
32
+ super().__init__()
33
+
34
+ # all conv layers have stride 1. an avgpool is performed after the second convolution when stride > 1
35
+ self.conv1 = nn.Conv2d(inplanes, planes, 1, bias=False)
36
+ self.bn1 = nn.BatchNorm2d(planes)
37
+ self.relu1 = nn.ReLU(inplace=True)
38
+
39
+ self.conv2 = nn.Conv2d(planes, planes, 3, padding=1, bias=False)
40
+ self.bn2 = nn.BatchNorm2d(planes)
41
+ self.relu2 = nn.ReLU(inplace=True)
42
+
43
+ self.avgpool = nn.AvgPool2d(stride) if stride > 1 else nn.Identity()
44
+
45
+ self.conv3 = nn.Conv2d(planes, planes * self.expansion, 1, bias=False)
46
+ self.bn3 = nn.BatchNorm2d(planes * self.expansion)
47
+ self.relu3 = nn.ReLU(inplace=True)
48
+
49
+ self.downsample = None
50
+ self.stride = stride
51
+
52
+ if stride > 1 or inplanes != planes * Bottleneck.expansion:
53
+ # downsampling layer is prepended with an avgpool, and the subsequent convolution has stride 1
54
+ self.downsample = nn.Sequential(OrderedDict([
55
+ ("-1", nn.AvgPool2d(stride)),
56
+ ("0", nn.Conv2d(inplanes, planes * self.expansion, 1, stride=1, bias=False)),
57
+ ("1", nn.BatchNorm2d(planes * self.expansion))
58
+ ]))
59
+
60
+ def forward(self, x: torch.Tensor):
61
+ identity = x
62
+
63
+ out = self.relu1(self.bn1(self.conv1(x)))
64
+ out = self.relu2(self.bn2(self.conv2(out)))
65
+ out = self.avgpool(out)
66
+ out = self.bn3(self.conv3(out))
67
+
68
+ if self.downsample is not None:
69
+ identity = self.downsample(x)
70
+
71
+ out += identity
72
+ out = self.relu3(out)
73
+ return out
74
+
75
+
76
+ class AttentionPool2d(nn.Module):
77
+ def __init__(self, spacial_dim: int, embed_dim: int, num_heads: int, output_dim: int = None):
78
+ super().__init__()
79
+ self.positional_embedding = nn.Parameter(torch.randn(spacial_dim ** 2 + 1, embed_dim) / embed_dim ** 0.5)
80
+ self.k_proj = nn.Linear(embed_dim, embed_dim)
81
+ self.q_proj = nn.Linear(embed_dim, embed_dim)
82
+ self.v_proj = nn.Linear(embed_dim, embed_dim)
83
+ self.c_proj = nn.Linear(embed_dim, output_dim or embed_dim)
84
+ self.num_heads = num_heads
85
+
86
+ def forward(self, x, H, W):
87
+ x = x.reshape(x.shape[0], x.shape[1], x.shape[2] * x.shape[3]).permute(2, 0, 1) # NCHW -> (HW)NC
88
+ x = torch.cat([x.mean(dim=0, keepdim=True), x], dim=0) # (HW+1)NC
89
+ self.positional_embedding_new = upsample_pos_emb(self.positional_embedding, (H//32,W//32))
90
+ x = x + self.positional_embedding_new[:, None, :].to(x.dtype) # (HW+1)NC
91
+ x, attn_weight = F.multi_head_attention_forward(
92
+ query=x, key=x, value=x,
93
+ embed_dim_to_check=x.shape[-1],
94
+ num_heads=self.num_heads,
95
+ q_proj_weight=self.q_proj.weight,
96
+ k_proj_weight=self.k_proj.weight,
97
+ v_proj_weight=self.v_proj.weight,
98
+ in_proj_weight=None,
99
+ in_proj_bias=torch.cat([self.q_proj.bias, self.k_proj.bias, self.v_proj.bias]),
100
+ bias_k=None,
101
+ bias_v=None,
102
+ add_zero_attn=False,
103
+ dropout_p=0,
104
+ out_proj_weight=self.c_proj.weight,
105
+ out_proj_bias=self.c_proj.bias,
106
+ use_separate_proj_weight=True,
107
+ training=self.training,
108
+ need_weights=False
109
+ )
110
+ return x[0]
111
+
112
+
113
+ class ModifiedResNet(nn.Module):
114
+ """
115
+ A ResNet class that is similar to torchvision's but contains the following changes:
116
+ - There are now 3 "stem" convolutions as opposed to 1, with an average pool instead of a max pool.
117
+ - Performs anti-aliasing strided convolutions, where an avgpool is prepended to convolutions with stride > 1
118
+ - The final pooling layer is a QKV attention instead of an average pool
119
+ """
120
+
121
+ def __init__(self, layers, output_dim, heads, input_resolution=224, width=64):
122
+ super().__init__()
123
+ self.output_dim = output_dim
124
+ self.input_resolution = input_resolution
125
+
126
+ # the 3-layer stem
127
+ self.conv1 = nn.Conv2d(3, width // 2, kernel_size=3, stride=2, padding=1, bias=False)
128
+ self.bn1 = nn.BatchNorm2d(width // 2)
129
+ self.relu1 = nn.ReLU(inplace=True)
130
+ self.conv2 = nn.Conv2d(width // 2, width // 2, kernel_size=3, padding=1, bias=False)
131
+ self.bn2 = nn.BatchNorm2d(width // 2)
132
+ self.relu2 = nn.ReLU(inplace=True)
133
+ self.conv3 = nn.Conv2d(width // 2, width, kernel_size=3, padding=1, bias=False)
134
+ self.bn3 = nn.BatchNorm2d(width)
135
+ self.relu3 = nn.ReLU(inplace=True)
136
+ self.avgpool = nn.AvgPool2d(2)
137
+
138
+ # residual layers
139
+ self._inplanes = width # this is a *mutable* variable used during construction
140
+ self.layer1 = self._make_layer(width, layers[0])
141
+ self.layer2 = self._make_layer(width * 2, layers[1], stride=2)
142
+ self.layer3 = self._make_layer(width * 4, layers[2], stride=2)
143
+ self.layer4 = self._make_layer(width * 8, layers[3], stride=2)
144
+
145
+ embed_dim = width * 32 # the ResNet feature dimension
146
+ self.attnpool = AttentionPool2d(input_resolution // 32, embed_dim, heads, output_dim)
147
+
148
+ def _make_layer(self, planes, blocks, stride=1):
149
+ layers = [Bottleneck(self._inplanes, planes, stride)]
150
+
151
+ self._inplanes = planes * Bottleneck.expansion
152
+ for _ in range(1, blocks):
153
+ layers.append(Bottleneck(self._inplanes, planes))
154
+
155
+ return nn.Sequential(*layers)
156
+
157
+ def forward(self, x, H, W):
158
+ def stem(x):
159
+ x = self.relu1(self.bn1(self.conv1(x)))
160
+ x = self.relu2(self.bn2(self.conv2(x)))
161
+ x = self.relu3(self.bn3(self.conv3(x)))
162
+ x = self.avgpool(x)
163
+ return x
164
+
165
+ x = x.type(self.conv1.weight.dtype)
166
+ x = stem(x)
167
+ x = self.layer1(x)
168
+ x = self.layer2(x)
169
+ x = self.layer3(x)
170
+ x = self.layer4(x)#(1,,2048, 7, 7)
171
+ x_pooled = self.attnpool(x, H, W)
172
+
173
+ return x_pooled
174
+
175
+
176
+ class LayerNorm(nn.LayerNorm):
177
+ """Subclass torch's LayerNorm to handle fp16."""
178
+
179
+ def forward(self, x: torch.Tensor):
180
+ orig_type = x.dtype
181
+ ret = super().forward(x.type(torch.float32))
182
+ return ret.type(orig_type)
183
+
184
+
185
+ class QuickGELU(nn.Module):
186
+ def forward(self, x: torch.Tensor):
187
+ return x * torch.sigmoid(1.702 * x)
188
+
189
+
190
+ class ResidualAttentionBlock(nn.Module):
191
+ def __init__(self, d_model: int, n_head: int, attn_mask: torch.Tensor = None):
192
+ super().__init__()
193
+
194
+ self.attn = nn.MultiheadAttention(d_model, n_head)
195
+ self.ln_1 = LayerNorm(d_model)
196
+ self.mlp = nn.Sequential(OrderedDict([
197
+ ("c_fc", nn.Linear(d_model, d_model * 4)),
198
+ ("gelu", QuickGELU()),
199
+ ("c_proj", nn.Linear(d_model * 4, d_model))
200
+ ]))
201
+ self.ln_2 = LayerNorm(d_model)
202
+ self.attn_mask = attn_mask
203
+
204
+ def attention(self, x: torch.Tensor):
205
+ self.attn_mask = self.attn_mask.to(dtype=x.dtype, device=x.device) if self.attn_mask is not None else None
206
+ return self.attn(x, x, x, need_weights=True, attn_mask=self.attn_mask)#[0]
207
+
208
+ def forward(self, x: torch.Tensor):
209
+ attn_output, attn_weight = self.attention(self.ln_1(x))#(L,N,E) (N,L,L)
210
+ x = x + attn_output
211
+ x = x + self.mlp(self.ln_2(x))
212
+ return x, attn_weight
213
+
214
+
215
+
216
+ class Transformer(nn.Module):
217
+ def __init__(self, width: int, layers: int, heads: int, attn_mask: torch.Tensor = None):
218
+ super().__init__()
219
+ self.width = width
220
+ self.layers = layers
221
+ self.resblocks = nn.Sequential(*[ResidualAttentionBlock(width, heads, attn_mask) for _ in range(layers)])
222
+
223
+ def forward(self, x: torch.Tensor):
224
+ attn_weights = []
225
+ with torch.no_grad():
226
+ layers = self.layers if x.shape[0] == 77 else self.layers-1
227
+ for i in range(layers):
228
+ x, attn_weight = self.resblocks[i](x)
229
+ attn_weights.append(attn_weight)
230
+ '''
231
+ for i in range(self.layers-1, self.layers):
232
+ x, attn_weight = self.resblocks[i](x)
233
+ attn_weights.append(attn_weight)
234
+ #feature_map_list.append(x)
235
+ '''
236
+ return x, attn_weights
237
+
238
+
239
+ class VisionTransformer(nn.Module):
240
+ def __init__(self, input_resolution: int, patch_size: int, width: int, layers: int, heads: int, output_dim: int):
241
+ super().__init__()
242
+ self.input_resolution = input_resolution
243
+ self.output_dim = output_dim
244
+ self.conv1 = nn.Conv2d(in_channels=3, out_channels=width, kernel_size=patch_size, stride=patch_size, bias=False)
245
+
246
+ scale = width ** -0.5
247
+ self.class_embedding = nn.Parameter(scale * torch.randn(width))
248
+ self.positional_embedding = nn.Parameter(scale * torch.randn((input_resolution // patch_size) ** 2 + 1, width))
249
+ self.ln_pre = LayerNorm(width)
250
+
251
+ self.transformer = Transformer(width, layers, heads)
252
+
253
+ self.ln_post = LayerNorm(width)
254
+ self.proj = nn.Parameter(scale * torch.randn(width, output_dim))
255
+ self.patch_size = patch_size
256
+
257
+ def forward(self, x: torch.Tensor, H, W):
258
+
259
+ self.positional_embedding_new = upsample_pos_emb(self.positional_embedding, (H//16,W//16))
260
+ x = self.conv1(x) # shape = [*, width, grid, grid]
261
+ x = x.reshape(x.shape[0], x.shape[1], -1) # shape = [*, width, grid ** 2]
262
+ x = x.permute(0, 2, 1) # shape = [*, grid ** 2, width]
263
+ x = torch.cat([self.class_embedding.to(x.dtype) + torch.zeros(x.shape[0], 1, x.shape[-1], dtype=x.dtype, device=x.device), x], dim=1) # shape = [*, grid ** 2 + 1, width]
264
+ x = x + self.positional_embedding_new.to(x.dtype)
265
+ x = self.ln_pre(x)
266
+
267
+ x = x.permute(1, 0, 2) # NLD -> LND
268
+ x, attn_weight = self.transformer(x)
269
+ '''
270
+ x = x.permute(1, 0, 2) # LND -> NLD
271
+
272
+ x = self.ln_post(x)
273
+ #x = x[:, 0, :]
274
+ #x = x[:,1:,:]
275
+ x = torch.mean(x[:,1:,:],dim=1)
276
+ #feature_map_list.append(x)
277
+
278
+ if self.proj is not None:
279
+ x = x @ self.proj
280
+ '''
281
+
282
+ return x, attn_weight#cls_attn
283
+
284
+
285
+ class CLIP(nn.Module):
286
+ def __init__(self,
287
+ embed_dim: int,
288
+ # vision
289
+ image_resolution: int,
290
+ vision_layers: Union[Tuple[int, int, int, int], int],
291
+ vision_width: int,
292
+ vision_patch_size: int,
293
+ # text
294
+ context_length: int,
295
+ vocab_size: int,
296
+ transformer_width: int,
297
+ transformer_heads: int,
298
+ transformer_layers: int
299
+ ):
300
+ super().__init__()
301
+
302
+ self.context_length = context_length
303
+
304
+ if isinstance(vision_layers, (tuple, list)):
305
+ vision_heads = vision_width * 32 // 64
306
+ self.visual = ModifiedResNet(
307
+ layers=vision_layers,
308
+ output_dim=embed_dim,
309
+ heads=vision_heads,
310
+ input_resolution=image_resolution,
311
+ width=vision_width
312
+ )
313
+ else:
314
+ vision_heads = vision_width // 64
315
+ self.visual = VisionTransformer(
316
+ input_resolution=image_resolution,
317
+ patch_size=vision_patch_size,
318
+ width=vision_width,
319
+ layers=vision_layers,
320
+ heads=vision_heads,
321
+ output_dim=embed_dim
322
+ )
323
+
324
+ self.transformer = Transformer(
325
+ width=transformer_width,
326
+ layers=transformer_layers,
327
+ heads=transformer_heads,
328
+ attn_mask=self.build_attention_mask()
329
+ )
330
+
331
+ self.vocab_size = vocab_size
332
+ self.token_embedding = nn.Embedding(vocab_size, transformer_width)
333
+ self.positional_embedding = nn.Parameter(torch.empty(self.context_length, transformer_width))
334
+ self.ln_final = LayerNorm(transformer_width)
335
+
336
+ self.text_projection = nn.Parameter(torch.empty(transformer_width, embed_dim))
337
+ self.logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / 0.07))
338
+
339
+ self.initialize_parameters()
340
+
341
+ def initialize_parameters(self):
342
+ nn.init.normal_(self.token_embedding.weight, std=0.02)
343
+ nn.init.normal_(self.positional_embedding, std=0.01)
344
+
345
+ if isinstance(self.visual, ModifiedResNet):
346
+ if self.visual.attnpool is not None:
347
+ std = self.visual.attnpool.c_proj.in_features ** -0.5
348
+ nn.init.normal_(self.visual.attnpool.q_proj.weight, std=std)
349
+ nn.init.normal_(self.visual.attnpool.k_proj.weight, std=std)
350
+ nn.init.normal_(self.visual.attnpool.v_proj.weight, std=std)
351
+ nn.init.normal_(self.visual.attnpool.c_proj.weight, std=std)
352
+
353
+ for resnet_block in [self.visual.layer1, self.visual.layer2, self.visual.layer3, self.visual.layer4]:
354
+ for name, param in resnet_block.named_parameters():
355
+ if name.endswith("bn3.weight"):
356
+ nn.init.zeros_(param)
357
+
358
+ proj_std = (self.transformer.width ** -0.5) * ((2 * self.transformer.layers) ** -0.5)
359
+ attn_std = self.transformer.width ** -0.5
360
+ fc_std = (2 * self.transformer.width) ** -0.5
361
+ for block in self.transformer.resblocks:
362
+ nn.init.normal_(block.attn.in_proj_weight, std=attn_std)
363
+ nn.init.normal_(block.attn.out_proj.weight, std=proj_std)
364
+ nn.init.normal_(block.mlp.c_fc.weight, std=fc_std)
365
+ nn.init.normal_(block.mlp.c_proj.weight, std=proj_std)
366
+
367
+ if self.text_projection is not None:
368
+ nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)
369
+
370
+ def build_attention_mask(self):
371
+ # lazily create causal attention mask, with full attention between the vision tokens
372
+ # pytorch uses additive attention mask; fill with -inf
373
+ mask = torch.empty(self.context_length, self.context_length)
374
+ mask.fill_(float("-inf"))
375
+ mask.triu_(1) # zero out the lower diagonal
376
+ return mask
377
+
378
+ @property
379
+ def dtype(self):
380
+ return self.visual.conv1.weight.dtype
381
+
382
+ def encode_image(self, image, H, W):
383
+ return self.visual(image.type(self.dtype), H, W)
384
+
385
+ def encode_text(self, text):
386
+ x = self.token_embedding(text).type(self.dtype) # [batch_size, n_ctx, d_model]
387
+
388
+ x = x + self.positional_embedding.type(self.dtype)
389
+ x = x.permute(1, 0, 2) # NLD -> LND
390
+ x, attn_weight = self.transformer(x)
391
+ x = x.permute(1, 0, 2) # LND -> NLD
392
+ x = self.ln_final(x).type(self.dtype)
393
+
394
+ # x.shape = [batch_size, n_ctx, transformer.width]
395
+ # take features from the eot embedding (eot_token is the highest number in each sequence)
396
+ x = x[torch.arange(x.shape[0]), text.argmax(dim=-1)] @ self.text_projection
397
+
398
+ return x
399
+
400
+ def forward_last_layer(self, image_features, text_features):
401
+ x, attn_weight = self.visual.transformer.resblocks[self.visual.transformer.layers-1](image_features)
402
+ x = x.permute(1, 0, 2) # LND -> NLD
403
+
404
+ x = self.visual.ln_post(x)
405
+ x = torch.mean(x[:, 1:, :], dim=1)
406
+
407
+ if self.visual.proj is not None:
408
+ x = x @ self.visual.proj
409
+
410
+ image_features = x
411
+
412
+ # normalized features
413
+ image_features = image_features / image_features.norm(dim=1, keepdim=True)
414
+ text_features = text_features / text_features.norm(dim=1, keepdim=True)
415
+ # cosine similarity as logits
416
+ logit_scale = self.logit_scale.exp()
417
+ logits_per_image = logit_scale * image_features @ text_features.t()
418
+
419
+ # shape = [global_batch_size, global_batch_size]
420
+ logits_per_image = logits_per_image.softmax(dim=-1)
421
+
422
+ return logits_per_image, attn_weight
423
+
424
+
425
+
426
+
427
+ def forward(self, image, text):
428
+ image_features, feature_map, cls_attn = self.encode_image(image)
429
+ with torch.no_grad():
430
+ text_features = self.encode_text(text)
431
+
432
+ # normalized features
433
+ image_features = image_features / image_features.norm(dim=1, keepdim=True)
434
+ text_features = text_features / text_features.norm(dim=1, keepdim=True)
435
+
436
+ # cosine similarity as logits
437
+ logit_scale = self.logit_scale.exp()
438
+ logits_per_image = logit_scale * image_features @ text_features.t()
439
+ #logits_per_text = logits_per_image.t()
440
+
441
+ # shape = [global_batch_size, global_batch_size]
442
+ return logits_per_image, logits_per_text
443
+
444
+
445
+ def convert_weights(model: nn.Module):
446
+ """Convert applicable model parameters to fp16"""
447
+
448
+ def _convert_weights_to_fp16(l):
449
+ if isinstance(l, (nn.Conv1d, nn.Conv2d, nn.Linear)):
450
+ l.weight.data = l.weight.data.half()
451
+ if l.bias is not None:
452
+ l.bias.data = l.bias.data.half()
453
+
454
+ if isinstance(l, nn.MultiheadAttention):
455
+ for attr in [*[f"{s}_proj_weight" for s in ["in", "q", "k", "v"]], "in_proj_bias", "bias_k", "bias_v"]:
456
+ tensor = getattr(l, attr)
457
+ if tensor is not None:
458
+ tensor.data = tensor.data.half()
459
+
460
+ for name in ["text_projection", "proj"]:
461
+ if hasattr(l, name):
462
+ attr = getattr(l, name)
463
+ if attr is not None:
464
+ attr.data = attr.data.half()
465
+
466
+ model.apply(_convert_weights_to_fp16)
467
+
468
+
469
+ def build_model(state_dict: dict):
470
+ vit = "visual.proj" in state_dict
471
+ '''
472
+ inv_freq = 1. / (10000 ** (torch.arange(0, 2048, 2, dtype=torch.float) / 2048))
473
+ position = torch.arange(50, dtype=torch.float)
474
+ sinusoid_inp = torch.einsum('i,j -> ij', position, inv_freq)
475
+ embeddings = torch.cat((sinusoid_inp.sin(), sinusoid_inp.cos()), dim=-1) / 2048 ** 0.5
476
+ state_dict["visual.attnpool.positional_embedding"] = embeddings
477
+ '''
478
+ #state_dict["visual.positional_embedding"] = upsample_pos_emb(state_dict["visual.positional_embedding"], 28)
479
+
480
+
481
+ if vit:
482
+ vision_width = state_dict["visual.conv1.weight"].shape[0]
483
+ vision_layers = len([k for k in state_dict.keys() if k.startswith("visual.") and k.endswith(".attn.in_proj_weight")])
484
+ vision_patch_size = state_dict["visual.conv1.weight"].shape[-1]
485
+ grid_size = round((state_dict["visual.positional_embedding"].shape[0] - 1) ** 0.5)
486
+ image_resolution = vision_patch_size * grid_size
487
+ else:
488
+ counts: list = [len(set(k.split(".")[2] for k in state_dict if k.startswith(f"visual.layer{b}"))) for b in [1, 2, 3, 4]]
489
+ vision_layers = tuple(counts)
490
+ vision_width = state_dict["visual.layer1.0.conv1.weight"].shape[0]
491
+ output_width = round((state_dict["visual.attnpool.positional_embedding"].shape[0] - 1) ** 0.5)
492
+ vision_patch_size = None
493
+ assert output_width ** 2 + 1 == state_dict["visual.attnpool.positional_embedding"].shape[0]
494
+ image_resolution = output_width * 32
495
+
496
+ embed_dim = state_dict["text_projection"].shape[1]
497
+ context_length = state_dict["positional_embedding"].shape[0]
498
+ vocab_size = state_dict["token_embedding.weight"].shape[0]
499
+ transformer_width = state_dict["ln_final.weight"].shape[0]
500
+ transformer_heads = transformer_width // 64
501
+ transformer_layers = len(set(k.split(".")[2] for k in state_dict if k.startswith(f"transformer.resblocks")))
502
+
503
+ model = CLIP(
504
+ embed_dim,
505
+ image_resolution, vision_layers, vision_width, vision_patch_size,
506
+ context_length, vocab_size, transformer_width, transformer_heads, transformer_layers
507
+ )
508
+
509
+ for key in ["input_resolution", "context_length", "vocab_size"]:
510
+ if key in state_dict:
511
+ del state_dict[key]
512
+
513
+
514
+
515
+ convert_weights(model)
516
+ model.load_state_dict(state_dict)
517
+ return model.eval()
models/dsp/CAM/clip/simple_tokenizer.py ADDED
@@ -0,0 +1,132 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import gzip
2
+ import html
3
+ import os
4
+ from functools import lru_cache
5
+
6
+ import ftfy
7
+ import regex as re
8
+
9
+
10
+ @lru_cache()
11
+ def default_bpe():
12
+ return os.path.join(os.path.dirname(os.path.abspath(__file__)), "bpe_simple_vocab_16e6.txt.gz")
13
+
14
+
15
+ @lru_cache()
16
+ def bytes_to_unicode():
17
+ """
18
+ Returns list of utf-8 byte and a corresponding list of unicode strings.
19
+ The reversible bpe codes work on unicode strings.
20
+ This means you need a large # of unicode characters in your vocab if you want to avoid UNKs.
21
+ When you're at something like a 10B token dataset you end up needing around 5K for decent coverage.
22
+ This is a signficant percentage of your normal, say, 32K bpe vocab.
23
+ To avoid that, we want lookup tables between utf-8 bytes and unicode strings.
24
+ And avoids mapping to whitespace/control characters the bpe code barfs on.
25
+ """
26
+ bs = list(range(ord("!"), ord("~")+1))+list(range(ord("¡"), ord("¬")+1))+list(range(ord("®"), ord("ÿ")+1))
27
+ cs = bs[:]
28
+ n = 0
29
+ for b in range(2**8):
30
+ if b not in bs:
31
+ bs.append(b)
32
+ cs.append(2**8+n)
33
+ n += 1
34
+ cs = [chr(n) for n in cs]
35
+ return dict(zip(bs, cs))
36
+
37
+
38
+ def get_pairs(word):
39
+ """Return set of symbol pairs in a word.
40
+ Word is represented as tuple of symbols (symbols being variable-length strings).
41
+ """
42
+ pairs = set()
43
+ prev_char = word[0]
44
+ for char in word[1:]:
45
+ pairs.add((prev_char, char))
46
+ prev_char = char
47
+ return pairs
48
+
49
+
50
+ def basic_clean(text):
51
+ text = ftfy.fix_text(text)
52
+ text = html.unescape(html.unescape(text))
53
+ return text.strip()
54
+
55
+
56
+ def whitespace_clean(text):
57
+ text = re.sub(r'\s+', ' ', text)
58
+ text = text.strip()
59
+ return text
60
+
61
+
62
+ class SimpleTokenizer(object):
63
+ def __init__(self, bpe_path: str = default_bpe()):
64
+ self.byte_encoder = bytes_to_unicode()
65
+ self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
66
+ merges = gzip.open(bpe_path).read().decode("utf-8").split('\n')
67
+ merges = merges[1:49152-256-2+1]
68
+ merges = [tuple(merge.split()) for merge in merges]
69
+ vocab = list(bytes_to_unicode().values())
70
+ vocab = vocab + [v+'</w>' for v in vocab]
71
+ for merge in merges:
72
+ vocab.append(''.join(merge))
73
+ vocab.extend(['<|startoftext|>', '<|endoftext|>'])
74
+ self.encoder = dict(zip(vocab, range(len(vocab))))
75
+ self.decoder = {v: k for k, v in self.encoder.items()}
76
+ self.bpe_ranks = dict(zip(merges, range(len(merges))))
77
+ self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
78
+ self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", re.IGNORECASE)
79
+
80
+ def bpe(self, token):
81
+ if token in self.cache:
82
+ return self.cache[token]
83
+ word = tuple(token[:-1]) + ( token[-1] + '</w>',)
84
+ pairs = get_pairs(word)
85
+
86
+ if not pairs:
87
+ return token+'</w>'
88
+
89
+ while True:
90
+ bigram = min(pairs, key = lambda pair: self.bpe_ranks.get(pair, float('inf')))
91
+ if bigram not in self.bpe_ranks:
92
+ break
93
+ first, second = bigram
94
+ new_word = []
95
+ i = 0
96
+ while i < len(word):
97
+ try:
98
+ j = word.index(first, i)
99
+ new_word.extend(word[i:j])
100
+ i = j
101
+ except:
102
+ new_word.extend(word[i:])
103
+ break
104
+
105
+ if word[i] == first and i < len(word)-1 and word[i+1] == second:
106
+ new_word.append(first+second)
107
+ i += 2
108
+ else:
109
+ new_word.append(word[i])
110
+ i += 1
111
+ new_word = tuple(new_word)
112
+ word = new_word
113
+ if len(word) == 1:
114
+ break
115
+ else:
116
+ pairs = get_pairs(word)
117
+ word = ' '.join(word)
118
+ self.cache[token] = word
119
+ return word
120
+
121
+ def encode(self, text):
122
+ bpe_tokens = []
123
+ text = whitespace_clean(basic_clean(text)).lower()
124
+ for token in re.findall(self.pat, text):
125
+ token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
126
+ bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
127
+ return bpe_tokens
128
+
129
+ def decode(self, tokens):
130
+ text = ''.join([self.decoder[token] for token in tokens])
131
+ text = bytearray([self.byte_decoder[c] for c in text]).decode('utf-8', errors="replace").replace('</w>', ' ')
132
+ return text
models/dsp/CAM/pytorch_grad_cam/__init__.py ADDED
@@ -0,0 +1,14 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .grad_cam import GradCAM
2
+ from .ablation_layer import AblationLayer, AblationLayerVit, AblationLayerFasterRCNN
3
+ from .ablation_cam import AblationCAM
4
+ from .xgrad_cam import XGradCAM
5
+ from .grad_cam_plusplus import GradCAMPlusPlus
6
+ from .score_cam import ScoreCAM
7
+ from .layer_cam import LayerCAM
8
+ from .eigen_cam import EigenCAM
9
+ from .eigen_grad_cam import EigenGradCAM
10
+ from .fullgrad_cam import FullGrad
11
+ from .guided_backprop import GuidedBackpropReLUModel
12
+ from .activations_and_gradients import ActivationsAndGradients
13
+ from .utils import model_targets
14
+ from .utils import reshape_transforms
models/dsp/CAM/pytorch_grad_cam/ablation_cam.py ADDED
@@ -0,0 +1,134 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ import tqdm
4
+ from typing import Callable, List
5
+ from .base_cam import BaseCAM
6
+ from .utils.find_layers import replace_layer_recursive
7
+ from .ablation_layer import AblationLayer
8
+
9
+
10
+ """ Implementation of AblationCAM
11
+ https://openaccess.thecvf.com/content_WACV_2020/papers/Desai_Ablation-CAM_Visual_Explanations_for_Deep_Convolutional_Network_via_Gradient-free_Localization_WACV_2020_paper.pdf
12
+
13
+ Ablate individual activations, and then measure the drop in the target score.
14
+
15
+ In the current implementation, the target layer activations is cached, so it won't be re-computed.
16
+ However layers before it, if any, will not be cached.
17
+ This means that if the target layer is a large block, for example model.featuers (in vgg), there will
18
+ be a large save in run time.
19
+
20
+ Since we have to go over many channels and ablate them, and every channel ablation requires a forward pass,
21
+ it would be nice if we could avoid doing that for channels that won't contribute anwyay, making it much faster.
22
+ The parameter ratio_channels_to_ablate controls how many channels should be ablated, using an experimental method
23
+ (to be improved). The default 1.0 value means that all channels will be ablated.
24
+ """
25
+
26
+
27
+ class AblationCAM(BaseCAM):
28
+ def __init__(self,
29
+ model: torch.nn.Module,
30
+ target_layers: List[torch.nn.Module],
31
+ use_cuda: bool = False,
32
+ reshape_transform: Callable = None,
33
+ ablation_layer: torch.nn.Module = AblationLayer(),
34
+ batch_size: int = 32,
35
+ ratio_channels_to_ablate: float = 1.0) -> None:
36
+
37
+ super(AblationCAM, self).__init__(model,
38
+ target_layers,
39
+ use_cuda,
40
+ reshape_transform,
41
+ uses_gradients=False)
42
+ self.batch_size = batch_size
43
+ self.ablation_layer = ablation_layer
44
+ self.ratio_channels_to_ablate = ratio_channels_to_ablate
45
+
46
+ def save_activation(self, module, input, output) -> None:
47
+ """ Helper function to save the raw activations from the target layer """
48
+ self.activations = output
49
+
50
+ def assemble_ablation_scores(self,
51
+ new_scores: list,
52
+ original_score: float ,
53
+ ablated_channels: np.ndarray,
54
+ number_of_channels: int) -> np.ndarray:
55
+ """ Take the value from the channels that were ablated,
56
+ and just set the original score for the channels that were skipped """
57
+
58
+ index = 0
59
+ result = []
60
+ sorted_indices = np.argsort(ablated_channels)
61
+ ablated_channels = ablated_channels[sorted_indices]
62
+ new_scores = np.float32(new_scores)[sorted_indices]
63
+
64
+ for i in range(number_of_channels):
65
+ if index < len(ablated_channels) and ablated_channels[index] == i:
66
+ weight = new_scores[index]
67
+ index = index + 1
68
+ else:
69
+ weight = original_score
70
+ result.append(weight)
71
+
72
+ return result
73
+
74
+ def get_cam_weights(self,
75
+ input_tensor: torch.Tensor,
76
+ target_layer: torch.nn.Module,
77
+ targets: List[Callable],
78
+ activations: torch.Tensor,
79
+ grads: torch.Tensor) -> np.ndarray:
80
+
81
+ # Do a forward pass, compute the target scores, and cache the activations
82
+ handle = target_layer.register_forward_hook(self.save_activation)
83
+ with torch.no_grad():
84
+ outputs = self.model(input_tensor)
85
+ handle.remove()
86
+ original_scores = np.float32([target(output).cpu().item() for target, output in zip(targets, outputs)])
87
+
88
+ # Replace the layer with the ablation layer.
89
+ # When we finish, we will replace it back, so the original model is unchanged.
90
+ ablation_layer = self.ablation_layer
91
+ replace_layer_recursive(self.model, target_layer, ablation_layer)
92
+
93
+ number_of_channels = activations.shape[1]
94
+ weights = []
95
+ # This is a "gradient free" method, so we don't need gradients here.
96
+ with torch.no_grad():
97
+ # Loop over each of the batch images and ablate activations for it.
98
+ for batch_index, (target, tensor) in enumerate(zip(targets, input_tensor)):
99
+ new_scores = []
100
+ batch_tensor = tensor.repeat(self.batch_size, 1, 1, 1)
101
+
102
+ # Check which channels should be ablated. Normally this will be all channels,
103
+ # But we can also try to speed this up by using a low ratio_channels_to_ablate.
104
+ channels_to_ablate = ablation_layer.activations_to_be_ablated(activations[batch_index, :],
105
+ self.ratio_channels_to_ablate)
106
+ number_channels_to_ablate = len(channels_to_ablate)
107
+
108
+ for i in tqdm.tqdm(range(0, number_channels_to_ablate, self.batch_size)):
109
+ if i + self.batch_size > number_channels_to_ablate:
110
+ batch_tensor = batch_tensor[:(number_channels_to_ablate - i)]
111
+
112
+ # Change the state of the ablation layer so it ablates the next channels.
113
+ # TBD: Move this into the ablation layer forward pass.
114
+ ablation_layer.set_next_batch(input_batch_index=batch_index,
115
+ activations=self.activations,
116
+ num_channels_to_ablate=batch_tensor.size(0))
117
+ score = [target(o).cpu().item() for o in self.model(batch_tensor)]
118
+ new_scores.extend(score)
119
+ ablation_layer.indices = ablation_layer.indices[batch_tensor.size(0):]
120
+
121
+ new_scores = self.assemble_ablation_scores(new_scores,
122
+ original_scores[batch_index],
123
+ channels_to_ablate,
124
+ number_of_channels)
125
+ weights.extend(new_scores)
126
+
127
+ weights = np.float32(weights)
128
+ weights = weights.reshape(activations.shape[:2])
129
+ original_scores = original_scores[:, None]
130
+ weights = (original_scores - weights) / original_scores
131
+
132
+ # Replace the model back to the original state
133
+ replace_layer_recursive(self.model, ablation_layer, target_layer)
134
+ return weights
models/dsp/CAM/pytorch_grad_cam/ablation_cam_multilayer.py ADDED
@@ -0,0 +1,136 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+ import torch
4
+ import tqdm
5
+ from .base_cam import BaseCAM
6
+
7
+
8
+ class AblationLayer(torch.nn.Module):
9
+ def __init__(self, layer, reshape_transform, indices):
10
+ super(AblationLayer, self).__init__()
11
+
12
+ self.layer = layer
13
+ self.reshape_transform = reshape_transform
14
+ # The channels to zero out:
15
+ self.indices = indices
16
+
17
+ def forward(self, x):
18
+ self.__call__(x)
19
+
20
+ def __call__(self, x):
21
+ output = self.layer(x)
22
+
23
+ # Hack to work with ViT,
24
+ # Since the activation channels are last and not first like in CNNs
25
+ # Probably should remove it?
26
+ if self.reshape_transform is not None:
27
+ output = output.transpose(1, 2)
28
+
29
+ for i in range(output.size(0)):
30
+
31
+ # Commonly the minimum activation will be 0,
32
+ # And then it makes sense to zero it out.
33
+ # However depending on the architecture,
34
+ # If the values can be negative, we use very negative values
35
+ # to perform the ablation, deviating from the paper.
36
+ if torch.min(output) == 0:
37
+ output[i, self.indices[i], :] = 0
38
+ else:
39
+ ABLATION_VALUE = 1e5
40
+ output[i, self.indices[i], :] = torch.min(
41
+ output) - ABLATION_VALUE
42
+
43
+ if self.reshape_transform is not None:
44
+ output = output.transpose(2, 1)
45
+
46
+ return output
47
+
48
+
49
+ def replace_layer_recursive(model, old_layer, new_layer):
50
+ for name, layer in model._modules.items():
51
+ if layer == old_layer:
52
+ model._modules[name] = new_layer
53
+ return True
54
+ elif replace_layer_recursive(layer, old_layer, new_layer):
55
+ return True
56
+ return False
57
+
58
+
59
+ class AblationCAM(BaseCAM):
60
+ def __init__(self, model, target_layers, use_cuda=False,
61
+ reshape_transform=None):
62
+ super(AblationCAM, self).__init__(model, target_layers, use_cuda,
63
+ reshape_transform)
64
+
65
+ if len(target_layers) > 1:
66
+ print(
67
+ "Warning. You are usign Ablation CAM with more than 1 layers. "
68
+ "This is supported only if all layers have the same output shape")
69
+
70
+ def set_ablation_layers(self):
71
+ self.ablation_layers = []
72
+ for target_layer in self.target_layers:
73
+ ablation_layer = AblationLayer(target_layer,
74
+ self.reshape_transform, indices=[])
75
+ self.ablation_layers.append(ablation_layer)
76
+ replace_layer_recursive(self.model, target_layer, ablation_layer)
77
+
78
+ def unset_ablation_layers(self):
79
+ # replace the model back to the original state
80
+ for ablation_layer, target_layer in zip(
81
+ self.ablation_layers, self.target_layers):
82
+ replace_layer_recursive(self.model, ablation_layer, target_layer)
83
+
84
+ def set_ablation_layer_batch_indices(self, indices):
85
+ for ablation_layer in self.ablation_layers:
86
+ ablation_layer.indices = indices
87
+
88
+ def trim_ablation_layer_batch_indices(self, keep):
89
+ for ablation_layer in self.ablation_layers:
90
+ ablation_layer.indices = ablation_layer.indices[:keep]
91
+
92
+ def get_cam_weights(self,
93
+ input_tensor,
94
+ target_category,
95
+ activations,
96
+ grads):
97
+ with torch.no_grad():
98
+ outputs = self.model(input_tensor).cpu().numpy()
99
+ original_scores = []
100
+ for i in range(input_tensor.size(0)):
101
+ original_scores.append(outputs[i, target_category[i]])
102
+ original_scores = np.float32(original_scores)
103
+
104
+ self.set_ablation_layers()
105
+
106
+ if hasattr(self, "batch_size"):
107
+ BATCH_SIZE = self.batch_size
108
+ else:
109
+ BATCH_SIZE = 32
110
+
111
+ number_of_channels = activations.shape[1]
112
+ weights = []
113
+
114
+ with torch.no_grad():
115
+ # Iterate over the input batch
116
+ for tensor, category in zip(input_tensor, target_category):
117
+ batch_tensor = tensor.repeat(BATCH_SIZE, 1, 1, 1)
118
+ for i in tqdm.tqdm(range(0, number_of_channels, BATCH_SIZE)):
119
+ self.set_ablation_layer_batch_indices(
120
+ list(range(i, i + BATCH_SIZE)))
121
+
122
+ if i + BATCH_SIZE > number_of_channels:
123
+ keep = number_of_channels - i
124
+ batch_tensor = batch_tensor[:keep]
125
+ self.trim_ablation_layer_batch_indices(self, keep)
126
+ score = self.model(batch_tensor)[:, category].cpu().numpy()
127
+ weights.extend(score)
128
+
129
+ weights = np.float32(weights)
130
+ weights = weights.reshape(activations.shape[:2])
131
+ original_scores = original_scores[:, None]
132
+ weights = (original_scores - weights) / original_scores
133
+
134
+ # replace the model back to the original state
135
+ self.unset_ablation_layers()
136
+ return weights
models/dsp/CAM/pytorch_grad_cam/ablation_layer.py ADDED
@@ -0,0 +1,131 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ from collections import OrderedDict
3
+ import numpy as np
4
+ from .utils.svd_on_activations import get_2d_projection
5
+
6
+
7
+ class AblationLayer(torch.nn.Module):
8
+ def __init__(self):
9
+ super(AblationLayer, self).__init__()
10
+
11
+ def objectiveness_mask_from_svd(self, activations, threshold=0.01):
12
+ """ Experimental method to get a binary mask to compare if the activation is worth ablating.
13
+ The idea is to apply the EigenCAM method by doing PCA on the activations.
14
+ Then we create a binary mask by comparing to a low threshold.
15
+ Areas that are masked out, are probably not interesting anyway.
16
+ """
17
+
18
+ projection = get_2d_projection(activations[None, :])[0, :]
19
+ projection = np.abs(projection)
20
+ projection = projection - projection.min()
21
+ projection = projection / projection.max()
22
+ projection = projection > threshold
23
+ return projection
24
+
25
+ def activations_to_be_ablated(self, activations, ratio_channels_to_ablate=1.0):
26
+ """ Experimental method to get a binary mask to compare if the activation is worth ablating.
27
+ Create a binary CAM mask with objectiveness_mask_from_svd.
28
+ Score each Activation channel, by seeing how much of its values are inside the mask.
29
+ Then keep the top channels.
30
+
31
+ """
32
+ if ratio_channels_to_ablate == 1.0:
33
+ self.indices = np.int32(range(activations.shape[0]))
34
+ return self.indices
35
+
36
+ projection = self.objectiveness_mask_from_svd(activations)
37
+
38
+ scores = []
39
+ for channel in activations:
40
+ normalized = np.abs(channel)
41
+ normalized = normalized - normalized.min()
42
+ normalized = normalized / np.max(normalized)
43
+ score = (projection*normalized).sum() / normalized.sum()
44
+ scores.append(score)
45
+ scores = np.float32(scores)
46
+
47
+ indices = list(np.argsort(scores))
48
+ high_score_indices = indices[::-1][: int(len(indices) * ratio_channels_to_ablate)]
49
+ low_score_indices = indices[: int(len(indices) * ratio_channels_to_ablate)]
50
+ self.indices = np.int32(high_score_indices + low_score_indices)
51
+ return self.indices
52
+
53
+ def set_next_batch(self, input_batch_index, activations, num_channels_to_ablate):
54
+ """ This creates the next batch of activations from the layer.
55
+ Just take corresponding batch member from activations, and repeat it num_channels_to_ablate times.
56
+ """
57
+ self.activations = activations[input_batch_index, :, :, :].clone().unsqueeze(0).repeat(num_channels_to_ablate, 1, 1, 1)
58
+
59
+ def __call__(self, x):
60
+ output = self.activations
61
+ for i in range(output.size(0)):
62
+ # Commonly the minimum activation will be 0,
63
+ # And then it makes sense to zero it out.
64
+ # However depending on the architecture,
65
+ # If the values can be negative, we use very negative values
66
+ # to perform the ablation, deviating from the paper.
67
+ if torch.min(output) == 0:
68
+ output[i, self.indices[i], :] = 0
69
+ else:
70
+ ABLATION_VALUE = 1e7
71
+ output[i, self.indices[i], :] = torch.min(
72
+ output) - ABLATION_VALUE
73
+
74
+ return output
75
+
76
+
77
+ class AblationLayerVit(AblationLayer):
78
+ def __init__(self):
79
+ super(AblationLayerVit, self).__init__()
80
+
81
+ def __call__(self, x):
82
+ output = self.activations
83
+ output = output.transpose(1, 2)
84
+ for i in range(output.size(0)):
85
+
86
+ # Commonly the minimum activation will be 0,
87
+ # And then it makes sense to zero it out.
88
+ # However depending on the architecture,
89
+ # If the values can be negative, we use very negative values
90
+ # to perform the ablation, deviating from the paper.
91
+ if torch.min(output) == 0:
92
+ output[i, self.indices[i], :] = 0
93
+ else:
94
+ ABLATION_VALUE = 1e7
95
+ output[i, self.indices[i], :] = torch.min(
96
+ output) - ABLATION_VALUE
97
+
98
+ output = output.transpose(2, 1)
99
+
100
+ return output
101
+
102
+ def set_next_batch(self, input_batch_index, activations, num_channels_to_ablate):
103
+ """ This creates the next batch of activations from the layer.
104
+ Just take corresponding batch member from activations, and repeat it num_channels_to_ablate times.
105
+ """
106
+ self.activations = activations[input_batch_index, :, :].clone().unsqueeze(0).repeat(num_channels_to_ablate, 1, 1)
107
+
108
+
109
+
110
+ class AblationLayerFasterRCNN(AblationLayer):
111
+ def __init__(self):
112
+ super(AblationLayerFasterRCNN, self).__init__()
113
+
114
+ def set_next_batch(self, input_batch_index, activations, num_channels_to_ablate):
115
+ """ Extract the next batch member from activations,
116
+ and repeat it num_channels_to_ablate times.
117
+ """
118
+ self.activations = OrderedDict()
119
+ for key, value in activations.items():
120
+ fpn_activation = value[input_batch_index, :, :, :].clone().unsqueeze(0)
121
+ self.activations[key] = fpn_activation.repeat(num_channels_to_ablate, 1, 1, 1)
122
+
123
+ def __call__(self, x):
124
+ result = self.activations
125
+ layers = {0: '0', 1: '1', 2: '2', 3: '3', 4: 'pool'}
126
+ num_channels_to_ablate = result['pool'].size(0)
127
+ for i in range(num_channels_to_ablate):
128
+ pyramid_layer = int(self.indices[i]/256)
129
+ index_in_pyramid_layer = int(self.indices[i] % 256)
130
+ result[layers[pyramid_layer]][i, index_in_pyramid_layer, :, :] = -1000
131
+ return result
models/dsp/CAM/pytorch_grad_cam/activations_and_gradients.py ADDED
@@ -0,0 +1,55 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ class ActivationsAndGradients:
2
+ """ Class for extracting activations and
3
+ registering gradients from targetted intermediate layers """
4
+
5
+ def __init__(self, model, target_layers, reshape_transform):
6
+ self.model = model
7
+ self.gradients = []
8
+ self.activations = []
9
+ self.reshape_transform = reshape_transform
10
+ self.handles = []
11
+ for target_layer in target_layers:
12
+ self.handles.append(
13
+ target_layer.register_forward_hook(self.save_activation))
14
+ # Because of https://github.com/pytorch/pytorch/issues/61519,
15
+ # we don't use backward hook to record gradients.
16
+ self.handles.append(
17
+ target_layer.register_forward_hook(self.save_gradient))
18
+
19
+ def save_activation(self, module, input, output):
20
+ activation = output
21
+
22
+ if self.reshape_transform is not None:
23
+ activation = self.reshape_transform(activation, self.height, self.width)
24
+ # self.activations.append(activation.cpu().detach())
25
+ # self.activations.append(activation.detach()) # original detach
26
+ self.activations.append(activation)
27
+
28
+ def save_gradient(self, module, input, output):
29
+ if not hasattr(output, "requires_grad") or not output.requires_grad:
30
+ # You can only register hooks on tensor requires grad.
31
+ return
32
+
33
+ # Gradients are computed in reverse order
34
+ def _store_grad(grad):
35
+ if self.reshape_transform is not None:
36
+ grad = self.reshape_transform(grad, self.height, self.width)
37
+ # self.gradients = [grad.cpu().detach()] + self.gradients
38
+ # self.gradients = [grad.detach()] + self.gradients # original detach
39
+ self.gradients = [grad] + self.gradients
40
+
41
+ output.register_hook(_store_grad)
42
+
43
+ def __call__(self, x, H, W):
44
+ self.height = H // 16
45
+ self.width = W // 16
46
+ self.gradients = []
47
+ self.activations = []
48
+ if isinstance(x, list):
49
+ return self.model.forward_last_layer(x[0], x[1])
50
+ else:
51
+ return self.model(x)
52
+
53
+ def release(self):
54
+ for handle in self.handles:
55
+ handle.remove()
models/dsp/CAM/pytorch_grad_cam/base_cam.py ADDED
@@ -0,0 +1,227 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ import ttach as tta
4
+ from typing import Callable, List, Tuple
5
+ from .activations_and_gradients import ActivationsAndGradients
6
+ from .utils.svd_on_activations import get_2d_projection
7
+ from .utils.image import scale_cam_image, scale_cam_image_torch
8
+ from .utils.model_targets import ClassifierOutputTarget
9
+
10
+
11
+ class BaseCAM:
12
+ def __init__(self,
13
+ model: torch.nn.Module,
14
+ target_layers: List[torch.nn.Module],
15
+ use_cuda: bool = False,
16
+ reshape_transform: Callable = None,
17
+ compute_input_gradient: bool = False,
18
+ uses_gradients: bool = True) -> None:
19
+ self.model = model.eval()
20
+ self.target_layers = target_layers
21
+ self.cuda = use_cuda
22
+ # if self.cuda:
23
+ # self.model = model.cuda()
24
+ self.reshape_transform = reshape_transform
25
+ self.compute_input_gradient = compute_input_gradient
26
+ self.uses_gradients = uses_gradients
27
+ self.activations_and_grads = ActivationsAndGradients(
28
+ self.model, target_layers, reshape_transform)
29
+
30
+ """ Get a vector of weights for every channel in the target layer.
31
+ Methods that return weights channels,
32
+ will typically need to only implement this function. """
33
+
34
+ def set_device(self, device):
35
+ self.device = device
36
+
37
+ def get_cam_weights(self,
38
+ input_tensor: torch.Tensor,
39
+ target_layers: List[torch.nn.Module],
40
+ targets: List[torch.nn.Module],
41
+ activations: torch.Tensor,
42
+ grads: torch.Tensor) -> np.ndarray:
43
+ raise Exception("Not Implemented")
44
+
45
+ def get_cam_image(self,
46
+ input_tensor: torch.Tensor,
47
+ target_layer: torch.nn.Module,
48
+ targets: List[torch.nn.Module],
49
+ activations: torch.Tensor,
50
+ grads: torch.Tensor,
51
+ eigen_smooth: bool = False) -> np.ndarray:
52
+
53
+ weights = self.get_cam_weights(input_tensor,
54
+ target_layer,
55
+ targets,
56
+ activations,
57
+ grads)
58
+ weighted_activations = weights[:, :, None, None] * activations
59
+ if eigen_smooth:
60
+ cam = get_2d_projection(weighted_activations)
61
+ else:
62
+ cam = weighted_activations.sum(axis=1)
63
+ return cam
64
+
65
+ def forward(self,
66
+ input_tensor: torch.Tensor,
67
+ targets: List[torch.nn.Module],
68
+ target_size,
69
+ eigen_smooth: bool = False) -> np.ndarray:
70
+
71
+ if self.cuda:
72
+ if isinstance(input_tensor, list):
73
+ input_tensor = [t.to(self.device) if isinstance(t, torch.Tensor) else t for t in input_tensor]
74
+ else:
75
+ input_tensor = input_tensor.to(self.device)
76
+
77
+ if self.compute_input_gradient:
78
+ input_tensor = torch.autograd.Variable(input_tensor,
79
+ requires_grad=True)
80
+
81
+ W,H = self.get_target_width_height(input_tensor)
82
+ outputs = self.activations_and_grads(input_tensor,H,W)
83
+ if targets is None:
84
+ if isinstance(input_tensor, list):
85
+ target_categories = np.argmax(outputs[0].cpu().data.numpy(), axis=-1)
86
+ else:
87
+ target_categories = np.argmax(outputs.cpu().data.numpy(), axis=-1)
88
+ targets = [ClassifierOutputTarget(category) for category in target_categories]
89
+
90
+ if self.uses_gradients:
91
+ self.model.zero_grad()
92
+ if isinstance(input_tensor, list):
93
+ loss = sum([target(output[0]) for target, output in zip(targets, outputs)])
94
+ else:
95
+ loss = sum([target(output) for target, output in zip(targets, outputs)])
96
+ loss.backward(retain_graph=True)
97
+
98
+ # In most of the saliency attribution papers, the saliency is
99
+ # computed with a single target layer.
100
+ # Commonly it is the last convolutional layer.
101
+ # Here we support passing a list with multiple target layers.
102
+ # It will compute the saliency image for every image,
103
+ # and then aggregate them (with a default mean aggregation).
104
+ # This gives you more flexibility in case you just want to
105
+ # use all conv layers for example, all Batchnorm layers,
106
+ # or something else.
107
+ cam_per_layer = self.compute_cam_per_layer(input_tensor,
108
+ targets,
109
+ target_size,
110
+ eigen_smooth)
111
+ if isinstance(input_tensor, list):
112
+ return self.aggregate_multi_layers_torch(cam_per_layer), outputs[0], outputs[1]
113
+ else:
114
+ return self.aggregate_multi_layers(cam_per_layer), outputs
115
+
116
+ def get_target_width_height(self,
117
+ input_tensor: torch.Tensor) -> Tuple[int, int]:
118
+ if isinstance(input_tensor, list):
119
+ width, height = input_tensor[-1], input_tensor[-2]
120
+ return width, height
121
+
122
+ def compute_cam_per_layer(
123
+ self,
124
+ input_tensor: torch.Tensor,
125
+ targets: List[torch.nn.Module],
126
+ target_size,
127
+ eigen_smooth: bool) -> np.ndarray:
128
+ activations_list = [a#.cpu().data.numpy()
129
+ for a in self.activations_and_grads.activations]
130
+ grads_list = [g#.cpu().data.numpy()
131
+ for g in self.activations_and_grads.gradients]
132
+
133
+ cam_per_target_layer = []
134
+ # Loop over the saliency image from every layer
135
+ for i in range(len(self.target_layers)):
136
+ target_layer = self.target_layers[i]
137
+ layer_activations = None
138
+ layer_grads = None
139
+ if i < len(activations_list):
140
+ layer_activations = activations_list[i]
141
+ if i < len(grads_list):
142
+ layer_grads = grads_list[i]
143
+
144
+ cam = self.get_cam_image(input_tensor,
145
+ target_layer,
146
+ targets,
147
+ layer_activations,
148
+ layer_grads,
149
+ eigen_smooth)
150
+ cam = torch.clamp(cam, min=0).float()
151
+ scaled = scale_cam_image_torch(cam, target_size)
152
+ # cam = np.maximum(cam, 0).astype(np.float32)#float16->32
153
+ # scaled = scale_cam_image(cam, target_size)
154
+ cam_per_target_layer.append(scaled[:, None, :])
155
+
156
+ return cam_per_target_layer
157
+
158
+ def aggregate_multi_layers(self, cam_per_target_layer: np.ndarray) -> np.ndarray:
159
+ cam_per_target_layer = np.concatenate(cam_per_target_layer, axis=1)
160
+ cam_per_target_layer = np.maximum(cam_per_target_layer, 0)
161
+ result = np.mean(cam_per_target_layer, axis=1)
162
+ return scale_cam_image(result)
163
+
164
+ def aggregate_multi_layers_torch(self, cam_per_target_layer):
165
+ cam_per_target_layer = torch.cat(cam_per_target_layer, dim=1)
166
+ cam_per_target_layer = torch.clamp(cam_per_target_layer, min=0)
167
+ result = cam_per_target_layer.mean(dim=1, keepdim=True)
168
+ return scale_cam_image_torch(result)
169
+
170
+ def forward_augmentation_smoothing(self,
171
+ input_tensor: torch.Tensor,
172
+ targets: List[torch.nn.Module],
173
+ eigen_smooth: bool = False) -> np.ndarray:
174
+ transforms = tta.Compose(
175
+ [
176
+ tta.HorizontalFlip(),
177
+ tta.Multiply(factors=[0.9, 1, 1.1]),
178
+ ]
179
+ )
180
+ cams = []
181
+ for transform in transforms:
182
+ augmented_tensor = transform.augment_image(input_tensor)
183
+ cam = self.forward(augmented_tensor,
184
+ targets,
185
+ eigen_smooth)
186
+
187
+ # The ttach library expects a tensor of size BxCxHxW
188
+ cam = cam[:, None, :, :]
189
+ cam = torch.from_numpy(cam)
190
+ cam = transform.deaugment_mask(cam)
191
+
192
+ # Back to numpy float32, HxW
193
+ cam = cam.numpy()
194
+ cam = cam[:, 0, :, :]
195
+ cams.append(cam)
196
+
197
+ cam = np.mean(np.float32(cams), axis=0)
198
+ return cam
199
+
200
+ def __call__(self,
201
+ input_tensor: torch.Tensor,
202
+ targets: List[torch.nn.Module] = None,
203
+ target_size=None,
204
+ aug_smooth: bool = False,
205
+ eigen_smooth: bool = False) -> np.ndarray:
206
+
207
+ # Smooth the CAM result with test time augmentation
208
+ if aug_smooth is True:
209
+ return self.forward_augmentation_smoothing(
210
+ input_tensor, targets, eigen_smooth)
211
+
212
+ return self.forward(input_tensor,
213
+ targets, target_size,eigen_smooth)
214
+
215
+ def __del__(self):
216
+ self.activations_and_grads.release()
217
+
218
+ def __enter__(self):
219
+ return self
220
+
221
+ def __exit__(self, exc_type, exc_value, exc_tb):
222
+ self.activations_and_grads.release()
223
+ if isinstance(exc_value, IndexError):
224
+ # Handle IndexError here...
225
+ print(
226
+ f"An exception occurred in CAM with block: {exc_type}. Message: {exc_value}")
227
+ return True
models/dsp/CAM/pytorch_grad_cam/eigen_cam.py ADDED
@@ -0,0 +1,23 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_cam import BaseCAM
2
+ from .utils.svd_on_activations import get_2d_projection
3
+
4
+ # https://arxiv.org/abs/2008.00299
5
+
6
+
7
+ class EigenCAM(BaseCAM):
8
+ def __init__(self, model, target_layers, use_cuda=False,
9
+ reshape_transform=None):
10
+ super(EigenCAM, self).__init__(model,
11
+ target_layers,
12
+ use_cuda,
13
+ reshape_transform,
14
+ uses_gradients=False)
15
+
16
+ def get_cam_image(self,
17
+ input_tensor,
18
+ target_layer,
19
+ target_category,
20
+ activations,
21
+ grads,
22
+ eigen_smooth):
23
+ return get_2d_projection(activations)
models/dsp/CAM/pytorch_grad_cam/eigen_grad_cam.py ADDED
@@ -0,0 +1,21 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from .base_cam import BaseCAM
2
+ from .utils.svd_on_activations import get_2d_projection
3
+
4
+ # Like Eigen CAM: https://arxiv.org/abs/2008.00299
5
+ # But multiply the activations x gradients
6
+
7
+
8
+ class EigenGradCAM(BaseCAM):
9
+ def __init__(self, model, target_layers, use_cuda=False,
10
+ reshape_transform=None):
11
+ super(EigenGradCAM, self).__init__(model, target_layers, use_cuda,
12
+ reshape_transform)
13
+
14
+ def get_cam_image(self,
15
+ input_tensor,
16
+ target_layer,
17
+ target_category,
18
+ activations,
19
+ grads,
20
+ eigen_smooth):
21
+ return get_2d_projection(grads * activations)
models/dsp/CAM/pytorch_grad_cam/fullgrad_cam.py ADDED
@@ -0,0 +1,95 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from .base_cam import BaseCAM
4
+ from .utils.find_layers import find_layer_predicate_recursive
5
+ from .utils.svd_on_activations import get_2d_projection
6
+ from .utils.image import scale_accross_batch_and_channels, scale_cam_image
7
+
8
+ # https://arxiv.org/abs/1905.00780
9
+
10
+
11
+ class FullGrad(BaseCAM):
12
+ def __init__(self, model, target_layers, use_cuda=False,
13
+ reshape_transform=None):
14
+ if len(target_layers) > 0:
15
+ print(
16
+ "Warning: target_layers is ignored in FullGrad. All bias layers will be used instead")
17
+
18
+ def layer_with_2D_bias(layer):
19
+ bias_target_layers = [torch.nn.Conv2d, torch.nn.BatchNorm2d]
20
+ if type(layer) in bias_target_layers and layer.bias is not None:
21
+ return True
22
+ return False
23
+ target_layers = find_layer_predicate_recursive(
24
+ model, layer_with_2D_bias)
25
+ super(
26
+ FullGrad,
27
+ self).__init__(
28
+ model,
29
+ target_layers,
30
+ use_cuda,
31
+ reshape_transform,
32
+ compute_input_gradient=True)
33
+ self.bias_data = [self.get_bias_data(
34
+ layer).cpu().numpy() for layer in target_layers]
35
+
36
+ def get_bias_data(self, layer):
37
+ # Borrowed from official paper impl:
38
+ # https://github.com/idiap/fullgrad-saliency/blob/master/saliency/tensor_extractor.py#L47
39
+ if isinstance(layer, torch.nn.BatchNorm2d):
40
+ bias = - (layer.running_mean * layer.weight
41
+ / torch.sqrt(layer.running_var + layer.eps)) + layer.bias
42
+ return bias.data
43
+ else:
44
+ return layer.bias.data
45
+
46
+ def compute_cam_per_layer(
47
+ self,
48
+ input_tensor,
49
+ target_category,
50
+ eigen_smooth):
51
+ input_grad = input_tensor.grad.data.cpu().numpy()
52
+ grads_list = [g.cpu().data.numpy() for g in
53
+ self.activations_and_grads.gradients]
54
+ cam_per_target_layer = []
55
+ target_size = self.get_target_width_height(input_tensor)
56
+
57
+ gradient_multiplied_input = input_grad * input_tensor.data.cpu().numpy()
58
+ gradient_multiplied_input = np.abs(gradient_multiplied_input)
59
+ gradient_multiplied_input = scale_accross_batch_and_channels(
60
+ gradient_multiplied_input,
61
+ target_size)
62
+ cam_per_target_layer.append(gradient_multiplied_input)
63
+
64
+ # Loop over the saliency image from every layer
65
+ assert(len(self.bias_data) == len(grads_list))
66
+ for bias, grads in zip(self.bias_data, grads_list):
67
+ bias = bias[None, :, None, None]
68
+ # In the paper they take the absolute value,
69
+ # but possibily taking only the positive gradients will work
70
+ # better.
71
+ bias_grad = np.abs(bias * grads)
72
+ result = scale_accross_batch_and_channels(
73
+ bias_grad, target_size)
74
+ result = np.sum(result, axis=1)
75
+ cam_per_target_layer.append(result[:, None, :])
76
+ cam_per_target_layer = np.concatenate(cam_per_target_layer, axis=1)
77
+ if eigen_smooth:
78
+ # Resize to a smaller image, since this method typically has a very large number of channels,
79
+ # and then consumes a lot of memory
80
+ cam_per_target_layer = scale_accross_batch_and_channels(
81
+ cam_per_target_layer, (target_size[0] // 8, target_size[1] // 8))
82
+ cam_per_target_layer = get_2d_projection(cam_per_target_layer)
83
+ cam_per_target_layer = cam_per_target_layer[:, None, :, :]
84
+ cam_per_target_layer = scale_accross_batch_and_channels(
85
+ cam_per_target_layer,
86
+ target_size)
87
+ else:
88
+ cam_per_target_layer = np.sum(
89
+ cam_per_target_layer, axis=1)[:, None, :]
90
+
91
+ return cam_per_target_layer
92
+
93
+ def aggregate_multi_layers(self, cam_per_target_layer):
94
+ result = np.sum(cam_per_target_layer, axis=1)
95
+ return scale_cam_image(result)
models/dsp/CAM/pytorch_grad_cam/grad_cam.py ADDED
@@ -0,0 +1,24 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import numpy as np
3
+ from .base_cam import BaseCAM
4
+
5
+
6
+ class GradCAM(BaseCAM):
7
+ def __init__(self, model, target_layers, use_cuda=False,
8
+ reshape_transform=None):
9
+ super(
10
+ GradCAM,
11
+ self).__init__(
12
+ model,
13
+ target_layers,
14
+ use_cuda,
15
+ reshape_transform)
16
+
17
+ def get_cam_weights(self,
18
+ input_tensor,
19
+ target_layer,
20
+ target_category,
21
+ activations,
22
+ grads):
23
+ # return np.mean(grads, axis=(2, 3))
24
+ return torch.mean(grads, dim=(2, 3))
models/dsp/CAM/pytorch_grad_cam/grad_cam_plusplus.py ADDED
@@ -0,0 +1,32 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ from .base_cam import BaseCAM
3
+
4
+ # https://arxiv.org/abs/1710.11063
5
+
6
+
7
+ class GradCAMPlusPlus(BaseCAM):
8
+ def __init__(self, model, target_layers, use_cuda=False,
9
+ reshape_transform=None):
10
+ super(GradCAMPlusPlus, self).__init__(model, target_layers, use_cuda,
11
+ reshape_transform)
12
+
13
+ def get_cam_weights(self,
14
+ input_tensor,
15
+ target_layers,
16
+ target_category,
17
+ activations,
18
+ grads):
19
+ grads_power_2 = grads**2
20
+ grads_power_3 = grads_power_2 * grads
21
+ # Equation 19 in https://arxiv.org/abs/1710.11063
22
+ sum_activations = np.sum(activations, axis=(2, 3))
23
+ eps = 0.000001
24
+ aij = grads_power_2 / (2 * grads_power_2 +
25
+ sum_activations[:, :, None, None] * grads_power_3 + eps)
26
+ # Now bring back the ReLU from eq.7 in the paper,
27
+ # And zero out aijs where the activations are 0
28
+ aij = np.where(grads != 0, aij, 0)
29
+
30
+ weights = np.maximum(grads, 0) * aij
31
+ weights = np.sum(weights, axis=(2, 3))
32
+ return weights
models/dsp/CAM/pytorch_grad_cam/guided_backprop.py ADDED
@@ -0,0 +1,100 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ import torch
3
+ from torch.autograd import Function
4
+ from .utils.find_layers import replace_all_layer_type_recursive
5
+
6
+
7
+ class GuidedBackpropReLU(Function):
8
+ @staticmethod
9
+ def forward(self, input_img):
10
+ positive_mask = (input_img > 0).type_as(input_img)
11
+ output = torch.addcmul(
12
+ torch.zeros(
13
+ input_img.size()).type_as(input_img),
14
+ input_img,
15
+ positive_mask)
16
+ self.save_for_backward(input_img, output)
17
+ return output
18
+
19
+ @staticmethod
20
+ def backward(self, grad_output):
21
+ input_img, output = self.saved_tensors
22
+ grad_input = None
23
+
24
+ positive_mask_1 = (input_img > 0).type_as(grad_output)
25
+ positive_mask_2 = (grad_output > 0).type_as(grad_output)
26
+ grad_input = torch.addcmul(
27
+ torch.zeros(
28
+ input_img.size()).type_as(input_img),
29
+ torch.addcmul(
30
+ torch.zeros(
31
+ input_img.size()).type_as(input_img),
32
+ grad_output,
33
+ positive_mask_1),
34
+ positive_mask_2)
35
+ return grad_input
36
+
37
+
38
+ class GuidedBackpropReLUasModule(torch.nn.Module):
39
+ def __init__(self):
40
+ super(GuidedBackpropReLUasModule, self).__init__()
41
+
42
+ def forward(self, input_img):
43
+ return GuidedBackpropReLU.apply(input_img)
44
+
45
+
46
+ class GuidedBackpropReLUModel:
47
+ def __init__(self, model, use_cuda):
48
+ self.model = model
49
+ self.model.eval()
50
+ self.cuda = use_cuda
51
+ if self.cuda:
52
+ self.model = self.model.cuda()
53
+
54
+ def forward(self, input_img):
55
+ return self.model(input_img)
56
+
57
+ def recursive_replace_relu_with_guidedrelu(self, module_top):
58
+
59
+ for idx, module in module_top._modules.items():
60
+ self.recursive_replace_relu_with_guidedrelu(module)
61
+ if module.__class__.__name__ == 'ReLU':
62
+ module_top._modules[idx] = GuidedBackpropReLU.apply
63
+ print("b")
64
+
65
+ def recursive_replace_guidedrelu_with_relu(self, module_top):
66
+ try:
67
+ for idx, module in module_top._modules.items():
68
+ self.recursive_replace_guidedrelu_with_relu(module)
69
+ if module == GuidedBackpropReLU.apply:
70
+ module_top._modules[idx] = torch.nn.ReLU()
71
+ except BaseException:
72
+ pass
73
+
74
+ def __call__(self, input_img, target_category=None):
75
+ replace_all_layer_type_recursive(self.model,
76
+ torch.nn.ReLU,
77
+ GuidedBackpropReLUasModule())
78
+
79
+ if self.cuda:
80
+ input_img = input_img.cuda()
81
+
82
+ input_img = input_img.requires_grad_(True)
83
+
84
+ output = self.forward(input_img)
85
+
86
+ if target_category is None:
87
+ target_category = np.argmax(output.cpu().data.numpy())
88
+
89
+ loss = output[0, target_category]
90
+ loss.backward(retain_graph=True)
91
+
92
+ output = input_img.grad.cpu().data.numpy()
93
+ output = output[0, :, :, :]
94
+ output = output.transpose((1, 2, 0))
95
+
96
+ replace_all_layer_type_recursive(self.model,
97
+ GuidedBackpropReLUasModule,
98
+ torch.nn.ReLU())
99
+
100
+ return output
models/dsp/CAM/pytorch_grad_cam/layer_cam.py ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ from .base_cam import BaseCAM
3
+ from .utils.svd_on_activations import get_2d_projection
4
+
5
+ # https://ieeexplore.ieee.org/document/9462463
6
+
7
+
8
+ class LayerCAM(BaseCAM):
9
+ def __init__(
10
+ self,
11
+ model,
12
+ target_layers,
13
+ use_cuda=False,
14
+ reshape_transform=None):
15
+ super(
16
+ LayerCAM,
17
+ self).__init__(
18
+ model,
19
+ target_layers,
20
+ use_cuda,
21
+ reshape_transform)
22
+
23
+ def get_cam_image(self,
24
+ input_tensor,
25
+ target_layer,
26
+ target_category,
27
+ activations,
28
+ grads,
29
+ eigen_smooth):
30
+ spatial_weighted_activations = np.maximum(grads, 0) * activations
31
+
32
+ if eigen_smooth:
33
+ cam = get_2d_projection(spatial_weighted_activations)
34
+ else:
35
+ cam = spatial_weighted_activations.sum(axis=1)
36
+ return cam
models/dsp/CAM/pytorch_grad_cam/score_cam.py ADDED
@@ -0,0 +1,63 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import tqdm
3
+ from .base_cam import BaseCAM
4
+
5
+
6
+ class ScoreCAM(BaseCAM):
7
+ def __init__(
8
+ self,
9
+ model,
10
+ target_layers,
11
+ use_cuda=False,
12
+ reshape_transform=None):
13
+ super(ScoreCAM, self).__init__(model,
14
+ target_layers,
15
+ use_cuda,
16
+ reshape_transform=reshape_transform,
17
+ uses_gradients=False)
18
+
19
+ if len(target_layers) > 0:
20
+ print("Warning: You are using ScoreCAM with target layers, "
21
+ "however ScoreCAM will ignore them.")
22
+
23
+ def get_cam_weights(self,
24
+ input_tensor,
25
+ target_layer,
26
+ targets,
27
+ activations,
28
+ grads):
29
+ with torch.no_grad():
30
+ upsample = torch.nn.UpsamplingBilinear2d(
31
+ size=input_tensor.shape[-2:])
32
+ activation_tensor = torch.from_numpy(activations)
33
+ if self.cuda:
34
+ activation_tensor = activation_tensor.cuda()
35
+
36
+ upsampled = upsample(activation_tensor)
37
+
38
+ maxs = upsampled.view(upsampled.size(0),
39
+ upsampled.size(1), -1).max(dim=-1)[0]
40
+ mins = upsampled.view(upsampled.size(0),
41
+ upsampled.size(1), -1).min(dim=-1)[0]
42
+
43
+ maxs, mins = maxs[:, :, None, None], mins[:, :, None, None]
44
+ upsampled = (upsampled - mins) / (maxs - mins)
45
+
46
+ input_tensors = input_tensor[:, None,
47
+ :, :] * upsampled[:, :, None, :, :]
48
+
49
+ if hasattr(self, "batch_size"):
50
+ BATCH_SIZE = self.batch_size
51
+ else:
52
+ BATCH_SIZE = 16
53
+
54
+ scores = []
55
+ for target, tensor in zip(targets, input_tensors):
56
+ for i in tqdm.tqdm(range(0, tensor.size(0), BATCH_SIZE)):
57
+ batch = tensor[i: i + BATCH_SIZE, :]
58
+ outputs = [target(o).cpu().item() for o in self.model(batch)]
59
+ scores.extend(outputs)
60
+ scores = torch.Tensor(scores)
61
+ scores = scores.view(activations.shape[0], activations.shape[1])
62
+ weights = torch.nn.Softmax(dim=-1)(scores).numpy()
63
+ return weights
models/dsp/CAM/pytorch_grad_cam/utils/__init__.py ADDED
@@ -0,0 +1,4 @@
 
 
 
 
 
1
+ from .image import deprocess_image
2
+ from .svd_on_activations import get_2d_projection
3
+ from . import model_targets
4
+ from . import reshape_transforms
models/dsp/CAM/pytorch_grad_cam/utils/find_layers.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ def replace_layer_recursive(model, old_layer, new_layer):
2
+ for name, layer in model._modules.items():
3
+ if layer == old_layer:
4
+ model._modules[name] = new_layer
5
+ return True
6
+ elif replace_layer_recursive(layer, old_layer, new_layer):
7
+ return True
8
+ return False
9
+
10
+
11
+ def replace_all_layer_type_recursive(model, old_layer_type, new_layer):
12
+ for name, layer in model._modules.items():
13
+ if isinstance(layer, old_layer_type):
14
+ model._modules[name] = new_layer
15
+ replace_all_layer_type_recursive(layer, old_layer_type, new_layer)
16
+
17
+
18
+ def find_layer_types_recursive(model, layer_types):
19
+ def predicate(layer):
20
+ return type(layer) in layer_types
21
+ return find_layer_predicate_recursive(model, predicate)
22
+
23
+
24
+ def find_layer_predicate_recursive(model, predicate):
25
+ result = []
26
+ for name, layer in model._modules.items():
27
+ if predicate(layer):
28
+ result.append(layer)
29
+ result.extend(find_layer_predicate_recursive(layer, predicate))
30
+ return result
models/dsp/CAM/pytorch_grad_cam/utils/image.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import cv2
2
+ import numpy as np
3
+ import torch
4
+ import torch.nn.functional as F
5
+ from torchvision.transforms import Compose, Normalize, ToTensor
6
+
7
+
8
+ def preprocess_image(img: np.ndarray, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]) -> torch.Tensor:
9
+ preprocessing = Compose([
10
+ ToTensor(),
11
+ Normalize(mean=mean, std=std)
12
+ ])
13
+ return preprocessing(img.copy()).unsqueeze(0)
14
+
15
+
16
+ def deprocess_image(img):
17
+ """ see https://github.com/jacobgil/keras-grad-cam/blob/master/grad-cam.py#L65 """
18
+ img = img - np.mean(img)
19
+ img = img / (np.std(img) + 1e-5)
20
+ img = img * 0.1
21
+ img = img + 0.5
22
+ img = np.clip(img, 0, 1)
23
+ return np.uint8(img * 255)
24
+
25
+
26
+ def show_cam_on_image(img: np.ndarray,
27
+ mask: np.ndarray,
28
+ use_rgb: bool = False,
29
+ colormap: int = cv2.COLORMAP_JET) -> np.ndarray:
30
+ """ This function overlays the cam mask on the image as an heatmap.
31
+ By default the heatmap is in BGR format.
32
+
33
+ :param img: The base image in RGB or BGR format.
34
+ :param mask: The cam mask.
35
+ :param use_rgb: Whether to use an RGB or BGR heatmap, this should be set to True if 'img' is in RGB format.
36
+ :param colormap: The OpenCV colormap to be used.
37
+ :returns: The default image with the cam overlay.
38
+ """
39
+ heatmap = cv2.applyColorMap(np.uint8(255 * mask), colormap)
40
+ if use_rgb:
41
+ heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)
42
+ heatmap = np.float32(heatmap) / 255
43
+
44
+ if np.max(img) > 1:
45
+ raise Exception(
46
+ "The input image should np.float32 in the range [0, 1]")
47
+
48
+ cam = heatmap + img
49
+ cam = cam / np.max(cam)
50
+ return np.uint8(255 * cam)
51
+
52
+ def scale_cam_image(cam, target_size=None):
53
+ result = []
54
+ for img in cam:
55
+ img = img - np.min(img)
56
+ img = img / (1e-7 + np.max(img))
57
+ if target_size is not None:
58
+ img = cv2.resize(img, target_size)
59
+ result.append(img)
60
+ result = np.float32(result)
61
+
62
+ return result
63
+
64
+ def scale_cam_image_torch(cam: torch.Tensor, target_size=None):
65
+ if cam.ndim == 3:
66
+ cam = cam.unsqueeze(1)
67
+
68
+ # normalize per image
69
+ cam_min = cam.amin(dim=(2, 3), keepdim=True)
70
+ cam_max = cam.amax(dim=(2, 3), keepdim=True)
71
+ cam = cam - cam_min
72
+ cam = cam / (cam_max + 1e-7)
73
+
74
+ # resize if needed
75
+ if target_size is not None:
76
+ cam = F.interpolate(cam, size=target_size, mode="bilinear", align_corners=False)
77
+
78
+ return cam.squeeze(1)
79
+
80
+ def scale_accross_batch_and_channels(tensor, target_size):
81
+ batch_size, channel_size = tensor.shape[:2]
82
+ reshaped_tensor = tensor.reshape(
83
+ batch_size * channel_size, *tensor.shape[2:])
84
+ result = scale_cam_image(reshaped_tensor, target_size)
85
+ result = result.reshape(
86
+ batch_size,
87
+ channel_size,
88
+ target_size[1],
89
+ target_size[0])
90
+ return result