Upload folder using huggingface_hub
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +4 -0
- LICENSE +21 -0
- README.md +252 -0
- claim1.py +59 -0
- claim2.py +73 -0
- claim5.py +107 -0
- claim5_official.py +79 -0
- claim6.py +34 -0
- configs/dsp-dior.yaml +99 -0
- configs/dsp-exdark.yaml +91 -0
- configs/dsp-ruod.yaml +91 -0
- crossover.py +42 -0
- databuilders/__init__.py +1 -0
- databuilders/dior.py +113 -0
- datamodules/__init__.py +2 -0
- datamodules/loader.py +109 -0
- datamodules/ref_table.py +76 -0
- datamodules/transforms/__init__.py +1 -0
- datamodules/transforms/layout_transform.py +132 -0
- figs/Main.png +3 -0
- fonts/GILI____.TTF +0 -0
- fonts/Rainbow-Party-2.ttf +3 -0
- gicdm_core.py +216 -0
- infer.sh +174 -0
- main.py +12 -0
- models/__init__.py +1 -0
- models/dsp/CAM/__init__.py +1 -0
- models/dsp/CAM/cam_generator.py +133 -0
- models/dsp/CAM/clip/__init__.py +1 -0
- models/dsp/CAM/clip/bpe_simple_vocab_16e6.txt.gz +3 -0
- models/dsp/CAM/clip/clip.py +245 -0
- models/dsp/CAM/clip/model.py +517 -0
- models/dsp/CAM/clip/simple_tokenizer.py +132 -0
- models/dsp/CAM/pytorch_grad_cam/__init__.py +14 -0
- models/dsp/CAM/pytorch_grad_cam/ablation_cam.py +134 -0
- models/dsp/CAM/pytorch_grad_cam/ablation_cam_multilayer.py +136 -0
- models/dsp/CAM/pytorch_grad_cam/ablation_layer.py +131 -0
- models/dsp/CAM/pytorch_grad_cam/activations_and_gradients.py +55 -0
- models/dsp/CAM/pytorch_grad_cam/base_cam.py +227 -0
- models/dsp/CAM/pytorch_grad_cam/eigen_cam.py +23 -0
- models/dsp/CAM/pytorch_grad_cam/eigen_grad_cam.py +21 -0
- models/dsp/CAM/pytorch_grad_cam/fullgrad_cam.py +95 -0
- models/dsp/CAM/pytorch_grad_cam/grad_cam.py +24 -0
- models/dsp/CAM/pytorch_grad_cam/grad_cam_plusplus.py +32 -0
- models/dsp/CAM/pytorch_grad_cam/guided_backprop.py +100 -0
- models/dsp/CAM/pytorch_grad_cam/layer_cam.py +36 -0
- models/dsp/CAM/pytorch_grad_cam/score_cam.py +63 -0
- models/dsp/CAM/pytorch_grad_cam/utils/__init__.py +4 -0
- models/dsp/CAM/pytorch_grad_cam/utils/find_layers.py +30 -0
- 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 |
+

|
| 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
|
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
|