.
This commit is contained in:
+146
@@ -0,0 +1,146 @@
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
**/*.pyc
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
|
||||
# PyBuilder
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# pyenv
|
||||
.python-version
|
||||
|
||||
# celery beat schedule file
|
||||
celerybeat-schedule
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
|
||||
# custom
|
||||
data
|
||||
# data for pytest moved to http server
|
||||
# !tests/data
|
||||
.vscode
|
||||
.idea
|
||||
*.pkl
|
||||
*.pkl.json
|
||||
*.log.json
|
||||
work_dirs/
|
||||
logs/
|
||||
|
||||
# Pytorch
|
||||
*.pth
|
||||
*.pt
|
||||
|
||||
|
||||
# Visualization
|
||||
*.mp4
|
||||
*.png
|
||||
*.gif
|
||||
*.jpg
|
||||
*.obj
|
||||
*.ply
|
||||
!demo/resources/*
|
||||
|
||||
# Resources as exception
|
||||
!resources/*
|
||||
|
||||
# Loaded/Saved data files
|
||||
*.npz
|
||||
*.npy
|
||||
*.pickle
|
||||
|
||||
# MacOS
|
||||
*DS_Store*
|
||||
# git
|
||||
*.orig
|
||||
|
||||
env.sh
|
||||
@@ -0,0 +1,10 @@
|
||||
S-Lab License 1.0
|
||||
|
||||
Copyright 2023 S-Lab
|
||||
|
||||
Redistribution and use for non-commercial purpose in source and binary forms, with or without modification, are permitted provided that the following conditions are met:
|
||||
1. Redistributions of source code must retain the above copyright notice, this list of conditions and the following disclaimer.
|
||||
2. Redistributions in binary form must reproduce the above copyright notice, this list of conditions and the following disclaimer in the documentation and/or other materials provided with the distribution.
|
||||
3. Neither the name of the copyright holder nor the names of its contributors may be used to endorse or promote products derived from this software without specific prior written permission.
|
||||
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
4. In the event that redistribution and/or use for commercial purpose in source or binary forms, with or without modification is required, please contact the contributor(s) of the work.
|
||||
@@ -0,0 +1,202 @@
|
||||
<div align="center">
|
||||
|
||||
<h1>THIS IS CURRENTLY FOR PERSONAL USE ONLY AS I'M TOO LAZY TO USE GITHUB TOKEN ON OR UPLOAD FOLDER TO GOOGLE COLAB</h1>
|
||||
|
||||
<h1>ReMoDiffuse: Retrieval-Augmented Motion Diffusion Model</h1>
|
||||
|
||||
<div>
|
||||
<a href='https://mingyuan-zhang.github.io/' target='_blank'>Mingyuan Zhang</a><sup>1</sup> 
|
||||
<a href='https://gxyes.github.io/' target='_blank'>Xinying Guo</a><sup>1</sup> 
|
||||
<a href='https://scholar.google.com/citations?user=lSDISOcAAAAJ&hl=zh-CN' target='_blank'>Liang Pan</a><sup>1</sup> 
|
||||
<a href='https://caizhongang.github.io/' target='_blank'>Zhongang Cai</a><sup>1,2</sup> 
|
||||
<a href='https://hongfz16.github.io/' target='_blank'>Fangzhou Hong</a><sup>1</sup> 
|
||||
<a href='https://www.linkedin.com/in/huirong-li' target='_blank'>Huirong Li</a><sup>1</sup>  <br>
|
||||
<a href='https://yanglei.me/' target='_blank'>Lei Yang</a><sup>2</sup> 
|
||||
<a href='https://liuziwei7.github.io/' target='_blank'>Ziwei Liu</a><sup>1+</sup>
|
||||
</div>
|
||||
<div>
|
||||
<sup>1</sup>S-Lab, Nanyang Technological University 
|
||||
<sup>2</sup>SenseTime Research 
|
||||
</div>
|
||||
<div>
|
||||
<sup>+</sup>corresponding author
|
||||
</div>
|
||||
|
||||
|
||||
---
|
||||
|
||||
<h4 align="center">
|
||||
<a href="https://mingyuan-zhang.github.io/projects/ReMoDiffuse.html" target='_blank'>[Project Page]</a> •
|
||||
<a href="https://arxiv.org/abs/2304.01116" target='_blank'>[arXiv]</a> •
|
||||
<a href="https://youtu.be/wSddrIA_2p8" target='_blank'>[Video]</a> •
|
||||
<a href="https://colab.research.google.com/drive/1jztE7c8js3P4YFbw5cGNPJAsCVrreTov?usp=sharing" target='_blank'>[Colab Demo]</a> •
|
||||
<a href="https://huggingface.co/spaces/mingyuan/ReMoDiffuse" target='_blank'>[Hugging Face Demo]</a>
|
||||
<br> <br>
|
||||
Accepted to <a href="https://iccv2023.thecvf.com/" target="_blank"><strong>ICCV 2023</strong></a></h2>
|
||||
<img src="https://visitor-badge.laobi.icu/badge?page_id=mingyuan-zhang/ReMoDiffuse" width="8%" alt="visitor badge"/>
|
||||
</h4>
|
||||
|
||||
</div>
|
||||
|
||||
|
||||
>**Abstract:** 3D human motion generation is crucial for creative industry. Recent advances rely on generative models with domain knowledge for text-driven motion generation, leading to substantial progress in capturing common motions. However, the performance on more diverse motions remains unsatisfactory. In this work, we propose **ReMoDiffuse**, a diffusion-model-based motion generation framework that integrates a retrieval mechanism to refine the denoising process.
|
||||
|
||||
<div align="center">
|
||||
<tr>
|
||||
<img src="imgs/teaser.png" width="90%"/>
|
||||
<img src="imgs/pipeline.png" width="90%"/>
|
||||
</tr>
|
||||
</div>
|
||||
|
||||
>**Pipeline Overview:** ReMoDiffuse is a retrieval-augmented 3D human motion diffusion model. Benefiting from the extra knowledge from the retrieved samples, ReMoDiffuse is able to achieve high-fidelity on the given prompts. It contains three core components: a) **Hybrid Retrieval** database stores multi-modality features of each motion sequence. b) Semantics-modulated transformer incorporates several identical decoder layers, including a **Semantics-Modulated Attention (SMA)** layer and an FFN layer. The SMA layer will adaptively absorb knowledge from both retrived samples and the given prompts. c) **Condition Mxture** technique is proposed to better mix model's outputs under different combinations of conditions.
|
||||
|
||||
## Updates
|
||||
|
||||
[09/2023] Add a [🤗Hugging Face Demo](https://huggingface.co/spaces/mingyuan/ReMoDiffuse)!
|
||||
|
||||
[09/2023] Add a [Colab Demo](https://colab.research.google.com/drive/1jztE7c8js3P4YFbw5cGNPJAsCVrreTov?usp=sharing)! [](https://colab.research.google.com/drive/1jztE7c8js3P4YFbw5cGNPJAsCVrreTov?usp=sharing)
|
||||
|
||||
[09/2023] Release code for [ReMoDiffuse](https://mingyuan-zhang.github.io/projects/ReMoDiffuse.html) and [MotionDiffuse](https://mingyuan-zhang.github.io/projects/MotionDiffuse.html)
|
||||
|
||||
## Benchmark and Model Zoo
|
||||
|
||||
#### Supported methods
|
||||
|
||||
- [x] [MotionDiffuse](https://mingyuan-zhang.github.io/projects/ReMoDiffuse.html)
|
||||
- [x] [MDM](https://guytevet.github.io/mdm-page/)
|
||||
- [x] [ReMoDiffuse](https://mingyuan-zhang.github.io/projects/MotionDiffuse.html)
|
||||
|
||||
|
||||
## Citation
|
||||
|
||||
If you find our work useful for your research, please consider citing the paper:
|
||||
|
||||
```
|
||||
@article{zhang2023remodiffuse,
|
||||
title={ReMoDiffuse: Retrieval-Augmented Motion Diffusion Model},
|
||||
author={Zhang, Mingyuan and Guo, Xinying and Pan, Liang and Cai, Zhongang and Hong, Fangzhou and Li, Huirong and Yang, Lei and Liu, Ziwei},
|
||||
journal={arXiv preprint arXiv:2304.01116},
|
||||
year={2023}
|
||||
}
|
||||
@article{zhang2022motiondiffuse,
|
||||
title={MotionDiffuse: Text-Driven Human Motion Generation with Diffusion Model},
|
||||
author={Zhang, Mingyuan and Cai, Zhongang and Pan, Liang and Hong, Fangzhou and Guo, Xinying and Yang, Lei and Liu, Ziwei},
|
||||
journal={arXiv preprint arXiv:2208.15001},
|
||||
year={2022}
|
||||
}
|
||||
```
|
||||
|
||||
## Installation
|
||||
|
||||
```shell
|
||||
# Create Conda Environment
|
||||
conda create -n mogen python=3.9 -y
|
||||
conda activate mogen
|
||||
|
||||
# C++ Environment
|
||||
export PATH=/mnt/lustre/share/gcc/gcc-8.5.0/bin:$PATH
|
||||
export LD_LIBRARY_PATH=/mnt/lustre/share/gcc/gcc-8.5.0/lib:/mnt/lustre/share/gcc/gcc-8.5.0/lib64:/mnt/lustre/share/gcc/gmp-4.3.2/lib:/mnt/lustre/share/gcc/mpc-0.8.1/lib:/mnt/lustre/share/gcc/mpfr-2.4.2/lib:$LD_LIBRARY_PATH
|
||||
|
||||
# Install Pytorch
|
||||
conda install pytorch==1.12.1 torchvision==0.13.1 torchaudio==0.12.1 cudatoolkit=11.3 -c pytorch -y
|
||||
|
||||
# Install Pytorch3d
|
||||
conda install -c bottler nvidiacub -y
|
||||
conda install -c fvcore -c iopath -c conda-forge fvcore iopath -y
|
||||
conda install pytorch3d -c pytorch3d -y
|
||||
|
||||
# Install other requirements
|
||||
pip install -r requirements.txt
|
||||
```
|
||||
|
||||
## Data Preparation
|
||||
|
||||
Download data files from google drive [link](https://drive.google.com/drive/folders/13kwahiktQ2GMVKfVH3WT-VGAQ6JHbvUv?usp=sharing) or Baidu Netdisk [link](https://pan.baidu.com/s/1604jks-9PtBUtqCpQQmeEg)(access code: vprc). Unzipped all files and arrange them in the following file structure:
|
||||
|
||||
```text
|
||||
ReMoDiffuse
|
||||
├── mogen
|
||||
├── tools
|
||||
├── configs
|
||||
├── logs
|
||||
│ ├── motiondiffuse
|
||||
│ ├── remodiffuse
|
||||
│ └── mdm
|
||||
└── data
|
||||
├── database
|
||||
├── datasets
|
||||
├── evaluators
|
||||
└── glove
|
||||
```
|
||||
|
||||
## Training
|
||||
|
||||
### Training with a single / multiple GPUs
|
||||
|
||||
```shell
|
||||
PYTHONPATH=".":$PYTHONPATH python tools/train.py ${CONFIG_FILE} ${WORK_DIR} --no-validate
|
||||
```
|
||||
|
||||
**Note:** The provided config files are designed for training with 8 gpus. If you want to train on a single gpu, you can reduce the number of epochs to one-fourth of the original.
|
||||
|
||||
### Training with Slurm
|
||||
|
||||
```shell
|
||||
./tools/slurm_train.sh ${PARTITION} ${JOB_NAME} ${CONFIG_FILE} ${WORK_DIR} ${GPU_NUM} --no-validate
|
||||
```
|
||||
|
||||
Common optional arguments include:
|
||||
- `--resume-from ${CHECKPOINT_FILE}`: Resume from a previous checkpoint file.
|
||||
- `--no-validate`: Whether not to evaluate the checkpoint during training.
|
||||
|
||||
Example: using 8 GPUs to train ReMoDiffuse on a slurm cluster.
|
||||
```shell
|
||||
./tools/slurm_train.sh my_partition my_job configs/remodiffuse/remodiffuse_kit.py logs/remodiffuse_kit 8 --no-validate
|
||||
```
|
||||
|
||||
## Evaluation
|
||||
|
||||
### Evaluate with a single GPU / multiple GPUs
|
||||
|
||||
```shell
|
||||
PYTHONPATH=".":$PYTHONPATH python tools/test.py ${CONFIG} --work-dir=${WORK_DIR} ${CHECKPOINT}
|
||||
```
|
||||
|
||||
### Evaluate with slurm
|
||||
|
||||
```shell
|
||||
./tools/slurm_test.sh ${PARTITION} ${JOB_NAME} ${CONFIG} ${WORK_DIR} ${CHECKPOINT}
|
||||
```
|
||||
Example:
|
||||
```shell
|
||||
./tools/slurm_test.sh my_partition test_remodiffuse configs/remodiffuse/remodiffuse_kit.py logs/remodiffuse_kit logs/remodiffuse_kit/latest.pth
|
||||
```
|
||||
|
||||
**Note:** Run full evaluation for HumanML3D dataset is very slow. You can change `replication_times` in [human_ml3d_bs128.py](configs/_base_/datasets/human_ml3d_bs128.py) to $1$ for a quick evaluation.
|
||||
|
||||
## Visualization
|
||||
|
||||
```shell
|
||||
PYTHONPATH=".":$PYTHONPATH python tools/visualize.py ${CONFIG} ${CHECKPOINT} \
|
||||
--text ${TEXT} \
|
||||
--motion_length ${MOTION_LENGTH} \
|
||||
--out ${OUTPUT_ANIMATION_PATH} \
|
||||
--device cpu
|
||||
```
|
||||
|
||||
Example:
|
||||
```shell
|
||||
PYTHONPATH=".":$PYTHONPATH python tools/visualize.py \
|
||||
configs/remodiffuse/remodiffuse_t2m.py \
|
||||
logs/remodiffuse/remodiffuse_t2m/latest.pth \
|
||||
--text "a person is running quickly" \
|
||||
--motion_length 120 \
|
||||
--out "test.gif" \
|
||||
--device cpu
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
This study is supported by the Ministry of Education, Singapore, under its MOE AcRF Tier 2 (MOE-T2EP20221-0012), NTU NAP, and under the RIE2020 Industry Alignment Fund – Industry Collaboration Projects (IAF-ICP) Funding Initiative, as well as cash and in-kind contribution from the industry partner(s).
|
||||
|
||||
The visualization tool is developed on top of [Generating Diverse and Natural 3D Human Motions from Text](https://github.com/EricGuo5513/text-to-motion)
|
||||
+176
@@ -0,0 +1,176 @@
|
||||
import matplotlib
|
||||
from .mogen import digit_version
|
||||
assert digit_version(matplotlib.__version__) == digit_version("3.3.1"), "This extension requires matplotlib==3.3.1, otherwise the visualization won't work."
|
||||
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
EXTENSION_PATH = Path(__file__).parent
|
||||
sys.path.insert(0, str(EXTENSION_PATH.resolve()))
|
||||
|
||||
import custom_mmpkg.custom_mmcv as mmcv
|
||||
import numpy as np
|
||||
import torch
|
||||
from mogen.models import build_architecture
|
||||
from custom_mmpkg.custom_mmcv.runner import load_checkpoint
|
||||
from custom_mmpkg.custom_mmcv.parallel import MMDataParallel
|
||||
from mogen.utils.plot_utils import (
|
||||
recover_from_ric,
|
||||
plot_3d_motion,
|
||||
t2m_kinematic_chain
|
||||
)
|
||||
from scipy.ndimage import gaussian_filter
|
||||
from IPython.display import Image
|
||||
import comfy.model_management as model_management
|
||||
|
||||
import yaml
|
||||
CONFIGS = yaml.load((EXTENSION_PATH / "config.yml").resolve())
|
||||
DATASET_PATH = Path(CONFIGS.dataset)
|
||||
mean_path = DATASET_PATH / "mean.npy"
|
||||
std_path = DATASET_PATH / "std.npy"
|
||||
assert os.path.exists(mean_path) and os.path.exists(std_path), "Dataset folder not found or lost mean.npy, std.npy"
|
||||
mean = np.load(mean_path)
|
||||
std = np.load(std_path)
|
||||
|
||||
def create_mdm_model(config_path, ckpt_path):
|
||||
cfg = mmcv.Config.fromfile(config_path)
|
||||
mdm = build_architecture(cfg.model)
|
||||
load_checkpoint(mdm, ckpt_path, map_location='cpu')
|
||||
motion_module = motion_module.to(model_management.unet_offload_device())
|
||||
mdm.eval()
|
||||
return mdm
|
||||
|
||||
def is_model_available(mdm_config):
|
||||
return os.path.exists(mdm_config["config"]) and os.path.exists(mdm_config["ckpt"])
|
||||
|
||||
class MotionDiffModel(torch.nn.Module): #Anything beside CLIP (mdm.model)
|
||||
def __init__(self, **kwargs) -> None:
|
||||
super(MotionDiffModel).__init__()
|
||||
self.loss_recon = kwargs["loss_recon"]
|
||||
self.diffusion_train = kwargs["diffusion_train"]
|
||||
self.diffusion_test = kwargs["sampler"]
|
||||
|
||||
def forward(self, clip, cond_dict, **kwargs):
|
||||
motion, motion_mask = kwargs['motion'].float(), kwargs['motion_mask'].float()
|
||||
sample_idx = kwargs.get('sample_idx', None)
|
||||
clip_feat = kwargs.get('clip_feat', None)
|
||||
sampler = kwargs.get('sampler', 'ddpm')
|
||||
B, T = motion.shape[:2]
|
||||
|
||||
dim_pose = kwargs['motion'].shape[-1]
|
||||
model_kwargs = cond_dict
|
||||
model_kwargs['motion_mask'] = motion_mask
|
||||
model_kwargs['sample_idx'] = sample_idx
|
||||
inference_kwargs = kwargs.get('inference_kwargs', {})
|
||||
if sampler == 'ddpm':
|
||||
output = self.diffusion_test.p_sample_loop(
|
||||
clip,
|
||||
(B, T, dim_pose),
|
||||
clip_denoised=False,
|
||||
progress=False,
|
||||
model_kwargs=model_kwargs,
|
||||
**inference_kwargs
|
||||
)
|
||||
else:
|
||||
output = self.diffusion_test.ddim_sample_loop(
|
||||
clip,
|
||||
(B, T, dim_pose),
|
||||
clip_denoised=False,
|
||||
progress=False,
|
||||
model_kwargs=model_kwargs,
|
||||
eta=0,
|
||||
**inference_kwargs
|
||||
)
|
||||
if getattr(clip, "post_process") is not None:
|
||||
output = clip.post_process(output)
|
||||
results = kwargs
|
||||
results['pred_motion'] = output
|
||||
results = self.split_results(results)
|
||||
return results
|
||||
|
||||
class MotionDiffCLIP(torch.nn.Module):
|
||||
def __init__(self, model):
|
||||
super(MotionDiffCLIP).__init__()
|
||||
self.model = model
|
||||
|
||||
def forward(self, text):
|
||||
return self.model.get_precompute_condition(device=model_management.get_torch_device(), text=text)
|
||||
|
||||
class MotionDiffLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"mdm_name": (
|
||||
list(
|
||||
filter(CONFIGS.keys(), lambda key: (key != "dataset") and is_model_available(CONFIGS[key]))
|
||||
),
|
||||
{ "default": "remodiffuse" }
|
||||
)
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MD_MODEL", "MD_CLIP")
|
||||
CATEGORY = "Human Motion Diff"
|
||||
FUNCTION = "load_mdm"
|
||||
|
||||
def load_mdm(self, mdm_name):
|
||||
mdm = create_mdm_model(mdm_name)
|
||||
model = MotionDiffModel(loss_recon=mdm.loss_recon, diffusion_train=mdm.diffusion_train, diffusion_test=mdm.diffusion_test, sampler=mdm.sampler)
|
||||
clip = mdm.model
|
||||
return (model, clip)
|
||||
|
||||
class MotionDiffTextEncode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"clip": ("MD_CLIP", ),
|
||||
"text": ("STRING", {"multiline": True})
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MD_CONDITIONING",)
|
||||
CATEGORY = "Human Motion Diff"
|
||||
FUNCTION = "encode_text"
|
||||
|
||||
def encode_text(self, clip, text):
|
||||
return (clip(text), )
|
||||
|
||||
class MotionDiffSimpleSampler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"sampler_name": (["ddpm", "ddim"], ),
|
||||
"model": ("MD_MODEL", ),
|
||||
"clip": ("MD_CLIP", ),
|
||||
"cond": ("MD_CONDITIONING", )
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MOTION_DATA",)
|
||||
CATEGORY = "Human Motion Diff"
|
||||
FUNCTION = "sample"
|
||||
|
||||
def sample(self, sampler_name, model, clip, cond):
|
||||
device = model_management.get_torch_device()
|
||||
motion = torch.zeros(1, motion_length, 263).to(device)
|
||||
motion_mask = torch.ones(1, motion_length).to(device)
|
||||
motion_length = torch.Tensor([motion_length]).long().to(device)
|
||||
model = model.to(device)
|
||||
kwargs = {
|
||||
'motion': motion,
|
||||
'motion_mask': motion_mask,
|
||||
'motion_length': motion_length,
|
||||
'inference_kwargs': {},
|
||||
'sampler': sampler_name,
|
||||
|
||||
}
|
||||
|
||||
with torch.no_grad():
|
||||
output = model(clip, cond_dict=cond, **kwargs)[0]['pred_motion']
|
||||
pred_motion = output.cpu().detach().numpy()
|
||||
pred_motion = pred_motion * std + mean
|
||||
|
||||
return pred_motion
|
||||
+15
@@ -0,0 +1,15 @@
|
||||
#Please don't change the field name/key, only change the value
|
||||
|
||||
remodiffuse:
|
||||
config: "configs/remodiffuse/remodiffuse_t2m.py"
|
||||
ckpt: "logs/remodiffuse/remodiffuse_t2m/latest.pth"
|
||||
|
||||
motiondiffuse:
|
||||
config: "configs/motiondiffuse/motiondiffuse_t2m.py"
|
||||
ckpt: "logs/motiondiffuse/motiondiffuse_t2m/latest.pth"
|
||||
|
||||
mdm:
|
||||
config: "configs/mdm/mdm_t2m_official.py"
|
||||
ckpt: "logs/mdm/mdm_t2m/latest.pth"
|
||||
|
||||
dataset: "data/datasets/human_ml3d/"
|
||||
@@ -0,0 +1,60 @@
|
||||
# dataset settings
|
||||
data_keys = ['motion', 'motion_mask', 'motion_length', 'clip_feat']
|
||||
meta_keys = ['text', 'token']
|
||||
train_pipeline = [
|
||||
dict(
|
||||
type='Normalize',
|
||||
mean_path='data/datasets/human_ml3d/mean.npy',
|
||||
std_path='data/datasets/human_ml3d/std.npy'),
|
||||
dict(type='Crop', crop_size=196),
|
||||
dict(type='ToTensor', keys=data_keys),
|
||||
dict(type='Collect', keys=data_keys, meta_keys=meta_keys)
|
||||
]
|
||||
|
||||
data = dict(
|
||||
samples_per_gpu=128,
|
||||
workers_per_gpu=1,
|
||||
train=dict(
|
||||
type='RepeatDataset',
|
||||
dataset=dict(
|
||||
type='TextMotionDataset',
|
||||
dataset_name='human_ml3d',
|
||||
data_prefix='data',
|
||||
pipeline=train_pipeline,
|
||||
ann_file='train.txt',
|
||||
motion_dir='motions',
|
||||
text_dir='texts',
|
||||
token_dir='tokens',
|
||||
clip_feat_dir='clip_feats',
|
||||
),
|
||||
times=200
|
||||
),
|
||||
test=dict(
|
||||
type='TextMotionDataset',
|
||||
dataset_name='human_ml3d',
|
||||
data_prefix='data',
|
||||
pipeline=train_pipeline,
|
||||
ann_file='test.txt',
|
||||
motion_dir='motions',
|
||||
text_dir='texts',
|
||||
token_dir='tokens',
|
||||
clip_feat_dir='clip_feats',
|
||||
eval_cfg=dict(
|
||||
shuffle_indexes=True,
|
||||
replication_times=20,
|
||||
replication_reduction='statistics',
|
||||
text_encoder_name='human_ml3d',
|
||||
text_encoder_path='data/evaluators/human_ml3d/finest.tar',
|
||||
motion_encoder_name='human_ml3d',
|
||||
motion_encoder_path='data/evaluators/human_ml3d/finest.tar',
|
||||
metrics=[
|
||||
dict(type='R Precision', batch_size=32, top_k=3),
|
||||
dict(type='Matching Score', batch_size=32),
|
||||
dict(type='FID'),
|
||||
dict(type='Diversity', num_samples=300),
|
||||
dict(type='MultiModality', num_samples=100, num_repeats=30, num_picks=10)
|
||||
]
|
||||
),
|
||||
test_mode=True
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,60 @@
|
||||
# dataset settings
|
||||
data_keys = ['motion', 'motion_mask', 'motion_length', 'clip_feat']
|
||||
meta_keys = ['text', 'token']
|
||||
train_pipeline = [
|
||||
dict(type='Crop', crop_size=196),
|
||||
dict(
|
||||
type='Normalize',
|
||||
mean_path='data/datasets/kit_ml/mean.npy',
|
||||
std_path='data/datasets/kit_ml/std.npy'),
|
||||
dict(type='ToTensor', keys=data_keys),
|
||||
dict(type='Collect', keys=data_keys, meta_keys=meta_keys)
|
||||
]
|
||||
|
||||
data = dict(
|
||||
samples_per_gpu=128,
|
||||
workers_per_gpu=1,
|
||||
train=dict(
|
||||
type='RepeatDataset',
|
||||
dataset=dict(
|
||||
type='TextMotionDataset',
|
||||
dataset_name='kit_ml',
|
||||
data_prefix='data',
|
||||
pipeline=train_pipeline,
|
||||
ann_file='train.txt',
|
||||
motion_dir='motions',
|
||||
text_dir='texts',
|
||||
token_dir='tokens',
|
||||
clip_feat_dir='clip_feats',
|
||||
),
|
||||
times=100
|
||||
),
|
||||
test=dict(
|
||||
type='TextMotionDataset',
|
||||
dataset_name='kit_ml',
|
||||
data_prefix='data',
|
||||
pipeline=train_pipeline,
|
||||
ann_file='test.txt',
|
||||
motion_dir='motions',
|
||||
text_dir='texts',
|
||||
token_dir='tokens',
|
||||
clip_feat_dir='clip_feats',
|
||||
eval_cfg=dict(
|
||||
shuffle_indexes=True,
|
||||
replication_times=20,
|
||||
replication_reduction='statistics',
|
||||
text_encoder_name='kit_ml',
|
||||
text_encoder_path='data/evaluators/kit_ml/finest.tar',
|
||||
motion_encoder_name='kit_ml',
|
||||
motion_encoder_path='data/evaluators/kit_ml/finest.tar',
|
||||
metrics=[
|
||||
dict(type='R Precision', batch_size=32, top_k=3),
|
||||
dict(type='Matching Score', batch_size=32),
|
||||
dict(type='FID'),
|
||||
dict(type='Diversity', num_samples=300),
|
||||
dict(type='MultiModality', num_samples=50, num_repeats=30, num_picks=10)
|
||||
]
|
||||
),
|
||||
test_mode=True
|
||||
)
|
||||
)
|
||||
@@ -0,0 +1,67 @@
|
||||
_base_ = ['../_base_/datasets/human_ml3d_bs128.py']
|
||||
|
||||
# checkpoint saving
|
||||
checkpoint_config = dict(interval=1)
|
||||
|
||||
dist_params = dict(backend='nccl')
|
||||
log_level = 'INFO'
|
||||
load_from = None
|
||||
resume_from = None
|
||||
workflow = [('train', 1)]
|
||||
|
||||
# optimizer
|
||||
optimizer = dict(type='Adam', lr=1e-4)
|
||||
optimizer_config = dict(grad_clip=None)
|
||||
# learning policy
|
||||
lr_config = dict(policy='step', step=[])
|
||||
runner = dict(type='EpochBasedRunner', max_epochs=50)
|
||||
|
||||
log_config = dict(
|
||||
interval=50,
|
||||
hooks=[
|
||||
dict(type='TextLoggerHook'),
|
||||
# dict(type='TensorboardLoggerHook')
|
||||
])
|
||||
|
||||
input_feats = 263
|
||||
max_seq_len = 196
|
||||
latent_dim = 512
|
||||
time_embed_dim = 2048
|
||||
text_latent_dim = 256
|
||||
ff_size = 1024
|
||||
num_layers = 8
|
||||
num_heads = 4
|
||||
dropout = 0.1
|
||||
cond_mask_prob = 0.1
|
||||
# model settings
|
||||
model = dict(
|
||||
type='MotionDiffusion',
|
||||
model=dict(
|
||||
type='MDMTransformer',
|
||||
input_feats=input_feats,
|
||||
latent_dim=latent_dim,
|
||||
ff_size=ff_size,
|
||||
num_layers=num_layers,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim,
|
||||
cond_mask_prob=cond_mask_prob,
|
||||
guide_scale=2.5,
|
||||
clip_version='ViT-B/32',
|
||||
use_official_ckpt=True
|
||||
),
|
||||
loss_recon=dict(type='MSELoss', loss_weight=1, reduction='none'),
|
||||
diffusion_train=dict(
|
||||
beta_scheduler='cosine',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='start_x',
|
||||
model_var_type='fixed_small',
|
||||
),
|
||||
diffusion_test=dict(
|
||||
beta_scheduler='cosine',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='start_x',
|
||||
model_var_type='fixed_small',
|
||||
),
|
||||
inference_type='ddpm'
|
||||
)
|
||||
@@ -0,0 +1,89 @@
|
||||
_base_ = ['../_base_/datasets/kit_ml_bs128.py']
|
||||
|
||||
# checkpoint saving
|
||||
checkpoint_config = dict(interval=1)
|
||||
|
||||
dist_params = dict(backend='nccl')
|
||||
log_level = 'INFO'
|
||||
load_from = None
|
||||
resume_from = None
|
||||
workflow = [('train', 1)]
|
||||
|
||||
# optimizer
|
||||
optimizer = dict(type='Adam', lr=2e-4)
|
||||
optimizer_config = dict(grad_clip=None)
|
||||
# learning policy
|
||||
lr_config = dict(policy='step', step=[])
|
||||
runner = dict(type='EpochBasedRunner', max_epochs=50)
|
||||
|
||||
log_config = dict(
|
||||
interval=50,
|
||||
hooks=[
|
||||
dict(type='TextLoggerHook'),
|
||||
# dict(type='TensorboardLoggerHook')
|
||||
])
|
||||
|
||||
input_feats = 251
|
||||
max_seq_len = 196
|
||||
latent_dim = 512
|
||||
time_embed_dim = 2048
|
||||
text_latent_dim = 256
|
||||
ff_size = 1024
|
||||
num_heads = 8
|
||||
dropout = 0
|
||||
# model settings
|
||||
model = dict(
|
||||
type='MotionDiffusion',
|
||||
model=dict(
|
||||
type='MotionDiffuseTransformer',
|
||||
input_feats=input_feats,
|
||||
max_seq_len=max_seq_len,
|
||||
latent_dim=latent_dim,
|
||||
time_embed_dim=time_embed_dim,
|
||||
num_layers=8,
|
||||
sa_block_cfg=dict(
|
||||
type='EfficientSelfAttention',
|
||||
latent_dim=latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
ca_block_cfg=dict(
|
||||
type='EfficientCrossAttention',
|
||||
latent_dim=latent_dim,
|
||||
text_latent_dim=text_latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
ffn_cfg=dict(
|
||||
latent_dim=latent_dim,
|
||||
ffn_dim=ff_size,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
text_encoder=dict(
|
||||
pretrained_model='clip',
|
||||
latent_dim=text_latent_dim,
|
||||
num_layers=4,
|
||||
num_heads=4,
|
||||
ff_size=2048,
|
||||
dropout=dropout,
|
||||
use_text_proj=True
|
||||
)
|
||||
),
|
||||
loss_recon=dict(type='MSELoss', loss_weight=1, reduction='none'),
|
||||
diffusion_train=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='epsilon',
|
||||
model_var_type='fixed_small',
|
||||
),
|
||||
diffusion_test=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='epsilon',
|
||||
model_var_type='fixed_small',
|
||||
),
|
||||
inference_type='ddpm'
|
||||
)
|
||||
@@ -0,0 +1,90 @@
|
||||
_base_ = ['../_base_/datasets/human_ml3d_bs128.py']
|
||||
|
||||
# checkpoint saving
|
||||
checkpoint_config = dict(interval=1)
|
||||
|
||||
dist_params = dict(backend='nccl')
|
||||
log_level = 'INFO'
|
||||
load_from = None
|
||||
resume_from = None
|
||||
workflow = [('train', 1)]
|
||||
|
||||
# optimizer
|
||||
optimizer = dict(type='Adam', lr=2e-4)
|
||||
optimizer_config = dict(grad_clip=None)
|
||||
# learning policy
|
||||
lr_config = dict(policy='step', step=[])
|
||||
runner = dict(type='EpochBasedRunner', max_epochs=50)
|
||||
|
||||
log_config = dict(
|
||||
interval=50,
|
||||
hooks=[
|
||||
dict(type='TextLoggerHook'),
|
||||
# dict(type='TensorboardLoggerHook')
|
||||
])
|
||||
|
||||
input_feats = 263
|
||||
max_seq_len = 196
|
||||
latent_dim = 512
|
||||
time_embed_dim = 2048
|
||||
text_latent_dim = 256
|
||||
ff_size = 1024
|
||||
num_heads = 8
|
||||
dropout = 0
|
||||
# model settings
|
||||
model = dict(
|
||||
type='MotionDiffusion',
|
||||
model=dict(
|
||||
type='MotionDiffuseTransformer',
|
||||
input_feats=input_feats,
|
||||
max_seq_len=max_seq_len,
|
||||
latent_dim=latent_dim,
|
||||
time_embed_dim=time_embed_dim,
|
||||
num_layers=8,
|
||||
sa_block_cfg=dict(
|
||||
type='EfficientSelfAttention',
|
||||
latent_dim=latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
ca_block_cfg=dict(
|
||||
type='EfficientCrossAttention',
|
||||
latent_dim=latent_dim,
|
||||
text_latent_dim=text_latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
ffn_cfg=dict(
|
||||
latent_dim=latent_dim,
|
||||
ffn_dim=ff_size,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
text_encoder=dict(
|
||||
pretrained_model='clip',
|
||||
latent_dim=text_latent_dim,
|
||||
num_layers=4,
|
||||
num_heads=4,
|
||||
ff_size=2048,
|
||||
dropout=dropout,
|
||||
use_text_proj=True
|
||||
)
|
||||
),
|
||||
loss_recon=dict(type='MSELoss', loss_weight=1, reduction='none'),
|
||||
diffusion_train=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='epsilon',
|
||||
model_var_type='fixed_small',
|
||||
),
|
||||
diffusion_test=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='epsilon',
|
||||
model_var_type='fixed_small',
|
||||
),
|
||||
inference_type='ddpm'
|
||||
)
|
||||
data = dict(samples_per_gpu=128)
|
||||
@@ -0,0 +1,115 @@
|
||||
_base_ = ['../_base_/datasets/kit_ml_bs128.py']
|
||||
|
||||
# checkpoint saving
|
||||
checkpoint_config = dict(interval=1)
|
||||
|
||||
dist_params = dict(backend='nccl')
|
||||
log_level = 'INFO'
|
||||
load_from = None
|
||||
resume_from = None
|
||||
workflow = [('train', 1)]
|
||||
|
||||
# optimizer
|
||||
optimizer = dict(type='Adam', lr=2e-4)
|
||||
optimizer_config = dict(grad_clip=None)
|
||||
# learning policy
|
||||
lr_config = dict(policy='CosineAnnealing', min_lr_ratio=2e-5, by_epoch=False)
|
||||
runner = dict(type='EpochBasedRunner', max_epochs=20)
|
||||
|
||||
log_config = dict(
|
||||
interval=50,
|
||||
hooks=[
|
||||
dict(type='TextLoggerHook'),
|
||||
# dict(type='TensorboardLoggerHook')
|
||||
])
|
||||
|
||||
input_feats = 251
|
||||
max_seq_len = 196
|
||||
latent_dim = 512
|
||||
time_embed_dim = 2048
|
||||
text_latent_dim = 256
|
||||
ff_size = 1024
|
||||
num_heads = 8
|
||||
dropout = 0
|
||||
|
||||
# model settings
|
||||
model = dict(
|
||||
type='MotionDiffusion',
|
||||
model=dict(
|
||||
type='ReMoDiffuseTransformer',
|
||||
input_feats=input_feats,
|
||||
max_seq_len=max_seq_len,
|
||||
latent_dim=latent_dim,
|
||||
time_embed_dim=time_embed_dim,
|
||||
num_layers=4,
|
||||
ca_block_cfg=dict(
|
||||
type='SemanticsModulatedAttention',
|
||||
latent_dim=latent_dim,
|
||||
text_latent_dim=text_latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
ffn_cfg=dict(
|
||||
latent_dim=latent_dim,
|
||||
ffn_dim=ff_size,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
text_encoder=dict(
|
||||
pretrained_model='clip',
|
||||
latent_dim=text_latent_dim,
|
||||
num_layers=2,
|
||||
ff_size=2048,
|
||||
dropout=dropout,
|
||||
use_text_proj=False
|
||||
),
|
||||
retrieval_cfg=dict(
|
||||
num_retrieval=2,
|
||||
stride=4,
|
||||
num_layers=2,
|
||||
num_motion_layers=2,
|
||||
kinematic_coef=0.1,
|
||||
topk=2,
|
||||
retrieval_file='data/database/kit_text_train.npz',
|
||||
latent_dim=latent_dim,
|
||||
output_dim=latent_dim,
|
||||
max_seq_len=max_seq_len,
|
||||
num_heads=num_heads,
|
||||
ff_size=ff_size,
|
||||
dropout=dropout,
|
||||
ffn_cfg=dict(
|
||||
latent_dim=latent_dim,
|
||||
ffn_dim=ff_size,
|
||||
dropout=dropout,
|
||||
),
|
||||
sa_block_cfg=dict(
|
||||
type='EfficientSelfAttention',
|
||||
latent_dim=latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout
|
||||
),
|
||||
),
|
||||
scale_func_cfg=dict(
|
||||
coarse_scale=4.0,
|
||||
both_coef=0.78123,
|
||||
text_coef=0.39284,
|
||||
retr_coef=-0.12475
|
||||
)
|
||||
),
|
||||
loss_recon=dict(type='MSELoss', loss_weight=1, reduction='none'),
|
||||
diffusion_train=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='start_x',
|
||||
model_var_type='fixed_large',
|
||||
),
|
||||
diffusion_test=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='start_x',
|
||||
model_var_type='fixed_large',
|
||||
respace='15,15,8,6,6',
|
||||
),
|
||||
inference_type='ddim'
|
||||
)
|
||||
@@ -0,0 +1,115 @@
|
||||
_base_ = ['../_base_/datasets/human_ml3d_bs128.py']
|
||||
|
||||
# checkpoint saving
|
||||
checkpoint_config = dict(interval=1)
|
||||
|
||||
dist_params = dict(backend='nccl')
|
||||
log_level = 'INFO'
|
||||
load_from = None
|
||||
resume_from = None
|
||||
workflow = [('train', 1)]
|
||||
|
||||
# optimizer
|
||||
optimizer = dict(type='Adam', lr=2e-4)
|
||||
optimizer_config = dict(grad_clip=None)
|
||||
# learning policy
|
||||
lr_config = dict(policy='CosineAnnealing', min_lr_ratio=2e-5, by_epoch=False)
|
||||
runner = dict(type='EpochBasedRunner', max_epochs=40)
|
||||
|
||||
log_config = dict(
|
||||
interval=50,
|
||||
hooks=[
|
||||
dict(type='TextLoggerHook'),
|
||||
# dict(type='TensorboardLoggerHook')
|
||||
])
|
||||
|
||||
input_feats = 263
|
||||
max_seq_len = 196
|
||||
latent_dim = 512
|
||||
time_embed_dim = 2048
|
||||
text_latent_dim = 256
|
||||
ff_size = 1024
|
||||
num_heads = 8
|
||||
dropout = 0
|
||||
|
||||
# model settings
|
||||
model = dict(
|
||||
type='MotionDiffusion',
|
||||
model=dict(
|
||||
type='ReMoDiffuseTransformer',
|
||||
input_feats=input_feats,
|
||||
max_seq_len=max_seq_len,
|
||||
latent_dim=latent_dim,
|
||||
time_embed_dim=time_embed_dim,
|
||||
num_layers=4,
|
||||
ca_block_cfg=dict(
|
||||
type='SemanticsModulatedAttention',
|
||||
latent_dim=latent_dim,
|
||||
text_latent_dim=text_latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
ffn_cfg=dict(
|
||||
latent_dim=latent_dim,
|
||||
ffn_dim=ff_size,
|
||||
dropout=dropout,
|
||||
time_embed_dim=time_embed_dim
|
||||
),
|
||||
text_encoder=dict(
|
||||
pretrained_model='clip',
|
||||
latent_dim=text_latent_dim,
|
||||
num_layers=2,
|
||||
ff_size=2048,
|
||||
dropout=dropout,
|
||||
use_text_proj=False
|
||||
),
|
||||
retrieval_cfg=dict(
|
||||
num_retrieval=2,
|
||||
stride=4,
|
||||
num_layers=2,
|
||||
num_motion_layers=2,
|
||||
kinematic_coef=0.1,
|
||||
topk=2,
|
||||
retrieval_file='data/database/t2m_text_train.npz',
|
||||
latent_dim=latent_dim,
|
||||
output_dim=latent_dim,
|
||||
max_seq_len=max_seq_len,
|
||||
num_heads=num_heads,
|
||||
ff_size=ff_size,
|
||||
dropout=dropout,
|
||||
ffn_cfg=dict(
|
||||
latent_dim=latent_dim,
|
||||
ffn_dim=ff_size,
|
||||
dropout=dropout,
|
||||
),
|
||||
sa_block_cfg=dict(
|
||||
type='EfficientSelfAttention',
|
||||
latent_dim=latent_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout
|
||||
),
|
||||
),
|
||||
scale_func_cfg=dict(
|
||||
coarse_scale=6.5,
|
||||
both_coef=0.52351,
|
||||
text_coef=-0.28419,
|
||||
retr_coef=2.39872
|
||||
)
|
||||
),
|
||||
loss_recon=dict(type='MSELoss', loss_weight=1, reduction='none'),
|
||||
diffusion_train=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='start_x',
|
||||
model_var_type='fixed_large',
|
||||
),
|
||||
diffusion_test=dict(
|
||||
beta_scheduler='linear',
|
||||
diffusion_steps=1000,
|
||||
model_mean_type='start_x',
|
||||
model_var_type='fixed_large',
|
||||
respace='15,15,8,6,6',
|
||||
),
|
||||
inference_type='ddim'
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
#Dummy file ensuring this package will be recognized
|
||||
@@ -0,0 +1,15 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
# flake8: noqa
|
||||
from .arraymisc import *
|
||||
from .fileio import *
|
||||
from .image import *
|
||||
from .utils import *
|
||||
from .version import *
|
||||
from .video import *
|
||||
from .visualization import *
|
||||
|
||||
# The following modules are not imported to this level, so mmcv may be used
|
||||
# without PyTorch.
|
||||
# - runner
|
||||
# - parallel
|
||||
# - op
|
||||
@@ -0,0 +1,4 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .quantization import dequantize, quantize
|
||||
|
||||
__all__ = ['quantize', 'dequantize']
|
||||
@@ -0,0 +1,55 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import numpy as np
|
||||
|
||||
|
||||
def quantize(arr, min_val, max_val, levels, dtype=np.int64):
|
||||
"""Quantize an array of (-inf, inf) to [0, levels-1].
|
||||
|
||||
Args:
|
||||
arr (ndarray): Input array.
|
||||
min_val (scalar): Minimum value to be clipped.
|
||||
max_val (scalar): Maximum value to be clipped.
|
||||
levels (int): Quantization levels.
|
||||
dtype (np.type): The type of the quantized array.
|
||||
|
||||
Returns:
|
||||
tuple: Quantized array.
|
||||
"""
|
||||
if not (isinstance(levels, int) and levels > 1):
|
||||
raise ValueError(
|
||||
f'levels must be a positive integer, but got {levels}')
|
||||
if min_val >= max_val:
|
||||
raise ValueError(
|
||||
f'min_val ({min_val}) must be smaller than max_val ({max_val})')
|
||||
|
||||
arr = np.clip(arr, min_val, max_val) - min_val
|
||||
quantized_arr = np.minimum(
|
||||
np.floor(levels * arr / (max_val - min_val)).astype(dtype), levels - 1)
|
||||
|
||||
return quantized_arr
|
||||
|
||||
|
||||
def dequantize(arr, min_val, max_val, levels, dtype=np.float64):
|
||||
"""Dequantize an array.
|
||||
|
||||
Args:
|
||||
arr (ndarray): Input array.
|
||||
min_val (scalar): Minimum value to be clipped.
|
||||
max_val (scalar): Maximum value to be clipped.
|
||||
levels (int): Quantization levels.
|
||||
dtype (np.type): The type of the dequantized array.
|
||||
|
||||
Returns:
|
||||
tuple: Dequantized array.
|
||||
"""
|
||||
if not (isinstance(levels, int) and levels > 1):
|
||||
raise ValueError(
|
||||
f'levels must be a positive integer, but got {levels}')
|
||||
if min_val >= max_val:
|
||||
raise ValueError(
|
||||
f'min_val ({min_val}) must be smaller than max_val ({max_val})')
|
||||
|
||||
dequantized_arr = (arr + 0.5).astype(dtype) * (max_val -
|
||||
min_val) / levels + min_val
|
||||
|
||||
return dequantized_arr
|
||||
@@ -0,0 +1,41 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .alexnet import AlexNet
|
||||
# yapf: disable
|
||||
from .bricks import (ACTIVATION_LAYERS, CONV_LAYERS, NORM_LAYERS,
|
||||
PADDING_LAYERS, PLUGIN_LAYERS, UPSAMPLE_LAYERS,
|
||||
ContextBlock, Conv2d, Conv3d, ConvAWS2d, ConvModule,
|
||||
ConvTranspose2d, ConvTranspose3d, ConvWS2d,
|
||||
DepthwiseSeparableConvModule, GeneralizedAttention,
|
||||
HSigmoid, HSwish, Linear, MaxPool2d, MaxPool3d,
|
||||
NonLocal1d, NonLocal2d, NonLocal3d, Scale, Swish,
|
||||
build_activation_layer, build_conv_layer,
|
||||
build_norm_layer, build_padding_layer, build_plugin_layer,
|
||||
build_upsample_layer, conv_ws_2d, is_norm)
|
||||
from .builder import MODELS, build_model_from_cfg
|
||||
# yapf: enable
|
||||
from .resnet import ResNet, make_res_layer
|
||||
from .utils import (INITIALIZERS, Caffe2XavierInit, ConstantInit, KaimingInit,
|
||||
NormalInit, PretrainedInit, TruncNormalInit, UniformInit,
|
||||
XavierInit, bias_init_with_prob, caffe2_xavier_init,
|
||||
constant_init, fuse_conv_bn, get_model_complexity_info,
|
||||
initialize, kaiming_init, normal_init, trunc_normal_init,
|
||||
uniform_init, xavier_init)
|
||||
from .vgg import VGG, make_vgg_layer
|
||||
|
||||
__all__ = [
|
||||
'AlexNet', 'VGG', 'make_vgg_layer', 'ResNet', 'make_res_layer',
|
||||
'constant_init', 'xavier_init', 'normal_init', 'trunc_normal_init',
|
||||
'uniform_init', 'kaiming_init', 'caffe2_xavier_init',
|
||||
'bias_init_with_prob', 'ConvModule', 'build_activation_layer',
|
||||
'build_conv_layer', 'build_norm_layer', 'build_padding_layer',
|
||||
'build_upsample_layer', 'build_plugin_layer', 'is_norm', 'NonLocal1d',
|
||||
'NonLocal2d', 'NonLocal3d', 'ContextBlock', 'HSigmoid', 'Swish', 'HSwish',
|
||||
'GeneralizedAttention', 'ACTIVATION_LAYERS', 'CONV_LAYERS', 'NORM_LAYERS',
|
||||
'PADDING_LAYERS', 'UPSAMPLE_LAYERS', 'PLUGIN_LAYERS', 'Scale',
|
||||
'get_model_complexity_info', 'conv_ws_2d', 'ConvAWS2d', 'ConvWS2d',
|
||||
'fuse_conv_bn', 'DepthwiseSeparableConvModule', 'Linear', 'Conv2d',
|
||||
'ConvTranspose2d', 'MaxPool2d', 'ConvTranspose3d', 'MaxPool3d', 'Conv3d',
|
||||
'initialize', 'INITIALIZERS', 'ConstantInit', 'XavierInit', 'NormalInit',
|
||||
'TruncNormalInit', 'UniformInit', 'KaimingInit', 'PretrainedInit',
|
||||
'Caffe2XavierInit', 'MODELS', 'build_model_from_cfg'
|
||||
]
|
||||
@@ -0,0 +1,61 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import logging
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class AlexNet(nn.Module):
|
||||
"""AlexNet backbone.
|
||||
|
||||
Args:
|
||||
num_classes (int): number of classes for classification.
|
||||
"""
|
||||
|
||||
def __init__(self, num_classes=-1):
|
||||
super(AlexNet, self).__init__()
|
||||
self.num_classes = num_classes
|
||||
self.features = nn.Sequential(
|
||||
nn.Conv2d(3, 64, kernel_size=11, stride=4, padding=2),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.MaxPool2d(kernel_size=3, stride=2),
|
||||
nn.Conv2d(64, 192, kernel_size=5, padding=2),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.MaxPool2d(kernel_size=3, stride=2),
|
||||
nn.Conv2d(192, 384, kernel_size=3, padding=1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(384, 256, kernel_size=3, padding=1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Conv2d(256, 256, kernel_size=3, padding=1),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.MaxPool2d(kernel_size=3, stride=2),
|
||||
)
|
||||
if self.num_classes > 0:
|
||||
self.classifier = nn.Sequential(
|
||||
nn.Dropout(),
|
||||
nn.Linear(256 * 6 * 6, 4096),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Dropout(),
|
||||
nn.Linear(4096, 4096),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(4096, num_classes),
|
||||
)
|
||||
|
||||
def init_weights(self, pretrained=None):
|
||||
if isinstance(pretrained, str):
|
||||
logger = logging.getLogger()
|
||||
from ..runner import load_checkpoint
|
||||
load_checkpoint(self, pretrained, strict=False, logger=logger)
|
||||
elif pretrained is None:
|
||||
# use default initializer
|
||||
pass
|
||||
else:
|
||||
raise TypeError('pretrained must be a str or None')
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
x = self.features(x)
|
||||
if self.num_classes > 0:
|
||||
x = x.view(x.size(0), 256 * 6 * 6)
|
||||
x = self.classifier(x)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,35 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .activation import build_activation_layer
|
||||
from .context_block import ContextBlock
|
||||
from .conv import build_conv_layer
|
||||
from .conv2d_adaptive_padding import Conv2dAdaptivePadding
|
||||
from .conv_module import ConvModule
|
||||
from .conv_ws import ConvAWS2d, ConvWS2d, conv_ws_2d
|
||||
from .depthwise_separable_conv_module import DepthwiseSeparableConvModule
|
||||
from .drop import Dropout, DropPath
|
||||
from .generalized_attention import GeneralizedAttention
|
||||
from .hsigmoid import HSigmoid
|
||||
from .hswish import HSwish
|
||||
from .non_local import NonLocal1d, NonLocal2d, NonLocal3d
|
||||
from .norm import build_norm_layer, is_norm
|
||||
from .padding import build_padding_layer
|
||||
from .plugin import build_plugin_layer
|
||||
from .registry import (ACTIVATION_LAYERS, CONV_LAYERS, NORM_LAYERS,
|
||||
PADDING_LAYERS, PLUGIN_LAYERS, UPSAMPLE_LAYERS)
|
||||
from .scale import Scale
|
||||
from .swish import Swish
|
||||
from .upsample import build_upsample_layer
|
||||
from .wrappers import (Conv2d, Conv3d, ConvTranspose2d, ConvTranspose3d,
|
||||
Linear, MaxPool2d, MaxPool3d)
|
||||
|
||||
__all__ = [
|
||||
'ConvModule', 'build_activation_layer', 'build_conv_layer',
|
||||
'build_norm_layer', 'build_padding_layer', 'build_upsample_layer',
|
||||
'build_plugin_layer', 'is_norm', 'HSigmoid', 'HSwish', 'NonLocal1d',
|
||||
'NonLocal2d', 'NonLocal3d', 'ContextBlock', 'GeneralizedAttention',
|
||||
'ACTIVATION_LAYERS', 'CONV_LAYERS', 'NORM_LAYERS', 'PADDING_LAYERS',
|
||||
'UPSAMPLE_LAYERS', 'PLUGIN_LAYERS', 'Scale', 'ConvAWS2d', 'ConvWS2d',
|
||||
'conv_ws_2d', 'DepthwiseSeparableConvModule', 'Swish', 'Linear',
|
||||
'Conv2dAdaptivePadding', 'Conv2d', 'ConvTranspose2d', 'MaxPool2d',
|
||||
'ConvTranspose3d', 'MaxPool3d', 'Conv3d', 'Dropout', 'DropPath'
|
||||
]
|
||||
@@ -0,0 +1,92 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import TORCH_VERSION, build_from_cfg, digit_version
|
||||
from .registry import ACTIVATION_LAYERS
|
||||
|
||||
for module in [
|
||||
nn.ReLU, nn.LeakyReLU, nn.PReLU, nn.RReLU, nn.ReLU6, nn.ELU,
|
||||
nn.Sigmoid, nn.Tanh
|
||||
]:
|
||||
ACTIVATION_LAYERS.register_module(module=module)
|
||||
|
||||
|
||||
@ACTIVATION_LAYERS.register_module(name='Clip')
|
||||
@ACTIVATION_LAYERS.register_module()
|
||||
class Clamp(nn.Module):
|
||||
"""Clamp activation layer.
|
||||
|
||||
This activation function is to clamp the feature map value within
|
||||
:math:`[min, max]`. More details can be found in ``torch.clamp()``.
|
||||
|
||||
Args:
|
||||
min (Number | optional): Lower-bound of the range to be clamped to.
|
||||
Default to -1.
|
||||
max (Number | optional): Upper-bound of the range to be clamped to.
|
||||
Default to 1.
|
||||
"""
|
||||
|
||||
def __init__(self, min=-1., max=1.):
|
||||
super(Clamp, self).__init__()
|
||||
self.min = min
|
||||
self.max = max
|
||||
|
||||
def forward(self, x):
|
||||
"""Forward function.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): The input tensor.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Clamped tensor.
|
||||
"""
|
||||
return torch.clamp(x, min=self.min, max=self.max)
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
r"""Applies the Gaussian Error Linear Units function:
|
||||
|
||||
.. math::
|
||||
\text{GELU}(x) = x * \Phi(x)
|
||||
where :math:`\Phi(x)` is the Cumulative Distribution Function for
|
||||
Gaussian Distribution.
|
||||
|
||||
Shape:
|
||||
- Input: :math:`(N, *)` where `*` means, any number of additional
|
||||
dimensions
|
||||
- Output: :math:`(N, *)`, same shape as the input
|
||||
|
||||
.. image:: scripts/activation_images/GELU.png
|
||||
|
||||
Examples::
|
||||
|
||||
>>> m = nn.GELU()
|
||||
>>> input = torch.randn(2)
|
||||
>>> output = m(input)
|
||||
"""
|
||||
|
||||
def forward(self, input):
|
||||
return F.gelu(input)
|
||||
|
||||
|
||||
if (TORCH_VERSION == 'parrots'
|
||||
or digit_version(TORCH_VERSION) < digit_version('1.4')):
|
||||
ACTIVATION_LAYERS.register_module(module=GELU)
|
||||
else:
|
||||
ACTIVATION_LAYERS.register_module(module=nn.GELU)
|
||||
|
||||
|
||||
def build_activation_layer(cfg):
|
||||
"""Build activation layer.
|
||||
|
||||
Args:
|
||||
cfg (dict): The activation layer config, which should contain:
|
||||
- type (str): Layer type.
|
||||
- layer args: Args needed to instantiate an activation layer.
|
||||
|
||||
Returns:
|
||||
nn.Module: Created activation layer.
|
||||
"""
|
||||
return build_from_cfg(cfg, ACTIVATION_LAYERS)
|
||||
@@ -0,0 +1,125 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from ..utils import constant_init, kaiming_init
|
||||
from .registry import PLUGIN_LAYERS
|
||||
|
||||
|
||||
def last_zero_init(m):
|
||||
if isinstance(m, nn.Sequential):
|
||||
constant_init(m[-1], val=0)
|
||||
else:
|
||||
constant_init(m, val=0)
|
||||
|
||||
|
||||
@PLUGIN_LAYERS.register_module()
|
||||
class ContextBlock(nn.Module):
|
||||
"""ContextBlock module in GCNet.
|
||||
|
||||
See 'GCNet: Non-local Networks Meet Squeeze-Excitation Networks and Beyond'
|
||||
(https://arxiv.org/abs/1904.11492) for details.
|
||||
|
||||
Args:
|
||||
in_channels (int): Channels of the input feature map.
|
||||
ratio (float): Ratio of channels of transform bottleneck
|
||||
pooling_type (str): Pooling method for context modeling.
|
||||
Options are 'att' and 'avg', stand for attention pooling and
|
||||
average pooling respectively. Default: 'att'.
|
||||
fusion_types (Sequence[str]): Fusion method for feature fusion,
|
||||
Options are 'channels_add', 'channel_mul', stand for channelwise
|
||||
addition and multiplication respectively. Default: ('channel_add',)
|
||||
"""
|
||||
|
||||
_abbr_ = 'context_block'
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
ratio,
|
||||
pooling_type='att',
|
||||
fusion_types=('channel_add', )):
|
||||
super(ContextBlock, self).__init__()
|
||||
assert pooling_type in ['avg', 'att']
|
||||
assert isinstance(fusion_types, (list, tuple))
|
||||
valid_fusion_types = ['channel_add', 'channel_mul']
|
||||
assert all([f in valid_fusion_types for f in fusion_types])
|
||||
assert len(fusion_types) > 0, 'at least one fusion should be used'
|
||||
self.in_channels = in_channels
|
||||
self.ratio = ratio
|
||||
self.planes = int(in_channels * ratio)
|
||||
self.pooling_type = pooling_type
|
||||
self.fusion_types = fusion_types
|
||||
if pooling_type == 'att':
|
||||
self.conv_mask = nn.Conv2d(in_channels, 1, kernel_size=1)
|
||||
self.softmax = nn.Softmax(dim=2)
|
||||
else:
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
if 'channel_add' in fusion_types:
|
||||
self.channel_add_conv = nn.Sequential(
|
||||
nn.Conv2d(self.in_channels, self.planes, kernel_size=1),
|
||||
nn.LayerNorm([self.planes, 1, 1]),
|
||||
nn.ReLU(inplace=True), # yapf: disable
|
||||
nn.Conv2d(self.planes, self.in_channels, kernel_size=1))
|
||||
else:
|
||||
self.channel_add_conv = None
|
||||
if 'channel_mul' in fusion_types:
|
||||
self.channel_mul_conv = nn.Sequential(
|
||||
nn.Conv2d(self.in_channels, self.planes, kernel_size=1),
|
||||
nn.LayerNorm([self.planes, 1, 1]),
|
||||
nn.ReLU(inplace=True), # yapf: disable
|
||||
nn.Conv2d(self.planes, self.in_channels, kernel_size=1))
|
||||
else:
|
||||
self.channel_mul_conv = None
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self):
|
||||
if self.pooling_type == 'att':
|
||||
kaiming_init(self.conv_mask, mode='fan_in')
|
||||
self.conv_mask.inited = True
|
||||
|
||||
if self.channel_add_conv is not None:
|
||||
last_zero_init(self.channel_add_conv)
|
||||
if self.channel_mul_conv is not None:
|
||||
last_zero_init(self.channel_mul_conv)
|
||||
|
||||
def spatial_pool(self, x):
|
||||
batch, channel, height, width = x.size()
|
||||
if self.pooling_type == 'att':
|
||||
input_x = x
|
||||
# [N, C, H * W]
|
||||
input_x = input_x.view(batch, channel, height * width)
|
||||
# [N, 1, C, H * W]
|
||||
input_x = input_x.unsqueeze(1)
|
||||
# [N, 1, H, W]
|
||||
context_mask = self.conv_mask(x)
|
||||
# [N, 1, H * W]
|
||||
context_mask = context_mask.view(batch, 1, height * width)
|
||||
# [N, 1, H * W]
|
||||
context_mask = self.softmax(context_mask)
|
||||
# [N, 1, H * W, 1]
|
||||
context_mask = context_mask.unsqueeze(-1)
|
||||
# [N, 1, C, 1]
|
||||
context = torch.matmul(input_x, context_mask)
|
||||
# [N, C, 1, 1]
|
||||
context = context.view(batch, channel, 1, 1)
|
||||
else:
|
||||
# [N, C, 1, 1]
|
||||
context = self.avg_pool(x)
|
||||
|
||||
return context
|
||||
|
||||
def forward(self, x):
|
||||
# [N, C, 1, 1]
|
||||
context = self.spatial_pool(x)
|
||||
|
||||
out = x
|
||||
if self.channel_mul_conv is not None:
|
||||
# [N, C, 1, 1]
|
||||
channel_mul_term = torch.sigmoid(self.channel_mul_conv(context))
|
||||
out = out * channel_mul_term
|
||||
if self.channel_add_conv is not None:
|
||||
# [N, C, 1, 1]
|
||||
channel_add_term = self.channel_add_conv(context)
|
||||
out = out + channel_add_term
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,44 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from torch import nn
|
||||
|
||||
from .registry import CONV_LAYERS
|
||||
|
||||
CONV_LAYERS.register_module('Conv1d', module=nn.Conv1d)
|
||||
CONV_LAYERS.register_module('Conv2d', module=nn.Conv2d)
|
||||
CONV_LAYERS.register_module('Conv3d', module=nn.Conv3d)
|
||||
CONV_LAYERS.register_module('Conv', module=nn.Conv2d)
|
||||
|
||||
|
||||
def build_conv_layer(cfg, *args, **kwargs):
|
||||
"""Build convolution layer.
|
||||
|
||||
Args:
|
||||
cfg (None or dict): The conv layer config, which should contain:
|
||||
- type (str): Layer type.
|
||||
- layer args: Args needed to instantiate an conv layer.
|
||||
args (argument list): Arguments passed to the `__init__`
|
||||
method of the corresponding conv layer.
|
||||
kwargs (keyword arguments): Keyword arguments passed to the `__init__`
|
||||
method of the corresponding conv layer.
|
||||
|
||||
Returns:
|
||||
nn.Module: Created conv layer.
|
||||
"""
|
||||
if cfg is None:
|
||||
cfg_ = dict(type='Conv2d')
|
||||
else:
|
||||
if not isinstance(cfg, dict):
|
||||
raise TypeError('cfg must be a dict')
|
||||
if 'type' not in cfg:
|
||||
raise KeyError('the cfg dict must contain the key "type"')
|
||||
cfg_ = cfg.copy()
|
||||
|
||||
layer_type = cfg_.pop('type')
|
||||
if layer_type not in CONV_LAYERS:
|
||||
raise KeyError(f'Unrecognized norm type {layer_type}')
|
||||
else:
|
||||
conv_layer = CONV_LAYERS.get(layer_type)
|
||||
|
||||
layer = conv_layer(*args, **kwargs, **cfg_)
|
||||
|
||||
return layer
|
||||
@@ -0,0 +1,62 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import math
|
||||
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
|
||||
from .registry import CONV_LAYERS
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module()
|
||||
class Conv2dAdaptivePadding(nn.Conv2d):
|
||||
"""Implementation of 2D convolution in tensorflow with `padding` as "same",
|
||||
which applies padding to input (if needed) so that input image gets fully
|
||||
covered by filter and stride you specified. For stride 1, this will ensure
|
||||
that output image size is same as input. For stride of 2, output dimensions
|
||||
will be half, for example.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of channels in the input image
|
||||
out_channels (int): Number of channels produced by the convolution
|
||||
kernel_size (int or tuple): Size of the convolving kernel
|
||||
stride (int or tuple, optional): Stride of the convolution. Default: 1
|
||||
padding (int or tuple, optional): Zero-padding added to both sides of
|
||||
the input. Default: 0
|
||||
dilation (int or tuple, optional): Spacing between kernel elements.
|
||||
Default: 1
|
||||
groups (int, optional): Number of blocked connections from input
|
||||
channels to output channels. Default: 1
|
||||
bias (bool, optional): If ``True``, adds a learnable bias to the
|
||||
output. Default: ``True``
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=True):
|
||||
super().__init__(in_channels, out_channels, kernel_size, stride, 0,
|
||||
dilation, groups, bias)
|
||||
|
||||
def forward(self, x):
|
||||
img_h, img_w = x.size()[-2:]
|
||||
kernel_h, kernel_w = self.weight.size()[-2:]
|
||||
stride_h, stride_w = self.stride
|
||||
output_h = math.ceil(img_h / stride_h)
|
||||
output_w = math.ceil(img_w / stride_w)
|
||||
pad_h = (
|
||||
max((output_h - 1) * self.stride[0] +
|
||||
(kernel_h - 1) * self.dilation[0] + 1 - img_h, 0))
|
||||
pad_w = (
|
||||
max((output_w - 1) * self.stride[1] +
|
||||
(kernel_w - 1) * self.dilation[1] + 1 - img_w, 0))
|
||||
if pad_h > 0 or pad_w > 0:
|
||||
x = F.pad(x, [
|
||||
pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2
|
||||
])
|
||||
return F.conv2d(x, self.weight, self.bias, self.stride, self.padding,
|
||||
self.dilation, self.groups)
|
||||
@@ -0,0 +1,206 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import warnings
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import _BatchNorm, _InstanceNorm
|
||||
from ..utils import constant_init, kaiming_init
|
||||
from .activation import build_activation_layer
|
||||
from .conv import build_conv_layer
|
||||
from .norm import build_norm_layer
|
||||
from .padding import build_padding_layer
|
||||
from .registry import PLUGIN_LAYERS
|
||||
|
||||
|
||||
@PLUGIN_LAYERS.register_module()
|
||||
class ConvModule(nn.Module):
|
||||
"""A conv block that bundles conv/norm/activation layers.
|
||||
|
||||
This block simplifies the usage of convolution layers, which are commonly
|
||||
used with a norm layer (e.g., BatchNorm) and activation layer (e.g., ReLU).
|
||||
It is based upon three build methods: `build_conv_layer()`,
|
||||
`build_norm_layer()` and `build_activation_layer()`.
|
||||
|
||||
Besides, we add some additional features in this module.
|
||||
1. Automatically set `bias` of the conv layer.
|
||||
2. Spectral norm is supported.
|
||||
3. More padding modes are supported. Before PyTorch 1.5, nn.Conv2d only
|
||||
supports zero and circular padding, and we add "reflect" padding mode.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of channels in the input feature map.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
out_channels (int): Number of channels produced by the convolution.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
kernel_size (int | tuple[int]): Size of the convolving kernel.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
stride (int | tuple[int]): Stride of the convolution.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
padding (int | tuple[int]): Zero-padding added to both sides of
|
||||
the input. Same as that in ``nn._ConvNd``.
|
||||
dilation (int | tuple[int]): Spacing between kernel elements.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
groups (int): Number of blocked connections from input channels to
|
||||
output channels. Same as that in ``nn._ConvNd``.
|
||||
bias (bool | str): If specified as `auto`, it will be decided by the
|
||||
norm_cfg. Bias will be set as True if `norm_cfg` is None, otherwise
|
||||
False. Default: "auto".
|
||||
conv_cfg (dict): Config dict for convolution layer. Default: None,
|
||||
which means using conv2d.
|
||||
norm_cfg (dict): Config dict for normalization layer. Default: None.
|
||||
act_cfg (dict): Config dict for activation layer.
|
||||
Default: dict(type='ReLU').
|
||||
inplace (bool): Whether to use inplace mode for activation.
|
||||
Default: True.
|
||||
with_spectral_norm (bool): Whether use spectral norm in conv module.
|
||||
Default: False.
|
||||
padding_mode (str): If the `padding_mode` has not been supported by
|
||||
current `Conv2d` in PyTorch, we will use our own padding layer
|
||||
instead. Currently, we support ['zeros', 'circular'] with official
|
||||
implementation and ['reflect'] with our own implementation.
|
||||
Default: 'zeros'.
|
||||
order (tuple[str]): The order of conv/norm/activation layers. It is a
|
||||
sequence of "conv", "norm" and "act". Common examples are
|
||||
("conv", "norm", "act") and ("act", "conv", "norm").
|
||||
Default: ('conv', 'norm', 'act').
|
||||
"""
|
||||
|
||||
_abbr_ = 'conv_block'
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias='auto',
|
||||
conv_cfg=None,
|
||||
norm_cfg=None,
|
||||
act_cfg=dict(type='ReLU'),
|
||||
inplace=True,
|
||||
with_spectral_norm=False,
|
||||
padding_mode='zeros',
|
||||
order=('conv', 'norm', 'act')):
|
||||
super(ConvModule, self).__init__()
|
||||
assert conv_cfg is None or isinstance(conv_cfg, dict)
|
||||
assert norm_cfg is None or isinstance(norm_cfg, dict)
|
||||
assert act_cfg is None or isinstance(act_cfg, dict)
|
||||
official_padding_mode = ['zeros', 'circular']
|
||||
self.conv_cfg = conv_cfg
|
||||
self.norm_cfg = norm_cfg
|
||||
self.act_cfg = act_cfg
|
||||
self.inplace = inplace
|
||||
self.with_spectral_norm = with_spectral_norm
|
||||
self.with_explicit_padding = padding_mode not in official_padding_mode
|
||||
self.order = order
|
||||
assert isinstance(self.order, tuple) and len(self.order) == 3
|
||||
assert set(order) == set(['conv', 'norm', 'act'])
|
||||
|
||||
self.with_norm = norm_cfg is not None
|
||||
self.with_activation = act_cfg is not None
|
||||
# if the conv layer is before a norm layer, bias is unnecessary.
|
||||
if bias == 'auto':
|
||||
bias = not self.with_norm
|
||||
self.with_bias = bias
|
||||
|
||||
if self.with_explicit_padding:
|
||||
pad_cfg = dict(type=padding_mode)
|
||||
self.padding_layer = build_padding_layer(pad_cfg, padding)
|
||||
|
||||
# reset padding to 0 for conv module
|
||||
conv_padding = 0 if self.with_explicit_padding else padding
|
||||
# build convolution layer
|
||||
self.conv = build_conv_layer(
|
||||
conv_cfg,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
padding=conv_padding,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
bias=bias)
|
||||
# export the attributes of self.conv to a higher level for convenience
|
||||
self.in_channels = self.conv.in_channels
|
||||
self.out_channels = self.conv.out_channels
|
||||
self.kernel_size = self.conv.kernel_size
|
||||
self.stride = self.conv.stride
|
||||
self.padding = padding
|
||||
self.dilation = self.conv.dilation
|
||||
self.transposed = self.conv.transposed
|
||||
self.output_padding = self.conv.output_padding
|
||||
self.groups = self.conv.groups
|
||||
|
||||
if self.with_spectral_norm:
|
||||
self.conv = nn.utils.spectral_norm(self.conv)
|
||||
|
||||
# build normalization layers
|
||||
if self.with_norm:
|
||||
# norm layer is after conv layer
|
||||
if order.index('norm') > order.index('conv'):
|
||||
norm_channels = out_channels
|
||||
else:
|
||||
norm_channels = in_channels
|
||||
self.norm_name, norm = build_norm_layer(norm_cfg, norm_channels)
|
||||
self.add_module(self.norm_name, norm)
|
||||
if self.with_bias:
|
||||
if isinstance(norm, (_BatchNorm, _InstanceNorm)):
|
||||
warnings.warn(
|
||||
'Unnecessary conv bias before batch/instance norm')
|
||||
else:
|
||||
self.norm_name = None
|
||||
|
||||
# build activation layer
|
||||
if self.with_activation:
|
||||
act_cfg_ = act_cfg.copy()
|
||||
# nn.Tanh has no 'inplace' argument
|
||||
if act_cfg_['type'] not in [
|
||||
'Tanh', 'PReLU', 'Sigmoid', 'HSigmoid', 'Swish'
|
||||
]:
|
||||
act_cfg_.setdefault('inplace', inplace)
|
||||
self.activate = build_activation_layer(act_cfg_)
|
||||
|
||||
# Use msra init by default
|
||||
self.init_weights()
|
||||
|
||||
@property
|
||||
def norm(self):
|
||||
if self.norm_name:
|
||||
return getattr(self, self.norm_name)
|
||||
else:
|
||||
return None
|
||||
|
||||
def init_weights(self):
|
||||
# 1. It is mainly for customized conv layers with their own
|
||||
# initialization manners by calling their own ``init_weights()``,
|
||||
# and we do not want ConvModule to override the initialization.
|
||||
# 2. For customized conv layers without their own initialization
|
||||
# manners (that is, they don't have their own ``init_weights()``)
|
||||
# and PyTorch's conv layers, they will be initialized by
|
||||
# this method with default ``kaiming_init``.
|
||||
# Note: For PyTorch's conv layers, they will be overwritten by our
|
||||
# initialization implementation using default ``kaiming_init``.
|
||||
if not hasattr(self.conv, 'init_weights'):
|
||||
if self.with_activation and self.act_cfg['type'] == 'LeakyReLU':
|
||||
nonlinearity = 'leaky_relu'
|
||||
a = self.act_cfg.get('negative_slope', 0.01)
|
||||
else:
|
||||
nonlinearity = 'relu'
|
||||
a = 0
|
||||
kaiming_init(self.conv, a=a, nonlinearity=nonlinearity)
|
||||
if self.with_norm:
|
||||
constant_init(self.norm, 1, bias=0)
|
||||
|
||||
def forward(self, x, activate=True, norm=True):
|
||||
for layer in self.order:
|
||||
if layer == 'conv':
|
||||
if self.with_explicit_padding:
|
||||
x = self.padding_layer(x)
|
||||
x = self.conv(x)
|
||||
elif layer == 'norm' and norm and self.with_norm:
|
||||
x = self.norm(x)
|
||||
elif layer == 'act' and activate and self.with_activation:
|
||||
x = self.activate(x)
|
||||
return x
|
||||
@@ -0,0 +1,148 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .registry import CONV_LAYERS
|
||||
|
||||
|
||||
def conv_ws_2d(input,
|
||||
weight,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
eps=1e-5):
|
||||
c_in = weight.size(0)
|
||||
weight_flat = weight.view(c_in, -1)
|
||||
mean = weight_flat.mean(dim=1, keepdim=True).view(c_in, 1, 1, 1)
|
||||
std = weight_flat.std(dim=1, keepdim=True).view(c_in, 1, 1, 1)
|
||||
weight = (weight - mean) / (std + eps)
|
||||
return F.conv2d(input, weight, bias, stride, padding, dilation, groups)
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module('ConvWS')
|
||||
class ConvWS2d(nn.Conv2d):
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=True,
|
||||
eps=1e-5):
|
||||
super(ConvWS2d, self).__init__(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
bias=bias)
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, x):
|
||||
return conv_ws_2d(x, self.weight, self.bias, self.stride, self.padding,
|
||||
self.dilation, self.groups, self.eps)
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module(name='ConvAWS')
|
||||
class ConvAWS2d(nn.Conv2d):
|
||||
"""AWS (Adaptive Weight Standardization)
|
||||
|
||||
This is a variant of Weight Standardization
|
||||
(https://arxiv.org/pdf/1903.10520.pdf)
|
||||
It is used in DetectoRS to avoid NaN
|
||||
(https://arxiv.org/pdf/2006.02334.pdf)
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of channels in the input image
|
||||
out_channels (int): Number of channels produced by the convolution
|
||||
kernel_size (int or tuple): Size of the conv kernel
|
||||
stride (int or tuple, optional): Stride of the convolution. Default: 1
|
||||
padding (int or tuple, optional): Zero-padding added to both sides of
|
||||
the input. Default: 0
|
||||
dilation (int or tuple, optional): Spacing between kernel elements.
|
||||
Default: 1
|
||||
groups (int, optional): Number of blocked connections from input
|
||||
channels to output channels. Default: 1
|
||||
bias (bool, optional): If set True, adds a learnable bias to the
|
||||
output. Default: True
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=True):
|
||||
super().__init__(
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=groups,
|
||||
bias=bias)
|
||||
self.register_buffer('weight_gamma',
|
||||
torch.ones(self.out_channels, 1, 1, 1))
|
||||
self.register_buffer('weight_beta',
|
||||
torch.zeros(self.out_channels, 1, 1, 1))
|
||||
|
||||
def _get_weight(self, weight):
|
||||
weight_flat = weight.view(weight.size(0), -1)
|
||||
mean = weight_flat.mean(dim=1).view(-1, 1, 1, 1)
|
||||
std = torch.sqrt(weight_flat.var(dim=1) + 1e-5).view(-1, 1, 1, 1)
|
||||
weight = (weight - mean) / std
|
||||
weight = self.weight_gamma * weight + self.weight_beta
|
||||
return weight
|
||||
|
||||
def forward(self, x):
|
||||
weight = self._get_weight(self.weight)
|
||||
return F.conv2d(x, weight, self.bias, self.stride, self.padding,
|
||||
self.dilation, self.groups)
|
||||
|
||||
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict,
|
||||
missing_keys, unexpected_keys, error_msgs):
|
||||
"""Override default load function.
|
||||
|
||||
AWS overrides the function _load_from_state_dict to recover
|
||||
weight_gamma and weight_beta if they are missing. If weight_gamma and
|
||||
weight_beta are found in the checkpoint, this function will return
|
||||
after super()._load_from_state_dict. Otherwise, it will compute the
|
||||
mean and std of the pretrained weights and store them in weight_beta
|
||||
and weight_gamma.
|
||||
"""
|
||||
|
||||
self.weight_gamma.data.fill_(-1)
|
||||
local_missing_keys = []
|
||||
super()._load_from_state_dict(state_dict, prefix, local_metadata,
|
||||
strict, local_missing_keys,
|
||||
unexpected_keys, error_msgs)
|
||||
if self.weight_gamma.data.mean() > 0:
|
||||
for k in local_missing_keys:
|
||||
missing_keys.append(k)
|
||||
return
|
||||
weight = self.weight.data
|
||||
weight_flat = weight.view(weight.size(0), -1)
|
||||
mean = weight_flat.mean(dim=1).view(-1, 1, 1, 1)
|
||||
std = torch.sqrt(weight_flat.var(dim=1) + 1e-5).view(-1, 1, 1, 1)
|
||||
self.weight_beta.data.copy_(mean)
|
||||
self.weight_gamma.data.copy_(std)
|
||||
missing_gamma_beta = [
|
||||
k for k in local_missing_keys
|
||||
if k.endswith('weight_gamma') or k.endswith('weight_beta')
|
||||
]
|
||||
for k in missing_gamma_beta:
|
||||
local_missing_keys.remove(k)
|
||||
for k in local_missing_keys:
|
||||
missing_keys.append(k)
|
||||
@@ -0,0 +1,96 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch.nn as nn
|
||||
|
||||
from .conv_module import ConvModule
|
||||
|
||||
|
||||
class DepthwiseSeparableConvModule(nn.Module):
|
||||
"""Depthwise separable convolution module.
|
||||
|
||||
See https://arxiv.org/pdf/1704.04861.pdf for details.
|
||||
|
||||
This module can replace a ConvModule with the conv block replaced by two
|
||||
conv block: depthwise conv block and pointwise conv block. The depthwise
|
||||
conv block contains depthwise-conv/norm/activation layers. The pointwise
|
||||
conv block contains pointwise-conv/norm/activation layers. It should be
|
||||
noted that there will be norm/activation layer in the depthwise conv block
|
||||
if `norm_cfg` and `act_cfg` are specified.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of channels in the input feature map.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
out_channels (int): Number of channels produced by the convolution.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
kernel_size (int | tuple[int]): Size of the convolving kernel.
|
||||
Same as that in ``nn._ConvNd``.
|
||||
stride (int | tuple[int]): Stride of the convolution.
|
||||
Same as that in ``nn._ConvNd``. Default: 1.
|
||||
padding (int | tuple[int]): Zero-padding added to both sides of
|
||||
the input. Same as that in ``nn._ConvNd``. Default: 0.
|
||||
dilation (int | tuple[int]): Spacing between kernel elements.
|
||||
Same as that in ``nn._ConvNd``. Default: 1.
|
||||
norm_cfg (dict): Default norm config for both depthwise ConvModule and
|
||||
pointwise ConvModule. Default: None.
|
||||
act_cfg (dict): Default activation config for both depthwise ConvModule
|
||||
and pointwise ConvModule. Default: dict(type='ReLU').
|
||||
dw_norm_cfg (dict): Norm config of depthwise ConvModule. If it is
|
||||
'default', it will be the same as `norm_cfg`. Default: 'default'.
|
||||
dw_act_cfg (dict): Activation config of depthwise ConvModule. If it is
|
||||
'default', it will be the same as `act_cfg`. Default: 'default'.
|
||||
pw_norm_cfg (dict): Norm config of pointwise ConvModule. If it is
|
||||
'default', it will be the same as `norm_cfg`. Default: 'default'.
|
||||
pw_act_cfg (dict): Activation config of pointwise ConvModule. If it is
|
||||
'default', it will be the same as `act_cfg`. Default: 'default'.
|
||||
kwargs (optional): Other shared arguments for depthwise and pointwise
|
||||
ConvModule. See ConvModule for ref.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
norm_cfg=None,
|
||||
act_cfg=dict(type='ReLU'),
|
||||
dw_norm_cfg='default',
|
||||
dw_act_cfg='default',
|
||||
pw_norm_cfg='default',
|
||||
pw_act_cfg='default',
|
||||
**kwargs):
|
||||
super(DepthwiseSeparableConvModule, self).__init__()
|
||||
assert 'groups' not in kwargs, 'groups should not be specified'
|
||||
|
||||
# if norm/activation config of depthwise/pointwise ConvModule is not
|
||||
# specified, use default config.
|
||||
dw_norm_cfg = dw_norm_cfg if dw_norm_cfg != 'default' else norm_cfg
|
||||
dw_act_cfg = dw_act_cfg if dw_act_cfg != 'default' else act_cfg
|
||||
pw_norm_cfg = pw_norm_cfg if pw_norm_cfg != 'default' else norm_cfg
|
||||
pw_act_cfg = pw_act_cfg if pw_act_cfg != 'default' else act_cfg
|
||||
|
||||
# depthwise convolution
|
||||
self.depthwise_conv = ConvModule(
|
||||
in_channels,
|
||||
in_channels,
|
||||
kernel_size,
|
||||
stride=stride,
|
||||
padding=padding,
|
||||
dilation=dilation,
|
||||
groups=in_channels,
|
||||
norm_cfg=dw_norm_cfg,
|
||||
act_cfg=dw_act_cfg,
|
||||
**kwargs)
|
||||
|
||||
self.pointwise_conv = ConvModule(
|
||||
in_channels,
|
||||
out_channels,
|
||||
1,
|
||||
norm_cfg=pw_norm_cfg,
|
||||
act_cfg=pw_act_cfg,
|
||||
**kwargs)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.depthwise_conv(x)
|
||||
x = self.pointwise_conv(x)
|
||||
return x
|
||||
@@ -0,0 +1,65 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from custom_mmpkg.custom_mmcv import build_from_cfg
|
||||
from .registry import DROPOUT_LAYERS
|
||||
|
||||
|
||||
def drop_path(x, drop_prob=0., training=False):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of
|
||||
residual blocks).
|
||||
|
||||
We follow the implementation
|
||||
https://github.com/rwightman/pytorch-image-models/blob/a2727c1bf78ba0d7b5727f5f95e37fb7f8866b1f/timm/models/layers/drop.py # noqa: E501
|
||||
"""
|
||||
if drop_prob == 0. or not training:
|
||||
return x
|
||||
keep_prob = 1 - drop_prob
|
||||
# handle tensors with different dimensions, not just 4D tensors.
|
||||
shape = (x.shape[0], ) + (1, ) * (x.ndim - 1)
|
||||
random_tensor = keep_prob + torch.rand(
|
||||
shape, dtype=x.dtype, device=x.device)
|
||||
output = x.div(keep_prob) * random_tensor.floor()
|
||||
return output
|
||||
|
||||
|
||||
@DROPOUT_LAYERS.register_module()
|
||||
class DropPath(nn.Module):
|
||||
"""Drop paths (Stochastic Depth) per sample (when applied in main path of
|
||||
residual blocks).
|
||||
|
||||
We follow the implementation
|
||||
https://github.com/rwightman/pytorch-image-models/blob/a2727c1bf78ba0d7b5727f5f95e37fb7f8866b1f/timm/models/layers/drop.py # noqa: E501
|
||||
|
||||
Args:
|
||||
drop_prob (float): Probability of the path to be zeroed. Default: 0.1
|
||||
"""
|
||||
|
||||
def __init__(self, drop_prob=0.1):
|
||||
super(DropPath, self).__init__()
|
||||
self.drop_prob = drop_prob
|
||||
|
||||
def forward(self, x):
|
||||
return drop_path(x, self.drop_prob, self.training)
|
||||
|
||||
|
||||
@DROPOUT_LAYERS.register_module()
|
||||
class Dropout(nn.Dropout):
|
||||
"""A wrapper for ``torch.nn.Dropout``, We rename the ``p`` of
|
||||
``torch.nn.Dropout`` to ``drop_prob`` so as to be consistent with
|
||||
``DropPath``
|
||||
|
||||
Args:
|
||||
drop_prob (float): Probability of the elements to be
|
||||
zeroed. Default: 0.5.
|
||||
inplace (bool): Do the operation inplace or not. Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self, drop_prob=0.5, inplace=False):
|
||||
super().__init__(p=drop_prob, inplace=inplace)
|
||||
|
||||
|
||||
def build_dropout(cfg, default_args=None):
|
||||
"""Builder for drop out layers."""
|
||||
return build_from_cfg(cfg, DROPOUT_LAYERS, default_args)
|
||||
@@ -0,0 +1,412 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import math
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..utils import kaiming_init
|
||||
from .registry import PLUGIN_LAYERS
|
||||
|
||||
|
||||
@PLUGIN_LAYERS.register_module()
|
||||
class GeneralizedAttention(nn.Module):
|
||||
"""GeneralizedAttention module.
|
||||
|
||||
See 'An Empirical Study of Spatial Attention Mechanisms in Deep Networks'
|
||||
(https://arxiv.org/abs/1711.07971) for details.
|
||||
|
||||
Args:
|
||||
in_channels (int): Channels of the input feature map.
|
||||
spatial_range (int): The spatial range. -1 indicates no spatial range
|
||||
constraint. Default: -1.
|
||||
num_heads (int): The head number of empirical_attention module.
|
||||
Default: 9.
|
||||
position_embedding_dim (int): The position embedding dimension.
|
||||
Default: -1.
|
||||
position_magnitude (int): A multiplier acting on coord difference.
|
||||
Default: 1.
|
||||
kv_stride (int): The feature stride acting on key/value feature map.
|
||||
Default: 2.
|
||||
q_stride (int): The feature stride acting on query feature map.
|
||||
Default: 1.
|
||||
attention_type (str): A binary indicator string for indicating which
|
||||
items in generalized empirical_attention module are used.
|
||||
Default: '1111'.
|
||||
|
||||
- '1000' indicates 'query and key content' (appr - appr) item,
|
||||
- '0100' indicates 'query content and relative position'
|
||||
(appr - position) item,
|
||||
- '0010' indicates 'key content only' (bias - appr) item,
|
||||
- '0001' indicates 'relative position only' (bias - position) item.
|
||||
"""
|
||||
|
||||
_abbr_ = 'gen_attention_block'
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
spatial_range=-1,
|
||||
num_heads=9,
|
||||
position_embedding_dim=-1,
|
||||
position_magnitude=1,
|
||||
kv_stride=2,
|
||||
q_stride=1,
|
||||
attention_type='1111'):
|
||||
|
||||
super(GeneralizedAttention, self).__init__()
|
||||
|
||||
# hard range means local range for non-local operation
|
||||
self.position_embedding_dim = (
|
||||
position_embedding_dim
|
||||
if position_embedding_dim > 0 else in_channels)
|
||||
|
||||
self.position_magnitude = position_magnitude
|
||||
self.num_heads = num_heads
|
||||
self.in_channels = in_channels
|
||||
self.spatial_range = spatial_range
|
||||
self.kv_stride = kv_stride
|
||||
self.q_stride = q_stride
|
||||
self.attention_type = [bool(int(_)) for _ in attention_type]
|
||||
self.qk_embed_dim = in_channels // num_heads
|
||||
out_c = self.qk_embed_dim * num_heads
|
||||
|
||||
if self.attention_type[0] or self.attention_type[1]:
|
||||
self.query_conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_c,
|
||||
kernel_size=1,
|
||||
bias=False)
|
||||
self.query_conv.kaiming_init = True
|
||||
|
||||
if self.attention_type[0] or self.attention_type[2]:
|
||||
self.key_conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=out_c,
|
||||
kernel_size=1,
|
||||
bias=False)
|
||||
self.key_conv.kaiming_init = True
|
||||
|
||||
self.v_dim = in_channels // num_heads
|
||||
self.value_conv = nn.Conv2d(
|
||||
in_channels=in_channels,
|
||||
out_channels=self.v_dim * num_heads,
|
||||
kernel_size=1,
|
||||
bias=False)
|
||||
self.value_conv.kaiming_init = True
|
||||
|
||||
if self.attention_type[1] or self.attention_type[3]:
|
||||
self.appr_geom_fc_x = nn.Linear(
|
||||
self.position_embedding_dim // 2, out_c, bias=False)
|
||||
self.appr_geom_fc_x.kaiming_init = True
|
||||
|
||||
self.appr_geom_fc_y = nn.Linear(
|
||||
self.position_embedding_dim // 2, out_c, bias=False)
|
||||
self.appr_geom_fc_y.kaiming_init = True
|
||||
|
||||
if self.attention_type[2]:
|
||||
stdv = 1.0 / math.sqrt(self.qk_embed_dim * 2)
|
||||
appr_bias_value = -2 * stdv * torch.rand(out_c) + stdv
|
||||
self.appr_bias = nn.Parameter(appr_bias_value)
|
||||
|
||||
if self.attention_type[3]:
|
||||
stdv = 1.0 / math.sqrt(self.qk_embed_dim * 2)
|
||||
geom_bias_value = -2 * stdv * torch.rand(out_c) + stdv
|
||||
self.geom_bias = nn.Parameter(geom_bias_value)
|
||||
|
||||
self.proj_conv = nn.Conv2d(
|
||||
in_channels=self.v_dim * num_heads,
|
||||
out_channels=in_channels,
|
||||
kernel_size=1,
|
||||
bias=True)
|
||||
self.proj_conv.kaiming_init = True
|
||||
self.gamma = nn.Parameter(torch.zeros(1))
|
||||
|
||||
if self.spatial_range >= 0:
|
||||
# only works when non local is after 3*3 conv
|
||||
if in_channels == 256:
|
||||
max_len = 84
|
||||
elif in_channels == 512:
|
||||
max_len = 42
|
||||
|
||||
max_len_kv = int((max_len - 1.0) / self.kv_stride + 1)
|
||||
local_constraint_map = np.ones(
|
||||
(max_len, max_len, max_len_kv, max_len_kv), dtype=np.int)
|
||||
for iy in range(max_len):
|
||||
for ix in range(max_len):
|
||||
local_constraint_map[
|
||||
iy, ix,
|
||||
max((iy - self.spatial_range) //
|
||||
self.kv_stride, 0):min((iy + self.spatial_range +
|
||||
1) // self.kv_stride +
|
||||
1, max_len),
|
||||
max((ix - self.spatial_range) //
|
||||
self.kv_stride, 0):min((ix + self.spatial_range +
|
||||
1) // self.kv_stride +
|
||||
1, max_len)] = 0
|
||||
|
||||
self.local_constraint_map = nn.Parameter(
|
||||
torch.from_numpy(local_constraint_map).byte(),
|
||||
requires_grad=False)
|
||||
|
||||
if self.q_stride > 1:
|
||||
self.q_downsample = nn.AvgPool2d(
|
||||
kernel_size=1, stride=self.q_stride)
|
||||
else:
|
||||
self.q_downsample = None
|
||||
|
||||
if self.kv_stride > 1:
|
||||
self.kv_downsample = nn.AvgPool2d(
|
||||
kernel_size=1, stride=self.kv_stride)
|
||||
else:
|
||||
self.kv_downsample = None
|
||||
|
||||
self.init_weights()
|
||||
|
||||
def get_position_embedding(self,
|
||||
h,
|
||||
w,
|
||||
h_kv,
|
||||
w_kv,
|
||||
q_stride,
|
||||
kv_stride,
|
||||
device,
|
||||
dtype,
|
||||
feat_dim,
|
||||
wave_length=1000):
|
||||
# the default type of Tensor is float32, leading to type mismatch
|
||||
# in fp16 mode. Cast it to support fp16 mode.
|
||||
h_idxs = torch.linspace(0, h - 1, h).to(device=device, dtype=dtype)
|
||||
h_idxs = h_idxs.view((h, 1)) * q_stride
|
||||
|
||||
w_idxs = torch.linspace(0, w - 1, w).to(device=device, dtype=dtype)
|
||||
w_idxs = w_idxs.view((w, 1)) * q_stride
|
||||
|
||||
h_kv_idxs = torch.linspace(0, h_kv - 1, h_kv).to(
|
||||
device=device, dtype=dtype)
|
||||
h_kv_idxs = h_kv_idxs.view((h_kv, 1)) * kv_stride
|
||||
|
||||
w_kv_idxs = torch.linspace(0, w_kv - 1, w_kv).to(
|
||||
device=device, dtype=dtype)
|
||||
w_kv_idxs = w_kv_idxs.view((w_kv, 1)) * kv_stride
|
||||
|
||||
# (h, h_kv, 1)
|
||||
h_diff = h_idxs.unsqueeze(1) - h_kv_idxs.unsqueeze(0)
|
||||
h_diff *= self.position_magnitude
|
||||
|
||||
# (w, w_kv, 1)
|
||||
w_diff = w_idxs.unsqueeze(1) - w_kv_idxs.unsqueeze(0)
|
||||
w_diff *= self.position_magnitude
|
||||
|
||||
feat_range = torch.arange(0, feat_dim / 4).to(
|
||||
device=device, dtype=dtype)
|
||||
|
||||
dim_mat = torch.Tensor([wave_length]).to(device=device, dtype=dtype)
|
||||
dim_mat = dim_mat**((4. / feat_dim) * feat_range)
|
||||
dim_mat = dim_mat.view((1, 1, -1))
|
||||
|
||||
embedding_x = torch.cat(
|
||||
((w_diff / dim_mat).sin(), (w_diff / dim_mat).cos()), dim=2)
|
||||
|
||||
embedding_y = torch.cat(
|
||||
((h_diff / dim_mat).sin(), (h_diff / dim_mat).cos()), dim=2)
|
||||
|
||||
return embedding_x, embedding_y
|
||||
|
||||
def forward(self, x_input):
|
||||
num_heads = self.num_heads
|
||||
|
||||
# use empirical_attention
|
||||
if self.q_downsample is not None:
|
||||
x_q = self.q_downsample(x_input)
|
||||
else:
|
||||
x_q = x_input
|
||||
n, _, h, w = x_q.shape
|
||||
|
||||
if self.kv_downsample is not None:
|
||||
x_kv = self.kv_downsample(x_input)
|
||||
else:
|
||||
x_kv = x_input
|
||||
_, _, h_kv, w_kv = x_kv.shape
|
||||
|
||||
if self.attention_type[0] or self.attention_type[1]:
|
||||
proj_query = self.query_conv(x_q).view(
|
||||
(n, num_heads, self.qk_embed_dim, h * w))
|
||||
proj_query = proj_query.permute(0, 1, 3, 2)
|
||||
|
||||
if self.attention_type[0] or self.attention_type[2]:
|
||||
proj_key = self.key_conv(x_kv).view(
|
||||
(n, num_heads, self.qk_embed_dim, h_kv * w_kv))
|
||||
|
||||
if self.attention_type[1] or self.attention_type[3]:
|
||||
position_embed_x, position_embed_y = self.get_position_embedding(
|
||||
h, w, h_kv, w_kv, self.q_stride, self.kv_stride,
|
||||
x_input.device, x_input.dtype, self.position_embedding_dim)
|
||||
# (n, num_heads, w, w_kv, dim)
|
||||
position_feat_x = self.appr_geom_fc_x(position_embed_x).\
|
||||
view(1, w, w_kv, num_heads, self.qk_embed_dim).\
|
||||
permute(0, 3, 1, 2, 4).\
|
||||
repeat(n, 1, 1, 1, 1)
|
||||
|
||||
# (n, num_heads, h, h_kv, dim)
|
||||
position_feat_y = self.appr_geom_fc_y(position_embed_y).\
|
||||
view(1, h, h_kv, num_heads, self.qk_embed_dim).\
|
||||
permute(0, 3, 1, 2, 4).\
|
||||
repeat(n, 1, 1, 1, 1)
|
||||
|
||||
position_feat_x /= math.sqrt(2)
|
||||
position_feat_y /= math.sqrt(2)
|
||||
|
||||
# accelerate for saliency only
|
||||
if (np.sum(self.attention_type) == 1) and self.attention_type[2]:
|
||||
appr_bias = self.appr_bias.\
|
||||
view(1, num_heads, 1, self.qk_embed_dim).\
|
||||
repeat(n, 1, 1, 1)
|
||||
|
||||
energy = torch.matmul(appr_bias, proj_key).\
|
||||
view(n, num_heads, 1, h_kv * w_kv)
|
||||
|
||||
h = 1
|
||||
w = 1
|
||||
else:
|
||||
# (n, num_heads, h*w, h_kv*w_kv), query before key, 540mb for
|
||||
if not self.attention_type[0]:
|
||||
energy = torch.zeros(
|
||||
n,
|
||||
num_heads,
|
||||
h,
|
||||
w,
|
||||
h_kv,
|
||||
w_kv,
|
||||
dtype=x_input.dtype,
|
||||
device=x_input.device)
|
||||
|
||||
# attention_type[0]: appr - appr
|
||||
# attention_type[1]: appr - position
|
||||
# attention_type[2]: bias - appr
|
||||
# attention_type[3]: bias - position
|
||||
if self.attention_type[0] or self.attention_type[2]:
|
||||
if self.attention_type[0] and self.attention_type[2]:
|
||||
appr_bias = self.appr_bias.\
|
||||
view(1, num_heads, 1, self.qk_embed_dim)
|
||||
energy = torch.matmul(proj_query + appr_bias, proj_key).\
|
||||
view(n, num_heads, h, w, h_kv, w_kv)
|
||||
|
||||
elif self.attention_type[0]:
|
||||
energy = torch.matmul(proj_query, proj_key).\
|
||||
view(n, num_heads, h, w, h_kv, w_kv)
|
||||
|
||||
elif self.attention_type[2]:
|
||||
appr_bias = self.appr_bias.\
|
||||
view(1, num_heads, 1, self.qk_embed_dim).\
|
||||
repeat(n, 1, 1, 1)
|
||||
|
||||
energy += torch.matmul(appr_bias, proj_key).\
|
||||
view(n, num_heads, 1, 1, h_kv, w_kv)
|
||||
|
||||
if self.attention_type[1] or self.attention_type[3]:
|
||||
if self.attention_type[1] and self.attention_type[3]:
|
||||
geom_bias = self.geom_bias.\
|
||||
view(1, num_heads, 1, self.qk_embed_dim)
|
||||
|
||||
proj_query_reshape = (proj_query + geom_bias).\
|
||||
view(n, num_heads, h, w, self.qk_embed_dim)
|
||||
|
||||
energy_x = torch.matmul(
|
||||
proj_query_reshape.permute(0, 1, 3, 2, 4),
|
||||
position_feat_x.permute(0, 1, 2, 4, 3))
|
||||
energy_x = energy_x.\
|
||||
permute(0, 1, 3, 2, 4).unsqueeze(4)
|
||||
|
||||
energy_y = torch.matmul(
|
||||
proj_query_reshape,
|
||||
position_feat_y.permute(0, 1, 2, 4, 3))
|
||||
energy_y = energy_y.unsqueeze(5)
|
||||
|
||||
energy += energy_x + energy_y
|
||||
|
||||
elif self.attention_type[1]:
|
||||
proj_query_reshape = proj_query.\
|
||||
view(n, num_heads, h, w, self.qk_embed_dim)
|
||||
proj_query_reshape = proj_query_reshape.\
|
||||
permute(0, 1, 3, 2, 4)
|
||||
position_feat_x_reshape = position_feat_x.\
|
||||
permute(0, 1, 2, 4, 3)
|
||||
position_feat_y_reshape = position_feat_y.\
|
||||
permute(0, 1, 2, 4, 3)
|
||||
|
||||
energy_x = torch.matmul(proj_query_reshape,
|
||||
position_feat_x_reshape)
|
||||
energy_x = energy_x.permute(0, 1, 3, 2, 4).unsqueeze(4)
|
||||
|
||||
energy_y = torch.matmul(proj_query_reshape,
|
||||
position_feat_y_reshape)
|
||||
energy_y = energy_y.unsqueeze(5)
|
||||
|
||||
energy += energy_x + energy_y
|
||||
|
||||
elif self.attention_type[3]:
|
||||
geom_bias = self.geom_bias.\
|
||||
view(1, num_heads, self.qk_embed_dim, 1).\
|
||||
repeat(n, 1, 1, 1)
|
||||
|
||||
position_feat_x_reshape = position_feat_x.\
|
||||
view(n, num_heads, w*w_kv, self.qk_embed_dim)
|
||||
|
||||
position_feat_y_reshape = position_feat_y.\
|
||||
view(n, num_heads, h * h_kv, self.qk_embed_dim)
|
||||
|
||||
energy_x = torch.matmul(position_feat_x_reshape, geom_bias)
|
||||
energy_x = energy_x.view(n, num_heads, 1, w, 1, w_kv)
|
||||
|
||||
energy_y = torch.matmul(position_feat_y_reshape, geom_bias)
|
||||
energy_y = energy_y.view(n, num_heads, h, 1, h_kv, 1)
|
||||
|
||||
energy += energy_x + energy_y
|
||||
|
||||
energy = energy.view(n, num_heads, h * w, h_kv * w_kv)
|
||||
|
||||
if self.spatial_range >= 0:
|
||||
cur_local_constraint_map = \
|
||||
self.local_constraint_map[:h, :w, :h_kv, :w_kv].\
|
||||
contiguous().\
|
||||
view(1, 1, h*w, h_kv*w_kv)
|
||||
|
||||
energy = energy.masked_fill_(cur_local_constraint_map,
|
||||
float('-inf'))
|
||||
|
||||
attention = F.softmax(energy, 3)
|
||||
|
||||
proj_value = self.value_conv(x_kv)
|
||||
proj_value_reshape = proj_value.\
|
||||
view((n, num_heads, self.v_dim, h_kv * w_kv)).\
|
||||
permute(0, 1, 3, 2)
|
||||
|
||||
out = torch.matmul(attention, proj_value_reshape).\
|
||||
permute(0, 1, 3, 2).\
|
||||
contiguous().\
|
||||
view(n, self.v_dim * self.num_heads, h, w)
|
||||
|
||||
out = self.proj_conv(out)
|
||||
|
||||
# output is downsampled, upsample back to input size
|
||||
if self.q_downsample is not None:
|
||||
out = F.interpolate(
|
||||
out,
|
||||
size=x_input.shape[2:],
|
||||
mode='bilinear',
|
||||
align_corners=False)
|
||||
|
||||
out = self.gamma * out + x_input
|
||||
return out
|
||||
|
||||
def init_weights(self):
|
||||
for m in self.modules():
|
||||
if hasattr(m, 'kaiming_init') and m.kaiming_init:
|
||||
kaiming_init(
|
||||
m,
|
||||
mode='fan_in',
|
||||
nonlinearity='leaky_relu',
|
||||
bias=0,
|
||||
distribution='uniform',
|
||||
a=1)
|
||||
@@ -0,0 +1,34 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch.nn as nn
|
||||
|
||||
from .registry import ACTIVATION_LAYERS
|
||||
|
||||
|
||||
@ACTIVATION_LAYERS.register_module()
|
||||
class HSigmoid(nn.Module):
|
||||
"""Hard Sigmoid Module. Apply the hard sigmoid function:
|
||||
Hsigmoid(x) = min(max((x + bias) / divisor, min_value), max_value)
|
||||
Default: Hsigmoid(x) = min(max((x + 1) / 2, 0), 1)
|
||||
|
||||
Args:
|
||||
bias (float): Bias of the input feature map. Default: 1.0.
|
||||
divisor (float): Divisor of the input feature map. Default: 2.0.
|
||||
min_value (float): Lower bound value. Default: 0.0.
|
||||
max_value (float): Upper bound value. Default: 1.0.
|
||||
|
||||
Returns:
|
||||
Tensor: The output tensor.
|
||||
"""
|
||||
|
||||
def __init__(self, bias=1.0, divisor=2.0, min_value=0.0, max_value=1.0):
|
||||
super(HSigmoid, self).__init__()
|
||||
self.bias = bias
|
||||
self.divisor = divisor
|
||||
assert self.divisor != 0
|
||||
self.min_value = min_value
|
||||
self.max_value = max_value
|
||||
|
||||
def forward(self, x):
|
||||
x = (x + self.bias) / self.divisor
|
||||
|
||||
return x.clamp_(self.min_value, self.max_value)
|
||||
@@ -0,0 +1,29 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch.nn as nn
|
||||
|
||||
from .registry import ACTIVATION_LAYERS
|
||||
|
||||
|
||||
@ACTIVATION_LAYERS.register_module()
|
||||
class HSwish(nn.Module):
|
||||
"""Hard Swish Module.
|
||||
|
||||
This module applies the hard swish function:
|
||||
|
||||
.. math::
|
||||
Hswish(x) = x * ReLU6(x + 3) / 6
|
||||
|
||||
Args:
|
||||
inplace (bool): can optionally do the operation in-place.
|
||||
Default: False.
|
||||
|
||||
Returns:
|
||||
Tensor: The output tensor.
|
||||
"""
|
||||
|
||||
def __init__(self, inplace=False):
|
||||
super(HSwish, self).__init__()
|
||||
self.act = nn.ReLU6(inplace)
|
||||
|
||||
def forward(self, x):
|
||||
return x * self.act(x + 3) / 6
|
||||
@@ -0,0 +1,306 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from abc import ABCMeta
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from ..utils import constant_init, normal_init
|
||||
from .conv_module import ConvModule
|
||||
from .registry import PLUGIN_LAYERS
|
||||
|
||||
|
||||
class _NonLocalNd(nn.Module, metaclass=ABCMeta):
|
||||
"""Basic Non-local module.
|
||||
|
||||
This module is proposed in
|
||||
"Non-local Neural Networks"
|
||||
Paper reference: https://arxiv.org/abs/1711.07971
|
||||
Code reference: https://github.com/AlexHex7/Non-local_pytorch
|
||||
|
||||
Args:
|
||||
in_channels (int): Channels of the input feature map.
|
||||
reduction (int): Channel reduction ratio. Default: 2.
|
||||
use_scale (bool): Whether to scale pairwise_weight by
|
||||
`1/sqrt(inter_channels)` when the mode is `embedded_gaussian`.
|
||||
Default: True.
|
||||
conv_cfg (None | dict): The config dict for convolution layers.
|
||||
If not specified, it will use `nn.Conv2d` for convolution layers.
|
||||
Default: None.
|
||||
norm_cfg (None | dict): The config dict for normalization layers.
|
||||
Default: None. (This parameter is only applicable to conv_out.)
|
||||
mode (str): Options are `gaussian`, `concatenation`,
|
||||
`embedded_gaussian` and `dot_product`. Default: embedded_gaussian.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
reduction=2,
|
||||
use_scale=True,
|
||||
conv_cfg=None,
|
||||
norm_cfg=None,
|
||||
mode='embedded_gaussian',
|
||||
**kwargs):
|
||||
super(_NonLocalNd, self).__init__()
|
||||
self.in_channels = in_channels
|
||||
self.reduction = reduction
|
||||
self.use_scale = use_scale
|
||||
self.inter_channels = max(in_channels // reduction, 1)
|
||||
self.mode = mode
|
||||
|
||||
if mode not in [
|
||||
'gaussian', 'embedded_gaussian', 'dot_product', 'concatenation'
|
||||
]:
|
||||
raise ValueError("Mode should be in 'gaussian', 'concatenation', "
|
||||
f"'embedded_gaussian' or 'dot_product', but got "
|
||||
f'{mode} instead.')
|
||||
|
||||
# g, theta, phi are defaulted as `nn.ConvNd`.
|
||||
# Here we use ConvModule for potential usage.
|
||||
self.g = ConvModule(
|
||||
self.in_channels,
|
||||
self.inter_channels,
|
||||
kernel_size=1,
|
||||
conv_cfg=conv_cfg,
|
||||
act_cfg=None)
|
||||
self.conv_out = ConvModule(
|
||||
self.inter_channels,
|
||||
self.in_channels,
|
||||
kernel_size=1,
|
||||
conv_cfg=conv_cfg,
|
||||
norm_cfg=norm_cfg,
|
||||
act_cfg=None)
|
||||
|
||||
if self.mode != 'gaussian':
|
||||
self.theta = ConvModule(
|
||||
self.in_channels,
|
||||
self.inter_channels,
|
||||
kernel_size=1,
|
||||
conv_cfg=conv_cfg,
|
||||
act_cfg=None)
|
||||
self.phi = ConvModule(
|
||||
self.in_channels,
|
||||
self.inter_channels,
|
||||
kernel_size=1,
|
||||
conv_cfg=conv_cfg,
|
||||
act_cfg=None)
|
||||
|
||||
if self.mode == 'concatenation':
|
||||
self.concat_project = ConvModule(
|
||||
self.inter_channels * 2,
|
||||
1,
|
||||
kernel_size=1,
|
||||
stride=1,
|
||||
padding=0,
|
||||
bias=False,
|
||||
act_cfg=dict(type='ReLU'))
|
||||
|
||||
self.init_weights(**kwargs)
|
||||
|
||||
def init_weights(self, std=0.01, zeros_init=True):
|
||||
if self.mode != 'gaussian':
|
||||
for m in [self.g, self.theta, self.phi]:
|
||||
normal_init(m.conv, std=std)
|
||||
else:
|
||||
normal_init(self.g.conv, std=std)
|
||||
if zeros_init:
|
||||
if self.conv_out.norm_cfg is None:
|
||||
constant_init(self.conv_out.conv, 0)
|
||||
else:
|
||||
constant_init(self.conv_out.norm, 0)
|
||||
else:
|
||||
if self.conv_out.norm_cfg is None:
|
||||
normal_init(self.conv_out.conv, std=std)
|
||||
else:
|
||||
normal_init(self.conv_out.norm, std=std)
|
||||
|
||||
def gaussian(self, theta_x, phi_x):
|
||||
# NonLocal1d pairwise_weight: [N, H, H]
|
||||
# NonLocal2d pairwise_weight: [N, HxW, HxW]
|
||||
# NonLocal3d pairwise_weight: [N, TxHxW, TxHxW]
|
||||
pairwise_weight = torch.matmul(theta_x, phi_x)
|
||||
pairwise_weight = pairwise_weight.softmax(dim=-1)
|
||||
return pairwise_weight
|
||||
|
||||
def embedded_gaussian(self, theta_x, phi_x):
|
||||
# NonLocal1d pairwise_weight: [N, H, H]
|
||||
# NonLocal2d pairwise_weight: [N, HxW, HxW]
|
||||
# NonLocal3d pairwise_weight: [N, TxHxW, TxHxW]
|
||||
pairwise_weight = torch.matmul(theta_x, phi_x)
|
||||
if self.use_scale:
|
||||
# theta_x.shape[-1] is `self.inter_channels`
|
||||
pairwise_weight /= theta_x.shape[-1]**0.5
|
||||
pairwise_weight = pairwise_weight.softmax(dim=-1)
|
||||
return pairwise_weight
|
||||
|
||||
def dot_product(self, theta_x, phi_x):
|
||||
# NonLocal1d pairwise_weight: [N, H, H]
|
||||
# NonLocal2d pairwise_weight: [N, HxW, HxW]
|
||||
# NonLocal3d pairwise_weight: [N, TxHxW, TxHxW]
|
||||
pairwise_weight = torch.matmul(theta_x, phi_x)
|
||||
pairwise_weight /= pairwise_weight.shape[-1]
|
||||
return pairwise_weight
|
||||
|
||||
def concatenation(self, theta_x, phi_x):
|
||||
# NonLocal1d pairwise_weight: [N, H, H]
|
||||
# NonLocal2d pairwise_weight: [N, HxW, HxW]
|
||||
# NonLocal3d pairwise_weight: [N, TxHxW, TxHxW]
|
||||
h = theta_x.size(2)
|
||||
w = phi_x.size(3)
|
||||
theta_x = theta_x.repeat(1, 1, 1, w)
|
||||
phi_x = phi_x.repeat(1, 1, h, 1)
|
||||
|
||||
concat_feature = torch.cat([theta_x, phi_x], dim=1)
|
||||
pairwise_weight = self.concat_project(concat_feature)
|
||||
n, _, h, w = pairwise_weight.size()
|
||||
pairwise_weight = pairwise_weight.view(n, h, w)
|
||||
pairwise_weight /= pairwise_weight.shape[-1]
|
||||
|
||||
return pairwise_weight
|
||||
|
||||
def forward(self, x):
|
||||
# Assume `reduction = 1`, then `inter_channels = C`
|
||||
# or `inter_channels = C` when `mode="gaussian"`
|
||||
|
||||
# NonLocal1d x: [N, C, H]
|
||||
# NonLocal2d x: [N, C, H, W]
|
||||
# NonLocal3d x: [N, C, T, H, W]
|
||||
n = x.size(0)
|
||||
|
||||
# NonLocal1d g_x: [N, H, C]
|
||||
# NonLocal2d g_x: [N, HxW, C]
|
||||
# NonLocal3d g_x: [N, TxHxW, C]
|
||||
g_x = self.g(x).view(n, self.inter_channels, -1)
|
||||
g_x = g_x.permute(0, 2, 1)
|
||||
|
||||
# NonLocal1d theta_x: [N, H, C], phi_x: [N, C, H]
|
||||
# NonLocal2d theta_x: [N, HxW, C], phi_x: [N, C, HxW]
|
||||
# NonLocal3d theta_x: [N, TxHxW, C], phi_x: [N, C, TxHxW]
|
||||
if self.mode == 'gaussian':
|
||||
theta_x = x.view(n, self.in_channels, -1)
|
||||
theta_x = theta_x.permute(0, 2, 1)
|
||||
if self.sub_sample:
|
||||
phi_x = self.phi(x).view(n, self.in_channels, -1)
|
||||
else:
|
||||
phi_x = x.view(n, self.in_channels, -1)
|
||||
elif self.mode == 'concatenation':
|
||||
theta_x = self.theta(x).view(n, self.inter_channels, -1, 1)
|
||||
phi_x = self.phi(x).view(n, self.inter_channels, 1, -1)
|
||||
else:
|
||||
theta_x = self.theta(x).view(n, self.inter_channels, -1)
|
||||
theta_x = theta_x.permute(0, 2, 1)
|
||||
phi_x = self.phi(x).view(n, self.inter_channels, -1)
|
||||
|
||||
pairwise_func = getattr(self, self.mode)
|
||||
# NonLocal1d pairwise_weight: [N, H, H]
|
||||
# NonLocal2d pairwise_weight: [N, HxW, HxW]
|
||||
# NonLocal3d pairwise_weight: [N, TxHxW, TxHxW]
|
||||
pairwise_weight = pairwise_func(theta_x, phi_x)
|
||||
|
||||
# NonLocal1d y: [N, H, C]
|
||||
# NonLocal2d y: [N, HxW, C]
|
||||
# NonLocal3d y: [N, TxHxW, C]
|
||||
y = torch.matmul(pairwise_weight, g_x)
|
||||
# NonLocal1d y: [N, C, H]
|
||||
# NonLocal2d y: [N, C, H, W]
|
||||
# NonLocal3d y: [N, C, T, H, W]
|
||||
y = y.permute(0, 2, 1).contiguous().reshape(n, self.inter_channels,
|
||||
*x.size()[2:])
|
||||
|
||||
output = x + self.conv_out(y)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class NonLocal1d(_NonLocalNd):
|
||||
"""1D Non-local module.
|
||||
|
||||
Args:
|
||||
in_channels (int): Same as `NonLocalND`.
|
||||
sub_sample (bool): Whether to apply max pooling after pairwise
|
||||
function (Note that the `sub_sample` is applied on spatial only).
|
||||
Default: False.
|
||||
conv_cfg (None | dict): Same as `NonLocalND`.
|
||||
Default: dict(type='Conv1d').
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
sub_sample=False,
|
||||
conv_cfg=dict(type='Conv1d'),
|
||||
**kwargs):
|
||||
super(NonLocal1d, self).__init__(
|
||||
in_channels, conv_cfg=conv_cfg, **kwargs)
|
||||
|
||||
self.sub_sample = sub_sample
|
||||
|
||||
if sub_sample:
|
||||
max_pool_layer = nn.MaxPool1d(kernel_size=2)
|
||||
self.g = nn.Sequential(self.g, max_pool_layer)
|
||||
if self.mode != 'gaussian':
|
||||
self.phi = nn.Sequential(self.phi, max_pool_layer)
|
||||
else:
|
||||
self.phi = max_pool_layer
|
||||
|
||||
|
||||
@PLUGIN_LAYERS.register_module()
|
||||
class NonLocal2d(_NonLocalNd):
|
||||
"""2D Non-local module.
|
||||
|
||||
Args:
|
||||
in_channels (int): Same as `NonLocalND`.
|
||||
sub_sample (bool): Whether to apply max pooling after pairwise
|
||||
function (Note that the `sub_sample` is applied on spatial only).
|
||||
Default: False.
|
||||
conv_cfg (None | dict): Same as `NonLocalND`.
|
||||
Default: dict(type='Conv2d').
|
||||
"""
|
||||
|
||||
_abbr_ = 'nonlocal_block'
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
sub_sample=False,
|
||||
conv_cfg=dict(type='Conv2d'),
|
||||
**kwargs):
|
||||
super(NonLocal2d, self).__init__(
|
||||
in_channels, conv_cfg=conv_cfg, **kwargs)
|
||||
|
||||
self.sub_sample = sub_sample
|
||||
|
||||
if sub_sample:
|
||||
max_pool_layer = nn.MaxPool2d(kernel_size=(2, 2))
|
||||
self.g = nn.Sequential(self.g, max_pool_layer)
|
||||
if self.mode != 'gaussian':
|
||||
self.phi = nn.Sequential(self.phi, max_pool_layer)
|
||||
else:
|
||||
self.phi = max_pool_layer
|
||||
|
||||
|
||||
class NonLocal3d(_NonLocalNd):
|
||||
"""3D Non-local module.
|
||||
|
||||
Args:
|
||||
in_channels (int): Same as `NonLocalND`.
|
||||
sub_sample (bool): Whether to apply max pooling after pairwise
|
||||
function (Note that the `sub_sample` is applied on spatial only).
|
||||
Default: False.
|
||||
conv_cfg (None | dict): Same as `NonLocalND`.
|
||||
Default: dict(type='Conv3d').
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
sub_sample=False,
|
||||
conv_cfg=dict(type='Conv3d'),
|
||||
**kwargs):
|
||||
super(NonLocal3d, self).__init__(
|
||||
in_channels, conv_cfg=conv_cfg, **kwargs)
|
||||
self.sub_sample = sub_sample
|
||||
|
||||
if sub_sample:
|
||||
max_pool_layer = nn.MaxPool3d(kernel_size=(1, 2, 2))
|
||||
self.g = nn.Sequential(self.g, max_pool_layer)
|
||||
if self.mode != 'gaussian':
|
||||
self.phi = nn.Sequential(self.phi, max_pool_layer)
|
||||
else:
|
||||
self.phi = max_pool_layer
|
||||
@@ -0,0 +1,144 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import inspect
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import is_tuple_of
|
||||
from custom_mmpkg.custom_mmcv.utils.parrots_wrapper import SyncBatchNorm, _BatchNorm, _InstanceNorm
|
||||
from .registry import NORM_LAYERS
|
||||
|
||||
NORM_LAYERS.register_module('BN', module=nn.BatchNorm2d)
|
||||
NORM_LAYERS.register_module('BN1d', module=nn.BatchNorm1d)
|
||||
NORM_LAYERS.register_module('BN2d', module=nn.BatchNorm2d)
|
||||
NORM_LAYERS.register_module('BN3d', module=nn.BatchNorm3d)
|
||||
NORM_LAYERS.register_module('SyncBN', module=SyncBatchNorm)
|
||||
NORM_LAYERS.register_module('GN', module=nn.GroupNorm)
|
||||
NORM_LAYERS.register_module('LN', module=nn.LayerNorm)
|
||||
NORM_LAYERS.register_module('IN', module=nn.InstanceNorm2d)
|
||||
NORM_LAYERS.register_module('IN1d', module=nn.InstanceNorm1d)
|
||||
NORM_LAYERS.register_module('IN2d', module=nn.InstanceNorm2d)
|
||||
NORM_LAYERS.register_module('IN3d', module=nn.InstanceNorm3d)
|
||||
|
||||
|
||||
def infer_abbr(class_type):
|
||||
"""Infer abbreviation from the class name.
|
||||
|
||||
When we build a norm layer with `build_norm_layer()`, we want to preserve
|
||||
the norm type in variable names, e.g, self.bn1, self.gn. This method will
|
||||
infer the abbreviation to map class types to abbreviations.
|
||||
|
||||
Rule 1: If the class has the property "_abbr_", return the property.
|
||||
Rule 2: If the parent class is _BatchNorm, GroupNorm, LayerNorm or
|
||||
InstanceNorm, the abbreviation of this layer will be "bn", "gn", "ln" and
|
||||
"in" respectively.
|
||||
Rule 3: If the class name contains "batch", "group", "layer" or "instance",
|
||||
the abbreviation of this layer will be "bn", "gn", "ln" and "in"
|
||||
respectively.
|
||||
Rule 4: Otherwise, the abbreviation falls back to "norm".
|
||||
|
||||
Args:
|
||||
class_type (type): The norm layer type.
|
||||
|
||||
Returns:
|
||||
str: The inferred abbreviation.
|
||||
"""
|
||||
if not inspect.isclass(class_type):
|
||||
raise TypeError(
|
||||
f'class_type must be a type, but got {type(class_type)}')
|
||||
if hasattr(class_type, '_abbr_'):
|
||||
return class_type._abbr_
|
||||
if issubclass(class_type, _InstanceNorm): # IN is a subclass of BN
|
||||
return 'in'
|
||||
elif issubclass(class_type, _BatchNorm):
|
||||
return 'bn'
|
||||
elif issubclass(class_type, nn.GroupNorm):
|
||||
return 'gn'
|
||||
elif issubclass(class_type, nn.LayerNorm):
|
||||
return 'ln'
|
||||
else:
|
||||
class_name = class_type.__name__.lower()
|
||||
if 'batch' in class_name:
|
||||
return 'bn'
|
||||
elif 'group' in class_name:
|
||||
return 'gn'
|
||||
elif 'layer' in class_name:
|
||||
return 'ln'
|
||||
elif 'instance' in class_name:
|
||||
return 'in'
|
||||
else:
|
||||
return 'norm_layer'
|
||||
|
||||
|
||||
def build_norm_layer(cfg, num_features, postfix=''):
|
||||
"""Build normalization layer.
|
||||
|
||||
Args:
|
||||
cfg (dict): The norm layer config, which should contain:
|
||||
|
||||
- type (str): Layer type.
|
||||
- layer args: Args needed to instantiate a norm layer.
|
||||
- requires_grad (bool, optional): Whether stop gradient updates.
|
||||
num_features (int): Number of input channels.
|
||||
postfix (int | str): The postfix to be appended into norm abbreviation
|
||||
to create named layer.
|
||||
|
||||
Returns:
|
||||
(str, nn.Module): The first element is the layer name consisting of
|
||||
abbreviation and postfix, e.g., bn1, gn. The second element is the
|
||||
created norm layer.
|
||||
"""
|
||||
if not isinstance(cfg, dict):
|
||||
raise TypeError('cfg must be a dict')
|
||||
if 'type' not in cfg:
|
||||
raise KeyError('the cfg dict must contain the key "type"')
|
||||
cfg_ = cfg.copy()
|
||||
|
||||
layer_type = cfg_.pop('type')
|
||||
if layer_type not in NORM_LAYERS:
|
||||
raise KeyError(f'Unrecognized norm type {layer_type}')
|
||||
|
||||
norm_layer = NORM_LAYERS.get(layer_type)
|
||||
abbr = infer_abbr(norm_layer)
|
||||
|
||||
assert isinstance(postfix, (int, str))
|
||||
name = abbr + str(postfix)
|
||||
|
||||
requires_grad = cfg_.pop('requires_grad', True)
|
||||
cfg_.setdefault('eps', 1e-5)
|
||||
if layer_type != 'GN':
|
||||
layer = norm_layer(num_features, **cfg_)
|
||||
if layer_type == 'SyncBN' and hasattr(layer, '_specify_ddp_gpu_num'):
|
||||
layer._specify_ddp_gpu_num(1)
|
||||
else:
|
||||
assert 'num_groups' in cfg_
|
||||
layer = norm_layer(num_channels=num_features, **cfg_)
|
||||
|
||||
for param in layer.parameters():
|
||||
param.requires_grad = requires_grad
|
||||
|
||||
return name, layer
|
||||
|
||||
|
||||
def is_norm(layer, exclude=None):
|
||||
"""Check if a layer is a normalization layer.
|
||||
|
||||
Args:
|
||||
layer (nn.Module): The layer to be checked.
|
||||
exclude (type | tuple[type]): Types to be excluded.
|
||||
|
||||
Returns:
|
||||
bool: Whether the layer is a norm layer.
|
||||
"""
|
||||
if exclude is not None:
|
||||
if not isinstance(exclude, tuple):
|
||||
exclude = (exclude, )
|
||||
if not is_tuple_of(exclude, type):
|
||||
raise TypeError(
|
||||
f'"exclude" must be either None or type or a tuple of types, '
|
||||
f'but got {type(exclude)}: {exclude}')
|
||||
|
||||
if exclude and isinstance(layer, exclude):
|
||||
return False
|
||||
|
||||
all_norm_bases = (_BatchNorm, _InstanceNorm, nn.GroupNorm, nn.LayerNorm)
|
||||
return isinstance(layer, all_norm_bases)
|
||||
@@ -0,0 +1,36 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch.nn as nn
|
||||
|
||||
from .registry import PADDING_LAYERS
|
||||
|
||||
PADDING_LAYERS.register_module('zero', module=nn.ZeroPad2d)
|
||||
PADDING_LAYERS.register_module('reflect', module=nn.ReflectionPad2d)
|
||||
PADDING_LAYERS.register_module('replicate', module=nn.ReplicationPad2d)
|
||||
|
||||
|
||||
def build_padding_layer(cfg, *args, **kwargs):
|
||||
"""Build padding layer.
|
||||
|
||||
Args:
|
||||
cfg (None or dict): The padding layer config, which should contain:
|
||||
- type (str): Layer type.
|
||||
- layer args: Args needed to instantiate a padding layer.
|
||||
|
||||
Returns:
|
||||
nn.Module: Created padding layer.
|
||||
"""
|
||||
if not isinstance(cfg, dict):
|
||||
raise TypeError('cfg must be a dict')
|
||||
if 'type' not in cfg:
|
||||
raise KeyError('the cfg dict must contain the key "type"')
|
||||
|
||||
cfg_ = cfg.copy()
|
||||
padding_type = cfg_.pop('type')
|
||||
if padding_type not in PADDING_LAYERS:
|
||||
raise KeyError(f'Unrecognized padding type {padding_type}.')
|
||||
else:
|
||||
padding_layer = PADDING_LAYERS.get(padding_type)
|
||||
|
||||
layer = padding_layer(*args, **kwargs, **cfg_)
|
||||
|
||||
return layer
|
||||
@@ -0,0 +1,88 @@
|
||||
import inspect
|
||||
import platform
|
||||
|
||||
from .registry import PLUGIN_LAYERS
|
||||
|
||||
if platform.system() == 'Windows':
|
||||
import regex as re
|
||||
else:
|
||||
import re
|
||||
|
||||
|
||||
def infer_abbr(class_type):
|
||||
"""Infer abbreviation from the class name.
|
||||
|
||||
This method will infer the abbreviation to map class types to
|
||||
abbreviations.
|
||||
|
||||
Rule 1: If the class has the property "abbr", return the property.
|
||||
Rule 2: Otherwise, the abbreviation falls back to snake case of class
|
||||
name, e.g. the abbreviation of ``FancyBlock`` will be ``fancy_block``.
|
||||
|
||||
Args:
|
||||
class_type (type): The norm layer type.
|
||||
|
||||
Returns:
|
||||
str: The inferred abbreviation.
|
||||
"""
|
||||
|
||||
def camel2snack(word):
|
||||
"""Convert camel case word into snack case.
|
||||
|
||||
Modified from `inflection lib
|
||||
<https://inflection.readthedocs.io/en/latest/#inflection.underscore>`_.
|
||||
|
||||
Example::
|
||||
|
||||
>>> camel2snack("FancyBlock")
|
||||
'fancy_block'
|
||||
"""
|
||||
|
||||
word = re.sub(r'([A-Z]+)([A-Z][a-z])', r'\1_\2', word)
|
||||
word = re.sub(r'([a-z\d])([A-Z])', r'\1_\2', word)
|
||||
word = word.replace('-', '_')
|
||||
return word.lower()
|
||||
|
||||
if not inspect.isclass(class_type):
|
||||
raise TypeError(
|
||||
f'class_type must be a type, but got {type(class_type)}')
|
||||
if hasattr(class_type, '_abbr_'):
|
||||
return class_type._abbr_
|
||||
else:
|
||||
return camel2snack(class_type.__name__)
|
||||
|
||||
|
||||
def build_plugin_layer(cfg, postfix='', **kwargs):
|
||||
"""Build plugin layer.
|
||||
|
||||
Args:
|
||||
cfg (None or dict): cfg should contain:
|
||||
type (str): identify plugin layer type.
|
||||
layer args: args needed to instantiate a plugin layer.
|
||||
postfix (int, str): appended into norm abbreviation to
|
||||
create named layer. Default: ''.
|
||||
|
||||
Returns:
|
||||
tuple[str, nn.Module]:
|
||||
name (str): abbreviation + postfix
|
||||
layer (nn.Module): created plugin layer
|
||||
"""
|
||||
if not isinstance(cfg, dict):
|
||||
raise TypeError('cfg must be a dict')
|
||||
if 'type' not in cfg:
|
||||
raise KeyError('the cfg dict must contain the key "type"')
|
||||
cfg_ = cfg.copy()
|
||||
|
||||
layer_type = cfg_.pop('type')
|
||||
if layer_type not in PLUGIN_LAYERS:
|
||||
raise KeyError(f'Unrecognized plugin type {layer_type}')
|
||||
|
||||
plugin_layer = PLUGIN_LAYERS.get(layer_type)
|
||||
abbr = infer_abbr(plugin_layer)
|
||||
|
||||
assert isinstance(postfix, (int, str))
|
||||
name = abbr + str(postfix)
|
||||
|
||||
layer = plugin_layer(**kwargs, **cfg_)
|
||||
|
||||
return name, layer
|
||||
@@ -0,0 +1,16 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from custom_mmpkg.custom_mmcv.utils import Registry
|
||||
|
||||
CONV_LAYERS = Registry('conv layer')
|
||||
NORM_LAYERS = Registry('norm layer')
|
||||
ACTIVATION_LAYERS = Registry('activation layer')
|
||||
PADDING_LAYERS = Registry('padding layer')
|
||||
UPSAMPLE_LAYERS = Registry('upsample layer')
|
||||
PLUGIN_LAYERS = Registry('plugin layer')
|
||||
|
||||
DROPOUT_LAYERS = Registry('drop out layers')
|
||||
POSITIONAL_ENCODING = Registry('position encoding')
|
||||
ATTENTION = Registry('attention')
|
||||
FEEDFORWARD_NETWORK = Registry('feed-forward Network')
|
||||
TRANSFORMER_LAYER = Registry('transformerLayer')
|
||||
TRANSFORMER_LAYER_SEQUENCE = Registry('transformer-layers sequence')
|
||||
@@ -0,0 +1,21 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class Scale(nn.Module):
|
||||
"""A learnable scale parameter.
|
||||
|
||||
This layer scales the input by a learnable factor. It multiplies a
|
||||
learnable scale parameter of shape (1,) with input of any shape.
|
||||
|
||||
Args:
|
||||
scale (float): Initial value of scale factor. Default: 1.0
|
||||
"""
|
||||
|
||||
def __init__(self, scale=1.0):
|
||||
super(Scale, self).__init__()
|
||||
self.scale = nn.Parameter(torch.tensor(scale, dtype=torch.float))
|
||||
|
||||
def forward(self, x):
|
||||
return x * self.scale
|
||||
@@ -0,0 +1,25 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from .registry import ACTIVATION_LAYERS
|
||||
|
||||
|
||||
@ACTIVATION_LAYERS.register_module()
|
||||
class Swish(nn.Module):
|
||||
"""Swish Module.
|
||||
|
||||
This module applies the swish function:
|
||||
|
||||
.. math::
|
||||
Swish(x) = x * Sigmoid(x)
|
||||
|
||||
Returns:
|
||||
Tensor: The output tensor.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super(Swish, self).__init__()
|
||||
|
||||
def forward(self, x):
|
||||
return x * torch.sigmoid(x)
|
||||
@@ -0,0 +1,595 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import copy
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from custom_mmpkg.custom_mmcv import ConfigDict, deprecated_api_warning
|
||||
from custom_mmpkg.custom_mmcv.cnn import Linear, build_activation_layer, build_norm_layer
|
||||
from custom_mmpkg.custom_mmcv.runner.base_module import BaseModule, ModuleList, Sequential
|
||||
from custom_mmpkg.custom_mmcv.utils import build_from_cfg
|
||||
from .drop import build_dropout
|
||||
from .registry import (ATTENTION, FEEDFORWARD_NETWORK, POSITIONAL_ENCODING,
|
||||
TRANSFORMER_LAYER, TRANSFORMER_LAYER_SEQUENCE)
|
||||
|
||||
# Avoid BC-breaking of importing MultiScaleDeformableAttention from this file
|
||||
try:
|
||||
from custom_mmpkg.custom_mmcv.ops.multi_scale_deform_attn import MultiScaleDeformableAttention # noqa F401
|
||||
warnings.warn(
|
||||
ImportWarning(
|
||||
'``MultiScaleDeformableAttention`` has been moved to '
|
||||
'``mmcv.ops.multi_scale_deform_attn``, please change original path ' # noqa E501
|
||||
'``from custom_mmpkg.custom_mmcv.cnn.bricks.transformer import MultiScaleDeformableAttention`` ' # noqa E501
|
||||
'to ``from custom_mmpkg.custom_mmcv.ops.multi_scale_deform_attn import MultiScaleDeformableAttention`` ' # noqa E501
|
||||
))
|
||||
|
||||
except ImportError:
|
||||
warnings.warn('Fail to import ``MultiScaleDeformableAttention`` from '
|
||||
'``mmcv.ops.multi_scale_deform_attn``, '
|
||||
'You should install ``mmcv-full`` if you need this module. ')
|
||||
|
||||
|
||||
def build_positional_encoding(cfg, default_args=None):
|
||||
"""Builder for Position Encoding."""
|
||||
return build_from_cfg(cfg, POSITIONAL_ENCODING, default_args)
|
||||
|
||||
|
||||
def build_attention(cfg, default_args=None):
|
||||
"""Builder for attention."""
|
||||
return build_from_cfg(cfg, ATTENTION, default_args)
|
||||
|
||||
|
||||
def build_feedforward_network(cfg, default_args=None):
|
||||
"""Builder for feed-forward network (FFN)."""
|
||||
return build_from_cfg(cfg, FEEDFORWARD_NETWORK, default_args)
|
||||
|
||||
|
||||
def build_transformer_layer(cfg, default_args=None):
|
||||
"""Builder for transformer layer."""
|
||||
return build_from_cfg(cfg, TRANSFORMER_LAYER, default_args)
|
||||
|
||||
|
||||
def build_transformer_layer_sequence(cfg, default_args=None):
|
||||
"""Builder for transformer encoder and transformer decoder."""
|
||||
return build_from_cfg(cfg, TRANSFORMER_LAYER_SEQUENCE, default_args)
|
||||
|
||||
|
||||
@ATTENTION.register_module()
|
||||
class MultiheadAttention(BaseModule):
|
||||
"""A wrapper for ``torch.nn.MultiheadAttention``.
|
||||
|
||||
This module implements MultiheadAttention with identity connection,
|
||||
and positional encoding is also passed as input.
|
||||
|
||||
Args:
|
||||
embed_dims (int): The embedding dimension.
|
||||
num_heads (int): Parallel attention heads.
|
||||
attn_drop (float): A Dropout layer on attn_output_weights.
|
||||
Default: 0.0.
|
||||
proj_drop (float): A Dropout layer after `nn.MultiheadAttention`.
|
||||
Default: 0.0.
|
||||
dropout_layer (obj:`ConfigDict`): The dropout_layer used
|
||||
when adding the shortcut.
|
||||
init_cfg (obj:`mmcv.ConfigDict`): The Config for initialization.
|
||||
Default: None.
|
||||
batch_first (bool): When it is True, Key, Query and Value are shape of
|
||||
(batch, n, embed_dim), otherwise (n, batch, embed_dim).
|
||||
Default to False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
embed_dims,
|
||||
num_heads,
|
||||
attn_drop=0.,
|
||||
proj_drop=0.,
|
||||
dropout_layer=dict(type='Dropout', drop_prob=0.),
|
||||
init_cfg=None,
|
||||
batch_first=False,
|
||||
**kwargs):
|
||||
super(MultiheadAttention, self).__init__(init_cfg)
|
||||
if 'dropout' in kwargs:
|
||||
warnings.warn('The arguments `dropout` in MultiheadAttention '
|
||||
'has been deprecated, now you can separately '
|
||||
'set `attn_drop`(float), proj_drop(float), '
|
||||
'and `dropout_layer`(dict) ')
|
||||
attn_drop = kwargs['dropout']
|
||||
dropout_layer['drop_prob'] = kwargs.pop('dropout')
|
||||
|
||||
self.embed_dims = embed_dims
|
||||
self.num_heads = num_heads
|
||||
self.batch_first = batch_first
|
||||
|
||||
self.attn = nn.MultiheadAttention(embed_dims, num_heads, attn_drop,
|
||||
**kwargs)
|
||||
|
||||
self.proj_drop = nn.Dropout(proj_drop)
|
||||
self.dropout_layer = build_dropout(
|
||||
dropout_layer) if dropout_layer else nn.Identity()
|
||||
|
||||
@deprecated_api_warning({'residual': 'identity'},
|
||||
cls_name='MultiheadAttention')
|
||||
def forward(self,
|
||||
query,
|
||||
key=None,
|
||||
value=None,
|
||||
identity=None,
|
||||
query_pos=None,
|
||||
key_pos=None,
|
||||
attn_mask=None,
|
||||
key_padding_mask=None,
|
||||
**kwargs):
|
||||
"""Forward function for `MultiheadAttention`.
|
||||
|
||||
**kwargs allow passing a more general data flow when combining
|
||||
with other operations in `transformerlayer`.
|
||||
|
||||
Args:
|
||||
query (Tensor): The input query with shape [num_queries, bs,
|
||||
embed_dims] if self.batch_first is False, else
|
||||
[bs, num_queries embed_dims].
|
||||
key (Tensor): The key tensor with shape [num_keys, bs,
|
||||
embed_dims] if self.batch_first is False, else
|
||||
[bs, num_keys, embed_dims] .
|
||||
If None, the ``query`` will be used. Defaults to None.
|
||||
value (Tensor): The value tensor with same shape as `key`.
|
||||
Same in `nn.MultiheadAttention.forward`. Defaults to None.
|
||||
If None, the `key` will be used.
|
||||
identity (Tensor): This tensor, with the same shape as x,
|
||||
will be used for the identity link.
|
||||
If None, `x` will be used. Defaults to None.
|
||||
query_pos (Tensor): The positional encoding for query, with
|
||||
the same shape as `x`. If not None, it will
|
||||
be added to `x` before forward function. Defaults to None.
|
||||
key_pos (Tensor): The positional encoding for `key`, with the
|
||||
same shape as `key`. Defaults to None. If not None, it will
|
||||
be added to `key` before forward function. If None, and
|
||||
`query_pos` has the same shape as `key`, then `query_pos`
|
||||
will be used for `key_pos`. Defaults to None.
|
||||
attn_mask (Tensor): ByteTensor mask with shape [num_queries,
|
||||
num_keys]. Same in `nn.MultiheadAttention.forward`.
|
||||
Defaults to None.
|
||||
key_padding_mask (Tensor): ByteTensor with shape [bs, num_keys].
|
||||
Defaults to None.
|
||||
|
||||
Returns:
|
||||
Tensor: forwarded results with shape
|
||||
[num_queries, bs, embed_dims]
|
||||
if self.batch_first is False, else
|
||||
[bs, num_queries embed_dims].
|
||||
"""
|
||||
|
||||
if key is None:
|
||||
key = query
|
||||
if value is None:
|
||||
value = key
|
||||
if identity is None:
|
||||
identity = query
|
||||
if key_pos is None:
|
||||
if query_pos is not None:
|
||||
# use query_pos if key_pos is not available
|
||||
if query_pos.shape == key.shape:
|
||||
key_pos = query_pos
|
||||
else:
|
||||
warnings.warn(f'position encoding of key is'
|
||||
f'missing in {self.__class__.__name__}.')
|
||||
if query_pos is not None:
|
||||
query = query + query_pos
|
||||
if key_pos is not None:
|
||||
key = key + key_pos
|
||||
|
||||
# Because the dataflow('key', 'query', 'value') of
|
||||
# ``torch.nn.MultiheadAttention`` is (num_query, batch,
|
||||
# embed_dims), We should adjust the shape of dataflow from
|
||||
# batch_first (batch, num_query, embed_dims) to num_query_first
|
||||
# (num_query ,batch, embed_dims), and recover ``attn_output``
|
||||
# from num_query_first to batch_first.
|
||||
if self.batch_first:
|
||||
query = query.transpose(0, 1)
|
||||
key = key.transpose(0, 1)
|
||||
value = value.transpose(0, 1)
|
||||
|
||||
out = self.attn(
|
||||
query=query,
|
||||
key=key,
|
||||
value=value,
|
||||
attn_mask=attn_mask,
|
||||
key_padding_mask=key_padding_mask)[0]
|
||||
|
||||
if self.batch_first:
|
||||
out = out.transpose(0, 1)
|
||||
|
||||
return identity + self.dropout_layer(self.proj_drop(out))
|
||||
|
||||
|
||||
@FEEDFORWARD_NETWORK.register_module()
|
||||
class FFN(BaseModule):
|
||||
"""Implements feed-forward networks (FFNs) with identity connection.
|
||||
|
||||
Args:
|
||||
embed_dims (int): The feature dimension. Same as
|
||||
`MultiheadAttention`. Defaults: 256.
|
||||
feedforward_channels (int): The hidden dimension of FFNs.
|
||||
Defaults: 1024.
|
||||
num_fcs (int, optional): The number of fully-connected layers in
|
||||
FFNs. Default: 2.
|
||||
act_cfg (dict, optional): The activation config for FFNs.
|
||||
Default: dict(type='ReLU')
|
||||
ffn_drop (float, optional): Probability of an element to be
|
||||
zeroed in FFN. Default 0.0.
|
||||
add_identity (bool, optional): Whether to add the
|
||||
identity connection. Default: `True`.
|
||||
dropout_layer (obj:`ConfigDict`): The dropout_layer used
|
||||
when adding the shortcut.
|
||||
init_cfg (obj:`mmcv.ConfigDict`): The Config for initialization.
|
||||
Default: None.
|
||||
"""
|
||||
|
||||
@deprecated_api_warning(
|
||||
{
|
||||
'dropout': 'ffn_drop',
|
||||
'add_residual': 'add_identity'
|
||||
},
|
||||
cls_name='FFN')
|
||||
def __init__(self,
|
||||
embed_dims=256,
|
||||
feedforward_channels=1024,
|
||||
num_fcs=2,
|
||||
act_cfg=dict(type='ReLU', inplace=True),
|
||||
ffn_drop=0.,
|
||||
dropout_layer=None,
|
||||
add_identity=True,
|
||||
init_cfg=None,
|
||||
**kwargs):
|
||||
super(FFN, self).__init__(init_cfg)
|
||||
assert num_fcs >= 2, 'num_fcs should be no less ' \
|
||||
f'than 2. got {num_fcs}.'
|
||||
self.embed_dims = embed_dims
|
||||
self.feedforward_channels = feedforward_channels
|
||||
self.num_fcs = num_fcs
|
||||
self.act_cfg = act_cfg
|
||||
self.activate = build_activation_layer(act_cfg)
|
||||
|
||||
layers = []
|
||||
in_channels = embed_dims
|
||||
for _ in range(num_fcs - 1):
|
||||
layers.append(
|
||||
Sequential(
|
||||
Linear(in_channels, feedforward_channels), self.activate,
|
||||
nn.Dropout(ffn_drop)))
|
||||
in_channels = feedforward_channels
|
||||
layers.append(Linear(feedforward_channels, embed_dims))
|
||||
layers.append(nn.Dropout(ffn_drop))
|
||||
self.layers = Sequential(*layers)
|
||||
self.dropout_layer = build_dropout(
|
||||
dropout_layer) if dropout_layer else torch.nn.Identity()
|
||||
self.add_identity = add_identity
|
||||
|
||||
@deprecated_api_warning({'residual': 'identity'}, cls_name='FFN')
|
||||
def forward(self, x, identity=None):
|
||||
"""Forward function for `FFN`.
|
||||
|
||||
The function would add x to the output tensor if residue is None.
|
||||
"""
|
||||
out = self.layers(x)
|
||||
if not self.add_identity:
|
||||
return self.dropout_layer(out)
|
||||
if identity is None:
|
||||
identity = x
|
||||
return identity + self.dropout_layer(out)
|
||||
|
||||
|
||||
@TRANSFORMER_LAYER.register_module()
|
||||
class BaseTransformerLayer(BaseModule):
|
||||
"""Base `TransformerLayer` for vision transformer.
|
||||
|
||||
It can be built from `mmcv.ConfigDict` and support more flexible
|
||||
customization, for example, using any number of `FFN or LN ` and
|
||||
use different kinds of `attention` by specifying a list of `ConfigDict`
|
||||
named `attn_cfgs`. It is worth mentioning that it supports `prenorm`
|
||||
when you specifying `norm` as the first element of `operation_order`.
|
||||
More details about the `prenorm`: `On Layer Normalization in the
|
||||
Transformer Architecture <https://arxiv.org/abs/2002.04745>`_ .
|
||||
|
||||
Args:
|
||||
attn_cfgs (list[`mmcv.ConfigDict`] | obj:`mmcv.ConfigDict` | None )):
|
||||
Configs for `self_attention` or `cross_attention` modules,
|
||||
The order of the configs in the list should be consistent with
|
||||
corresponding attentions in operation_order.
|
||||
If it is a dict, all of the attention modules in operation_order
|
||||
will be built with this config. Default: None.
|
||||
ffn_cfgs (list[`mmcv.ConfigDict`] | obj:`mmcv.ConfigDict` | None )):
|
||||
Configs for FFN, The order of the configs in the list should be
|
||||
consistent with corresponding ffn in operation_order.
|
||||
If it is a dict, all of the attention modules in operation_order
|
||||
will be built with this config.
|
||||
operation_order (tuple[str]): The execution order of operation
|
||||
in transformer. Such as ('self_attn', 'norm', 'ffn', 'norm').
|
||||
Support `prenorm` when you specifying first element as `norm`.
|
||||
Default:None.
|
||||
norm_cfg (dict): Config dict for normalization layer.
|
||||
Default: dict(type='LN').
|
||||
init_cfg (obj:`mmcv.ConfigDict`): The Config for initialization.
|
||||
Default: None.
|
||||
batch_first (bool): Key, Query and Value are shape
|
||||
of (batch, n, embed_dim)
|
||||
or (n, batch, embed_dim). Default to False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
attn_cfgs=None,
|
||||
ffn_cfgs=dict(
|
||||
type='FFN',
|
||||
embed_dims=256,
|
||||
feedforward_channels=1024,
|
||||
num_fcs=2,
|
||||
ffn_drop=0.,
|
||||
act_cfg=dict(type='ReLU', inplace=True),
|
||||
),
|
||||
operation_order=None,
|
||||
norm_cfg=dict(type='LN'),
|
||||
init_cfg=None,
|
||||
batch_first=False,
|
||||
**kwargs):
|
||||
|
||||
deprecated_args = dict(
|
||||
feedforward_channels='feedforward_channels',
|
||||
ffn_dropout='ffn_drop',
|
||||
ffn_num_fcs='num_fcs')
|
||||
for ori_name, new_name in deprecated_args.items():
|
||||
if ori_name in kwargs:
|
||||
warnings.warn(
|
||||
f'The arguments `{ori_name}` in BaseTransformerLayer '
|
||||
f'has been deprecated, now you should set `{new_name}` '
|
||||
f'and other FFN related arguments '
|
||||
f'to a dict named `ffn_cfgs`. ')
|
||||
ffn_cfgs[new_name] = kwargs[ori_name]
|
||||
|
||||
super(BaseTransformerLayer, self).__init__(init_cfg)
|
||||
|
||||
self.batch_first = batch_first
|
||||
|
||||
assert set(operation_order) & set(
|
||||
['self_attn', 'norm', 'ffn', 'cross_attn']) == \
|
||||
set(operation_order), f'The operation_order of' \
|
||||
f' {self.__class__.__name__} should ' \
|
||||
f'contains all four operation type ' \
|
||||
f"{['self_attn', 'norm', 'ffn', 'cross_attn']}"
|
||||
|
||||
num_attn = operation_order.count('self_attn') + operation_order.count(
|
||||
'cross_attn')
|
||||
if isinstance(attn_cfgs, dict):
|
||||
attn_cfgs = [copy.deepcopy(attn_cfgs) for _ in range(num_attn)]
|
||||
else:
|
||||
assert num_attn == len(attn_cfgs), f'The length ' \
|
||||
f'of attn_cfg {num_attn} is ' \
|
||||
f'not consistent with the number of attention' \
|
||||
f'in operation_order {operation_order}.'
|
||||
|
||||
self.num_attn = num_attn
|
||||
self.operation_order = operation_order
|
||||
self.norm_cfg = norm_cfg
|
||||
self.pre_norm = operation_order[0] == 'norm'
|
||||
self.attentions = ModuleList()
|
||||
|
||||
index = 0
|
||||
for operation_name in operation_order:
|
||||
if operation_name in ['self_attn', 'cross_attn']:
|
||||
if 'batch_first' in attn_cfgs[index]:
|
||||
assert self.batch_first == attn_cfgs[index]['batch_first']
|
||||
else:
|
||||
attn_cfgs[index]['batch_first'] = self.batch_first
|
||||
attention = build_attention(attn_cfgs[index])
|
||||
# Some custom attentions used as `self_attn`
|
||||
# or `cross_attn` can have different behavior.
|
||||
attention.operation_name = operation_name
|
||||
self.attentions.append(attention)
|
||||
index += 1
|
||||
|
||||
self.embed_dims = self.attentions[0].embed_dims
|
||||
|
||||
self.ffns = ModuleList()
|
||||
num_ffns = operation_order.count('ffn')
|
||||
if isinstance(ffn_cfgs, dict):
|
||||
ffn_cfgs = ConfigDict(ffn_cfgs)
|
||||
if isinstance(ffn_cfgs, dict):
|
||||
ffn_cfgs = [copy.deepcopy(ffn_cfgs) for _ in range(num_ffns)]
|
||||
assert len(ffn_cfgs) == num_ffns
|
||||
for ffn_index in range(num_ffns):
|
||||
if 'embed_dims' not in ffn_cfgs[ffn_index]:
|
||||
ffn_cfgs['embed_dims'] = self.embed_dims
|
||||
else:
|
||||
assert ffn_cfgs[ffn_index]['embed_dims'] == self.embed_dims
|
||||
self.ffns.append(
|
||||
build_feedforward_network(ffn_cfgs[ffn_index],
|
||||
dict(type='FFN')))
|
||||
|
||||
self.norms = ModuleList()
|
||||
num_norms = operation_order.count('norm')
|
||||
for _ in range(num_norms):
|
||||
self.norms.append(build_norm_layer(norm_cfg, self.embed_dims)[1])
|
||||
|
||||
def forward(self,
|
||||
query,
|
||||
key=None,
|
||||
value=None,
|
||||
query_pos=None,
|
||||
key_pos=None,
|
||||
attn_masks=None,
|
||||
query_key_padding_mask=None,
|
||||
key_padding_mask=None,
|
||||
**kwargs):
|
||||
"""Forward function for `TransformerDecoderLayer`.
|
||||
|
||||
**kwargs contains some specific arguments of attentions.
|
||||
|
||||
Args:
|
||||
query (Tensor): The input query with shape
|
||||
[num_queries, bs, embed_dims] if
|
||||
self.batch_first is False, else
|
||||
[bs, num_queries embed_dims].
|
||||
key (Tensor): The key tensor with shape [num_keys, bs,
|
||||
embed_dims] if self.batch_first is False, else
|
||||
[bs, num_keys, embed_dims] .
|
||||
value (Tensor): The value tensor with same shape as `key`.
|
||||
query_pos (Tensor): The positional encoding for `query`.
|
||||
Default: None.
|
||||
key_pos (Tensor): The positional encoding for `key`.
|
||||
Default: None.
|
||||
attn_masks (List[Tensor] | None): 2D Tensor used in
|
||||
calculation of corresponding attention. The length of
|
||||
it should equal to the number of `attention` in
|
||||
`operation_order`. Default: None.
|
||||
query_key_padding_mask (Tensor): ByteTensor for `query`, with
|
||||
shape [bs, num_queries]. Only used in `self_attn` layer.
|
||||
Defaults to None.
|
||||
key_padding_mask (Tensor): ByteTensor for `query`, with
|
||||
shape [bs, num_keys]. Default: None.
|
||||
|
||||
Returns:
|
||||
Tensor: forwarded results with shape [num_queries, bs, embed_dims].
|
||||
"""
|
||||
|
||||
norm_index = 0
|
||||
attn_index = 0
|
||||
ffn_index = 0
|
||||
identity = query
|
||||
if attn_masks is None:
|
||||
attn_masks = [None for _ in range(self.num_attn)]
|
||||
elif isinstance(attn_masks, torch.Tensor):
|
||||
attn_masks = [
|
||||
copy.deepcopy(attn_masks) for _ in range(self.num_attn)
|
||||
]
|
||||
warnings.warn(f'Use same attn_mask in all attentions in '
|
||||
f'{self.__class__.__name__} ')
|
||||
else:
|
||||
assert len(attn_masks) == self.num_attn, f'The length of ' \
|
||||
f'attn_masks {len(attn_masks)} must be equal ' \
|
||||
f'to the number of attention in ' \
|
||||
f'operation_order {self.num_attn}'
|
||||
|
||||
for layer in self.operation_order:
|
||||
if layer == 'self_attn':
|
||||
temp_key = temp_value = query
|
||||
query = self.attentions[attn_index](
|
||||
query,
|
||||
temp_key,
|
||||
temp_value,
|
||||
identity if self.pre_norm else None,
|
||||
query_pos=query_pos,
|
||||
key_pos=query_pos,
|
||||
attn_mask=attn_masks[attn_index],
|
||||
key_padding_mask=query_key_padding_mask,
|
||||
**kwargs)
|
||||
attn_index += 1
|
||||
identity = query
|
||||
|
||||
elif layer == 'norm':
|
||||
query = self.norms[norm_index](query)
|
||||
norm_index += 1
|
||||
|
||||
elif layer == 'cross_attn':
|
||||
query = self.attentions[attn_index](
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
identity if self.pre_norm else None,
|
||||
query_pos=query_pos,
|
||||
key_pos=key_pos,
|
||||
attn_mask=attn_masks[attn_index],
|
||||
key_padding_mask=key_padding_mask,
|
||||
**kwargs)
|
||||
attn_index += 1
|
||||
identity = query
|
||||
|
||||
elif layer == 'ffn':
|
||||
query = self.ffns[ffn_index](
|
||||
query, identity if self.pre_norm else None)
|
||||
ffn_index += 1
|
||||
|
||||
return query
|
||||
|
||||
|
||||
@TRANSFORMER_LAYER_SEQUENCE.register_module()
|
||||
class TransformerLayerSequence(BaseModule):
|
||||
"""Base class for TransformerEncoder and TransformerDecoder in vision
|
||||
transformer.
|
||||
|
||||
As base-class of Encoder and Decoder in vision transformer.
|
||||
Support customization such as specifying different kind
|
||||
of `transformer_layer` in `transformer_coder`.
|
||||
|
||||
Args:
|
||||
transformerlayer (list[obj:`mmcv.ConfigDict`] |
|
||||
obj:`mmcv.ConfigDict`): Config of transformerlayer
|
||||
in TransformerCoder. If it is obj:`mmcv.ConfigDict`,
|
||||
it would be repeated `num_layer` times to a
|
||||
list[`mmcv.ConfigDict`]. Default: None.
|
||||
num_layers (int): The number of `TransformerLayer`. Default: None.
|
||||
init_cfg (obj:`mmcv.ConfigDict`): The Config for initialization.
|
||||
Default: None.
|
||||
"""
|
||||
|
||||
def __init__(self, transformerlayers=None, num_layers=None, init_cfg=None):
|
||||
super(TransformerLayerSequence, self).__init__(init_cfg)
|
||||
if isinstance(transformerlayers, dict):
|
||||
transformerlayers = [
|
||||
copy.deepcopy(transformerlayers) for _ in range(num_layers)
|
||||
]
|
||||
else:
|
||||
assert isinstance(transformerlayers, list) and \
|
||||
len(transformerlayers) == num_layers
|
||||
self.num_layers = num_layers
|
||||
self.layers = ModuleList()
|
||||
for i in range(num_layers):
|
||||
self.layers.append(build_transformer_layer(transformerlayers[i]))
|
||||
self.embed_dims = self.layers[0].embed_dims
|
||||
self.pre_norm = self.layers[0].pre_norm
|
||||
|
||||
def forward(self,
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_pos=None,
|
||||
key_pos=None,
|
||||
attn_masks=None,
|
||||
query_key_padding_mask=None,
|
||||
key_padding_mask=None,
|
||||
**kwargs):
|
||||
"""Forward function for `TransformerCoder`.
|
||||
|
||||
Args:
|
||||
query (Tensor): Input query with shape
|
||||
`(num_queries, bs, embed_dims)`.
|
||||
key (Tensor): The key tensor with shape
|
||||
`(num_keys, bs, embed_dims)`.
|
||||
value (Tensor): The value tensor with shape
|
||||
`(num_keys, bs, embed_dims)`.
|
||||
query_pos (Tensor): The positional encoding for `query`.
|
||||
Default: None.
|
||||
key_pos (Tensor): The positional encoding for `key`.
|
||||
Default: None.
|
||||
attn_masks (List[Tensor], optional): Each element is 2D Tensor
|
||||
which is used in calculation of corresponding attention in
|
||||
operation_order. Default: None.
|
||||
query_key_padding_mask (Tensor): ByteTensor for `query`, with
|
||||
shape [bs, num_queries]. Only used in self-attention
|
||||
Default: None.
|
||||
key_padding_mask (Tensor): ByteTensor for `query`, with
|
||||
shape [bs, num_keys]. Default: None.
|
||||
|
||||
Returns:
|
||||
Tensor: results with shape [num_queries, bs, embed_dims].
|
||||
"""
|
||||
for layer in self.layers:
|
||||
query = layer(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
query_pos=query_pos,
|
||||
key_pos=key_pos,
|
||||
attn_masks=attn_masks,
|
||||
query_key_padding_mask=query_key_padding_mask,
|
||||
key_padding_mask=key_padding_mask,
|
||||
**kwargs)
|
||||
return query
|
||||
@@ -0,0 +1,84 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..utils import xavier_init
|
||||
from .registry import UPSAMPLE_LAYERS
|
||||
|
||||
UPSAMPLE_LAYERS.register_module('nearest', module=nn.Upsample)
|
||||
UPSAMPLE_LAYERS.register_module('bilinear', module=nn.Upsample)
|
||||
|
||||
|
||||
@UPSAMPLE_LAYERS.register_module(name='pixel_shuffle')
|
||||
class PixelShufflePack(nn.Module):
|
||||
"""Pixel Shuffle upsample layer.
|
||||
|
||||
This module packs `F.pixel_shuffle()` and a nn.Conv2d module together to
|
||||
achieve a simple upsampling with pixel shuffle.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of input channels.
|
||||
out_channels (int): Number of output channels.
|
||||
scale_factor (int): Upsample ratio.
|
||||
upsample_kernel (int): Kernel size of the conv layer to expand the
|
||||
channels.
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels, out_channels, scale_factor,
|
||||
upsample_kernel):
|
||||
super(PixelShufflePack, self).__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.scale_factor = scale_factor
|
||||
self.upsample_kernel = upsample_kernel
|
||||
self.upsample_conv = nn.Conv2d(
|
||||
self.in_channels,
|
||||
self.out_channels * scale_factor * scale_factor,
|
||||
self.upsample_kernel,
|
||||
padding=(self.upsample_kernel - 1) // 2)
|
||||
self.init_weights()
|
||||
|
||||
def init_weights(self):
|
||||
xavier_init(self.upsample_conv, distribution='uniform')
|
||||
|
||||
def forward(self, x):
|
||||
x = self.upsample_conv(x)
|
||||
x = F.pixel_shuffle(x, self.scale_factor)
|
||||
return x
|
||||
|
||||
|
||||
def build_upsample_layer(cfg, *args, **kwargs):
|
||||
"""Build upsample layer.
|
||||
|
||||
Args:
|
||||
cfg (dict): The upsample layer config, which should contain:
|
||||
|
||||
- type (str): Layer type.
|
||||
- scale_factor (int): Upsample ratio, which is not applicable to
|
||||
deconv.
|
||||
- layer args: Args needed to instantiate a upsample layer.
|
||||
args (argument list): Arguments passed to the ``__init__``
|
||||
method of the corresponding conv layer.
|
||||
kwargs (keyword arguments): Keyword arguments passed to the
|
||||
``__init__`` method of the corresponding conv layer.
|
||||
|
||||
Returns:
|
||||
nn.Module: Created upsample layer.
|
||||
"""
|
||||
if not isinstance(cfg, dict):
|
||||
raise TypeError(f'cfg must be a dict, but got {type(cfg)}')
|
||||
if 'type' not in cfg:
|
||||
raise KeyError(
|
||||
f'the cfg dict must contain the key "type", but got {cfg}')
|
||||
cfg_ = cfg.copy()
|
||||
|
||||
layer_type = cfg_.pop('type')
|
||||
if layer_type not in UPSAMPLE_LAYERS:
|
||||
raise KeyError(f'Unrecognized upsample type {layer_type}')
|
||||
else:
|
||||
upsample = UPSAMPLE_LAYERS.get(layer_type)
|
||||
|
||||
if upsample is nn.Upsample:
|
||||
cfg_['mode'] = layer_type
|
||||
layer = upsample(*args, **kwargs, **cfg_)
|
||||
return layer
|
||||
@@ -0,0 +1,180 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
r"""Modified from https://github.com/facebookresearch/detectron2/blob/master/detectron2/layers/wrappers.py # noqa: E501
|
||||
|
||||
Wrap some nn modules to support empty tensor input. Currently, these wrappers
|
||||
are mainly used in mask heads like fcn_mask_head and maskiou_heads since mask
|
||||
heads are trained on only positive RoIs.
|
||||
"""
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.nn.modules.utils import _pair, _triple
|
||||
|
||||
from .registry import CONV_LAYERS, UPSAMPLE_LAYERS
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
TORCH_VERSION = torch.__version__
|
||||
else:
|
||||
# torch.__version__ could be 1.3.1+cu92, we only need the first two
|
||||
# for comparison
|
||||
TORCH_VERSION = tuple(int(x) for x in torch.__version__.split('.')[:2])
|
||||
|
||||
|
||||
def obsolete_torch_version(torch_version, version_threshold):
|
||||
return torch_version == 'parrots' or torch_version <= version_threshold
|
||||
|
||||
|
||||
class NewEmptyTensorOp(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, x, new_shape):
|
||||
ctx.shape = x.shape
|
||||
return x.new_empty(new_shape)
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad):
|
||||
shape = ctx.shape
|
||||
return NewEmptyTensorOp.apply(grad, shape), None
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module('Conv', force=True)
|
||||
class Conv2d(nn.Conv2d):
|
||||
|
||||
def forward(self, x):
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 4)):
|
||||
out_shape = [x.shape[0], self.out_channels]
|
||||
for i, k, p, s, d in zip(x.shape[-2:], self.kernel_size,
|
||||
self.padding, self.stride, self.dilation):
|
||||
o = (i + 2 * p - (d * (k - 1) + 1)) // s + 1
|
||||
out_shape.append(o)
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
if self.training:
|
||||
# produce dummy gradient to avoid DDP warning.
|
||||
dummy = sum(x.view(-1)[0] for x in self.parameters()) * 0.0
|
||||
return empty + dummy
|
||||
else:
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module('Conv3d', force=True)
|
||||
class Conv3d(nn.Conv3d):
|
||||
|
||||
def forward(self, x):
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 4)):
|
||||
out_shape = [x.shape[0], self.out_channels]
|
||||
for i, k, p, s, d in zip(x.shape[-3:], self.kernel_size,
|
||||
self.padding, self.stride, self.dilation):
|
||||
o = (i + 2 * p - (d * (k - 1) + 1)) // s + 1
|
||||
out_shape.append(o)
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
if self.training:
|
||||
# produce dummy gradient to avoid DDP warning.
|
||||
dummy = sum(x.view(-1)[0] for x in self.parameters()) * 0.0
|
||||
return empty + dummy
|
||||
else:
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module()
|
||||
@CONV_LAYERS.register_module('deconv')
|
||||
@UPSAMPLE_LAYERS.register_module('deconv', force=True)
|
||||
class ConvTranspose2d(nn.ConvTranspose2d):
|
||||
|
||||
def forward(self, x):
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 4)):
|
||||
out_shape = [x.shape[0], self.out_channels]
|
||||
for i, k, p, s, d, op in zip(x.shape[-2:], self.kernel_size,
|
||||
self.padding, self.stride,
|
||||
self.dilation, self.output_padding):
|
||||
out_shape.append((i - 1) * s - 2 * p + (d * (k - 1) + 1) + op)
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
if self.training:
|
||||
# produce dummy gradient to avoid DDP warning.
|
||||
dummy = sum(x.view(-1)[0] for x in self.parameters()) * 0.0
|
||||
return empty + dummy
|
||||
else:
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module()
|
||||
@CONV_LAYERS.register_module('deconv3d')
|
||||
@UPSAMPLE_LAYERS.register_module('deconv3d', force=True)
|
||||
class ConvTranspose3d(nn.ConvTranspose3d):
|
||||
|
||||
def forward(self, x):
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 4)):
|
||||
out_shape = [x.shape[0], self.out_channels]
|
||||
for i, k, p, s, d, op in zip(x.shape[-3:], self.kernel_size,
|
||||
self.padding, self.stride,
|
||||
self.dilation, self.output_padding):
|
||||
out_shape.append((i - 1) * s - 2 * p + (d * (k - 1) + 1) + op)
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
if self.training:
|
||||
# produce dummy gradient to avoid DDP warning.
|
||||
dummy = sum(x.view(-1)[0] for x in self.parameters()) * 0.0
|
||||
return empty + dummy
|
||||
else:
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class MaxPool2d(nn.MaxPool2d):
|
||||
|
||||
def forward(self, x):
|
||||
# PyTorch 1.9 does not support empty tensor inference yet
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 9)):
|
||||
out_shape = list(x.shape[:2])
|
||||
for i, k, p, s, d in zip(x.shape[-2:], _pair(self.kernel_size),
|
||||
_pair(self.padding), _pair(self.stride),
|
||||
_pair(self.dilation)):
|
||||
o = (i + 2 * p - (d * (k - 1) + 1)) / s + 1
|
||||
o = math.ceil(o) if self.ceil_mode else math.floor(o)
|
||||
out_shape.append(o)
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class MaxPool3d(nn.MaxPool3d):
|
||||
|
||||
def forward(self, x):
|
||||
# PyTorch 1.9 does not support empty tensor inference yet
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 9)):
|
||||
out_shape = list(x.shape[:2])
|
||||
for i, k, p, s, d in zip(x.shape[-3:], _triple(self.kernel_size),
|
||||
_triple(self.padding),
|
||||
_triple(self.stride),
|
||||
_triple(self.dilation)):
|
||||
o = (i + 2 * p - (d * (k - 1) + 1)) / s + 1
|
||||
o = math.ceil(o) if self.ceil_mode else math.floor(o)
|
||||
out_shape.append(o)
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
|
||||
|
||||
class Linear(torch.nn.Linear):
|
||||
|
||||
def forward(self, x):
|
||||
# empty tensor forward of Linear layer is supported in Pytorch 1.6
|
||||
if x.numel() == 0 and obsolete_torch_version(TORCH_VERSION, (1, 5)):
|
||||
out_shape = [x.shape[0], self.out_features]
|
||||
empty = NewEmptyTensorOp.apply(x, out_shape)
|
||||
if self.training:
|
||||
# produce dummy gradient to avoid DDP warning.
|
||||
dummy = sum(x.view(-1)[0] for x in self.parameters()) * 0.0
|
||||
return empty + dummy
|
||||
else:
|
||||
return empty
|
||||
|
||||
return super().forward(x)
|
||||
@@ -0,0 +1,30 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from ..runner import Sequential
|
||||
from ..utils import Registry, build_from_cfg
|
||||
|
||||
|
||||
def build_model_from_cfg(cfg, registry, default_args=None):
|
||||
"""Build a PyTorch model from config dict(s). Different from
|
||||
``build_from_cfg``, if cfg is a list, a ``nn.Sequential`` will be built.
|
||||
|
||||
Args:
|
||||
cfg (dict, list[dict]): The config of modules, is is either a config
|
||||
dict or a list of config dicts. If cfg is a list, a
|
||||
the built modules will be wrapped with ``nn.Sequential``.
|
||||
registry (:obj:`Registry`): A registry the module belongs to.
|
||||
default_args (dict, optional): Default arguments to build the module.
|
||||
Defaults to None.
|
||||
|
||||
Returns:
|
||||
nn.Module: A built nn module.
|
||||
"""
|
||||
if isinstance(cfg, list):
|
||||
modules = [
|
||||
build_from_cfg(cfg_, registry, default_args) for cfg_ in cfg
|
||||
]
|
||||
return Sequential(*modules)
|
||||
else:
|
||||
return build_from_cfg(cfg, registry, default_args)
|
||||
|
||||
|
||||
MODELS = Registry('model', build_func=build_model_from_cfg)
|
||||
@@ -0,0 +1,316 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import logging
|
||||
|
||||
import torch.nn as nn
|
||||
import torch.utils.checkpoint as cp
|
||||
|
||||
from .utils import constant_init, kaiming_init
|
||||
|
||||
|
||||
def conv3x3(in_planes, out_planes, stride=1, dilation=1):
|
||||
"""3x3 convolution with padding."""
|
||||
return nn.Conv2d(
|
||||
in_planes,
|
||||
out_planes,
|
||||
kernel_size=3,
|
||||
stride=stride,
|
||||
padding=dilation,
|
||||
dilation=dilation,
|
||||
bias=False)
|
||||
|
||||
|
||||
class BasicBlock(nn.Module):
|
||||
expansion = 1
|
||||
|
||||
def __init__(self,
|
||||
inplanes,
|
||||
planes,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
downsample=None,
|
||||
style='pytorch',
|
||||
with_cp=False):
|
||||
super(BasicBlock, self).__init__()
|
||||
assert style in ['pytorch', 'caffe']
|
||||
self.conv1 = conv3x3(inplanes, planes, stride, dilation)
|
||||
self.bn1 = nn.BatchNorm2d(planes)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.conv2 = conv3x3(planes, planes)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
self.downsample = downsample
|
||||
self.stride = stride
|
||||
self.dilation = dilation
|
||||
assert not with_cp
|
||||
|
||||
def forward(self, x):
|
||||
residual = x
|
||||
|
||||
out = self.conv1(x)
|
||||
out = self.bn1(out)
|
||||
out = self.relu(out)
|
||||
|
||||
out = self.conv2(out)
|
||||
out = self.bn2(out)
|
||||
|
||||
if self.downsample is not None:
|
||||
residual = self.downsample(x)
|
||||
|
||||
out += residual
|
||||
out = self.relu(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
class Bottleneck(nn.Module):
|
||||
expansion = 4
|
||||
|
||||
def __init__(self,
|
||||
inplanes,
|
||||
planes,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
downsample=None,
|
||||
style='pytorch',
|
||||
with_cp=False):
|
||||
"""Bottleneck block.
|
||||
|
||||
If style is "pytorch", the stride-two layer is the 3x3 conv layer, if
|
||||
it is "caffe", the stride-two layer is the first 1x1 conv layer.
|
||||
"""
|
||||
super(Bottleneck, self).__init__()
|
||||
assert style in ['pytorch', 'caffe']
|
||||
if style == 'pytorch':
|
||||
conv1_stride = 1
|
||||
conv2_stride = stride
|
||||
else:
|
||||
conv1_stride = stride
|
||||
conv2_stride = 1
|
||||
self.conv1 = nn.Conv2d(
|
||||
inplanes, planes, kernel_size=1, stride=conv1_stride, bias=False)
|
||||
self.conv2 = nn.Conv2d(
|
||||
planes,
|
||||
planes,
|
||||
kernel_size=3,
|
||||
stride=conv2_stride,
|
||||
padding=dilation,
|
||||
dilation=dilation,
|
||||
bias=False)
|
||||
|
||||
self.bn1 = nn.BatchNorm2d(planes)
|
||||
self.bn2 = nn.BatchNorm2d(planes)
|
||||
self.conv3 = nn.Conv2d(
|
||||
planes, planes * self.expansion, kernel_size=1, bias=False)
|
||||
self.bn3 = nn.BatchNorm2d(planes * self.expansion)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.downsample = downsample
|
||||
self.stride = stride
|
||||
self.dilation = dilation
|
||||
self.with_cp = with_cp
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
def _inner_forward(x):
|
||||
residual = x
|
||||
|
||||
out = self.conv1(x)
|
||||
out = self.bn1(out)
|
||||
out = self.relu(out)
|
||||
|
||||
out = self.conv2(out)
|
||||
out = self.bn2(out)
|
||||
out = self.relu(out)
|
||||
|
||||
out = self.conv3(out)
|
||||
out = self.bn3(out)
|
||||
|
||||
if self.downsample is not None:
|
||||
residual = self.downsample(x)
|
||||
|
||||
out += residual
|
||||
|
||||
return out
|
||||
|
||||
if self.with_cp and x.requires_grad:
|
||||
out = cp.checkpoint(_inner_forward, x)
|
||||
else:
|
||||
out = _inner_forward(x)
|
||||
|
||||
out = self.relu(out)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def make_res_layer(block,
|
||||
inplanes,
|
||||
planes,
|
||||
blocks,
|
||||
stride=1,
|
||||
dilation=1,
|
||||
style='pytorch',
|
||||
with_cp=False):
|
||||
downsample = None
|
||||
if stride != 1 or inplanes != planes * block.expansion:
|
||||
downsample = nn.Sequential(
|
||||
nn.Conv2d(
|
||||
inplanes,
|
||||
planes * block.expansion,
|
||||
kernel_size=1,
|
||||
stride=stride,
|
||||
bias=False),
|
||||
nn.BatchNorm2d(planes * block.expansion),
|
||||
)
|
||||
|
||||
layers = []
|
||||
layers.append(
|
||||
block(
|
||||
inplanes,
|
||||
planes,
|
||||
stride,
|
||||
dilation,
|
||||
downsample,
|
||||
style=style,
|
||||
with_cp=with_cp))
|
||||
inplanes = planes * block.expansion
|
||||
for _ in range(1, blocks):
|
||||
layers.append(
|
||||
block(inplanes, planes, 1, dilation, style=style, with_cp=with_cp))
|
||||
|
||||
return nn.Sequential(*layers)
|
||||
|
||||
|
||||
class ResNet(nn.Module):
|
||||
"""ResNet backbone.
|
||||
|
||||
Args:
|
||||
depth (int): Depth of resnet, from {18, 34, 50, 101, 152}.
|
||||
num_stages (int): Resnet stages, normally 4.
|
||||
strides (Sequence[int]): Strides of the first block of each stage.
|
||||
dilations (Sequence[int]): Dilation of each stage.
|
||||
out_indices (Sequence[int]): Output from which stages.
|
||||
style (str): `pytorch` or `caffe`. If set to "pytorch", the stride-two
|
||||
layer is the 3x3 conv layer, otherwise the stride-two layer is
|
||||
the first 1x1 conv layer.
|
||||
frozen_stages (int): Stages to be frozen (all param fixed). -1 means
|
||||
not freezing any parameters.
|
||||
bn_eval (bool): Whether to set BN layers as eval mode, namely, freeze
|
||||
running stats (mean and var).
|
||||
bn_frozen (bool): Whether to freeze weight and bias of BN layers.
|
||||
with_cp (bool): Use checkpoint or not. Using checkpoint will save some
|
||||
memory while slowing down the training speed.
|
||||
"""
|
||||
|
||||
arch_settings = {
|
||||
18: (BasicBlock, (2, 2, 2, 2)),
|
||||
34: (BasicBlock, (3, 4, 6, 3)),
|
||||
50: (Bottleneck, (3, 4, 6, 3)),
|
||||
101: (Bottleneck, (3, 4, 23, 3)),
|
||||
152: (Bottleneck, (3, 8, 36, 3))
|
||||
}
|
||||
|
||||
def __init__(self,
|
||||
depth,
|
||||
num_stages=4,
|
||||
strides=(1, 2, 2, 2),
|
||||
dilations=(1, 1, 1, 1),
|
||||
out_indices=(0, 1, 2, 3),
|
||||
style='pytorch',
|
||||
frozen_stages=-1,
|
||||
bn_eval=True,
|
||||
bn_frozen=False,
|
||||
with_cp=False):
|
||||
super(ResNet, self).__init__()
|
||||
if depth not in self.arch_settings:
|
||||
raise KeyError(f'invalid depth {depth} for resnet')
|
||||
assert num_stages >= 1 and num_stages <= 4
|
||||
block, stage_blocks = self.arch_settings[depth]
|
||||
stage_blocks = stage_blocks[:num_stages]
|
||||
assert len(strides) == len(dilations) == num_stages
|
||||
assert max(out_indices) < num_stages
|
||||
|
||||
self.out_indices = out_indices
|
||||
self.style = style
|
||||
self.frozen_stages = frozen_stages
|
||||
self.bn_eval = bn_eval
|
||||
self.bn_frozen = bn_frozen
|
||||
self.with_cp = with_cp
|
||||
|
||||
self.inplanes = 64
|
||||
self.conv1 = nn.Conv2d(
|
||||
3, 64, kernel_size=7, stride=2, padding=3, bias=False)
|
||||
self.bn1 = nn.BatchNorm2d(64)
|
||||
self.relu = nn.ReLU(inplace=True)
|
||||
self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
|
||||
|
||||
self.res_layers = []
|
||||
for i, num_blocks in enumerate(stage_blocks):
|
||||
stride = strides[i]
|
||||
dilation = dilations[i]
|
||||
planes = 64 * 2**i
|
||||
res_layer = make_res_layer(
|
||||
block,
|
||||
self.inplanes,
|
||||
planes,
|
||||
num_blocks,
|
||||
stride=stride,
|
||||
dilation=dilation,
|
||||
style=self.style,
|
||||
with_cp=with_cp)
|
||||
self.inplanes = planes * block.expansion
|
||||
layer_name = f'layer{i + 1}'
|
||||
self.add_module(layer_name, res_layer)
|
||||
self.res_layers.append(layer_name)
|
||||
|
||||
self.feat_dim = block.expansion * 64 * 2**(len(stage_blocks) - 1)
|
||||
|
||||
def init_weights(self, pretrained=None):
|
||||
if isinstance(pretrained, str):
|
||||
logger = logging.getLogger()
|
||||
from ..runner import load_checkpoint
|
||||
load_checkpoint(self, pretrained, strict=False, logger=logger)
|
||||
elif pretrained is None:
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
kaiming_init(m)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
constant_init(m, 1)
|
||||
else:
|
||||
raise TypeError('pretrained must be a str or None')
|
||||
|
||||
def forward(self, x):
|
||||
x = self.conv1(x)
|
||||
x = self.bn1(x)
|
||||
x = self.relu(x)
|
||||
x = self.maxpool(x)
|
||||
outs = []
|
||||
for i, layer_name in enumerate(self.res_layers):
|
||||
res_layer = getattr(self, layer_name)
|
||||
x = res_layer(x)
|
||||
if i in self.out_indices:
|
||||
outs.append(x)
|
||||
if len(outs) == 1:
|
||||
return outs[0]
|
||||
else:
|
||||
return tuple(outs)
|
||||
|
||||
def train(self, mode=True):
|
||||
super(ResNet, self).train(mode)
|
||||
if self.bn_eval:
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.BatchNorm2d):
|
||||
m.eval()
|
||||
if self.bn_frozen:
|
||||
for params in m.parameters():
|
||||
params.requires_grad = False
|
||||
if mode and self.frozen_stages >= 0:
|
||||
for param in self.conv1.parameters():
|
||||
param.requires_grad = False
|
||||
for param in self.bn1.parameters():
|
||||
param.requires_grad = False
|
||||
self.bn1.eval()
|
||||
self.bn1.weight.requires_grad = False
|
||||
self.bn1.bias.requires_grad = False
|
||||
for i in range(1, self.frozen_stages + 1):
|
||||
mod = getattr(self, f'layer{i}')
|
||||
mod.eval()
|
||||
for param in mod.parameters():
|
||||
param.requires_grad = False
|
||||
@@ -0,0 +1,19 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .flops_counter import get_model_complexity_info
|
||||
from .fuse_conv_bn import fuse_conv_bn
|
||||
from .sync_bn import revert_sync_batchnorm
|
||||
from .weight_init import (INITIALIZERS, Caffe2XavierInit, ConstantInit,
|
||||
KaimingInit, NormalInit, PretrainedInit,
|
||||
TruncNormalInit, UniformInit, XavierInit,
|
||||
bias_init_with_prob, caffe2_xavier_init,
|
||||
constant_init, initialize, kaiming_init, normal_init,
|
||||
trunc_normal_init, uniform_init, xavier_init)
|
||||
|
||||
__all__ = [
|
||||
'get_model_complexity_info', 'bias_init_with_prob', 'caffe2_xavier_init',
|
||||
'constant_init', 'kaiming_init', 'normal_init', 'trunc_normal_init',
|
||||
'uniform_init', 'xavier_init', 'fuse_conv_bn', 'initialize',
|
||||
'INITIALIZERS', 'ConstantInit', 'XavierInit', 'NormalInit',
|
||||
'TruncNormalInit', 'UniformInit', 'KaimingInit', 'PretrainedInit',
|
||||
'Caffe2XavierInit', 'revert_sync_batchnorm'
|
||||
]
|
||||
@@ -0,0 +1,599 @@
|
||||
# Modified from flops-counter.pytorch by Vladislav Sovrasov
|
||||
# original repo: https://github.com/sovrasov/flops-counter.pytorch
|
||||
|
||||
# MIT License
|
||||
|
||||
# Copyright (c) 2018 Vladislav Sovrasov
|
||||
|
||||
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
# of this software and associated documentation files (the "Software"), to deal
|
||||
# in the Software without restriction, including without limitation the rights
|
||||
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
# copies of the Software, and to permit persons to whom the Software is
|
||||
# furnished to do so, subject to the following conditions:
|
||||
|
||||
# The above copyright notice and this permission notice shall be included in
|
||||
# all copies or substantial portions of the Software.
|
||||
|
||||
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
# SOFTWARE.
|
||||
|
||||
import sys
|
||||
from functools import partial
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import custom_mmpkg.custom_mmcv as mmcv
|
||||
|
||||
|
||||
def get_model_complexity_info(model,
|
||||
input_shape,
|
||||
print_per_layer_stat=True,
|
||||
as_strings=True,
|
||||
input_constructor=None,
|
||||
flush=False,
|
||||
ost=sys.stdout):
|
||||
"""Get complexity information of a model.
|
||||
|
||||
This method can calculate FLOPs and parameter counts of a model with
|
||||
corresponding input shape. It can also print complexity information for
|
||||
each layer in a model.
|
||||
|
||||
Supported layers are listed as below:
|
||||
- Convolutions: ``nn.Conv1d``, ``nn.Conv2d``, ``nn.Conv3d``.
|
||||
- Activations: ``nn.ReLU``, ``nn.PReLU``, ``nn.ELU``, ``nn.LeakyReLU``,
|
||||
``nn.ReLU6``.
|
||||
- Poolings: ``nn.MaxPool1d``, ``nn.MaxPool2d``, ``nn.MaxPool3d``,
|
||||
``nn.AvgPool1d``, ``nn.AvgPool2d``, ``nn.AvgPool3d``,
|
||||
``nn.AdaptiveMaxPool1d``, ``nn.AdaptiveMaxPool2d``,
|
||||
``nn.AdaptiveMaxPool3d``, ``nn.AdaptiveAvgPool1d``,
|
||||
``nn.AdaptiveAvgPool2d``, ``nn.AdaptiveAvgPool3d``.
|
||||
- BatchNorms: ``nn.BatchNorm1d``, ``nn.BatchNorm2d``,
|
||||
``nn.BatchNorm3d``, ``nn.GroupNorm``, ``nn.InstanceNorm1d``,
|
||||
``InstanceNorm2d``, ``InstanceNorm3d``, ``nn.LayerNorm``.
|
||||
- Linear: ``nn.Linear``.
|
||||
- Deconvolution: ``nn.ConvTranspose2d``.
|
||||
- Upsample: ``nn.Upsample``.
|
||||
|
||||
Args:
|
||||
model (nn.Module): The model for complexity calculation.
|
||||
input_shape (tuple): Input shape used for calculation.
|
||||
print_per_layer_stat (bool): Whether to print complexity information
|
||||
for each layer in a model. Default: True.
|
||||
as_strings (bool): Output FLOPs and params counts in a string form.
|
||||
Default: True.
|
||||
input_constructor (None | callable): If specified, it takes a callable
|
||||
method that generates input. otherwise, it will generate a random
|
||||
tensor with input shape to calculate FLOPs. Default: None.
|
||||
flush (bool): same as that in :func:`print`. Default: False.
|
||||
ost (stream): same as ``file`` param in :func:`print`.
|
||||
Default: sys.stdout.
|
||||
|
||||
Returns:
|
||||
tuple[float | str]: If ``as_strings`` is set to True, it will return
|
||||
FLOPs and parameter counts in a string format. otherwise, it will
|
||||
return those in a float number format.
|
||||
"""
|
||||
assert type(input_shape) is tuple
|
||||
assert len(input_shape) >= 1
|
||||
assert isinstance(model, nn.Module)
|
||||
flops_model = add_flops_counting_methods(model)
|
||||
flops_model.eval()
|
||||
flops_model.start_flops_count()
|
||||
if input_constructor:
|
||||
input = input_constructor(input_shape)
|
||||
_ = flops_model(**input)
|
||||
else:
|
||||
try:
|
||||
batch = torch.ones(()).new_empty(
|
||||
(1, *input_shape),
|
||||
dtype=next(flops_model.parameters()).dtype,
|
||||
device=next(flops_model.parameters()).device)
|
||||
except StopIteration:
|
||||
# Avoid StopIteration for models which have no parameters,
|
||||
# like `nn.Relu()`, `nn.AvgPool2d`, etc.
|
||||
batch = torch.ones(()).new_empty((1, *input_shape))
|
||||
|
||||
_ = flops_model(batch)
|
||||
|
||||
flops_count, params_count = flops_model.compute_average_flops_cost()
|
||||
if print_per_layer_stat:
|
||||
print_model_with_flops(
|
||||
flops_model, flops_count, params_count, ost=ost, flush=flush)
|
||||
flops_model.stop_flops_count()
|
||||
|
||||
if as_strings:
|
||||
return flops_to_string(flops_count), params_to_string(params_count)
|
||||
|
||||
return flops_count, params_count
|
||||
|
||||
|
||||
def flops_to_string(flops, units='GFLOPs', precision=2):
|
||||
"""Convert FLOPs number into a string.
|
||||
|
||||
Note that Here we take a multiply-add counts as one FLOP.
|
||||
|
||||
Args:
|
||||
flops (float): FLOPs number to be converted.
|
||||
units (str | None): Converted FLOPs units. Options are None, 'GFLOPs',
|
||||
'MFLOPs', 'KFLOPs', 'FLOPs'. If set to None, it will automatically
|
||||
choose the most suitable unit for FLOPs. Default: 'GFLOPs'.
|
||||
precision (int): Digit number after the decimal point. Default: 2.
|
||||
|
||||
Returns:
|
||||
str: The converted FLOPs number with units.
|
||||
|
||||
Examples:
|
||||
>>> flops_to_string(1e9)
|
||||
'1.0 GFLOPs'
|
||||
>>> flops_to_string(2e5, 'MFLOPs')
|
||||
'0.2 MFLOPs'
|
||||
>>> flops_to_string(3e-9, None)
|
||||
'3e-09 FLOPs'
|
||||
"""
|
||||
if units is None:
|
||||
if flops // 10**9 > 0:
|
||||
return str(round(flops / 10.**9, precision)) + ' GFLOPs'
|
||||
elif flops // 10**6 > 0:
|
||||
return str(round(flops / 10.**6, precision)) + ' MFLOPs'
|
||||
elif flops // 10**3 > 0:
|
||||
return str(round(flops / 10.**3, precision)) + ' KFLOPs'
|
||||
else:
|
||||
return str(flops) + ' FLOPs'
|
||||
else:
|
||||
if units == 'GFLOPs':
|
||||
return str(round(flops / 10.**9, precision)) + ' ' + units
|
||||
elif units == 'MFLOPs':
|
||||
return str(round(flops / 10.**6, precision)) + ' ' + units
|
||||
elif units == 'KFLOPs':
|
||||
return str(round(flops / 10.**3, precision)) + ' ' + units
|
||||
else:
|
||||
return str(flops) + ' FLOPs'
|
||||
|
||||
|
||||
def params_to_string(num_params, units=None, precision=2):
|
||||
"""Convert parameter number into a string.
|
||||
|
||||
Args:
|
||||
num_params (float): Parameter number to be converted.
|
||||
units (str | None): Converted FLOPs units. Options are None, 'M',
|
||||
'K' and ''. If set to None, it will automatically choose the most
|
||||
suitable unit for Parameter number. Default: None.
|
||||
precision (int): Digit number after the decimal point. Default: 2.
|
||||
|
||||
Returns:
|
||||
str: The converted parameter number with units.
|
||||
|
||||
Examples:
|
||||
>>> params_to_string(1e9)
|
||||
'1000.0 M'
|
||||
>>> params_to_string(2e5)
|
||||
'200.0 k'
|
||||
>>> params_to_string(3e-9)
|
||||
'3e-09'
|
||||
"""
|
||||
if units is None:
|
||||
if num_params // 10**6 > 0:
|
||||
return str(round(num_params / 10**6, precision)) + ' M'
|
||||
elif num_params // 10**3:
|
||||
return str(round(num_params / 10**3, precision)) + ' k'
|
||||
else:
|
||||
return str(num_params)
|
||||
else:
|
||||
if units == 'M':
|
||||
return str(round(num_params / 10.**6, precision)) + ' ' + units
|
||||
elif units == 'K':
|
||||
return str(round(num_params / 10.**3, precision)) + ' ' + units
|
||||
else:
|
||||
return str(num_params)
|
||||
|
||||
|
||||
def print_model_with_flops(model,
|
||||
total_flops,
|
||||
total_params,
|
||||
units='GFLOPs',
|
||||
precision=3,
|
||||
ost=sys.stdout,
|
||||
flush=False):
|
||||
"""Print a model with FLOPs for each layer.
|
||||
|
||||
Args:
|
||||
model (nn.Module): The model to be printed.
|
||||
total_flops (float): Total FLOPs of the model.
|
||||
total_params (float): Total parameter counts of the model.
|
||||
units (str | None): Converted FLOPs units. Default: 'GFLOPs'.
|
||||
precision (int): Digit number after the decimal point. Default: 3.
|
||||
ost (stream): same as `file` param in :func:`print`.
|
||||
Default: sys.stdout.
|
||||
flush (bool): same as that in :func:`print`. Default: False.
|
||||
|
||||
Example:
|
||||
>>> class ExampleModel(nn.Module):
|
||||
|
||||
>>> def __init__(self):
|
||||
>>> super().__init__()
|
||||
>>> self.conv1 = nn.Conv2d(3, 8, 3)
|
||||
>>> self.conv2 = nn.Conv2d(8, 256, 3)
|
||||
>>> self.conv3 = nn.Conv2d(256, 8, 3)
|
||||
>>> self.avg_pool = nn.AdaptiveAvgPool2d((1, 1))
|
||||
>>> self.flatten = nn.Flatten()
|
||||
>>> self.fc = nn.Linear(8, 1)
|
||||
|
||||
>>> def forward(self, x):
|
||||
>>> x = self.conv1(x)
|
||||
>>> x = self.conv2(x)
|
||||
>>> x = self.conv3(x)
|
||||
>>> x = self.avg_pool(x)
|
||||
>>> x = self.flatten(x)
|
||||
>>> x = self.fc(x)
|
||||
>>> return x
|
||||
|
||||
>>> model = ExampleModel()
|
||||
>>> x = (3, 16, 16)
|
||||
to print the complexity information state for each layer, you can use
|
||||
>>> get_model_complexity_info(model, x)
|
||||
or directly use
|
||||
>>> print_model_with_flops(model, 4579784.0, 37361)
|
||||
ExampleModel(
|
||||
0.037 M, 100.000% Params, 0.005 GFLOPs, 100.000% FLOPs,
|
||||
(conv1): Conv2d(0.0 M, 0.600% Params, 0.0 GFLOPs, 0.959% FLOPs, 3, 8, kernel_size=(3, 3), stride=(1, 1)) # noqa: E501
|
||||
(conv2): Conv2d(0.019 M, 50.020% Params, 0.003 GFLOPs, 58.760% FLOPs, 8, 256, kernel_size=(3, 3), stride=(1, 1))
|
||||
(conv3): Conv2d(0.018 M, 49.356% Params, 0.002 GFLOPs, 40.264% FLOPs, 256, 8, kernel_size=(3, 3), stride=(1, 1))
|
||||
(avg_pool): AdaptiveAvgPool2d(0.0 M, 0.000% Params, 0.0 GFLOPs, 0.017% FLOPs, output_size=(1, 1))
|
||||
(flatten): Flatten(0.0 M, 0.000% Params, 0.0 GFLOPs, 0.000% FLOPs, )
|
||||
(fc): Linear(0.0 M, 0.024% Params, 0.0 GFLOPs, 0.000% FLOPs, in_features=8, out_features=1, bias=True)
|
||||
)
|
||||
"""
|
||||
|
||||
def accumulate_params(self):
|
||||
if is_supported_instance(self):
|
||||
return self.__params__
|
||||
else:
|
||||
sum = 0
|
||||
for m in self.children():
|
||||
sum += m.accumulate_params()
|
||||
return sum
|
||||
|
||||
def accumulate_flops(self):
|
||||
if is_supported_instance(self):
|
||||
return self.__flops__ / model.__batch_counter__
|
||||
else:
|
||||
sum = 0
|
||||
for m in self.children():
|
||||
sum += m.accumulate_flops()
|
||||
return sum
|
||||
|
||||
def flops_repr(self):
|
||||
accumulated_num_params = self.accumulate_params()
|
||||
accumulated_flops_cost = self.accumulate_flops()
|
||||
return ', '.join([
|
||||
params_to_string(
|
||||
accumulated_num_params, units='M', precision=precision),
|
||||
'{:.3%} Params'.format(accumulated_num_params / total_params),
|
||||
flops_to_string(
|
||||
accumulated_flops_cost, units=units, precision=precision),
|
||||
'{:.3%} FLOPs'.format(accumulated_flops_cost / total_flops),
|
||||
self.original_extra_repr()
|
||||
])
|
||||
|
||||
def add_extra_repr(m):
|
||||
m.accumulate_flops = accumulate_flops.__get__(m)
|
||||
m.accumulate_params = accumulate_params.__get__(m)
|
||||
flops_extra_repr = flops_repr.__get__(m)
|
||||
if m.extra_repr != flops_extra_repr:
|
||||
m.original_extra_repr = m.extra_repr
|
||||
m.extra_repr = flops_extra_repr
|
||||
assert m.extra_repr != m.original_extra_repr
|
||||
|
||||
def del_extra_repr(m):
|
||||
if hasattr(m, 'original_extra_repr'):
|
||||
m.extra_repr = m.original_extra_repr
|
||||
del m.original_extra_repr
|
||||
if hasattr(m, 'accumulate_flops'):
|
||||
del m.accumulate_flops
|
||||
|
||||
model.apply(add_extra_repr)
|
||||
print(model, file=ost, flush=flush)
|
||||
model.apply(del_extra_repr)
|
||||
|
||||
|
||||
def get_model_parameters_number(model):
|
||||
"""Calculate parameter number of a model.
|
||||
|
||||
Args:
|
||||
model (nn.module): The model for parameter number calculation.
|
||||
|
||||
Returns:
|
||||
float: Parameter number of the model.
|
||||
"""
|
||||
num_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
|
||||
return num_params
|
||||
|
||||
|
||||
def add_flops_counting_methods(net_main_module):
|
||||
# adding additional methods to the existing module object,
|
||||
# this is done this way so that each function has access to self object
|
||||
net_main_module.start_flops_count = start_flops_count.__get__(
|
||||
net_main_module)
|
||||
net_main_module.stop_flops_count = stop_flops_count.__get__(
|
||||
net_main_module)
|
||||
net_main_module.reset_flops_count = reset_flops_count.__get__(
|
||||
net_main_module)
|
||||
net_main_module.compute_average_flops_cost = compute_average_flops_cost.__get__( # noqa: E501
|
||||
net_main_module)
|
||||
|
||||
net_main_module.reset_flops_count()
|
||||
|
||||
return net_main_module
|
||||
|
||||
|
||||
def compute_average_flops_cost(self):
|
||||
"""Compute average FLOPs cost.
|
||||
|
||||
A method to compute average FLOPs cost, which will be available after
|
||||
`add_flops_counting_methods()` is called on a desired net object.
|
||||
|
||||
Returns:
|
||||
float: Current mean flops consumption per image.
|
||||
"""
|
||||
batches_count = self.__batch_counter__
|
||||
flops_sum = 0
|
||||
for module in self.modules():
|
||||
if is_supported_instance(module):
|
||||
flops_sum += module.__flops__
|
||||
params_sum = get_model_parameters_number(self)
|
||||
return flops_sum / batches_count, params_sum
|
||||
|
||||
|
||||
def start_flops_count(self):
|
||||
"""Activate the computation of mean flops consumption per image.
|
||||
|
||||
A method to activate the computation of mean flops consumption per image.
|
||||
which will be available after ``add_flops_counting_methods()`` is called on
|
||||
a desired net object. It should be called before running the network.
|
||||
"""
|
||||
add_batch_counter_hook_function(self)
|
||||
|
||||
def add_flops_counter_hook_function(module):
|
||||
if is_supported_instance(module):
|
||||
if hasattr(module, '__flops_handle__'):
|
||||
return
|
||||
|
||||
else:
|
||||
handle = module.register_forward_hook(
|
||||
get_modules_mapping()[type(module)])
|
||||
|
||||
module.__flops_handle__ = handle
|
||||
|
||||
self.apply(partial(add_flops_counter_hook_function))
|
||||
|
||||
|
||||
def stop_flops_count(self):
|
||||
"""Stop computing the mean flops consumption per image.
|
||||
|
||||
A method to stop computing the mean flops consumption per image, which will
|
||||
be available after ``add_flops_counting_methods()`` is called on a desired
|
||||
net object. It can be called to pause the computation whenever.
|
||||
"""
|
||||
remove_batch_counter_hook_function(self)
|
||||
self.apply(remove_flops_counter_hook_function)
|
||||
|
||||
|
||||
def reset_flops_count(self):
|
||||
"""Reset statistics computed so far.
|
||||
|
||||
A method to Reset computed statistics, which will be available after
|
||||
`add_flops_counting_methods()` is called on a desired net object.
|
||||
"""
|
||||
add_batch_counter_variables_or_reset(self)
|
||||
self.apply(add_flops_counter_variable_or_reset)
|
||||
|
||||
|
||||
# ---- Internal functions
|
||||
def empty_flops_counter_hook(module, input, output):
|
||||
module.__flops__ += 0
|
||||
|
||||
|
||||
def upsample_flops_counter_hook(module, input, output):
|
||||
output_size = output[0]
|
||||
batch_size = output_size.shape[0]
|
||||
output_elements_count = batch_size
|
||||
for val in output_size.shape[1:]:
|
||||
output_elements_count *= val
|
||||
module.__flops__ += int(output_elements_count)
|
||||
|
||||
|
||||
def relu_flops_counter_hook(module, input, output):
|
||||
active_elements_count = output.numel()
|
||||
module.__flops__ += int(active_elements_count)
|
||||
|
||||
|
||||
def linear_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
output_last_dim = output.shape[
|
||||
-1] # pytorch checks dimensions, so here we don't care much
|
||||
module.__flops__ += int(np.prod(input.shape) * output_last_dim)
|
||||
|
||||
|
||||
def pool_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
module.__flops__ += int(np.prod(input.shape))
|
||||
|
||||
|
||||
def norm_flops_counter_hook(module, input, output):
|
||||
input = input[0]
|
||||
|
||||
batch_flops = np.prod(input.shape)
|
||||
if (getattr(module, 'affine', False)
|
||||
or getattr(module, 'elementwise_affine', False)):
|
||||
batch_flops *= 2
|
||||
module.__flops__ += int(batch_flops)
|
||||
|
||||
|
||||
def deconv_flops_counter_hook(conv_module, input, output):
|
||||
# Can have multiple inputs, getting the first one
|
||||
input = input[0]
|
||||
|
||||
batch_size = input.shape[0]
|
||||
input_height, input_width = input.shape[2:]
|
||||
|
||||
kernel_height, kernel_width = conv_module.kernel_size
|
||||
in_channels = conv_module.in_channels
|
||||
out_channels = conv_module.out_channels
|
||||
groups = conv_module.groups
|
||||
|
||||
filters_per_channel = out_channels // groups
|
||||
conv_per_position_flops = (
|
||||
kernel_height * kernel_width * in_channels * filters_per_channel)
|
||||
|
||||
active_elements_count = batch_size * input_height * input_width
|
||||
overall_conv_flops = conv_per_position_flops * active_elements_count
|
||||
bias_flops = 0
|
||||
if conv_module.bias is not None:
|
||||
output_height, output_width = output.shape[2:]
|
||||
bias_flops = out_channels * batch_size * output_height * output_height
|
||||
overall_flops = overall_conv_flops + bias_flops
|
||||
|
||||
conv_module.__flops__ += int(overall_flops)
|
||||
|
||||
|
||||
def conv_flops_counter_hook(conv_module, input, output):
|
||||
# Can have multiple inputs, getting the first one
|
||||
input = input[0]
|
||||
|
||||
batch_size = input.shape[0]
|
||||
output_dims = list(output.shape[2:])
|
||||
|
||||
kernel_dims = list(conv_module.kernel_size)
|
||||
in_channels = conv_module.in_channels
|
||||
out_channels = conv_module.out_channels
|
||||
groups = conv_module.groups
|
||||
|
||||
filters_per_channel = out_channels // groups
|
||||
conv_per_position_flops = int(
|
||||
np.prod(kernel_dims)) * in_channels * filters_per_channel
|
||||
|
||||
active_elements_count = batch_size * int(np.prod(output_dims))
|
||||
|
||||
overall_conv_flops = conv_per_position_flops * active_elements_count
|
||||
|
||||
bias_flops = 0
|
||||
|
||||
if conv_module.bias is not None:
|
||||
|
||||
bias_flops = out_channels * active_elements_count
|
||||
|
||||
overall_flops = overall_conv_flops + bias_flops
|
||||
|
||||
conv_module.__flops__ += int(overall_flops)
|
||||
|
||||
|
||||
def batch_counter_hook(module, input, output):
|
||||
batch_size = 1
|
||||
if len(input) > 0:
|
||||
# Can have multiple inputs, getting the first one
|
||||
input = input[0]
|
||||
batch_size = len(input)
|
||||
else:
|
||||
pass
|
||||
print('Warning! No positional inputs found for a module, '
|
||||
'assuming batch size is 1.')
|
||||
module.__batch_counter__ += batch_size
|
||||
|
||||
|
||||
def add_batch_counter_variables_or_reset(module):
|
||||
|
||||
module.__batch_counter__ = 0
|
||||
|
||||
|
||||
def add_batch_counter_hook_function(module):
|
||||
if hasattr(module, '__batch_counter_handle__'):
|
||||
return
|
||||
|
||||
handle = module.register_forward_hook(batch_counter_hook)
|
||||
module.__batch_counter_handle__ = handle
|
||||
|
||||
|
||||
def remove_batch_counter_hook_function(module):
|
||||
if hasattr(module, '__batch_counter_handle__'):
|
||||
module.__batch_counter_handle__.remove()
|
||||
del module.__batch_counter_handle__
|
||||
|
||||
|
||||
def add_flops_counter_variable_or_reset(module):
|
||||
if is_supported_instance(module):
|
||||
if hasattr(module, '__flops__') or hasattr(module, '__params__'):
|
||||
print('Warning: variables __flops__ or __params__ are already '
|
||||
'defined for the module' + type(module).__name__ +
|
||||
' ptflops can affect your code!')
|
||||
module.__flops__ = 0
|
||||
module.__params__ = get_model_parameters_number(module)
|
||||
|
||||
|
||||
def is_supported_instance(module):
|
||||
if type(module) in get_modules_mapping():
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def remove_flops_counter_hook_function(module):
|
||||
if is_supported_instance(module):
|
||||
if hasattr(module, '__flops_handle__'):
|
||||
module.__flops_handle__.remove()
|
||||
del module.__flops_handle__
|
||||
|
||||
|
||||
def get_modules_mapping():
|
||||
return {
|
||||
# convolutions
|
||||
nn.Conv1d: conv_flops_counter_hook,
|
||||
nn.Conv2d: conv_flops_counter_hook,
|
||||
mmcv.cnn.bricks.Conv2d: conv_flops_counter_hook,
|
||||
nn.Conv3d: conv_flops_counter_hook,
|
||||
mmcv.cnn.bricks.Conv3d: conv_flops_counter_hook,
|
||||
# activations
|
||||
nn.ReLU: relu_flops_counter_hook,
|
||||
nn.PReLU: relu_flops_counter_hook,
|
||||
nn.ELU: relu_flops_counter_hook,
|
||||
nn.LeakyReLU: relu_flops_counter_hook,
|
||||
nn.ReLU6: relu_flops_counter_hook,
|
||||
# poolings
|
||||
nn.MaxPool1d: pool_flops_counter_hook,
|
||||
nn.AvgPool1d: pool_flops_counter_hook,
|
||||
nn.AvgPool2d: pool_flops_counter_hook,
|
||||
nn.MaxPool2d: pool_flops_counter_hook,
|
||||
mmcv.cnn.bricks.MaxPool2d: pool_flops_counter_hook,
|
||||
nn.MaxPool3d: pool_flops_counter_hook,
|
||||
mmcv.cnn.bricks.MaxPool3d: pool_flops_counter_hook,
|
||||
nn.AvgPool3d: pool_flops_counter_hook,
|
||||
nn.AdaptiveMaxPool1d: pool_flops_counter_hook,
|
||||
nn.AdaptiveAvgPool1d: pool_flops_counter_hook,
|
||||
nn.AdaptiveMaxPool2d: pool_flops_counter_hook,
|
||||
nn.AdaptiveAvgPool2d: pool_flops_counter_hook,
|
||||
nn.AdaptiveMaxPool3d: pool_flops_counter_hook,
|
||||
nn.AdaptiveAvgPool3d: pool_flops_counter_hook,
|
||||
# normalizations
|
||||
nn.BatchNorm1d: norm_flops_counter_hook,
|
||||
nn.BatchNorm2d: norm_flops_counter_hook,
|
||||
nn.BatchNorm3d: norm_flops_counter_hook,
|
||||
nn.GroupNorm: norm_flops_counter_hook,
|
||||
nn.InstanceNorm1d: norm_flops_counter_hook,
|
||||
nn.InstanceNorm2d: norm_flops_counter_hook,
|
||||
nn.InstanceNorm3d: norm_flops_counter_hook,
|
||||
nn.LayerNorm: norm_flops_counter_hook,
|
||||
# FC
|
||||
nn.Linear: linear_flops_counter_hook,
|
||||
mmcv.cnn.bricks.Linear: linear_flops_counter_hook,
|
||||
# Upscale
|
||||
nn.Upsample: upsample_flops_counter_hook,
|
||||
# Deconvolution
|
||||
nn.ConvTranspose2d: deconv_flops_counter_hook,
|
||||
mmcv.cnn.bricks.ConvTranspose2d: deconv_flops_counter_hook,
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
def _fuse_conv_bn(conv, bn):
|
||||
"""Fuse conv and bn into one module.
|
||||
|
||||
Args:
|
||||
conv (nn.Module): Conv to be fused.
|
||||
bn (nn.Module): BN to be fused.
|
||||
|
||||
Returns:
|
||||
nn.Module: Fused module.
|
||||
"""
|
||||
conv_w = conv.weight
|
||||
conv_b = conv.bias if conv.bias is not None else torch.zeros_like(
|
||||
bn.running_mean)
|
||||
|
||||
factor = bn.weight / torch.sqrt(bn.running_var + bn.eps)
|
||||
conv.weight = nn.Parameter(conv_w *
|
||||
factor.reshape([conv.out_channels, 1, 1, 1]))
|
||||
conv.bias = nn.Parameter((conv_b - bn.running_mean) * factor + bn.bias)
|
||||
return conv
|
||||
|
||||
|
||||
def fuse_conv_bn(module):
|
||||
"""Recursively fuse conv and bn in a module.
|
||||
|
||||
During inference, the functionary of batch norm layers is turned off
|
||||
but only the mean and var alone channels are used, which exposes the
|
||||
chance to fuse it with the preceding conv layers to save computations and
|
||||
simplify network structures.
|
||||
|
||||
Args:
|
||||
module (nn.Module): Module to be fused.
|
||||
|
||||
Returns:
|
||||
nn.Module: Fused module.
|
||||
"""
|
||||
last_conv = None
|
||||
last_conv_name = None
|
||||
|
||||
for name, child in module.named_children():
|
||||
if isinstance(child,
|
||||
(nn.modules.batchnorm._BatchNorm, nn.SyncBatchNorm)):
|
||||
if last_conv is None: # only fuse BN that is after Conv
|
||||
continue
|
||||
fused_conv = _fuse_conv_bn(last_conv, child)
|
||||
module._modules[last_conv_name] = fused_conv
|
||||
# To reduce changes, set BN as Identity instead of deleting it.
|
||||
module._modules[name] = nn.Identity()
|
||||
last_conv = None
|
||||
elif isinstance(child, nn.Conv2d):
|
||||
last_conv = child
|
||||
last_conv_name = name
|
||||
else:
|
||||
fuse_conv_bn(child)
|
||||
return module
|
||||
@@ -0,0 +1,59 @@
|
||||
import torch
|
||||
|
||||
import custom_mmpkg.custom_mmcv as mmcv
|
||||
|
||||
|
||||
class _BatchNormXd(torch.nn.modules.batchnorm._BatchNorm):
|
||||
"""A general BatchNorm layer without input dimension check.
|
||||
|
||||
Reproduced from @kapily's work:
|
||||
(https://github.com/pytorch/pytorch/issues/41081#issuecomment-783961547)
|
||||
The only difference between BatchNorm1d, BatchNorm2d, BatchNorm3d, etc
|
||||
is `_check_input_dim` that is designed for tensor sanity checks.
|
||||
The check has been bypassed in this class for the convenience of converting
|
||||
SyncBatchNorm.
|
||||
"""
|
||||
|
||||
def _check_input_dim(self, input):
|
||||
return
|
||||
|
||||
|
||||
def revert_sync_batchnorm(module):
|
||||
"""Helper function to convert all `SyncBatchNorm` (SyncBN) and
|
||||
`mmcv.ops.sync_bn.SyncBatchNorm`(MMSyncBN) layers in the model to
|
||||
`BatchNormXd` layers.
|
||||
|
||||
Adapted from @kapily's work:
|
||||
(https://github.com/pytorch/pytorch/issues/41081#issuecomment-783961547)
|
||||
|
||||
Args:
|
||||
module (nn.Module): The module containing `SyncBatchNorm` layers.
|
||||
|
||||
Returns:
|
||||
module_output: The converted module with `BatchNormXd` layers.
|
||||
"""
|
||||
module_output = module
|
||||
module_checklist = [torch.nn.modules.batchnorm.SyncBatchNorm]
|
||||
if hasattr(mmcv, 'ops'):
|
||||
module_checklist.append(mmcv.ops.SyncBatchNorm)
|
||||
if isinstance(module, tuple(module_checklist)):
|
||||
module_output = _BatchNormXd(module.num_features, module.eps,
|
||||
module.momentum, module.affine,
|
||||
module.track_running_stats)
|
||||
if module.affine:
|
||||
# no_grad() may not be needed here but
|
||||
# just to be consistent with `convert_sync_batchnorm()`
|
||||
with torch.no_grad():
|
||||
module_output.weight = module.weight
|
||||
module_output.bias = module.bias
|
||||
module_output.running_mean = module.running_mean
|
||||
module_output.running_var = module.running_var
|
||||
module_output.num_batches_tracked = module.num_batches_tracked
|
||||
module_output.training = module.training
|
||||
# qconfig exists in quantized models
|
||||
if hasattr(module, 'qconfig'):
|
||||
module_output.qconfig = module.qconfig
|
||||
for name, child in module.named_children():
|
||||
module_output.add_module(name, revert_sync_batchnorm(child))
|
||||
del module
|
||||
return module_output
|
||||
@@ -0,0 +1,684 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import copy
|
||||
import math
|
||||
import warnings
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch import Tensor
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import Registry, build_from_cfg, get_logger, print_log
|
||||
|
||||
INITIALIZERS = Registry('initializer')
|
||||
|
||||
|
||||
def update_init_info(module, init_info):
|
||||
"""Update the `_params_init_info` in the module if the value of parameters
|
||||
are changed.
|
||||
|
||||
Args:
|
||||
module (obj:`nn.Module`): The module of PyTorch with a user-defined
|
||||
attribute `_params_init_info` which records the initialization
|
||||
information.
|
||||
init_info (str): The string that describes the initialization.
|
||||
"""
|
||||
assert hasattr(
|
||||
module,
|
||||
'_params_init_info'), f'Can not find `_params_init_info` in {module}'
|
||||
for name, param in module.named_parameters():
|
||||
|
||||
assert param in module._params_init_info, (
|
||||
f'Find a new :obj:`Parameter` '
|
||||
f'named `{name}` during executing the '
|
||||
f'`init_weights` of '
|
||||
f'`{module.__class__.__name__}`. '
|
||||
f'Please do not add or '
|
||||
f'replace parameters during executing '
|
||||
f'the `init_weights`. ')
|
||||
|
||||
# The parameter has been changed during executing the
|
||||
# `init_weights` of module
|
||||
mean_value = param.data.mean()
|
||||
if module._params_init_info[param]['tmp_mean_value'] != mean_value:
|
||||
module._params_init_info[param]['init_info'] = init_info
|
||||
module._params_init_info[param]['tmp_mean_value'] = mean_value
|
||||
|
||||
|
||||
def constant_init(module, val, bias=0):
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
nn.init.constant_(module.weight, val)
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
nn.init.constant_(module.bias, bias)
|
||||
|
||||
|
||||
def xavier_init(module, gain=1, bias=0, distribution='normal'):
|
||||
assert distribution in ['uniform', 'normal']
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
if distribution == 'uniform':
|
||||
nn.init.xavier_uniform_(module.weight, gain=gain)
|
||||
else:
|
||||
nn.init.xavier_normal_(module.weight, gain=gain)
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
nn.init.constant_(module.bias, bias)
|
||||
|
||||
|
||||
def normal_init(module, mean=0, std=1, bias=0):
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
nn.init.normal_(module.weight, mean, std)
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
nn.init.constant_(module.bias, bias)
|
||||
|
||||
|
||||
def trunc_normal_init(module: nn.Module,
|
||||
mean: float = 0,
|
||||
std: float = 1,
|
||||
a: float = -2,
|
||||
b: float = 2,
|
||||
bias: float = 0) -> None:
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
trunc_normal_(module.weight, mean, std, a, b) # type: ignore
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
nn.init.constant_(module.bias, bias) # type: ignore
|
||||
|
||||
|
||||
def uniform_init(module, a=0, b=1, bias=0):
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
nn.init.uniform_(module.weight, a, b)
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
nn.init.constant_(module.bias, bias)
|
||||
|
||||
|
||||
def kaiming_init(module,
|
||||
a=0,
|
||||
mode='fan_out',
|
||||
nonlinearity='relu',
|
||||
bias=0,
|
||||
distribution='normal'):
|
||||
assert distribution in ['uniform', 'normal']
|
||||
if hasattr(module, 'weight') and module.weight is not None:
|
||||
if distribution == 'uniform':
|
||||
nn.init.kaiming_uniform_(
|
||||
module.weight, a=a, mode=mode, nonlinearity=nonlinearity)
|
||||
else:
|
||||
nn.init.kaiming_normal_(
|
||||
module.weight, a=a, mode=mode, nonlinearity=nonlinearity)
|
||||
if hasattr(module, 'bias') and module.bias is not None:
|
||||
nn.init.constant_(module.bias, bias)
|
||||
|
||||
|
||||
def caffe2_xavier_init(module, bias=0):
|
||||
# `XavierFill` in Caffe2 corresponds to `kaiming_uniform_` in PyTorch
|
||||
# Acknowledgment to FAIR's internal code
|
||||
kaiming_init(
|
||||
module,
|
||||
a=1,
|
||||
mode='fan_in',
|
||||
nonlinearity='leaky_relu',
|
||||
bias=bias,
|
||||
distribution='uniform')
|
||||
|
||||
|
||||
def bias_init_with_prob(prior_prob):
|
||||
"""initialize conv/fc bias value according to a given probability value."""
|
||||
bias_init = float(-np.log((1 - prior_prob) / prior_prob))
|
||||
return bias_init
|
||||
|
||||
|
||||
def _get_bases_name(m):
|
||||
return [b.__name__ for b in m.__class__.__bases__]
|
||||
|
||||
|
||||
class BaseInit(object):
|
||||
|
||||
def __init__(self, *, bias=0, bias_prob=None, layer=None):
|
||||
self.wholemodule = False
|
||||
if not isinstance(bias, (int, float)):
|
||||
raise TypeError(f'bias must be a number, but got a {type(bias)}')
|
||||
|
||||
if bias_prob is not None:
|
||||
if not isinstance(bias_prob, float):
|
||||
raise TypeError(f'bias_prob type must be float, \
|
||||
but got {type(bias_prob)}')
|
||||
|
||||
if layer is not None:
|
||||
if not isinstance(layer, (str, list)):
|
||||
raise TypeError(f'layer must be a str or a list of str, \
|
||||
but got a {type(layer)}')
|
||||
else:
|
||||
layer = []
|
||||
|
||||
if bias_prob is not None:
|
||||
self.bias = bias_init_with_prob(bias_prob)
|
||||
else:
|
||||
self.bias = bias
|
||||
self.layer = [layer] if isinstance(layer, str) else layer
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Constant')
|
||||
class ConstantInit(BaseInit):
|
||||
"""Initialize module parameters with constant values.
|
||||
|
||||
Args:
|
||||
val (int | float): the value to fill the weights in the module with
|
||||
bias (int | float): the value to fill the bias. Defaults to 0.
|
||||
bias_prob (float, optional): the probability for bias initialization.
|
||||
Defaults to None.
|
||||
layer (str | list[str], optional): the layer will be initialized.
|
||||
Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(self, val, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.val = val
|
||||
|
||||
def __call__(self, module):
|
||||
|
||||
def init(m):
|
||||
if self.wholemodule:
|
||||
constant_init(m, self.val, self.bias)
|
||||
else:
|
||||
layername = m.__class__.__name__
|
||||
basesname = _get_bases_name(m)
|
||||
if len(set(self.layer) & set([layername] + basesname)):
|
||||
constant_init(m, self.val, self.bias)
|
||||
|
||||
module.apply(init)
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: val={self.val}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Xavier')
|
||||
class XavierInit(BaseInit):
|
||||
r"""Initialize module parameters with values according to the method
|
||||
described in `Understanding the difficulty of training deep feedforward
|
||||
neural networks - Glorot, X. & Bengio, Y. (2010).
|
||||
<http://proceedings.mlr.press/v9/glorot10a/glorot10a.pdf>`_
|
||||
|
||||
Args:
|
||||
gain (int | float): an optional scaling factor. Defaults to 1.
|
||||
bias (int | float): the value to fill the bias. Defaults to 0.
|
||||
bias_prob (float, optional): the probability for bias initialization.
|
||||
Defaults to None.
|
||||
distribution (str): distribution either be ``'normal'``
|
||||
or ``'uniform'``. Defaults to ``'normal'``.
|
||||
layer (str | list[str], optional): the layer will be initialized.
|
||||
Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(self, gain=1, distribution='normal', **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.gain = gain
|
||||
self.distribution = distribution
|
||||
|
||||
def __call__(self, module):
|
||||
|
||||
def init(m):
|
||||
if self.wholemodule:
|
||||
xavier_init(m, self.gain, self.bias, self.distribution)
|
||||
else:
|
||||
layername = m.__class__.__name__
|
||||
basesname = _get_bases_name(m)
|
||||
if len(set(self.layer) & set([layername] + basesname)):
|
||||
xavier_init(m, self.gain, self.bias, self.distribution)
|
||||
|
||||
module.apply(init)
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: gain={self.gain}, ' \
|
||||
f'distribution={self.distribution}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Normal')
|
||||
class NormalInit(BaseInit):
|
||||
r"""Initialize module parameters with the values drawn from the normal
|
||||
distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)`.
|
||||
|
||||
Args:
|
||||
mean (int | float):the mean of the normal distribution. Defaults to 0.
|
||||
std (int | float): the standard deviation of the normal distribution.
|
||||
Defaults to 1.
|
||||
bias (int | float): the value to fill the bias. Defaults to 0.
|
||||
bias_prob (float, optional): the probability for bias initialization.
|
||||
Defaults to None.
|
||||
layer (str | list[str], optional): the layer will be initialized.
|
||||
Defaults to None.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, mean=0, std=1, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
|
||||
def __call__(self, module):
|
||||
|
||||
def init(m):
|
||||
if self.wholemodule:
|
||||
normal_init(m, self.mean, self.std, self.bias)
|
||||
else:
|
||||
layername = m.__class__.__name__
|
||||
basesname = _get_bases_name(m)
|
||||
if len(set(self.layer) & set([layername] + basesname)):
|
||||
normal_init(m, self.mean, self.std, self.bias)
|
||||
|
||||
module.apply(init)
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: mean={self.mean},' \
|
||||
f' std={self.std}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='TruncNormal')
|
||||
class TruncNormalInit(BaseInit):
|
||||
r"""Initialize module parameters with the values drawn from the normal
|
||||
distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)` with values
|
||||
outside :math:`[a, b]`.
|
||||
|
||||
Args:
|
||||
mean (float): the mean of the normal distribution. Defaults to 0.
|
||||
std (float): the standard deviation of the normal distribution.
|
||||
Defaults to 1.
|
||||
a (float): The minimum cutoff value.
|
||||
b ( float): The maximum cutoff value.
|
||||
bias (float): the value to fill the bias. Defaults to 0.
|
||||
bias_prob (float, optional): the probability for bias initialization.
|
||||
Defaults to None.
|
||||
layer (str | list[str], optional): the layer will be initialized.
|
||||
Defaults to None.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
mean: float = 0,
|
||||
std: float = 1,
|
||||
a: float = -2,
|
||||
b: float = 2,
|
||||
**kwargs) -> None:
|
||||
super().__init__(**kwargs)
|
||||
self.mean = mean
|
||||
self.std = std
|
||||
self.a = a
|
||||
self.b = b
|
||||
|
||||
def __call__(self, module: nn.Module) -> None:
|
||||
|
||||
def init(m):
|
||||
if self.wholemodule:
|
||||
trunc_normal_init(m, self.mean, self.std, self.a, self.b,
|
||||
self.bias)
|
||||
else:
|
||||
layername = m.__class__.__name__
|
||||
basesname = _get_bases_name(m)
|
||||
if len(set(self.layer) & set([layername] + basesname)):
|
||||
trunc_normal_init(m, self.mean, self.std, self.a, self.b,
|
||||
self.bias)
|
||||
|
||||
module.apply(init)
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: a={self.a}, b={self.b},' \
|
||||
f' mean={self.mean}, std={self.std}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Uniform')
|
||||
class UniformInit(BaseInit):
|
||||
r"""Initialize module parameters with values drawn from the uniform
|
||||
distribution :math:`\mathcal{U}(a, b)`.
|
||||
|
||||
Args:
|
||||
a (int | float): the lower bound of the uniform distribution.
|
||||
Defaults to 0.
|
||||
b (int | float): the upper bound of the uniform distribution.
|
||||
Defaults to 1.
|
||||
bias (int | float): the value to fill the bias. Defaults to 0.
|
||||
bias_prob (float, optional): the probability for bias initialization.
|
||||
Defaults to None.
|
||||
layer (str | list[str], optional): the layer will be initialized.
|
||||
Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(self, a=0, b=1, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.a = a
|
||||
self.b = b
|
||||
|
||||
def __call__(self, module):
|
||||
|
||||
def init(m):
|
||||
if self.wholemodule:
|
||||
uniform_init(m, self.a, self.b, self.bias)
|
||||
else:
|
||||
layername = m.__class__.__name__
|
||||
basesname = _get_bases_name(m)
|
||||
if len(set(self.layer) & set([layername] + basesname)):
|
||||
uniform_init(m, self.a, self.b, self.bias)
|
||||
|
||||
module.apply(init)
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: a={self.a},' \
|
||||
f' b={self.b}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Kaiming')
|
||||
class KaimingInit(BaseInit):
|
||||
r"""Initialize module parameters with the values according to the method
|
||||
described in `Delving deep into rectifiers: Surpassing human-level
|
||||
performance on ImageNet classification - He, K. et al. (2015).
|
||||
<https://www.cv-foundation.org/openaccess/content_iccv_2015/
|
||||
papers/He_Delving_Deep_into_ICCV_2015_paper.pdf>`_
|
||||
|
||||
Args:
|
||||
a (int | float): the negative slope of the rectifier used after this
|
||||
layer (only used with ``'leaky_relu'``). Defaults to 0.
|
||||
mode (str): either ``'fan_in'`` or ``'fan_out'``. Choosing
|
||||
``'fan_in'`` preserves the magnitude of the variance of the weights
|
||||
in the forward pass. Choosing ``'fan_out'`` preserves the
|
||||
magnitudes in the backwards pass. Defaults to ``'fan_out'``.
|
||||
nonlinearity (str): the non-linear function (`nn.functional` name),
|
||||
recommended to use only with ``'relu'`` or ``'leaky_relu'`` .
|
||||
Defaults to 'relu'.
|
||||
bias (int | float): the value to fill the bias. Defaults to 0.
|
||||
bias_prob (float, optional): the probability for bias initialization.
|
||||
Defaults to None.
|
||||
distribution (str): distribution either be ``'normal'`` or
|
||||
``'uniform'``. Defaults to ``'normal'``.
|
||||
layer (str | list[str], optional): the layer will be initialized.
|
||||
Defaults to None.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
a=0,
|
||||
mode='fan_out',
|
||||
nonlinearity='relu',
|
||||
distribution='normal',
|
||||
**kwargs):
|
||||
super().__init__(**kwargs)
|
||||
self.a = a
|
||||
self.mode = mode
|
||||
self.nonlinearity = nonlinearity
|
||||
self.distribution = distribution
|
||||
|
||||
def __call__(self, module):
|
||||
|
||||
def init(m):
|
||||
if self.wholemodule:
|
||||
kaiming_init(m, self.a, self.mode, self.nonlinearity,
|
||||
self.bias, self.distribution)
|
||||
else:
|
||||
layername = m.__class__.__name__
|
||||
basesname = _get_bases_name(m)
|
||||
if len(set(self.layer) & set([layername] + basesname)):
|
||||
kaiming_init(m, self.a, self.mode, self.nonlinearity,
|
||||
self.bias, self.distribution)
|
||||
|
||||
module.apply(init)
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: a={self.a}, mode={self.mode}, ' \
|
||||
f'nonlinearity={self.nonlinearity}, ' \
|
||||
f'distribution ={self.distribution}, bias={self.bias}'
|
||||
return info
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Caffe2Xavier')
|
||||
class Caffe2XavierInit(KaimingInit):
|
||||
# `XavierFill` in Caffe2 corresponds to `kaiming_uniform_` in PyTorch
|
||||
# Acknowledgment to FAIR's internal code
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(
|
||||
a=1,
|
||||
mode='fan_in',
|
||||
nonlinearity='leaky_relu',
|
||||
distribution='uniform',
|
||||
**kwargs)
|
||||
|
||||
def __call__(self, module):
|
||||
super().__call__(module)
|
||||
|
||||
|
||||
@INITIALIZERS.register_module(name='Pretrained')
|
||||
class PretrainedInit(object):
|
||||
"""Initialize module by loading a pretrained model.
|
||||
|
||||
Args:
|
||||
checkpoint (str): the checkpoint file of the pretrained model should
|
||||
be load.
|
||||
prefix (str, optional): the prefix of a sub-module in the pretrained
|
||||
model. it is for loading a part of the pretrained model to
|
||||
initialize. For example, if we would like to only load the
|
||||
backbone of a detector model, we can set ``prefix='backbone.'``.
|
||||
Defaults to None.
|
||||
map_location (str): map tensors into proper locations.
|
||||
"""
|
||||
|
||||
def __init__(self, checkpoint, prefix=None, map_location=None):
|
||||
self.checkpoint = checkpoint
|
||||
self.prefix = prefix
|
||||
self.map_location = map_location
|
||||
|
||||
def __call__(self, module):
|
||||
from custom_mmpkg.custom_mmcv.runner import (_load_checkpoint_with_prefix, load_checkpoint,
|
||||
load_state_dict)
|
||||
logger = get_logger('mmcv')
|
||||
if self.prefix is None:
|
||||
print_log(f'load model from: {self.checkpoint}', logger=logger)
|
||||
load_checkpoint(
|
||||
module,
|
||||
self.checkpoint,
|
||||
map_location=self.map_location,
|
||||
strict=False,
|
||||
logger=logger)
|
||||
else:
|
||||
print_log(
|
||||
f'load {self.prefix} in model from: {self.checkpoint}',
|
||||
logger=logger)
|
||||
state_dict = _load_checkpoint_with_prefix(
|
||||
self.prefix, self.checkpoint, map_location=self.map_location)
|
||||
load_state_dict(module, state_dict, strict=False, logger=logger)
|
||||
|
||||
if hasattr(module, '_params_init_info'):
|
||||
update_init_info(module, init_info=self._get_init_info())
|
||||
|
||||
def _get_init_info(self):
|
||||
info = f'{self.__class__.__name__}: load from {self.checkpoint}'
|
||||
return info
|
||||
|
||||
|
||||
def _initialize(module, cfg, wholemodule=False):
|
||||
func = build_from_cfg(cfg, INITIALIZERS)
|
||||
# wholemodule flag is for override mode, there is no layer key in override
|
||||
# and initializer will give init values for the whole module with the name
|
||||
# in override.
|
||||
func.wholemodule = wholemodule
|
||||
func(module)
|
||||
|
||||
|
||||
def _initialize_override(module, override, cfg):
|
||||
if not isinstance(override, (dict, list)):
|
||||
raise TypeError(f'override must be a dict or a list of dict, \
|
||||
but got {type(override)}')
|
||||
|
||||
override = [override] if isinstance(override, dict) else override
|
||||
|
||||
for override_ in override:
|
||||
|
||||
cp_override = copy.deepcopy(override_)
|
||||
name = cp_override.pop('name', None)
|
||||
if name is None:
|
||||
raise ValueError('`override` must contain the key "name",'
|
||||
f'but got {cp_override}')
|
||||
# if override only has name key, it means use args in init_cfg
|
||||
if not cp_override:
|
||||
cp_override.update(cfg)
|
||||
# if override has name key and other args except type key, it will
|
||||
# raise error
|
||||
elif 'type' not in cp_override.keys():
|
||||
raise ValueError(
|
||||
f'`override` need "type" key, but got {cp_override}')
|
||||
|
||||
if hasattr(module, name):
|
||||
_initialize(getattr(module, name), cp_override, wholemodule=True)
|
||||
else:
|
||||
raise RuntimeError(f'module did not have attribute {name}, '
|
||||
f'but init_cfg is {cp_override}.')
|
||||
|
||||
|
||||
def initialize(module, init_cfg):
|
||||
"""Initialize a module.
|
||||
|
||||
Args:
|
||||
module (``torch.nn.Module``): the module will be initialized.
|
||||
init_cfg (dict | list[dict]): initialization configuration dict to
|
||||
define initializer. OpenMMLab has implemented 6 initializers
|
||||
including ``Constant``, ``Xavier``, ``Normal``, ``Uniform``,
|
||||
``Kaiming``, and ``Pretrained``.
|
||||
Example:
|
||||
>>> module = nn.Linear(2, 3, bias=True)
|
||||
>>> init_cfg = dict(type='Constant', layer='Linear', val =1 , bias =2)
|
||||
>>> initialize(module, init_cfg)
|
||||
|
||||
>>> module = nn.Sequential(nn.Conv1d(3, 1, 3), nn.Linear(1,2))
|
||||
>>> # define key ``'layer'`` for initializing layer with different
|
||||
>>> # configuration
|
||||
>>> init_cfg = [dict(type='Constant', layer='Conv1d', val=1),
|
||||
dict(type='Constant', layer='Linear', val=2)]
|
||||
>>> initialize(module, init_cfg)
|
||||
|
||||
>>> # define key``'override'`` to initialize some specific part in
|
||||
>>> # module
|
||||
>>> class FooNet(nn.Module):
|
||||
>>> def __init__(self):
|
||||
>>> super().__init__()
|
||||
>>> self.feat = nn.Conv2d(3, 16, 3)
|
||||
>>> self.reg = nn.Conv2d(16, 10, 3)
|
||||
>>> self.cls = nn.Conv2d(16, 5, 3)
|
||||
>>> model = FooNet()
|
||||
>>> init_cfg = dict(type='Constant', val=1, bias=2, layer='Conv2d',
|
||||
>>> override=dict(type='Constant', name='reg', val=3, bias=4))
|
||||
>>> initialize(model, init_cfg)
|
||||
|
||||
>>> model = ResNet(depth=50)
|
||||
>>> # Initialize weights with the pretrained model.
|
||||
>>> init_cfg = dict(type='Pretrained',
|
||||
checkpoint='torchvision://resnet50')
|
||||
>>> initialize(model, init_cfg)
|
||||
|
||||
>>> # Initialize weights of a sub-module with the specific part of
|
||||
>>> # a pretrained model by using "prefix".
|
||||
>>> url = 'http://download.openmmlab.com/mmdetection/v2.0/retinanet/'\
|
||||
>>> 'retinanet_r50_fpn_1x_coco/'\
|
||||
>>> 'retinanet_r50_fpn_1x_coco_20200130-c2398f9e.pth'
|
||||
>>> init_cfg = dict(type='Pretrained',
|
||||
checkpoint=url, prefix='backbone.')
|
||||
"""
|
||||
if not isinstance(init_cfg, (dict, list)):
|
||||
raise TypeError(f'init_cfg must be a dict or a list of dict, \
|
||||
but got {type(init_cfg)}')
|
||||
|
||||
if isinstance(init_cfg, dict):
|
||||
init_cfg = [init_cfg]
|
||||
|
||||
for cfg in init_cfg:
|
||||
# should deeply copy the original config because cfg may be used by
|
||||
# other modules, e.g., one init_cfg shared by multiple bottleneck
|
||||
# blocks, the expected cfg will be changed after pop and will change
|
||||
# the initialization behavior of other modules
|
||||
cp_cfg = copy.deepcopy(cfg)
|
||||
override = cp_cfg.pop('override', None)
|
||||
_initialize(module, cp_cfg)
|
||||
|
||||
if override is not None:
|
||||
cp_cfg.pop('layer', None)
|
||||
_initialize_override(module, override, cp_cfg)
|
||||
else:
|
||||
# All attributes in module have same initialization.
|
||||
pass
|
||||
|
||||
|
||||
def _no_grad_trunc_normal_(tensor: Tensor, mean: float, std: float, a: float,
|
||||
b: float) -> Tensor:
|
||||
# Method based on
|
||||
# https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf
|
||||
# Modified from
|
||||
# https://github.com/pytorch/pytorch/blob/master/torch/nn/init.py
|
||||
def norm_cdf(x):
|
||||
# Computes standard normal cumulative distribution function
|
||||
return (1. + math.erf(x / math.sqrt(2.))) / 2.
|
||||
|
||||
if (mean < a - 2 * std) or (mean > b + 2 * std):
|
||||
warnings.warn(
|
||||
'mean is more than 2 std from [a, b] in nn.init.trunc_normal_. '
|
||||
'The distribution of values may be incorrect.',
|
||||
stacklevel=2)
|
||||
|
||||
with torch.no_grad():
|
||||
# Values are generated by using a truncated uniform distribution and
|
||||
# then using the inverse CDF for the normal distribution.
|
||||
# Get upper and lower cdf values
|
||||
lower = norm_cdf((a - mean) / std)
|
||||
upper = norm_cdf((b - mean) / std)
|
||||
|
||||
# Uniformly fill tensor with values from [lower, upper], then translate
|
||||
# to [2lower-1, 2upper-1].
|
||||
tensor.uniform_(2 * lower - 1, 2 * upper - 1)
|
||||
|
||||
# Use inverse cdf transform for normal distribution to get truncated
|
||||
# standard normal
|
||||
tensor.erfinv_()
|
||||
|
||||
# Transform to proper mean, std
|
||||
tensor.mul_(std * math.sqrt(2.))
|
||||
tensor.add_(mean)
|
||||
|
||||
# Clamp to ensure it's in the proper range
|
||||
tensor.clamp_(min=a, max=b)
|
||||
return tensor
|
||||
|
||||
|
||||
def trunc_normal_(tensor: Tensor,
|
||||
mean: float = 0.,
|
||||
std: float = 1.,
|
||||
a: float = -2.,
|
||||
b: float = 2.) -> Tensor:
|
||||
r"""Fills the input Tensor with values drawn from a truncated
|
||||
normal distribution. The values are effectively drawn from the
|
||||
normal distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)`
|
||||
with values outside :math:`[a, b]` redrawn until they are within
|
||||
the bounds. The method used for generating the random values works
|
||||
best when :math:`a \leq \text{mean} \leq b`.
|
||||
|
||||
Modified from
|
||||
https://github.com/pytorch/pytorch/blob/master/torch/nn/init.py
|
||||
|
||||
Args:
|
||||
tensor (``torch.Tensor``): an n-dimensional `torch.Tensor`.
|
||||
mean (float): the mean of the normal distribution.
|
||||
std (float): the standard deviation of the normal distribution.
|
||||
a (float): the minimum cutoff value.
|
||||
b (float): the maximum cutoff value.
|
||||
"""
|
||||
return _no_grad_trunc_normal_(tensor, mean, std, a, b)
|
||||
@@ -0,0 +1,175 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import logging
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
from .utils import constant_init, kaiming_init, normal_init
|
||||
|
||||
|
||||
def conv3x3(in_planes, out_planes, dilation=1):
|
||||
"""3x3 convolution with padding."""
|
||||
return nn.Conv2d(
|
||||
in_planes,
|
||||
out_planes,
|
||||
kernel_size=3,
|
||||
padding=dilation,
|
||||
dilation=dilation)
|
||||
|
||||
|
||||
def make_vgg_layer(inplanes,
|
||||
planes,
|
||||
num_blocks,
|
||||
dilation=1,
|
||||
with_bn=False,
|
||||
ceil_mode=False):
|
||||
layers = []
|
||||
for _ in range(num_blocks):
|
||||
layers.append(conv3x3(inplanes, planes, dilation))
|
||||
if with_bn:
|
||||
layers.append(nn.BatchNorm2d(planes))
|
||||
layers.append(nn.ReLU(inplace=True))
|
||||
inplanes = planes
|
||||
layers.append(nn.MaxPool2d(kernel_size=2, stride=2, ceil_mode=ceil_mode))
|
||||
|
||||
return layers
|
||||
|
||||
|
||||
class VGG(nn.Module):
|
||||
"""VGG backbone.
|
||||
|
||||
Args:
|
||||
depth (int): Depth of vgg, from {11, 13, 16, 19}.
|
||||
with_bn (bool): Use BatchNorm or not.
|
||||
num_classes (int): number of classes for classification.
|
||||
num_stages (int): VGG stages, normally 5.
|
||||
dilations (Sequence[int]): Dilation of each stage.
|
||||
out_indices (Sequence[int]): Output from which stages.
|
||||
frozen_stages (int): Stages to be frozen (all param fixed). -1 means
|
||||
not freezing any parameters.
|
||||
bn_eval (bool): Whether to set BN layers as eval mode, namely, freeze
|
||||
running stats (mean and var).
|
||||
bn_frozen (bool): Whether to freeze weight and bias of BN layers.
|
||||
"""
|
||||
|
||||
arch_settings = {
|
||||
11: (1, 1, 2, 2, 2),
|
||||
13: (2, 2, 2, 2, 2),
|
||||
16: (2, 2, 3, 3, 3),
|
||||
19: (2, 2, 4, 4, 4)
|
||||
}
|
||||
|
||||
def __init__(self,
|
||||
depth,
|
||||
with_bn=False,
|
||||
num_classes=-1,
|
||||
num_stages=5,
|
||||
dilations=(1, 1, 1, 1, 1),
|
||||
out_indices=(0, 1, 2, 3, 4),
|
||||
frozen_stages=-1,
|
||||
bn_eval=True,
|
||||
bn_frozen=False,
|
||||
ceil_mode=False,
|
||||
with_last_pool=True):
|
||||
super(VGG, self).__init__()
|
||||
if depth not in self.arch_settings:
|
||||
raise KeyError(f'invalid depth {depth} for vgg')
|
||||
assert num_stages >= 1 and num_stages <= 5
|
||||
stage_blocks = self.arch_settings[depth]
|
||||
self.stage_blocks = stage_blocks[:num_stages]
|
||||
assert len(dilations) == num_stages
|
||||
assert max(out_indices) <= num_stages
|
||||
|
||||
self.num_classes = num_classes
|
||||
self.out_indices = out_indices
|
||||
self.frozen_stages = frozen_stages
|
||||
self.bn_eval = bn_eval
|
||||
self.bn_frozen = bn_frozen
|
||||
|
||||
self.inplanes = 3
|
||||
start_idx = 0
|
||||
vgg_layers = []
|
||||
self.range_sub_modules = []
|
||||
for i, num_blocks in enumerate(self.stage_blocks):
|
||||
num_modules = num_blocks * (2 + with_bn) + 1
|
||||
end_idx = start_idx + num_modules
|
||||
dilation = dilations[i]
|
||||
planes = 64 * 2**i if i < 4 else 512
|
||||
vgg_layer = make_vgg_layer(
|
||||
self.inplanes,
|
||||
planes,
|
||||
num_blocks,
|
||||
dilation=dilation,
|
||||
with_bn=with_bn,
|
||||
ceil_mode=ceil_mode)
|
||||
vgg_layers.extend(vgg_layer)
|
||||
self.inplanes = planes
|
||||
self.range_sub_modules.append([start_idx, end_idx])
|
||||
start_idx = end_idx
|
||||
if not with_last_pool:
|
||||
vgg_layers.pop(-1)
|
||||
self.range_sub_modules[-1][1] -= 1
|
||||
self.module_name = 'features'
|
||||
self.add_module(self.module_name, nn.Sequential(*vgg_layers))
|
||||
|
||||
if self.num_classes > 0:
|
||||
self.classifier = nn.Sequential(
|
||||
nn.Linear(512 * 7 * 7, 4096),
|
||||
nn.ReLU(True),
|
||||
nn.Dropout(),
|
||||
nn.Linear(4096, 4096),
|
||||
nn.ReLU(True),
|
||||
nn.Dropout(),
|
||||
nn.Linear(4096, num_classes),
|
||||
)
|
||||
|
||||
def init_weights(self, pretrained=None):
|
||||
if isinstance(pretrained, str):
|
||||
logger = logging.getLogger()
|
||||
from ..runner import load_checkpoint
|
||||
load_checkpoint(self, pretrained, strict=False, logger=logger)
|
||||
elif pretrained is None:
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
kaiming_init(m)
|
||||
elif isinstance(m, nn.BatchNorm2d):
|
||||
constant_init(m, 1)
|
||||
elif isinstance(m, nn.Linear):
|
||||
normal_init(m, std=0.01)
|
||||
else:
|
||||
raise TypeError('pretrained must be a str or None')
|
||||
|
||||
def forward(self, x):
|
||||
outs = []
|
||||
vgg_layers = getattr(self, self.module_name)
|
||||
for i in range(len(self.stage_blocks)):
|
||||
for j in range(*self.range_sub_modules[i]):
|
||||
vgg_layer = vgg_layers[j]
|
||||
x = vgg_layer(x)
|
||||
if i in self.out_indices:
|
||||
outs.append(x)
|
||||
if self.num_classes > 0:
|
||||
x = x.view(x.size(0), -1)
|
||||
x = self.classifier(x)
|
||||
outs.append(x)
|
||||
if len(outs) == 1:
|
||||
return outs[0]
|
||||
else:
|
||||
return tuple(outs)
|
||||
|
||||
def train(self, mode=True):
|
||||
super(VGG, self).train(mode)
|
||||
if self.bn_eval:
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.BatchNorm2d):
|
||||
m.eval()
|
||||
if self.bn_frozen:
|
||||
for params in m.parameters():
|
||||
params.requires_grad = False
|
||||
vgg_layers = getattr(self, self.module_name)
|
||||
if mode and self.frozen_stages >= 0:
|
||||
for i in range(self.frozen_stages):
|
||||
for j in range(*self.range_sub_modules[i]):
|
||||
mod = vgg_layers[j]
|
||||
mod.eval()
|
||||
for param in mod.parameters():
|
||||
param.requires_grad = False
|
||||
@@ -0,0 +1,8 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .test import (collect_results_cpu, collect_results_gpu, multi_gpu_test,
|
||||
single_gpu_test)
|
||||
|
||||
__all__ = [
|
||||
'collect_results_cpu', 'collect_results_gpu', 'multi_gpu_test',
|
||||
'single_gpu_test'
|
||||
]
|
||||
@@ -0,0 +1,202 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import os.path as osp
|
||||
import pickle
|
||||
import shutil
|
||||
import tempfile
|
||||
import time
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
import custom_mmpkg.custom_mmcv as mmcv
|
||||
from custom_mmpkg.custom_mmcv.runner import get_dist_info
|
||||
|
||||
|
||||
def single_gpu_test(model, data_loader):
|
||||
"""Test model with a single gpu.
|
||||
|
||||
This method tests model with a single gpu and displays test progress bar.
|
||||
|
||||
Args:
|
||||
model (nn.Module): Model to be tested.
|
||||
data_loader (nn.Dataloader): Pytorch data loader.
|
||||
|
||||
Returns:
|
||||
list: The prediction results.
|
||||
"""
|
||||
model.eval()
|
||||
results = []
|
||||
dataset = data_loader.dataset
|
||||
prog_bar = mmcv.ProgressBar(len(dataset))
|
||||
for data in data_loader:
|
||||
with torch.no_grad():
|
||||
result = model(return_loss=False, **data)
|
||||
results.extend(result)
|
||||
|
||||
# Assume result has the same length of batch_size
|
||||
# refer to https://github.com/open-mmlab/mmcv/issues/985
|
||||
batch_size = len(result)
|
||||
for _ in range(batch_size):
|
||||
prog_bar.update()
|
||||
return results
|
||||
|
||||
|
||||
def multi_gpu_test(model, data_loader, tmpdir=None, gpu_collect=False):
|
||||
"""Test model with multiple gpus.
|
||||
|
||||
This method tests model with multiple gpus and collects the results
|
||||
under two different modes: gpu and cpu modes. By setting
|
||||
``gpu_collect=True``, it encodes results to gpu tensors and use gpu
|
||||
communication for results collection. On cpu mode it saves the results on
|
||||
different gpus to ``tmpdir`` and collects them by the rank 0 worker.
|
||||
|
||||
Args:
|
||||
model (nn.Module): Model to be tested.
|
||||
data_loader (nn.Dataloader): Pytorch data loader.
|
||||
tmpdir (str): Path of directory to save the temporary results from
|
||||
different gpus under cpu mode.
|
||||
gpu_collect (bool): Option to use either gpu or cpu to collect results.
|
||||
|
||||
Returns:
|
||||
list: The prediction results.
|
||||
"""
|
||||
model.eval()
|
||||
results = []
|
||||
dataset = data_loader.dataset
|
||||
rank, world_size = get_dist_info()
|
||||
if rank == 0:
|
||||
prog_bar = mmcv.ProgressBar(len(dataset))
|
||||
time.sleep(2) # This line can prevent deadlock problem in some cases.
|
||||
for i, data in enumerate(data_loader):
|
||||
with torch.no_grad():
|
||||
result = model(return_loss=False, **data)
|
||||
results.extend(result)
|
||||
|
||||
if rank == 0:
|
||||
batch_size = len(result)
|
||||
batch_size_all = batch_size * world_size
|
||||
if batch_size_all + prog_bar.completed > len(dataset):
|
||||
batch_size_all = len(dataset) - prog_bar.completed
|
||||
for _ in range(batch_size_all):
|
||||
prog_bar.update()
|
||||
|
||||
# collect results from all ranks
|
||||
if gpu_collect:
|
||||
results = collect_results_gpu(results, len(dataset))
|
||||
else:
|
||||
results = collect_results_cpu(results, len(dataset), tmpdir)
|
||||
return results
|
||||
|
||||
|
||||
def collect_results_cpu(result_part, size, tmpdir=None):
|
||||
"""Collect results under cpu mode.
|
||||
|
||||
On cpu mode, this function will save the results on different gpus to
|
||||
``tmpdir`` and collect them by the rank 0 worker.
|
||||
|
||||
Args:
|
||||
result_part (list): Result list containing result parts
|
||||
to be collected.
|
||||
size (int): Size of the results, commonly equal to length of
|
||||
the results.
|
||||
tmpdir (str | None): temporal directory for collected results to
|
||||
store. If set to None, it will create a random temporal directory
|
||||
for it.
|
||||
|
||||
Returns:
|
||||
list: The collected results.
|
||||
"""
|
||||
rank, world_size = get_dist_info()
|
||||
# create a tmp dir if it is not specified
|
||||
if tmpdir is None:
|
||||
MAX_LEN = 512
|
||||
# 32 is whitespace
|
||||
dir_tensor = torch.full((MAX_LEN, ),
|
||||
32,
|
||||
dtype=torch.uint8,
|
||||
device='cuda')
|
||||
if rank == 0:
|
||||
mmcv.mkdir_or_exist('.dist_test')
|
||||
tmpdir = tempfile.mkdtemp(dir='.dist_test')
|
||||
tmpdir = torch.tensor(
|
||||
bytearray(tmpdir.encode()), dtype=torch.uint8, device='cuda')
|
||||
dir_tensor[:len(tmpdir)] = tmpdir
|
||||
dist.broadcast(dir_tensor, 0)
|
||||
tmpdir = dir_tensor.cpu().numpy().tobytes().decode().rstrip()
|
||||
else:
|
||||
mmcv.mkdir_or_exist(tmpdir)
|
||||
# dump the part result to the dir
|
||||
mmcv.dump(result_part, osp.join(tmpdir, f'part_{rank}.pkl'))
|
||||
dist.barrier()
|
||||
# collect all parts
|
||||
if rank != 0:
|
||||
return None
|
||||
else:
|
||||
# load results of all parts from tmp dir
|
||||
part_list = []
|
||||
for i in range(world_size):
|
||||
part_file = osp.join(tmpdir, f'part_{i}.pkl')
|
||||
part_result = mmcv.load(part_file)
|
||||
# When data is severely insufficient, an empty part_result
|
||||
# on a certain gpu could makes the overall outputs empty.
|
||||
if part_result:
|
||||
part_list.append(part_result)
|
||||
# sort the results
|
||||
ordered_results = []
|
||||
for res in zip(*part_list):
|
||||
ordered_results.extend(list(res))
|
||||
# the dataloader may pad some samples
|
||||
ordered_results = ordered_results[:size]
|
||||
# remove tmp dir
|
||||
shutil.rmtree(tmpdir)
|
||||
return ordered_results
|
||||
|
||||
|
||||
def collect_results_gpu(result_part, size):
|
||||
"""Collect results under gpu mode.
|
||||
|
||||
On gpu mode, this function will encode results to gpu tensors and use gpu
|
||||
communication for results collection.
|
||||
|
||||
Args:
|
||||
result_part (list): Result list containing result parts
|
||||
to be collected.
|
||||
size (int): Size of the results, commonly equal to length of
|
||||
the results.
|
||||
|
||||
Returns:
|
||||
list: The collected results.
|
||||
"""
|
||||
rank, world_size = get_dist_info()
|
||||
# dump result part to tensor with pickle
|
||||
part_tensor = torch.tensor(
|
||||
bytearray(pickle.dumps(result_part)), dtype=torch.uint8, device='cuda')
|
||||
# gather all result part tensor shape
|
||||
shape_tensor = torch.tensor(part_tensor.shape, device='cuda')
|
||||
shape_list = [shape_tensor.clone() for _ in range(world_size)]
|
||||
dist.all_gather(shape_list, shape_tensor)
|
||||
# padding result part tensor to max length
|
||||
shape_max = torch.tensor(shape_list).max()
|
||||
part_send = torch.zeros(shape_max, dtype=torch.uint8, device='cuda')
|
||||
part_send[:shape_tensor[0]] = part_tensor
|
||||
part_recv_list = [
|
||||
part_tensor.new_zeros(shape_max) for _ in range(world_size)
|
||||
]
|
||||
# gather all result part
|
||||
dist.all_gather(part_recv_list, part_send)
|
||||
|
||||
if rank == 0:
|
||||
part_list = []
|
||||
for recv, shape in zip(part_recv_list, shape_list):
|
||||
part_result = pickle.loads(recv[:shape[0]].cpu().numpy().tobytes())
|
||||
# When data is severely insufficient, an empty part_result
|
||||
# on a certain gpu could makes the overall outputs empty.
|
||||
if part_result:
|
||||
part_list.append(part_result)
|
||||
# sort the results
|
||||
ordered_results = []
|
||||
for res in zip(*part_list):
|
||||
ordered_results.extend(list(res))
|
||||
# the dataloader may pad some samples
|
||||
ordered_results = ordered_results[:size]
|
||||
return ordered_results
|
||||
@@ -0,0 +1,11 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .file_client import BaseStorageBackend, FileClient
|
||||
from .handlers import BaseFileHandler, JsonHandler, PickleHandler, YamlHandler
|
||||
from .io import dump, load, register_handler
|
||||
from .parse import dict_from_file, list_from_file
|
||||
|
||||
__all__ = [
|
||||
'BaseStorageBackend', 'FileClient', 'load', 'dump', 'register_handler',
|
||||
'BaseFileHandler', 'JsonHandler', 'PickleHandler', 'YamlHandler',
|
||||
'list_from_file', 'dict_from_file'
|
||||
]
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,7 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .base import BaseFileHandler
|
||||
from .json_handler import JsonHandler
|
||||
from .pickle_handler import PickleHandler
|
||||
from .yaml_handler import YamlHandler
|
||||
|
||||
__all__ = ['BaseFileHandler', 'JsonHandler', 'PickleHandler', 'YamlHandler']
|
||||
@@ -0,0 +1,30 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from abc import ABCMeta, abstractmethod
|
||||
|
||||
|
||||
class BaseFileHandler(metaclass=ABCMeta):
|
||||
# `str_like` is a flag to indicate whether the type of file object is
|
||||
# str-like object or bytes-like object. Pickle only processes bytes-like
|
||||
# objects but json only processes str-like object. If it is str-like
|
||||
# object, `StringIO` will be used to process the buffer.
|
||||
str_like = True
|
||||
|
||||
@abstractmethod
|
||||
def load_from_fileobj(self, file, **kwargs):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def dump_to_fileobj(self, obj, file, **kwargs):
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def dump_to_str(self, obj, **kwargs):
|
||||
pass
|
||||
|
||||
def load_from_path(self, filepath, mode='r', **kwargs):
|
||||
with open(filepath, mode) as f:
|
||||
return self.load_from_fileobj(f, **kwargs)
|
||||
|
||||
def dump_to_path(self, obj, filepath, mode='w', **kwargs):
|
||||
with open(filepath, mode) as f:
|
||||
self.dump_to_fileobj(obj, f, **kwargs)
|
||||
@@ -0,0 +1,36 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
|
||||
from .base import BaseFileHandler
|
||||
|
||||
|
||||
def set_default(obj):
|
||||
"""Set default json values for non-serializable values.
|
||||
|
||||
It helps convert ``set``, ``range`` and ``np.ndarray`` data types to list.
|
||||
It also converts ``np.generic`` (including ``np.int32``, ``np.float32``,
|
||||
etc.) into plain numbers of plain python built-in types.
|
||||
"""
|
||||
if isinstance(obj, (set, range)):
|
||||
return list(obj)
|
||||
elif isinstance(obj, np.ndarray):
|
||||
return obj.tolist()
|
||||
elif isinstance(obj, np.generic):
|
||||
return obj.item()
|
||||
raise TypeError(f'{type(obj)} is unsupported for json dump')
|
||||
|
||||
|
||||
class JsonHandler(BaseFileHandler):
|
||||
|
||||
def load_from_fileobj(self, file):
|
||||
return json.load(file)
|
||||
|
||||
def dump_to_fileobj(self, obj, file, **kwargs):
|
||||
kwargs.setdefault('default', set_default)
|
||||
json.dump(obj, file, **kwargs)
|
||||
|
||||
def dump_to_str(self, obj, **kwargs):
|
||||
kwargs.setdefault('default', set_default)
|
||||
return json.dumps(obj, **kwargs)
|
||||
@@ -0,0 +1,28 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import pickle
|
||||
|
||||
from .base import BaseFileHandler
|
||||
|
||||
|
||||
class PickleHandler(BaseFileHandler):
|
||||
|
||||
str_like = False
|
||||
|
||||
def load_from_fileobj(self, file, **kwargs):
|
||||
return pickle.load(file, **kwargs)
|
||||
|
||||
def load_from_path(self, filepath, **kwargs):
|
||||
return super(PickleHandler, self).load_from_path(
|
||||
filepath, mode='rb', **kwargs)
|
||||
|
||||
def dump_to_str(self, obj, **kwargs):
|
||||
kwargs.setdefault('protocol', 2)
|
||||
return pickle.dumps(obj, **kwargs)
|
||||
|
||||
def dump_to_fileobj(self, obj, file, **kwargs):
|
||||
kwargs.setdefault('protocol', 2)
|
||||
pickle.dump(obj, file, **kwargs)
|
||||
|
||||
def dump_to_path(self, obj, filepath, **kwargs):
|
||||
super(PickleHandler, self).dump_to_path(
|
||||
obj, filepath, mode='wb', **kwargs)
|
||||
@@ -0,0 +1,24 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import yaml
|
||||
|
||||
try:
|
||||
from yaml import CLoader as Loader, CDumper as Dumper
|
||||
except ImportError:
|
||||
from yaml import Loader, Dumper
|
||||
|
||||
from .base import BaseFileHandler # isort:skip
|
||||
|
||||
|
||||
class YamlHandler(BaseFileHandler):
|
||||
|
||||
def load_from_fileobj(self, file, **kwargs):
|
||||
kwargs.setdefault('Loader', Loader)
|
||||
return yaml.load(file, **kwargs)
|
||||
|
||||
def dump_to_fileobj(self, obj, file, **kwargs):
|
||||
kwargs.setdefault('Dumper', Dumper)
|
||||
yaml.dump(obj, file, **kwargs)
|
||||
|
||||
def dump_to_str(self, obj, **kwargs):
|
||||
kwargs.setdefault('Dumper', Dumper)
|
||||
return yaml.dump(obj, **kwargs)
|
||||
@@ -0,0 +1,151 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from io import BytesIO, StringIO
|
||||
from pathlib import Path
|
||||
|
||||
from ..utils import is_list_of, is_str
|
||||
from .file_client import FileClient
|
||||
from .handlers import BaseFileHandler, JsonHandler, PickleHandler, YamlHandler
|
||||
|
||||
file_handlers = {
|
||||
'json': JsonHandler(),
|
||||
'yaml': YamlHandler(),
|
||||
'yml': YamlHandler(),
|
||||
'pickle': PickleHandler(),
|
||||
'pkl': PickleHandler()
|
||||
}
|
||||
|
||||
|
||||
def load(file, file_format=None, file_client_args=None, **kwargs):
|
||||
"""Load data from json/yaml/pickle files.
|
||||
|
||||
This method provides a unified api for loading data from serialized files.
|
||||
|
||||
Note:
|
||||
In v1.3.16 and later, ``load`` supports loading data from serialized
|
||||
files those can be storaged in different backends.
|
||||
|
||||
Args:
|
||||
file (str or :obj:`Path` or file-like object): Filename or a file-like
|
||||
object.
|
||||
file_format (str, optional): If not specified, the file format will be
|
||||
inferred from the file extension, otherwise use the specified one.
|
||||
Currently supported formats include "json", "yaml/yml" and
|
||||
"pickle/pkl".
|
||||
file_client_args (dict, optional): Arguments to instantiate a
|
||||
FileClient. See :class:`mmcv.fileio.FileClient` for details.
|
||||
Default: None.
|
||||
|
||||
Examples:
|
||||
>>> load('/path/of/your/file') # file is storaged in disk
|
||||
>>> load('https://path/of/your/file') # file is storaged in Internet
|
||||
>>> load('s3://path/of/your/file') # file is storaged in petrel
|
||||
|
||||
Returns:
|
||||
The content from the file.
|
||||
"""
|
||||
if isinstance(file, Path):
|
||||
file = str(file)
|
||||
if file_format is None and is_str(file):
|
||||
file_format = file.split('.')[-1]
|
||||
if file_format not in file_handlers:
|
||||
raise TypeError(f'Unsupported format: {file_format}')
|
||||
|
||||
handler = file_handlers[file_format]
|
||||
if is_str(file):
|
||||
file_client = FileClient.infer_client(file_client_args, file)
|
||||
if handler.str_like:
|
||||
with StringIO(file_client.get_text(file)) as f:
|
||||
obj = handler.load_from_fileobj(f, **kwargs)
|
||||
else:
|
||||
with BytesIO(file_client.get(file)) as f:
|
||||
obj = handler.load_from_fileobj(f, **kwargs)
|
||||
elif hasattr(file, 'read'):
|
||||
obj = handler.load_from_fileobj(file, **kwargs)
|
||||
else:
|
||||
raise TypeError('"file" must be a filepath str or a file-object')
|
||||
return obj
|
||||
|
||||
|
||||
def dump(obj, file=None, file_format=None, file_client_args=None, **kwargs):
|
||||
"""Dump data to json/yaml/pickle strings or files.
|
||||
|
||||
This method provides a unified api for dumping data as strings or to files,
|
||||
and also supports custom arguments for each file format.
|
||||
|
||||
Note:
|
||||
In v1.3.16 and later, ``dump`` supports dumping data as strings or to
|
||||
files which is saved to different backends.
|
||||
|
||||
Args:
|
||||
obj (any): The python object to be dumped.
|
||||
file (str or :obj:`Path` or file-like object, optional): If not
|
||||
specified, then the object is dumped to a str, otherwise to a file
|
||||
specified by the filename or file-like object.
|
||||
file_format (str, optional): Same as :func:`load`.
|
||||
file_client_args (dict, optional): Arguments to instantiate a
|
||||
FileClient. See :class:`mmcv.fileio.FileClient` for details.
|
||||
Default: None.
|
||||
|
||||
Examples:
|
||||
>>> dump('hello world', '/path/of/your/file') # disk
|
||||
>>> dump('hello world', 's3://path/of/your/file') # ceph or petrel
|
||||
|
||||
Returns:
|
||||
bool: True for success, False otherwise.
|
||||
"""
|
||||
if isinstance(file, Path):
|
||||
file = str(file)
|
||||
if file_format is None:
|
||||
if is_str(file):
|
||||
file_format = file.split('.')[-1]
|
||||
elif file is None:
|
||||
raise ValueError(
|
||||
'file_format must be specified since file is None')
|
||||
if file_format not in file_handlers:
|
||||
raise TypeError(f'Unsupported format: {file_format}')
|
||||
|
||||
handler = file_handlers[file_format]
|
||||
if file is None:
|
||||
return handler.dump_to_str(obj, **kwargs)
|
||||
elif is_str(file):
|
||||
file_client = FileClient.infer_client(file_client_args, file)
|
||||
if handler.str_like:
|
||||
with StringIO() as f:
|
||||
handler.dump_to_fileobj(obj, f, **kwargs)
|
||||
file_client.put_text(f.getvalue(), file)
|
||||
else:
|
||||
with BytesIO() as f:
|
||||
handler.dump_to_fileobj(obj, f, **kwargs)
|
||||
file_client.put(f.getvalue(), file)
|
||||
elif hasattr(file, 'write'):
|
||||
handler.dump_to_fileobj(obj, file, **kwargs)
|
||||
else:
|
||||
raise TypeError('"file" must be a filename str or a file-object')
|
||||
|
||||
|
||||
def _register_handler(handler, file_formats):
|
||||
"""Register a handler for some file extensions.
|
||||
|
||||
Args:
|
||||
handler (:obj:`BaseFileHandler`): Handler to be registered.
|
||||
file_formats (str or list[str]): File formats to be handled by this
|
||||
handler.
|
||||
"""
|
||||
if not isinstance(handler, BaseFileHandler):
|
||||
raise TypeError(
|
||||
f'handler must be a child of BaseFileHandler, not {type(handler)}')
|
||||
if isinstance(file_formats, str):
|
||||
file_formats = [file_formats]
|
||||
if not is_list_of(file_formats, str):
|
||||
raise TypeError('file_formats must be a str or a list of str')
|
||||
for ext in file_formats:
|
||||
file_handlers[ext] = handler
|
||||
|
||||
|
||||
def register_handler(file_formats, **kwargs):
|
||||
|
||||
def wrap(cls):
|
||||
_register_handler(cls(**kwargs), file_formats)
|
||||
return cls
|
||||
|
||||
return wrap
|
||||
@@ -0,0 +1,97 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
|
||||
from io import StringIO
|
||||
|
||||
from .file_client import FileClient
|
||||
|
||||
|
||||
def list_from_file(filename,
|
||||
prefix='',
|
||||
offset=0,
|
||||
max_num=0,
|
||||
encoding='utf-8',
|
||||
file_client_args=None):
|
||||
"""Load a text file and parse the content as a list of strings.
|
||||
|
||||
Note:
|
||||
In v1.3.16 and later, ``list_from_file`` supports loading a text file
|
||||
which can be storaged in different backends and parsing the content as
|
||||
a list for strings.
|
||||
|
||||
Args:
|
||||
filename (str): Filename.
|
||||
prefix (str): The prefix to be inserted to the beginning of each item.
|
||||
offset (int): The offset of lines.
|
||||
max_num (int): The maximum number of lines to be read,
|
||||
zeros and negatives mean no limitation.
|
||||
encoding (str): Encoding used to open the file. Default utf-8.
|
||||
file_client_args (dict, optional): Arguments to instantiate a
|
||||
FileClient. See :class:`mmcv.fileio.FileClient` for details.
|
||||
Default: None.
|
||||
|
||||
Examples:
|
||||
>>> list_from_file('/path/of/your/file') # disk
|
||||
['hello', 'world']
|
||||
>>> list_from_file('s3://path/of/your/file') # ceph or petrel
|
||||
['hello', 'world']
|
||||
|
||||
Returns:
|
||||
list[str]: A list of strings.
|
||||
"""
|
||||
cnt = 0
|
||||
item_list = []
|
||||
file_client = FileClient.infer_client(file_client_args, filename)
|
||||
with StringIO(file_client.get_text(filename, encoding)) as f:
|
||||
for _ in range(offset):
|
||||
f.readline()
|
||||
for line in f:
|
||||
if 0 < max_num <= cnt:
|
||||
break
|
||||
item_list.append(prefix + line.rstrip('\n\r'))
|
||||
cnt += 1
|
||||
return item_list
|
||||
|
||||
|
||||
def dict_from_file(filename,
|
||||
key_type=str,
|
||||
encoding='utf-8',
|
||||
file_client_args=None):
|
||||
"""Load a text file and parse the content as a dict.
|
||||
|
||||
Each line of the text file will be two or more columns split by
|
||||
whitespaces or tabs. The first column will be parsed as dict keys, and
|
||||
the following columns will be parsed as dict values.
|
||||
|
||||
Note:
|
||||
In v1.3.16 and later, ``dict_from_file`` supports loading a text file
|
||||
which can be storaged in different backends and parsing the content as
|
||||
a dict.
|
||||
|
||||
Args:
|
||||
filename(str): Filename.
|
||||
key_type(type): Type of the dict keys. str is user by default and
|
||||
type conversion will be performed if specified.
|
||||
encoding (str): Encoding used to open the file. Default utf-8.
|
||||
file_client_args (dict, optional): Arguments to instantiate a
|
||||
FileClient. See :class:`mmcv.fileio.FileClient` for details.
|
||||
Default: None.
|
||||
|
||||
Examples:
|
||||
>>> dict_from_file('/path/of/your/file') # disk
|
||||
{'key1': 'value1', 'key2': 'value2'}
|
||||
>>> dict_from_file('s3://path/of/your/file') # ceph or petrel
|
||||
{'key1': 'value1', 'key2': 'value2'}
|
||||
|
||||
Returns:
|
||||
dict: The parsed contents.
|
||||
"""
|
||||
mapping = {}
|
||||
file_client = FileClient.infer_client(file_client_args, filename)
|
||||
with StringIO(file_client.get_text(filename, encoding)) as f:
|
||||
for line in f:
|
||||
items = line.rstrip('\n').split()
|
||||
assert len(items) >= 2
|
||||
key = key_type(items[0])
|
||||
val = items[1:] if len(items) > 2 else items[1]
|
||||
mapping[key] = val
|
||||
return mapping
|
||||
@@ -0,0 +1,28 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .colorspace import (bgr2gray, bgr2hls, bgr2hsv, bgr2rgb, bgr2ycbcr,
|
||||
gray2bgr, gray2rgb, hls2bgr, hsv2bgr, imconvert,
|
||||
rgb2bgr, rgb2gray, rgb2ycbcr, ycbcr2bgr, ycbcr2rgb)
|
||||
from .geometric import (cutout, imcrop, imflip, imflip_, impad,
|
||||
impad_to_multiple, imrescale, imresize, imresize_like,
|
||||
imresize_to_multiple, imrotate, imshear, imtranslate,
|
||||
rescale_size)
|
||||
from .io import imfrombytes, imread, imwrite, supported_backends, use_backend
|
||||
from .misc import tensor2imgs
|
||||
from .photometric import (adjust_brightness, adjust_color, adjust_contrast,
|
||||
adjust_lighting, adjust_sharpness, auto_contrast,
|
||||
clahe, imdenormalize, imequalize, iminvert,
|
||||
imnormalize, imnormalize_, lut_transform, posterize,
|
||||
solarize)
|
||||
|
||||
__all__ = [
|
||||
'bgr2gray', 'bgr2hls', 'bgr2hsv', 'bgr2rgb', 'gray2bgr', 'gray2rgb',
|
||||
'hls2bgr', 'hsv2bgr', 'imconvert', 'rgb2bgr', 'rgb2gray', 'imrescale',
|
||||
'imresize', 'imresize_like', 'imresize_to_multiple', 'rescale_size',
|
||||
'imcrop', 'imflip', 'imflip_', 'impad', 'impad_to_multiple', 'imrotate',
|
||||
'imfrombytes', 'imread', 'imwrite', 'supported_backends', 'use_backend',
|
||||
'imdenormalize', 'imnormalize', 'imnormalize_', 'iminvert', 'posterize',
|
||||
'solarize', 'rgb2ycbcr', 'bgr2ycbcr', 'ycbcr2rgb', 'ycbcr2bgr',
|
||||
'tensor2imgs', 'imshear', 'imtranslate', 'adjust_color', 'imequalize',
|
||||
'adjust_brightness', 'adjust_contrast', 'lut_transform', 'clahe',
|
||||
'adjust_sharpness', 'auto_contrast', 'cutout', 'adjust_lighting'
|
||||
]
|
||||
@@ -0,0 +1,306 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
|
||||
def imconvert(img, src, dst):
|
||||
"""Convert an image from the src colorspace to dst colorspace.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
src (str): The source colorspace, e.g., 'rgb', 'hsv'.
|
||||
dst (str): The destination colorspace, e.g., 'rgb', 'hsv'.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted image.
|
||||
"""
|
||||
code = getattr(cv2, f'COLOR_{src.upper()}2{dst.upper()}')
|
||||
out_img = cv2.cvtColor(img, code)
|
||||
return out_img
|
||||
|
||||
|
||||
def bgr2gray(img, keepdim=False):
|
||||
"""Convert a BGR image to grayscale image.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
keepdim (bool): If False (by default), then return the grayscale image
|
||||
with 2 dims, otherwise 3 dims.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted grayscale image.
|
||||
"""
|
||||
out_img = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)
|
||||
if keepdim:
|
||||
out_img = out_img[..., None]
|
||||
return out_img
|
||||
|
||||
|
||||
def rgb2gray(img, keepdim=False):
|
||||
"""Convert a RGB image to grayscale image.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
keepdim (bool): If False (by default), then return the grayscale image
|
||||
with 2 dims, otherwise 3 dims.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted grayscale image.
|
||||
"""
|
||||
out_img = cv2.cvtColor(img, cv2.COLOR_RGB2GRAY)
|
||||
if keepdim:
|
||||
out_img = out_img[..., None]
|
||||
return out_img
|
||||
|
||||
|
||||
def gray2bgr(img):
|
||||
"""Convert a grayscale image to BGR image.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted BGR image.
|
||||
"""
|
||||
img = img[..., None] if img.ndim == 2 else img
|
||||
out_img = cv2.cvtColor(img, cv2.COLOR_GRAY2BGR)
|
||||
return out_img
|
||||
|
||||
|
||||
def gray2rgb(img):
|
||||
"""Convert a grayscale image to RGB image.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted RGB image.
|
||||
"""
|
||||
img = img[..., None] if img.ndim == 2 else img
|
||||
out_img = cv2.cvtColor(img, cv2.COLOR_GRAY2RGB)
|
||||
return out_img
|
||||
|
||||
|
||||
def _convert_input_type_range(img):
|
||||
"""Convert the type and range of the input image.
|
||||
|
||||
It converts the input image to np.float32 type and range of [0, 1].
|
||||
It is mainly used for pre-processing the input image in colorspace
|
||||
conversion functions such as rgb2ycbcr and ycbcr2rgb.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image. It accepts:
|
||||
1. np.uint8 type with range [0, 255];
|
||||
2. np.float32 type with range [0, 1].
|
||||
|
||||
Returns:
|
||||
(ndarray): The converted image with type of np.float32 and range of
|
||||
[0, 1].
|
||||
"""
|
||||
img_type = img.dtype
|
||||
img = img.astype(np.float32)
|
||||
if img_type == np.float32:
|
||||
pass
|
||||
elif img_type == np.uint8:
|
||||
img /= 255.
|
||||
else:
|
||||
raise TypeError('The img type should be np.float32 or np.uint8, '
|
||||
f'but got {img_type}')
|
||||
return img
|
||||
|
||||
|
||||
def _convert_output_type_range(img, dst_type):
|
||||
"""Convert the type and range of the image according to dst_type.
|
||||
|
||||
It converts the image to desired type and range. If `dst_type` is np.uint8,
|
||||
images will be converted to np.uint8 type with range [0, 255]. If
|
||||
`dst_type` is np.float32, it converts the image to np.float32 type with
|
||||
range [0, 1].
|
||||
It is mainly used for post-processing images in colorspace conversion
|
||||
functions such as rgb2ycbcr and ycbcr2rgb.
|
||||
|
||||
Args:
|
||||
img (ndarray): The image to be converted with np.float32 type and
|
||||
range [0, 255].
|
||||
dst_type (np.uint8 | np.float32): If dst_type is np.uint8, it
|
||||
converts the image to np.uint8 type with range [0, 255]. If
|
||||
dst_type is np.float32, it converts the image to np.float32 type
|
||||
with range [0, 1].
|
||||
|
||||
Returns:
|
||||
(ndarray): The converted image with desired type and range.
|
||||
"""
|
||||
if dst_type not in (np.uint8, np.float32):
|
||||
raise TypeError('The dst_type should be np.float32 or np.uint8, '
|
||||
f'but got {dst_type}')
|
||||
if dst_type == np.uint8:
|
||||
img = img.round()
|
||||
else:
|
||||
img /= 255.
|
||||
return img.astype(dst_type)
|
||||
|
||||
|
||||
def rgb2ycbcr(img, y_only=False):
|
||||
"""Convert a RGB image to YCbCr image.
|
||||
|
||||
This function produces the same results as Matlab's `rgb2ycbcr` function.
|
||||
It implements the ITU-R BT.601 conversion for standard-definition
|
||||
television. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
|
||||
|
||||
It differs from a similar function in cv2.cvtColor: `RGB <-> YCrCb`.
|
||||
In OpenCV, it implements a JPEG conversion. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image. It accepts:
|
||||
1. np.uint8 type with range [0, 255];
|
||||
2. np.float32 type with range [0, 1].
|
||||
y_only (bool): Whether to only return Y channel. Default: False.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted YCbCr image. The output image has the same type
|
||||
and range as input image.
|
||||
"""
|
||||
img_type = img.dtype
|
||||
img = _convert_input_type_range(img)
|
||||
if y_only:
|
||||
out_img = np.dot(img, [65.481, 128.553, 24.966]) + 16.0
|
||||
else:
|
||||
out_img = np.matmul(
|
||||
img, [[65.481, -37.797, 112.0], [128.553, -74.203, -93.786],
|
||||
[24.966, 112.0, -18.214]]) + [16, 128, 128]
|
||||
out_img = _convert_output_type_range(out_img, img_type)
|
||||
return out_img
|
||||
|
||||
|
||||
def bgr2ycbcr(img, y_only=False):
|
||||
"""Convert a BGR image to YCbCr image.
|
||||
|
||||
The bgr version of rgb2ycbcr.
|
||||
It implements the ITU-R BT.601 conversion for standard-definition
|
||||
television. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
|
||||
|
||||
It differs from a similar function in cv2.cvtColor: `BGR <-> YCrCb`.
|
||||
In OpenCV, it implements a JPEG conversion. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image. It accepts:
|
||||
1. np.uint8 type with range [0, 255];
|
||||
2. np.float32 type with range [0, 1].
|
||||
y_only (bool): Whether to only return Y channel. Default: False.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted YCbCr image. The output image has the same type
|
||||
and range as input image.
|
||||
"""
|
||||
img_type = img.dtype
|
||||
img = _convert_input_type_range(img)
|
||||
if y_only:
|
||||
out_img = np.dot(img, [24.966, 128.553, 65.481]) + 16.0
|
||||
else:
|
||||
out_img = np.matmul(
|
||||
img, [[24.966, 112.0, -18.214], [128.553, -74.203, -93.786],
|
||||
[65.481, -37.797, 112.0]]) + [16, 128, 128]
|
||||
out_img = _convert_output_type_range(out_img, img_type)
|
||||
return out_img
|
||||
|
||||
|
||||
def ycbcr2rgb(img):
|
||||
"""Convert a YCbCr image to RGB image.
|
||||
|
||||
This function produces the same results as Matlab's ycbcr2rgb function.
|
||||
It implements the ITU-R BT.601 conversion for standard-definition
|
||||
television. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
|
||||
|
||||
It differs from a similar function in cv2.cvtColor: `YCrCb <-> RGB`.
|
||||
In OpenCV, it implements a JPEG conversion. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image. It accepts:
|
||||
1. np.uint8 type with range [0, 255];
|
||||
2. np.float32 type with range [0, 1].
|
||||
|
||||
Returns:
|
||||
ndarray: The converted RGB image. The output image has the same type
|
||||
and range as input image.
|
||||
"""
|
||||
img_type = img.dtype
|
||||
img = _convert_input_type_range(img) * 255
|
||||
out_img = np.matmul(img, [[0.00456621, 0.00456621, 0.00456621],
|
||||
[0, -0.00153632, 0.00791071],
|
||||
[0.00625893, -0.00318811, 0]]) * 255.0 + [
|
||||
-222.921, 135.576, -276.836
|
||||
]
|
||||
out_img = _convert_output_type_range(out_img, img_type)
|
||||
return out_img
|
||||
|
||||
|
||||
def ycbcr2bgr(img):
|
||||
"""Convert a YCbCr image to BGR image.
|
||||
|
||||
The bgr version of ycbcr2rgb.
|
||||
It implements the ITU-R BT.601 conversion for standard-definition
|
||||
television. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#ITU-R_BT.601_conversion.
|
||||
|
||||
It differs from a similar function in cv2.cvtColor: `YCrCb <-> BGR`.
|
||||
In OpenCV, it implements a JPEG conversion. See more details in
|
||||
https://en.wikipedia.org/wiki/YCbCr#JPEG_conversion.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image. It accepts:
|
||||
1. np.uint8 type with range [0, 255];
|
||||
2. np.float32 type with range [0, 1].
|
||||
|
||||
Returns:
|
||||
ndarray: The converted BGR image. The output image has the same type
|
||||
and range as input image.
|
||||
"""
|
||||
img_type = img.dtype
|
||||
img = _convert_input_type_range(img) * 255
|
||||
out_img = np.matmul(img, [[0.00456621, 0.00456621, 0.00456621],
|
||||
[0.00791071, -0.00153632, 0],
|
||||
[0, -0.00318811, 0.00625893]]) * 255.0 + [
|
||||
-276.836, 135.576, -222.921
|
||||
]
|
||||
out_img = _convert_output_type_range(out_img, img_type)
|
||||
return out_img
|
||||
|
||||
|
||||
def convert_color_factory(src, dst):
|
||||
|
||||
code = getattr(cv2, f'COLOR_{src.upper()}2{dst.upper()}')
|
||||
|
||||
def convert_color(img):
|
||||
out_img = cv2.cvtColor(img, code)
|
||||
return out_img
|
||||
|
||||
convert_color.__doc__ = f"""Convert a {src.upper()} image to {dst.upper()}
|
||||
image.
|
||||
|
||||
Args:
|
||||
img (ndarray or str): The input image.
|
||||
|
||||
Returns:
|
||||
ndarray: The converted {dst.upper()} image.
|
||||
"""
|
||||
|
||||
return convert_color
|
||||
|
||||
|
||||
bgr2rgb = convert_color_factory('bgr', 'rgb')
|
||||
|
||||
rgb2bgr = convert_color_factory('rgb', 'bgr')
|
||||
|
||||
bgr2hsv = convert_color_factory('bgr', 'hsv')
|
||||
|
||||
hsv2bgr = convert_color_factory('hsv', 'bgr')
|
||||
|
||||
bgr2hls = convert_color_factory('bgr', 'hls')
|
||||
|
||||
hls2bgr = convert_color_factory('hls', 'bgr')
|
||||
@@ -0,0 +1,728 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import numbers
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from ..utils import to_2tuple
|
||||
from .io import imread_backend
|
||||
|
||||
try:
|
||||
from PIL import Image
|
||||
except ImportError:
|
||||
Image = None
|
||||
|
||||
|
||||
def _scale_size(size, scale):
|
||||
"""Rescale a size by a ratio.
|
||||
|
||||
Args:
|
||||
size (tuple[int]): (w, h).
|
||||
scale (float | tuple(float)): Scaling factor.
|
||||
|
||||
Returns:
|
||||
tuple[int]: scaled size.
|
||||
"""
|
||||
if isinstance(scale, (float, int)):
|
||||
scale = (scale, scale)
|
||||
w, h = size
|
||||
return int(w * float(scale[0]) + 0.5), int(h * float(scale[1]) + 0.5)
|
||||
|
||||
|
||||
cv2_interp_codes = {
|
||||
'nearest': cv2.INTER_NEAREST,
|
||||
'bilinear': cv2.INTER_LINEAR,
|
||||
'bicubic': cv2.INTER_CUBIC,
|
||||
'area': cv2.INTER_AREA,
|
||||
'lanczos': cv2.INTER_LANCZOS4
|
||||
}
|
||||
|
||||
if Image is not None:
|
||||
pillow_interp_codes = {
|
||||
'nearest': Image.NEAREST,
|
||||
'bilinear': Image.BILINEAR,
|
||||
'bicubic': Image.BICUBIC,
|
||||
'box': Image.BOX,
|
||||
'lanczos': Image.LANCZOS,
|
||||
'hamming': Image.HAMMING
|
||||
}
|
||||
|
||||
|
||||
def imresize(img,
|
||||
size,
|
||||
return_scale=False,
|
||||
interpolation='bilinear',
|
||||
out=None,
|
||||
backend=None):
|
||||
"""Resize image to a given size.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
size (tuple[int]): Target size (w, h).
|
||||
return_scale (bool): Whether to return `w_scale` and `h_scale`.
|
||||
interpolation (str): Interpolation method, accepted values are
|
||||
"nearest", "bilinear", "bicubic", "area", "lanczos" for 'cv2'
|
||||
backend, "nearest", "bilinear" for 'pillow' backend.
|
||||
out (ndarray): The output destination.
|
||||
backend (str | None): The image resize backend type. Options are `cv2`,
|
||||
`pillow`, `None`. If backend is None, the global imread_backend
|
||||
specified by ``mmcv.use_backend()`` will be used. Default: None.
|
||||
|
||||
Returns:
|
||||
tuple | ndarray: (`resized_img`, `w_scale`, `h_scale`) or
|
||||
`resized_img`.
|
||||
"""
|
||||
h, w = img.shape[:2]
|
||||
if backend is None:
|
||||
backend = imread_backend
|
||||
if backend not in ['cv2', 'pillow']:
|
||||
raise ValueError(f'backend: {backend} is not supported for resize.'
|
||||
f"Supported backends are 'cv2', 'pillow'")
|
||||
|
||||
if backend == 'pillow':
|
||||
assert img.dtype == np.uint8, 'Pillow backend only support uint8 type'
|
||||
pil_image = Image.fromarray(img)
|
||||
pil_image = pil_image.resize(size, pillow_interp_codes[interpolation])
|
||||
resized_img = np.array(pil_image)
|
||||
else:
|
||||
resized_img = cv2.resize(
|
||||
img, size, dst=out, interpolation=cv2_interp_codes[interpolation])
|
||||
if not return_scale:
|
||||
return resized_img
|
||||
else:
|
||||
w_scale = size[0] / w
|
||||
h_scale = size[1] / h
|
||||
return resized_img, w_scale, h_scale
|
||||
|
||||
|
||||
def imresize_to_multiple(img,
|
||||
divisor,
|
||||
size=None,
|
||||
scale_factor=None,
|
||||
keep_ratio=False,
|
||||
return_scale=False,
|
||||
interpolation='bilinear',
|
||||
out=None,
|
||||
backend=None):
|
||||
"""Resize image according to a given size or scale factor and then rounds
|
||||
up the the resized or rescaled image size to the nearest value that can be
|
||||
divided by the divisor.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
divisor (int | tuple): Resized image size will be a multiple of
|
||||
divisor. If divisor is a tuple, divisor should be
|
||||
(w_divisor, h_divisor).
|
||||
size (None | int | tuple[int]): Target size (w, h). Default: None.
|
||||
scale_factor (None | float | tuple[float]): Multiplier for spatial
|
||||
size. Should match input size if it is a tuple and the 2D style is
|
||||
(w_scale_factor, h_scale_factor). Default: None.
|
||||
keep_ratio (bool): Whether to keep the aspect ratio when resizing the
|
||||
image. Default: False.
|
||||
return_scale (bool): Whether to return `w_scale` and `h_scale`.
|
||||
interpolation (str): Interpolation method, accepted values are
|
||||
"nearest", "bilinear", "bicubic", "area", "lanczos" for 'cv2'
|
||||
backend, "nearest", "bilinear" for 'pillow' backend.
|
||||
out (ndarray): The output destination.
|
||||
backend (str | None): The image resize backend type. Options are `cv2`,
|
||||
`pillow`, `None`. If backend is None, the global imread_backend
|
||||
specified by ``mmcv.use_backend()`` will be used. Default: None.
|
||||
|
||||
Returns:
|
||||
tuple | ndarray: (`resized_img`, `w_scale`, `h_scale`) or
|
||||
`resized_img`.
|
||||
"""
|
||||
h, w = img.shape[:2]
|
||||
if size is not None and scale_factor is not None:
|
||||
raise ValueError('only one of size or scale_factor should be defined')
|
||||
elif size is None and scale_factor is None:
|
||||
raise ValueError('one of size or scale_factor should be defined')
|
||||
elif size is not None:
|
||||
size = to_2tuple(size)
|
||||
if keep_ratio:
|
||||
size = rescale_size((w, h), size, return_scale=False)
|
||||
else:
|
||||
size = _scale_size((w, h), scale_factor)
|
||||
|
||||
divisor = to_2tuple(divisor)
|
||||
size = tuple([int(np.ceil(s / d)) * d for s, d in zip(size, divisor)])
|
||||
resized_img, w_scale, h_scale = imresize(
|
||||
img,
|
||||
size,
|
||||
return_scale=True,
|
||||
interpolation=interpolation,
|
||||
out=out,
|
||||
backend=backend)
|
||||
if return_scale:
|
||||
return resized_img, w_scale, h_scale
|
||||
else:
|
||||
return resized_img
|
||||
|
||||
|
||||
def imresize_like(img,
|
||||
dst_img,
|
||||
return_scale=False,
|
||||
interpolation='bilinear',
|
||||
backend=None):
|
||||
"""Resize image to the same size of a given image.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
dst_img (ndarray): The target image.
|
||||
return_scale (bool): Whether to return `w_scale` and `h_scale`.
|
||||
interpolation (str): Same as :func:`resize`.
|
||||
backend (str | None): Same as :func:`resize`.
|
||||
|
||||
Returns:
|
||||
tuple or ndarray: (`resized_img`, `w_scale`, `h_scale`) or
|
||||
`resized_img`.
|
||||
"""
|
||||
h, w = dst_img.shape[:2]
|
||||
return imresize(img, (w, h), return_scale, interpolation, backend=backend)
|
||||
|
||||
|
||||
def rescale_size(old_size, scale, return_scale=False):
|
||||
"""Calculate the new size to be rescaled to.
|
||||
|
||||
Args:
|
||||
old_size (tuple[int]): The old size (w, h) of image.
|
||||
scale (float | tuple[int]): The scaling factor or maximum size.
|
||||
If it is a float number, then the image will be rescaled by this
|
||||
factor, else if it is a tuple of 2 integers, then the image will
|
||||
be rescaled as large as possible within the scale.
|
||||
return_scale (bool): Whether to return the scaling factor besides the
|
||||
rescaled image size.
|
||||
|
||||
Returns:
|
||||
tuple[int]: The new rescaled image size.
|
||||
"""
|
||||
w, h = old_size
|
||||
if isinstance(scale, (float, int)):
|
||||
if scale <= 0:
|
||||
raise ValueError(f'Invalid scale {scale}, must be positive.')
|
||||
scale_factor = scale
|
||||
elif isinstance(scale, tuple):
|
||||
max_long_edge = max(scale)
|
||||
max_short_edge = min(scale)
|
||||
scale_factor = min(max_long_edge / max(h, w),
|
||||
max_short_edge / min(h, w))
|
||||
else:
|
||||
raise TypeError(
|
||||
f'Scale must be a number or tuple of int, but got {type(scale)}')
|
||||
|
||||
new_size = _scale_size((w, h), scale_factor)
|
||||
|
||||
if return_scale:
|
||||
return new_size, scale_factor
|
||||
else:
|
||||
return new_size
|
||||
|
||||
|
||||
def imrescale(img,
|
||||
scale,
|
||||
return_scale=False,
|
||||
interpolation='bilinear',
|
||||
backend=None):
|
||||
"""Resize image while keeping the aspect ratio.
|
||||
|
||||
Args:
|
||||
img (ndarray): The input image.
|
||||
scale (float | tuple[int]): The scaling factor or maximum size.
|
||||
If it is a float number, then the image will be rescaled by this
|
||||
factor, else if it is a tuple of 2 integers, then the image will
|
||||
be rescaled as large as possible within the scale.
|
||||
return_scale (bool): Whether to return the scaling factor besides the
|
||||
rescaled image.
|
||||
interpolation (str): Same as :func:`resize`.
|
||||
backend (str | None): Same as :func:`resize`.
|
||||
|
||||
Returns:
|
||||
ndarray: The rescaled image.
|
||||
"""
|
||||
h, w = img.shape[:2]
|
||||
new_size, scale_factor = rescale_size((w, h), scale, return_scale=True)
|
||||
rescaled_img = imresize(
|
||||
img, new_size, interpolation=interpolation, backend=backend)
|
||||
if return_scale:
|
||||
return rescaled_img, scale_factor
|
||||
else:
|
||||
return rescaled_img
|
||||
|
||||
|
||||
def imflip(img, direction='horizontal'):
|
||||
"""Flip an image horizontally or vertically.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be flipped.
|
||||
direction (str): The flip direction, either "horizontal" or
|
||||
"vertical" or "diagonal".
|
||||
|
||||
Returns:
|
||||
ndarray: The flipped image.
|
||||
"""
|
||||
assert direction in ['horizontal', 'vertical', 'diagonal']
|
||||
if direction == 'horizontal':
|
||||
return np.flip(img, axis=1)
|
||||
elif direction == 'vertical':
|
||||
return np.flip(img, axis=0)
|
||||
else:
|
||||
return np.flip(img, axis=(0, 1))
|
||||
|
||||
|
||||
def imflip_(img, direction='horizontal'):
|
||||
"""Inplace flip an image horizontally or vertically.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be flipped.
|
||||
direction (str): The flip direction, either "horizontal" or
|
||||
"vertical" or "diagonal".
|
||||
|
||||
Returns:
|
||||
ndarray: The flipped image (inplace).
|
||||
"""
|
||||
assert direction in ['horizontal', 'vertical', 'diagonal']
|
||||
if direction == 'horizontal':
|
||||
return cv2.flip(img, 1, img)
|
||||
elif direction == 'vertical':
|
||||
return cv2.flip(img, 0, img)
|
||||
else:
|
||||
return cv2.flip(img, -1, img)
|
||||
|
||||
|
||||
def imrotate(img,
|
||||
angle,
|
||||
center=None,
|
||||
scale=1.0,
|
||||
border_value=0,
|
||||
interpolation='bilinear',
|
||||
auto_bound=False):
|
||||
"""Rotate an image.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be rotated.
|
||||
angle (float): Rotation angle in degrees, positive values mean
|
||||
clockwise rotation.
|
||||
center (tuple[float], optional): Center point (w, h) of the rotation in
|
||||
the source image. If not specified, the center of the image will be
|
||||
used.
|
||||
scale (float): Isotropic scale factor.
|
||||
border_value (int): Border value.
|
||||
interpolation (str): Same as :func:`resize`.
|
||||
auto_bound (bool): Whether to adjust the image size to cover the whole
|
||||
rotated image.
|
||||
|
||||
Returns:
|
||||
ndarray: The rotated image.
|
||||
"""
|
||||
if center is not None and auto_bound:
|
||||
raise ValueError('`auto_bound` conflicts with `center`')
|
||||
h, w = img.shape[:2]
|
||||
if center is None:
|
||||
center = ((w - 1) * 0.5, (h - 1) * 0.5)
|
||||
assert isinstance(center, tuple)
|
||||
|
||||
matrix = cv2.getRotationMatrix2D(center, -angle, scale)
|
||||
if auto_bound:
|
||||
cos = np.abs(matrix[0, 0])
|
||||
sin = np.abs(matrix[0, 1])
|
||||
new_w = h * sin + w * cos
|
||||
new_h = h * cos + w * sin
|
||||
matrix[0, 2] += (new_w - w) * 0.5
|
||||
matrix[1, 2] += (new_h - h) * 0.5
|
||||
w = int(np.round(new_w))
|
||||
h = int(np.round(new_h))
|
||||
rotated = cv2.warpAffine(
|
||||
img,
|
||||
matrix, (w, h),
|
||||
flags=cv2_interp_codes[interpolation],
|
||||
borderValue=border_value)
|
||||
return rotated
|
||||
|
||||
|
||||
def bbox_clip(bboxes, img_shape):
|
||||
"""Clip bboxes to fit the image shape.
|
||||
|
||||
Args:
|
||||
bboxes (ndarray): Shape (..., 4*k)
|
||||
img_shape (tuple[int]): (height, width) of the image.
|
||||
|
||||
Returns:
|
||||
ndarray: Clipped bboxes.
|
||||
"""
|
||||
assert bboxes.shape[-1] % 4 == 0
|
||||
cmin = np.empty(bboxes.shape[-1], dtype=bboxes.dtype)
|
||||
cmin[0::2] = img_shape[1] - 1
|
||||
cmin[1::2] = img_shape[0] - 1
|
||||
clipped_bboxes = np.maximum(np.minimum(bboxes, cmin), 0)
|
||||
return clipped_bboxes
|
||||
|
||||
|
||||
def bbox_scaling(bboxes, scale, clip_shape=None):
|
||||
"""Scaling bboxes w.r.t the box center.
|
||||
|
||||
Args:
|
||||
bboxes (ndarray): Shape(..., 4).
|
||||
scale (float): Scaling factor.
|
||||
clip_shape (tuple[int], optional): If specified, bboxes that exceed the
|
||||
boundary will be clipped according to the given shape (h, w).
|
||||
|
||||
Returns:
|
||||
ndarray: Scaled bboxes.
|
||||
"""
|
||||
if float(scale) == 1.0:
|
||||
scaled_bboxes = bboxes.copy()
|
||||
else:
|
||||
w = bboxes[..., 2] - bboxes[..., 0] + 1
|
||||
h = bboxes[..., 3] - bboxes[..., 1] + 1
|
||||
dw = (w * (scale - 1)) * 0.5
|
||||
dh = (h * (scale - 1)) * 0.5
|
||||
scaled_bboxes = bboxes + np.stack((-dw, -dh, dw, dh), axis=-1)
|
||||
if clip_shape is not None:
|
||||
return bbox_clip(scaled_bboxes, clip_shape)
|
||||
else:
|
||||
return scaled_bboxes
|
||||
|
||||
|
||||
def imcrop(img, bboxes, scale=1.0, pad_fill=None):
|
||||
"""Crop image patches.
|
||||
|
||||
3 steps: scale the bboxes -> clip bboxes -> crop and pad.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be cropped.
|
||||
bboxes (ndarray): Shape (k, 4) or (4, ), location of cropped bboxes.
|
||||
scale (float, optional): Scale ratio of bboxes, the default value
|
||||
1.0 means no padding.
|
||||
pad_fill (Number | list[Number]): Value to be filled for padding.
|
||||
Default: None, which means no padding.
|
||||
|
||||
Returns:
|
||||
list[ndarray] | ndarray: The cropped image patches.
|
||||
"""
|
||||
chn = 1 if img.ndim == 2 else img.shape[2]
|
||||
if pad_fill is not None:
|
||||
if isinstance(pad_fill, (int, float)):
|
||||
pad_fill = [pad_fill for _ in range(chn)]
|
||||
assert len(pad_fill) == chn
|
||||
|
||||
_bboxes = bboxes[None, ...] if bboxes.ndim == 1 else bboxes
|
||||
scaled_bboxes = bbox_scaling(_bboxes, scale).astype(np.int32)
|
||||
clipped_bbox = bbox_clip(scaled_bboxes, img.shape)
|
||||
|
||||
patches = []
|
||||
for i in range(clipped_bbox.shape[0]):
|
||||
x1, y1, x2, y2 = tuple(clipped_bbox[i, :])
|
||||
if pad_fill is None:
|
||||
patch = img[y1:y2 + 1, x1:x2 + 1, ...]
|
||||
else:
|
||||
_x1, _y1, _x2, _y2 = tuple(scaled_bboxes[i, :])
|
||||
if chn == 1:
|
||||
patch_shape = (_y2 - _y1 + 1, _x2 - _x1 + 1)
|
||||
else:
|
||||
patch_shape = (_y2 - _y1 + 1, _x2 - _x1 + 1, chn)
|
||||
patch = np.array(
|
||||
pad_fill, dtype=img.dtype) * np.ones(
|
||||
patch_shape, dtype=img.dtype)
|
||||
x_start = 0 if _x1 >= 0 else -_x1
|
||||
y_start = 0 if _y1 >= 0 else -_y1
|
||||
w = x2 - x1 + 1
|
||||
h = y2 - y1 + 1
|
||||
patch[y_start:y_start + h, x_start:x_start + w,
|
||||
...] = img[y1:y1 + h, x1:x1 + w, ...]
|
||||
patches.append(patch)
|
||||
|
||||
if bboxes.ndim == 1:
|
||||
return patches[0]
|
||||
else:
|
||||
return patches
|
||||
|
||||
|
||||
def impad(img,
|
||||
*,
|
||||
shape=None,
|
||||
padding=None,
|
||||
pad_val=0,
|
||||
padding_mode='constant'):
|
||||
"""Pad the given image to a certain shape or pad on all sides with
|
||||
specified padding mode and padding value.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be padded.
|
||||
shape (tuple[int]): Expected padding shape (h, w). Default: None.
|
||||
padding (int or tuple[int]): Padding on each border. If a single int is
|
||||
provided this is used to pad all borders. If tuple of length 2 is
|
||||
provided this is the padding on left/right and top/bottom
|
||||
respectively. If a tuple of length 4 is provided this is the
|
||||
padding for the left, top, right and bottom borders respectively.
|
||||
Default: None. Note that `shape` and `padding` can not be both
|
||||
set.
|
||||
pad_val (Number | Sequence[Number]): Values to be filled in padding
|
||||
areas when padding_mode is 'constant'. Default: 0.
|
||||
padding_mode (str): Type of padding. Should be: constant, edge,
|
||||
reflect or symmetric. Default: constant.
|
||||
|
||||
- constant: pads with a constant value, this value is specified
|
||||
with pad_val.
|
||||
- edge: pads with the last value at the edge of the image.
|
||||
- reflect: pads with reflection of image without repeating the
|
||||
last value on the edge. For example, padding [1, 2, 3, 4]
|
||||
with 2 elements on both sides in reflect mode will result
|
||||
in [3, 2, 1, 2, 3, 4, 3, 2].
|
||||
- symmetric: pads with reflection of image repeating the last
|
||||
value on the edge. For example, padding [1, 2, 3, 4] with
|
||||
2 elements on both sides in symmetric mode will result in
|
||||
[2, 1, 1, 2, 3, 4, 4, 3]
|
||||
|
||||
Returns:
|
||||
ndarray: The padded image.
|
||||
"""
|
||||
|
||||
assert (shape is not None) ^ (padding is not None)
|
||||
if shape is not None:
|
||||
padding = (0, 0, shape[1] - img.shape[1], shape[0] - img.shape[0])
|
||||
|
||||
# check pad_val
|
||||
if isinstance(pad_val, tuple):
|
||||
assert len(pad_val) == img.shape[-1]
|
||||
elif not isinstance(pad_val, numbers.Number):
|
||||
raise TypeError('pad_val must be a int or a tuple. '
|
||||
f'But received {type(pad_val)}')
|
||||
|
||||
# check padding
|
||||
if isinstance(padding, tuple) and len(padding) in [2, 4]:
|
||||
if len(padding) == 2:
|
||||
padding = (padding[0], padding[1], padding[0], padding[1])
|
||||
elif isinstance(padding, numbers.Number):
|
||||
padding = (padding, padding, padding, padding)
|
||||
else:
|
||||
raise ValueError('Padding must be a int or a 2, or 4 element tuple.'
|
||||
f'But received {padding}')
|
||||
|
||||
# check padding mode
|
||||
assert padding_mode in ['constant', 'edge', 'reflect', 'symmetric']
|
||||
|
||||
border_type = {
|
||||
'constant': cv2.BORDER_CONSTANT,
|
||||
'edge': cv2.BORDER_REPLICATE,
|
||||
'reflect': cv2.BORDER_REFLECT_101,
|
||||
'symmetric': cv2.BORDER_REFLECT
|
||||
}
|
||||
img = cv2.copyMakeBorder(
|
||||
img,
|
||||
padding[1],
|
||||
padding[3],
|
||||
padding[0],
|
||||
padding[2],
|
||||
border_type[padding_mode],
|
||||
value=pad_val)
|
||||
|
||||
return img
|
||||
|
||||
|
||||
def impad_to_multiple(img, divisor, pad_val=0):
|
||||
"""Pad an image to ensure each edge to be multiple to some number.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be padded.
|
||||
divisor (int): Padded image edges will be multiple to divisor.
|
||||
pad_val (Number | Sequence[Number]): Same as :func:`impad`.
|
||||
|
||||
Returns:
|
||||
ndarray: The padded image.
|
||||
"""
|
||||
pad_h = int(np.ceil(img.shape[0] / divisor)) * divisor
|
||||
pad_w = int(np.ceil(img.shape[1] / divisor)) * divisor
|
||||
return impad(img, shape=(pad_h, pad_w), pad_val=pad_val)
|
||||
|
||||
|
||||
def cutout(img, shape, pad_val=0):
|
||||
"""Randomly cut out a rectangle from the original img.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be cutout.
|
||||
shape (int | tuple[int]): Expected cutout shape (h, w). If given as a
|
||||
int, the value will be used for both h and w.
|
||||
pad_val (int | float | tuple[int | float]): Values to be filled in the
|
||||
cut area. Defaults to 0.
|
||||
|
||||
Returns:
|
||||
ndarray: The cutout image.
|
||||
"""
|
||||
|
||||
channels = 1 if img.ndim == 2 else img.shape[2]
|
||||
if isinstance(shape, int):
|
||||
cut_h, cut_w = shape, shape
|
||||
else:
|
||||
assert isinstance(shape, tuple) and len(shape) == 2, \
|
||||
f'shape must be a int or a tuple with length 2, but got type ' \
|
||||
f'{type(shape)} instead.'
|
||||
cut_h, cut_w = shape
|
||||
if isinstance(pad_val, (int, float)):
|
||||
pad_val = tuple([pad_val] * channels)
|
||||
elif isinstance(pad_val, tuple):
|
||||
assert len(pad_val) == channels, \
|
||||
'Expected the num of elements in tuple equals the channels' \
|
||||
'of input image. Found {} vs {}'.format(
|
||||
len(pad_val), channels)
|
||||
else:
|
||||
raise TypeError(f'Invalid type {type(pad_val)} for `pad_val`')
|
||||
|
||||
img_h, img_w = img.shape[:2]
|
||||
y0 = np.random.uniform(img_h)
|
||||
x0 = np.random.uniform(img_w)
|
||||
|
||||
y1 = int(max(0, y0 - cut_h / 2.))
|
||||
x1 = int(max(0, x0 - cut_w / 2.))
|
||||
y2 = min(img_h, y1 + cut_h)
|
||||
x2 = min(img_w, x1 + cut_w)
|
||||
|
||||
if img.ndim == 2:
|
||||
patch_shape = (y2 - y1, x2 - x1)
|
||||
else:
|
||||
patch_shape = (y2 - y1, x2 - x1, channels)
|
||||
|
||||
img_cutout = img.copy()
|
||||
patch = np.array(
|
||||
pad_val, dtype=img.dtype) * np.ones(
|
||||
patch_shape, dtype=img.dtype)
|
||||
img_cutout[y1:y2, x1:x2, ...] = patch
|
||||
|
||||
return img_cutout
|
||||
|
||||
|
||||
def _get_shear_matrix(magnitude, direction='horizontal'):
|
||||
"""Generate the shear matrix for transformation.
|
||||
|
||||
Args:
|
||||
magnitude (int | float): The magnitude used for shear.
|
||||
direction (str): The flip direction, either "horizontal"
|
||||
or "vertical".
|
||||
|
||||
Returns:
|
||||
ndarray: The shear matrix with dtype float32.
|
||||
"""
|
||||
if direction == 'horizontal':
|
||||
shear_matrix = np.float32([[1, magnitude, 0], [0, 1, 0]])
|
||||
elif direction == 'vertical':
|
||||
shear_matrix = np.float32([[1, 0, 0], [magnitude, 1, 0]])
|
||||
return shear_matrix
|
||||
|
||||
|
||||
def imshear(img,
|
||||
magnitude,
|
||||
direction='horizontal',
|
||||
border_value=0,
|
||||
interpolation='bilinear'):
|
||||
"""Shear an image.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be sheared with format (h, w)
|
||||
or (h, w, c).
|
||||
magnitude (int | float): The magnitude used for shear.
|
||||
direction (str): The flip direction, either "horizontal"
|
||||
or "vertical".
|
||||
border_value (int | tuple[int]): Value used in case of a
|
||||
constant border.
|
||||
interpolation (str): Same as :func:`resize`.
|
||||
|
||||
Returns:
|
||||
ndarray: The sheared image.
|
||||
"""
|
||||
assert direction in ['horizontal',
|
||||
'vertical'], f'Invalid direction: {direction}'
|
||||
height, width = img.shape[:2]
|
||||
if img.ndim == 2:
|
||||
channels = 1
|
||||
elif img.ndim == 3:
|
||||
channels = img.shape[-1]
|
||||
if isinstance(border_value, int):
|
||||
border_value = tuple([border_value] * channels)
|
||||
elif isinstance(border_value, tuple):
|
||||
assert len(border_value) == channels, \
|
||||
'Expected the num of elements in tuple equals the channels' \
|
||||
'of input image. Found {} vs {}'.format(
|
||||
len(border_value), channels)
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Invalid type {type(border_value)} for `border_value`')
|
||||
shear_matrix = _get_shear_matrix(magnitude, direction)
|
||||
sheared = cv2.warpAffine(
|
||||
img,
|
||||
shear_matrix,
|
||||
(width, height),
|
||||
# Note case when the number elements in `border_value`
|
||||
# greater than 3 (e.g. shearing masks whose channels large
|
||||
# than 3) will raise TypeError in `cv2.warpAffine`.
|
||||
# Here simply slice the first 3 values in `border_value`.
|
||||
borderValue=border_value[:3],
|
||||
flags=cv2_interp_codes[interpolation])
|
||||
return sheared
|
||||
|
||||
|
||||
def _get_translate_matrix(offset, direction='horizontal'):
|
||||
"""Generate the translate matrix.
|
||||
|
||||
Args:
|
||||
offset (int | float): The offset used for translate.
|
||||
direction (str): The translate direction, either
|
||||
"horizontal" or "vertical".
|
||||
|
||||
Returns:
|
||||
ndarray: The translate matrix with dtype float32.
|
||||
"""
|
||||
if direction == 'horizontal':
|
||||
translate_matrix = np.float32([[1, 0, offset], [0, 1, 0]])
|
||||
elif direction == 'vertical':
|
||||
translate_matrix = np.float32([[1, 0, 0], [0, 1, offset]])
|
||||
return translate_matrix
|
||||
|
||||
|
||||
def imtranslate(img,
|
||||
offset,
|
||||
direction='horizontal',
|
||||
border_value=0,
|
||||
interpolation='bilinear'):
|
||||
"""Translate an image.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be translated with format
|
||||
(h, w) or (h, w, c).
|
||||
offset (int | float): The offset used for translate.
|
||||
direction (str): The translate direction, either "horizontal"
|
||||
or "vertical".
|
||||
border_value (int | tuple[int]): Value used in case of a
|
||||
constant border.
|
||||
interpolation (str): Same as :func:`resize`.
|
||||
|
||||
Returns:
|
||||
ndarray: The translated image.
|
||||
"""
|
||||
assert direction in ['horizontal',
|
||||
'vertical'], f'Invalid direction: {direction}'
|
||||
height, width = img.shape[:2]
|
||||
if img.ndim == 2:
|
||||
channels = 1
|
||||
elif img.ndim == 3:
|
||||
channels = img.shape[-1]
|
||||
if isinstance(border_value, int):
|
||||
border_value = tuple([border_value] * channels)
|
||||
elif isinstance(border_value, tuple):
|
||||
assert len(border_value) == channels, \
|
||||
'Expected the num of elements in tuple equals the channels' \
|
||||
'of input image. Found {} vs {}'.format(
|
||||
len(border_value), channels)
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Invalid type {type(border_value)} for `border_value`.')
|
||||
translate_matrix = _get_translate_matrix(offset, direction)
|
||||
translated = cv2.warpAffine(
|
||||
img,
|
||||
translate_matrix,
|
||||
(width, height),
|
||||
# Note case when the number elements in `border_value`
|
||||
# greater than 3 (e.g. translating masks whose channels
|
||||
# large than 3) will raise TypeError in `cv2.warpAffine`.
|
||||
# Here simply slice the first 3 values in `border_value`.
|
||||
borderValue=border_value[:3],
|
||||
flags=cv2_interp_codes[interpolation])
|
||||
return translated
|
||||
@@ -0,0 +1,258 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import io
|
||||
import os.path as osp
|
||||
from pathlib import Path
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
from cv2 import (IMREAD_COLOR, IMREAD_GRAYSCALE, IMREAD_IGNORE_ORIENTATION,
|
||||
IMREAD_UNCHANGED)
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import check_file_exist, is_str, mkdir_or_exist
|
||||
|
||||
try:
|
||||
from turbojpeg import TJCS_RGB, TJPF_BGR, TJPF_GRAY, TurboJPEG
|
||||
except ImportError:
|
||||
TJCS_RGB = TJPF_GRAY = TJPF_BGR = TurboJPEG = None
|
||||
|
||||
try:
|
||||
from PIL import Image, ImageOps
|
||||
except ImportError:
|
||||
Image = None
|
||||
|
||||
try:
|
||||
import tifffile
|
||||
except ImportError:
|
||||
tifffile = None
|
||||
|
||||
jpeg = None
|
||||
supported_backends = ['cv2', 'turbojpeg', 'pillow', 'tifffile']
|
||||
|
||||
imread_flags = {
|
||||
'color': IMREAD_COLOR,
|
||||
'grayscale': IMREAD_GRAYSCALE,
|
||||
'unchanged': IMREAD_UNCHANGED,
|
||||
'color_ignore_orientation': IMREAD_IGNORE_ORIENTATION | IMREAD_COLOR,
|
||||
'grayscale_ignore_orientation':
|
||||
IMREAD_IGNORE_ORIENTATION | IMREAD_GRAYSCALE
|
||||
}
|
||||
|
||||
imread_backend = 'cv2'
|
||||
|
||||
|
||||
def use_backend(backend):
|
||||
"""Select a backend for image decoding.
|
||||
|
||||
Args:
|
||||
backend (str): The image decoding backend type. Options are `cv2`,
|
||||
`pillow`, `turbojpeg` (see https://github.com/lilohuang/PyTurboJPEG)
|
||||
and `tifffile`. `turbojpeg` is faster but it only supports `.jpeg`
|
||||
file format.
|
||||
"""
|
||||
assert backend in supported_backends
|
||||
global imread_backend
|
||||
imread_backend = backend
|
||||
if imread_backend == 'turbojpeg':
|
||||
if TurboJPEG is None:
|
||||
raise ImportError('`PyTurboJPEG` is not installed')
|
||||
global jpeg
|
||||
if jpeg is None:
|
||||
jpeg = TurboJPEG()
|
||||
elif imread_backend == 'pillow':
|
||||
if Image is None:
|
||||
raise ImportError('`Pillow` is not installed')
|
||||
elif imread_backend == 'tifffile':
|
||||
if tifffile is None:
|
||||
raise ImportError('`tifffile` is not installed')
|
||||
|
||||
|
||||
def _jpegflag(flag='color', channel_order='bgr'):
|
||||
channel_order = channel_order.lower()
|
||||
if channel_order not in ['rgb', 'bgr']:
|
||||
raise ValueError('channel order must be either "rgb" or "bgr"')
|
||||
|
||||
if flag == 'color':
|
||||
if channel_order == 'bgr':
|
||||
return TJPF_BGR
|
||||
elif channel_order == 'rgb':
|
||||
return TJCS_RGB
|
||||
elif flag == 'grayscale':
|
||||
return TJPF_GRAY
|
||||
else:
|
||||
raise ValueError('flag must be "color" or "grayscale"')
|
||||
|
||||
|
||||
def _pillow2array(img, flag='color', channel_order='bgr'):
|
||||
"""Convert a pillow image to numpy array.
|
||||
|
||||
Args:
|
||||
img (:obj:`PIL.Image.Image`): The image loaded using PIL
|
||||
flag (str): Flags specifying the color type of a loaded image,
|
||||
candidates are 'color', 'grayscale' and 'unchanged'.
|
||||
Default to 'color'.
|
||||
channel_order (str): The channel order of the output image array,
|
||||
candidates are 'bgr' and 'rgb'. Default to 'bgr'.
|
||||
|
||||
Returns:
|
||||
np.ndarray: The converted numpy array
|
||||
"""
|
||||
channel_order = channel_order.lower()
|
||||
if channel_order not in ['rgb', 'bgr']:
|
||||
raise ValueError('channel order must be either "rgb" or "bgr"')
|
||||
|
||||
if flag == 'unchanged':
|
||||
array = np.array(img)
|
||||
if array.ndim >= 3 and array.shape[2] >= 3: # color image
|
||||
array[:, :, :3] = array[:, :, (2, 1, 0)] # RGB to BGR
|
||||
else:
|
||||
# Handle exif orientation tag
|
||||
if flag in ['color', 'grayscale']:
|
||||
img = ImageOps.exif_transpose(img)
|
||||
# If the image mode is not 'RGB', convert it to 'RGB' first.
|
||||
if img.mode != 'RGB':
|
||||
if img.mode != 'LA':
|
||||
# Most formats except 'LA' can be directly converted to RGB
|
||||
img = img.convert('RGB')
|
||||
else:
|
||||
# When the mode is 'LA', the default conversion will fill in
|
||||
# the canvas with black, which sometimes shadows black objects
|
||||
# in the foreground.
|
||||
#
|
||||
# Therefore, a random color (124, 117, 104) is used for canvas
|
||||
img_rgba = img.convert('RGBA')
|
||||
img = Image.new('RGB', img_rgba.size, (124, 117, 104))
|
||||
img.paste(img_rgba, mask=img_rgba.split()[3]) # 3 is alpha
|
||||
if flag in ['color', 'color_ignore_orientation']:
|
||||
array = np.array(img)
|
||||
if channel_order != 'rgb':
|
||||
array = array[:, :, ::-1] # RGB to BGR
|
||||
elif flag in ['grayscale', 'grayscale_ignore_orientation']:
|
||||
img = img.convert('L')
|
||||
array = np.array(img)
|
||||
else:
|
||||
raise ValueError(
|
||||
'flag must be "color", "grayscale", "unchanged", '
|
||||
f'"color_ignore_orientation" or "grayscale_ignore_orientation"'
|
||||
f' but got {flag}')
|
||||
return array
|
||||
|
||||
|
||||
def imread(img_or_path, flag='color', channel_order='bgr', backend=None):
|
||||
"""Read an image.
|
||||
|
||||
Args:
|
||||
img_or_path (ndarray or str or Path): Either a numpy array or str or
|
||||
pathlib.Path. If it is a numpy array (loaded image), then
|
||||
it will be returned as is.
|
||||
flag (str): Flags specifying the color type of a loaded image,
|
||||
candidates are `color`, `grayscale`, `unchanged`,
|
||||
`color_ignore_orientation` and `grayscale_ignore_orientation`.
|
||||
By default, `cv2` and `pillow` backend would rotate the image
|
||||
according to its EXIF info unless called with `unchanged` or
|
||||
`*_ignore_orientation` flags. `turbojpeg` and `tifffile` backend
|
||||
always ignore image's EXIF info regardless of the flag.
|
||||
The `turbojpeg` backend only supports `color` and `grayscale`.
|
||||
channel_order (str): Order of channel, candidates are `bgr` and `rgb`.
|
||||
backend (str | None): The image decoding backend type. Options are
|
||||
`cv2`, `pillow`, `turbojpeg`, `tifffile`, `None`.
|
||||
If backend is None, the global imread_backend specified by
|
||||
``mmcv.use_backend()`` will be used. Default: None.
|
||||
|
||||
Returns:
|
||||
ndarray: Loaded image array.
|
||||
"""
|
||||
|
||||
if backend is None:
|
||||
backend = imread_backend
|
||||
if backend not in supported_backends:
|
||||
raise ValueError(f'backend: {backend} is not supported. Supported '
|
||||
"backends are 'cv2', 'turbojpeg', 'pillow'")
|
||||
if isinstance(img_or_path, Path):
|
||||
img_or_path = str(img_or_path)
|
||||
|
||||
if isinstance(img_or_path, np.ndarray):
|
||||
return img_or_path
|
||||
elif is_str(img_or_path):
|
||||
check_file_exist(img_or_path,
|
||||
f'img file does not exist: {img_or_path}')
|
||||
if backend == 'turbojpeg':
|
||||
with open(img_or_path, 'rb') as in_file:
|
||||
img = jpeg.decode(in_file.read(),
|
||||
_jpegflag(flag, channel_order))
|
||||
if img.shape[-1] == 1:
|
||||
img = img[:, :, 0]
|
||||
return img
|
||||
elif backend == 'pillow':
|
||||
img = Image.open(img_or_path)
|
||||
img = _pillow2array(img, flag, channel_order)
|
||||
return img
|
||||
elif backend == 'tifffile':
|
||||
img = tifffile.imread(img_or_path)
|
||||
return img
|
||||
else:
|
||||
flag = imread_flags[flag] if is_str(flag) else flag
|
||||
img = cv2.imread(img_or_path, flag)
|
||||
if flag == IMREAD_COLOR and channel_order == 'rgb':
|
||||
cv2.cvtColor(img, cv2.COLOR_BGR2RGB, img)
|
||||
return img
|
||||
else:
|
||||
raise TypeError('"img" must be a numpy array or a str or '
|
||||
'a pathlib.Path object')
|
||||
|
||||
|
||||
def imfrombytes(content, flag='color', channel_order='bgr', backend=None):
|
||||
"""Read an image from bytes.
|
||||
|
||||
Args:
|
||||
content (bytes): Image bytes got from files or other streams.
|
||||
flag (str): Same as :func:`imread`.
|
||||
backend (str | None): The image decoding backend type. Options are
|
||||
`cv2`, `pillow`, `turbojpeg`, `None`. If backend is None, the
|
||||
global imread_backend specified by ``mmcv.use_backend()`` will be
|
||||
used. Default: None.
|
||||
|
||||
Returns:
|
||||
ndarray: Loaded image array.
|
||||
"""
|
||||
|
||||
if backend is None:
|
||||
backend = imread_backend
|
||||
if backend not in supported_backends:
|
||||
raise ValueError(f'backend: {backend} is not supported. Supported '
|
||||
"backends are 'cv2', 'turbojpeg', 'pillow'")
|
||||
if backend == 'turbojpeg':
|
||||
img = jpeg.decode(content, _jpegflag(flag, channel_order))
|
||||
if img.shape[-1] == 1:
|
||||
img = img[:, :, 0]
|
||||
return img
|
||||
elif backend == 'pillow':
|
||||
buff = io.BytesIO(content)
|
||||
img = Image.open(buff)
|
||||
img = _pillow2array(img, flag, channel_order)
|
||||
return img
|
||||
else:
|
||||
img_np = np.frombuffer(content, np.uint8)
|
||||
flag = imread_flags[flag] if is_str(flag) else flag
|
||||
img = cv2.imdecode(img_np, flag)
|
||||
if flag == IMREAD_COLOR and channel_order == 'rgb':
|
||||
cv2.cvtColor(img, cv2.COLOR_BGR2RGB, img)
|
||||
return img
|
||||
|
||||
|
||||
def imwrite(img, file_path, params=None, auto_mkdir=True):
|
||||
"""Write image to file.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image array to be written.
|
||||
file_path (str): Image file path.
|
||||
params (None or list): Same as opencv :func:`imwrite` interface.
|
||||
auto_mkdir (bool): If the parent folder of `file_path` does not exist,
|
||||
whether to create it automatically.
|
||||
|
||||
Returns:
|
||||
bool: Successful or not.
|
||||
"""
|
||||
if auto_mkdir:
|
||||
dir_name = osp.abspath(osp.dirname(file_path))
|
||||
mkdir_or_exist(dir_name)
|
||||
return cv2.imwrite(file_path, img, params)
|
||||
@@ -0,0 +1,44 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import numpy as np
|
||||
|
||||
import custom_mmpkg.custom_mmcv as mmcv
|
||||
|
||||
try:
|
||||
import torch
|
||||
except ImportError:
|
||||
torch = None
|
||||
|
||||
|
||||
def tensor2imgs(tensor, mean=(0, 0, 0), std=(1, 1, 1), to_rgb=True):
|
||||
"""Convert tensor to 3-channel images.
|
||||
|
||||
Args:
|
||||
tensor (torch.Tensor): Tensor that contains multiple images, shape (
|
||||
N, C, H, W).
|
||||
mean (tuple[float], optional): Mean of images. Defaults to (0, 0, 0).
|
||||
std (tuple[float], optional): Standard deviation of images.
|
||||
Defaults to (1, 1, 1).
|
||||
to_rgb (bool, optional): Whether the tensor was converted to RGB
|
||||
format in the first place. If so, convert it back to BGR.
|
||||
Defaults to True.
|
||||
|
||||
Returns:
|
||||
list[np.ndarray]: A list that contains multiple images.
|
||||
"""
|
||||
|
||||
if torch is None:
|
||||
raise RuntimeError('pytorch is not installed')
|
||||
assert torch.is_tensor(tensor) and tensor.ndim == 4
|
||||
assert len(mean) == 3
|
||||
assert len(std) == 3
|
||||
|
||||
num_imgs = tensor.size(0)
|
||||
mean = np.array(mean, dtype=np.float32)
|
||||
std = np.array(std, dtype=np.float32)
|
||||
imgs = []
|
||||
for img_id in range(num_imgs):
|
||||
img = tensor[img_id, ...].cpu().numpy().transpose(1, 2, 0)
|
||||
img = mmcv.imdenormalize(
|
||||
img, mean, std, to_bgr=to_rgb).astype(np.uint8)
|
||||
imgs.append(np.ascontiguousarray(img))
|
||||
return imgs
|
||||
@@ -0,0 +1,428 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import cv2
|
||||
import numpy as np
|
||||
|
||||
from ..utils import is_tuple_of
|
||||
from .colorspace import bgr2gray, gray2bgr
|
||||
|
||||
|
||||
def imnormalize(img, mean, std, to_rgb=True):
|
||||
"""Normalize an image with mean and std.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be normalized.
|
||||
mean (ndarray): The mean to be used for normalize.
|
||||
std (ndarray): The std to be used for normalize.
|
||||
to_rgb (bool): Whether to convert to rgb.
|
||||
|
||||
Returns:
|
||||
ndarray: The normalized image.
|
||||
"""
|
||||
img = img.copy().astype(np.float32)
|
||||
return imnormalize_(img, mean, std, to_rgb)
|
||||
|
||||
|
||||
def imnormalize_(img, mean, std, to_rgb=True):
|
||||
"""Inplace normalize an image with mean and std.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be normalized.
|
||||
mean (ndarray): The mean to be used for normalize.
|
||||
std (ndarray): The std to be used for normalize.
|
||||
to_rgb (bool): Whether to convert to rgb.
|
||||
|
||||
Returns:
|
||||
ndarray: The normalized image.
|
||||
"""
|
||||
# cv2 inplace normalization does not accept uint8
|
||||
assert img.dtype != np.uint8
|
||||
mean = np.float64(mean.reshape(1, -1))
|
||||
stdinv = 1 / np.float64(std.reshape(1, -1))
|
||||
if to_rgb:
|
||||
cv2.cvtColor(img, cv2.COLOR_BGR2RGB, img) # inplace
|
||||
cv2.subtract(img, mean, img) # inplace
|
||||
cv2.multiply(img, stdinv, img) # inplace
|
||||
return img
|
||||
|
||||
|
||||
def imdenormalize(img, mean, std, to_bgr=True):
|
||||
assert img.dtype != np.uint8
|
||||
mean = mean.reshape(1, -1).astype(np.float64)
|
||||
std = std.reshape(1, -1).astype(np.float64)
|
||||
img = cv2.multiply(img, std) # make a copy
|
||||
cv2.add(img, mean, img) # inplace
|
||||
if to_bgr:
|
||||
cv2.cvtColor(img, cv2.COLOR_RGB2BGR, img) # inplace
|
||||
return img
|
||||
|
||||
|
||||
def iminvert(img):
|
||||
"""Invert (negate) an image.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be inverted.
|
||||
|
||||
Returns:
|
||||
ndarray: The inverted image.
|
||||
"""
|
||||
return np.full_like(img, 255) - img
|
||||
|
||||
|
||||
def solarize(img, thr=128):
|
||||
"""Solarize an image (invert all pixel values above a threshold)
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be solarized.
|
||||
thr (int): Threshold for solarizing (0 - 255).
|
||||
|
||||
Returns:
|
||||
ndarray: The solarized image.
|
||||
"""
|
||||
img = np.where(img < thr, img, 255 - img)
|
||||
return img
|
||||
|
||||
|
||||
def posterize(img, bits):
|
||||
"""Posterize an image (reduce the number of bits for each color channel)
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be posterized.
|
||||
bits (int): Number of bits (1 to 8) to use for posterizing.
|
||||
|
||||
Returns:
|
||||
ndarray: The posterized image.
|
||||
"""
|
||||
shift = 8 - bits
|
||||
img = np.left_shift(np.right_shift(img, shift), shift)
|
||||
return img
|
||||
|
||||
|
||||
def adjust_color(img, alpha=1, beta=None, gamma=0):
|
||||
r"""It blends the source image and its gray image:
|
||||
|
||||
.. math::
|
||||
output = img * alpha + gray\_img * beta + gamma
|
||||
|
||||
Args:
|
||||
img (ndarray): The input source image.
|
||||
alpha (int | float): Weight for the source image. Default 1.
|
||||
beta (int | float): Weight for the converted gray image.
|
||||
If None, it's assigned the value (1 - `alpha`).
|
||||
gamma (int | float): Scalar added to each sum.
|
||||
Same as :func:`cv2.addWeighted`. Default 0.
|
||||
|
||||
Returns:
|
||||
ndarray: Colored image which has the same size and dtype as input.
|
||||
"""
|
||||
gray_img = bgr2gray(img)
|
||||
gray_img = np.tile(gray_img[..., None], [1, 1, 3])
|
||||
if beta is None:
|
||||
beta = 1 - alpha
|
||||
colored_img = cv2.addWeighted(img, alpha, gray_img, beta, gamma)
|
||||
if not colored_img.dtype == np.uint8:
|
||||
# Note when the dtype of `img` is not the default `np.uint8`
|
||||
# (e.g. np.float32), the value in `colored_img` got from cv2
|
||||
# is not guaranteed to be in range [0, 255], so here clip
|
||||
# is needed.
|
||||
colored_img = np.clip(colored_img, 0, 255)
|
||||
return colored_img
|
||||
|
||||
|
||||
def imequalize(img):
|
||||
"""Equalize the image histogram.
|
||||
|
||||
This function applies a non-linear mapping to the input image,
|
||||
in order to create a uniform distribution of grayscale values
|
||||
in the output image.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be equalized.
|
||||
|
||||
Returns:
|
||||
ndarray: The equalized image.
|
||||
"""
|
||||
|
||||
def _scale_channel(im, c):
|
||||
"""Scale the data in the corresponding channel."""
|
||||
im = im[:, :, c]
|
||||
# Compute the histogram of the image channel.
|
||||
histo = np.histogram(im, 256, (0, 255))[0]
|
||||
# For computing the step, filter out the nonzeros.
|
||||
nonzero_histo = histo[histo > 0]
|
||||
step = (np.sum(nonzero_histo) - nonzero_histo[-1]) // 255
|
||||
if not step:
|
||||
lut = np.array(range(256))
|
||||
else:
|
||||
# Compute the cumulative sum, shifted by step // 2
|
||||
# and then normalized by step.
|
||||
lut = (np.cumsum(histo) + (step // 2)) // step
|
||||
# Shift lut, prepending with 0.
|
||||
lut = np.concatenate([[0], lut[:-1]], 0)
|
||||
# handle potential integer overflow
|
||||
lut[lut > 255] = 255
|
||||
# If step is zero, return the original image.
|
||||
# Otherwise, index from lut.
|
||||
return np.where(np.equal(step, 0), im, lut[im])
|
||||
|
||||
# Scales each channel independently and then stacks
|
||||
# the result.
|
||||
s1 = _scale_channel(img, 0)
|
||||
s2 = _scale_channel(img, 1)
|
||||
s3 = _scale_channel(img, 2)
|
||||
equalized_img = np.stack([s1, s2, s3], axis=-1)
|
||||
return equalized_img.astype(img.dtype)
|
||||
|
||||
|
||||
def adjust_brightness(img, factor=1.):
|
||||
"""Adjust image brightness.
|
||||
|
||||
This function controls the brightness of an image. An
|
||||
enhancement factor of 0.0 gives a black image.
|
||||
A factor of 1.0 gives the original image. This function
|
||||
blends the source image and the degenerated black image:
|
||||
|
||||
.. math::
|
||||
output = img * factor + degenerated * (1 - factor)
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be brightened.
|
||||
factor (float): A value controls the enhancement.
|
||||
Factor 1.0 returns the original image, lower
|
||||
factors mean less color (brightness, contrast,
|
||||
etc), and higher values more. Default 1.
|
||||
|
||||
Returns:
|
||||
ndarray: The brightened image.
|
||||
"""
|
||||
degenerated = np.zeros_like(img)
|
||||
# Note manually convert the dtype to np.float32, to
|
||||
# achieve as close results as PIL.ImageEnhance.Brightness.
|
||||
# Set beta=1-factor, and gamma=0
|
||||
brightened_img = cv2.addWeighted(
|
||||
img.astype(np.float32), factor, degenerated.astype(np.float32),
|
||||
1 - factor, 0)
|
||||
brightened_img = np.clip(brightened_img, 0, 255)
|
||||
return brightened_img.astype(img.dtype)
|
||||
|
||||
|
||||
def adjust_contrast(img, factor=1.):
|
||||
"""Adjust image contrast.
|
||||
|
||||
This function controls the contrast of an image. An
|
||||
enhancement factor of 0.0 gives a solid grey
|
||||
image. A factor of 1.0 gives the original image. It
|
||||
blends the source image and the degenerated mean image:
|
||||
|
||||
.. math::
|
||||
output = img * factor + degenerated * (1 - factor)
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be contrasted. BGR order.
|
||||
factor (float): Same as :func:`mmcv.adjust_brightness`.
|
||||
|
||||
Returns:
|
||||
ndarray: The contrasted image.
|
||||
"""
|
||||
gray_img = bgr2gray(img)
|
||||
hist = np.histogram(gray_img, 256, (0, 255))[0]
|
||||
mean = round(np.sum(gray_img) / np.sum(hist))
|
||||
degenerated = (np.ones_like(img[..., 0]) * mean).astype(img.dtype)
|
||||
degenerated = gray2bgr(degenerated)
|
||||
contrasted_img = cv2.addWeighted(
|
||||
img.astype(np.float32), factor, degenerated.astype(np.float32),
|
||||
1 - factor, 0)
|
||||
contrasted_img = np.clip(contrasted_img, 0, 255)
|
||||
return contrasted_img.astype(img.dtype)
|
||||
|
||||
|
||||
def auto_contrast(img, cutoff=0):
|
||||
"""Auto adjust image contrast.
|
||||
|
||||
This function maximize (normalize) image contrast by first removing cutoff
|
||||
percent of the lightest and darkest pixels from the histogram and remapping
|
||||
the image so that the darkest pixel becomes black (0), and the lightest
|
||||
becomes white (255).
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be contrasted. BGR order.
|
||||
cutoff (int | float | tuple): The cutoff percent of the lightest and
|
||||
darkest pixels to be removed. If given as tuple, it shall be
|
||||
(low, high). Otherwise, the single value will be used for both.
|
||||
Defaults to 0.
|
||||
|
||||
Returns:
|
||||
ndarray: The contrasted image.
|
||||
"""
|
||||
|
||||
def _auto_contrast_channel(im, c, cutoff):
|
||||
im = im[:, :, c]
|
||||
# Compute the histogram of the image channel.
|
||||
histo = np.histogram(im, 256, (0, 255))[0]
|
||||
# Remove cut-off percent pixels from histo
|
||||
histo_sum = np.cumsum(histo)
|
||||
cut_low = histo_sum[-1] * cutoff[0] // 100
|
||||
cut_high = histo_sum[-1] - histo_sum[-1] * cutoff[1] // 100
|
||||
histo_sum = np.clip(histo_sum, cut_low, cut_high) - cut_low
|
||||
histo = np.concatenate([[histo_sum[0]], np.diff(histo_sum)], 0)
|
||||
|
||||
# Compute mapping
|
||||
low, high = np.nonzero(histo)[0][0], np.nonzero(histo)[0][-1]
|
||||
# If all the values have been cut off, return the origin img
|
||||
if low >= high:
|
||||
return im
|
||||
scale = 255.0 / (high - low)
|
||||
offset = -low * scale
|
||||
lut = np.array(range(256))
|
||||
lut = lut * scale + offset
|
||||
lut = np.clip(lut, 0, 255)
|
||||
return lut[im]
|
||||
|
||||
if isinstance(cutoff, (int, float)):
|
||||
cutoff = (cutoff, cutoff)
|
||||
else:
|
||||
assert isinstance(cutoff, tuple), 'cutoff must be of type int, ' \
|
||||
f'float or tuple, but got {type(cutoff)} instead.'
|
||||
# Auto adjusts contrast for each channel independently and then stacks
|
||||
# the result.
|
||||
s1 = _auto_contrast_channel(img, 0, cutoff)
|
||||
s2 = _auto_contrast_channel(img, 1, cutoff)
|
||||
s3 = _auto_contrast_channel(img, 2, cutoff)
|
||||
contrasted_img = np.stack([s1, s2, s3], axis=-1)
|
||||
return contrasted_img.astype(img.dtype)
|
||||
|
||||
|
||||
def adjust_sharpness(img, factor=1., kernel=None):
|
||||
"""Adjust image sharpness.
|
||||
|
||||
This function controls the sharpness of an image. An
|
||||
enhancement factor of 0.0 gives a blurred image. A
|
||||
factor of 1.0 gives the original image. And a factor
|
||||
of 2.0 gives a sharpened image. It blends the source
|
||||
image and the degenerated mean image:
|
||||
|
||||
.. math::
|
||||
output = img * factor + degenerated * (1 - factor)
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be sharpened. BGR order.
|
||||
factor (float): Same as :func:`mmcv.adjust_brightness`.
|
||||
kernel (np.ndarray, optional): Filter kernel to be applied on the img
|
||||
to obtain the degenerated img. Defaults to None.
|
||||
|
||||
Note:
|
||||
No value sanity check is enforced on the kernel set by users. So with
|
||||
an inappropriate kernel, the ``adjust_sharpness`` may fail to perform
|
||||
the function its name indicates but end up performing whatever
|
||||
transform determined by the kernel.
|
||||
|
||||
Returns:
|
||||
ndarray: The sharpened image.
|
||||
"""
|
||||
|
||||
if kernel is None:
|
||||
# adopted from PIL.ImageFilter.SMOOTH
|
||||
kernel = np.array([[1., 1., 1.], [1., 5., 1.], [1., 1., 1.]]) / 13
|
||||
assert isinstance(kernel, np.ndarray), \
|
||||
f'kernel must be of type np.ndarray, but got {type(kernel)} instead.'
|
||||
assert kernel.ndim == 2, \
|
||||
f'kernel must have a dimension of 2, but got {kernel.ndim} instead.'
|
||||
|
||||
degenerated = cv2.filter2D(img, -1, kernel)
|
||||
sharpened_img = cv2.addWeighted(
|
||||
img.astype(np.float32), factor, degenerated.astype(np.float32),
|
||||
1 - factor, 0)
|
||||
sharpened_img = np.clip(sharpened_img, 0, 255)
|
||||
return sharpened_img.astype(img.dtype)
|
||||
|
||||
|
||||
def adjust_lighting(img, eigval, eigvec, alphastd=0.1, to_rgb=True):
|
||||
"""AlexNet-style PCA jitter.
|
||||
|
||||
This data augmentation is proposed in `ImageNet Classification with Deep
|
||||
Convolutional Neural Networks
|
||||
<https://dl.acm.org/doi/pdf/10.1145/3065386>`_.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be adjusted lighting. BGR order.
|
||||
eigval (ndarray): the eigenvalue of the convariance matrix of pixel
|
||||
values, respectively.
|
||||
eigvec (ndarray): the eigenvector of the convariance matrix of pixel
|
||||
values, respectively.
|
||||
alphastd (float): The standard deviation for distribution of alpha.
|
||||
Defaults to 0.1
|
||||
to_rgb (bool): Whether to convert img to rgb.
|
||||
|
||||
Returns:
|
||||
ndarray: The adjusted image.
|
||||
"""
|
||||
assert isinstance(eigval, np.ndarray) and isinstance(eigvec, np.ndarray), \
|
||||
f'eigval and eigvec should both be of type np.ndarray, got ' \
|
||||
f'{type(eigval)} and {type(eigvec)} instead.'
|
||||
|
||||
assert eigval.ndim == 1 and eigvec.ndim == 2
|
||||
assert eigvec.shape == (3, eigval.shape[0])
|
||||
n_eigval = eigval.shape[0]
|
||||
assert isinstance(alphastd, float), 'alphastd should be of type float, ' \
|
||||
f'got {type(alphastd)} instead.'
|
||||
|
||||
img = img.copy().astype(np.float32)
|
||||
if to_rgb:
|
||||
cv2.cvtColor(img, cv2.COLOR_BGR2RGB, img) # inplace
|
||||
|
||||
alpha = np.random.normal(0, alphastd, n_eigval)
|
||||
alter = eigvec \
|
||||
* np.broadcast_to(alpha.reshape(1, n_eigval), (3, n_eigval)) \
|
||||
* np.broadcast_to(eigval.reshape(1, n_eigval), (3, n_eigval))
|
||||
alter = np.broadcast_to(alter.sum(axis=1).reshape(1, 1, 3), img.shape)
|
||||
img_adjusted = img + alter
|
||||
return img_adjusted
|
||||
|
||||
|
||||
def lut_transform(img, lut_table):
|
||||
"""Transform array by look-up table.
|
||||
|
||||
The function lut_transform fills the output array with values from the
|
||||
look-up table. Indices of the entries are taken from the input array.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be transformed.
|
||||
lut_table (ndarray): look-up table of 256 elements; in case of
|
||||
multi-channel input array, the table should either have a single
|
||||
channel (in this case the same table is used for all channels) or
|
||||
the same number of channels as in the input array.
|
||||
|
||||
Returns:
|
||||
ndarray: The transformed image.
|
||||
"""
|
||||
assert isinstance(img, np.ndarray)
|
||||
assert 0 <= np.min(img) and np.max(img) <= 255
|
||||
assert isinstance(lut_table, np.ndarray)
|
||||
assert lut_table.shape == (256, )
|
||||
|
||||
return cv2.LUT(np.array(img, dtype=np.uint8), lut_table)
|
||||
|
||||
|
||||
def clahe(img, clip_limit=40.0, tile_grid_size=(8, 8)):
|
||||
"""Use CLAHE method to process the image.
|
||||
|
||||
See `ZUIDERVELD,K. Contrast Limited Adaptive Histogram Equalization[J].
|
||||
Graphics Gems, 1994:474-485.` for more information.
|
||||
|
||||
Args:
|
||||
img (ndarray): Image to be processed.
|
||||
clip_limit (float): Threshold for contrast limiting. Default: 40.0.
|
||||
tile_grid_size (tuple[int]): Size of grid for histogram equalization.
|
||||
Input image will be divided into equally sized rectangular tiles.
|
||||
It defines the number of tiles in row and column. Default: (8, 8).
|
||||
|
||||
Returns:
|
||||
ndarray: The processed image.
|
||||
"""
|
||||
assert isinstance(img, np.ndarray)
|
||||
assert img.ndim == 2
|
||||
assert isinstance(clip_limit, (float, int))
|
||||
assert is_tuple_of(tile_grid_size, int)
|
||||
assert len(tile_grid_size) == 2
|
||||
|
||||
clahe = cv2.createCLAHE(clip_limit, tile_grid_size)
|
||||
return clahe.apply(np.array(img, dtype=np.uint8))
|
||||
@@ -0,0 +1,6 @@
|
||||
{
|
||||
"resnet50_caffe": "detectron/resnet50_caffe",
|
||||
"resnet50_caffe_bgr": "detectron2/resnet50_caffe_bgr",
|
||||
"resnet101_caffe": "detectron/resnet101_caffe",
|
||||
"resnet101_caffe_bgr": "detectron2/resnet101_caffe_bgr"
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"vgg11": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg11_batch256_imagenet_20210208-4271cd6c.pth",
|
||||
"vgg13": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg13_batch256_imagenet_20210208-4d1d6080.pth",
|
||||
"vgg16": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg16_batch256_imagenet_20210208-db26f1a5.pth",
|
||||
"vgg19": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg19_batch256_imagenet_20210208-e6920e4a.pth",
|
||||
"vgg11_bn": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg11_bn_batch256_imagenet_20210207-f244902c.pth",
|
||||
"vgg13_bn": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg13_bn_batch256_imagenet_20210207-1a8b7864.pth",
|
||||
"vgg16_bn": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg16_bn_batch256_imagenet_20210208-7e55cd29.pth",
|
||||
"vgg19_bn": "https://download.openmmlab.com/mmclassification/v0/vgg/vgg19_bn_batch256_imagenet_20210208-da620c4f.pth",
|
||||
"resnet18": "https://download.openmmlab.com/mmclassification/v0/resnet/resnet18_batch256_imagenet_20200708-34ab8f90.pth",
|
||||
"resnet34": "https://download.openmmlab.com/mmclassification/v0/resnet/resnet34_batch256_imagenet_20200708-32ffb4f7.pth",
|
||||
"resnet50": "https://download.openmmlab.com/mmclassification/v0/resnet/resnet50_batch256_imagenet_20200708-cfb998bf.pth",
|
||||
"resnet101": "https://download.openmmlab.com/mmclassification/v0/resnet/resnet101_batch256_imagenet_20200708-753f3608.pth",
|
||||
"resnet152": "https://download.openmmlab.com/mmclassification/v0/resnet/resnet152_batch256_imagenet_20200708-ec25b1f9.pth",
|
||||
"resnet50_v1d": "https://download.openmmlab.com/mmclassification/v0/resnet/resnetv1d50_batch256_imagenet_20200708-1ad0ce94.pth",
|
||||
"resnet101_v1d": "https://download.openmmlab.com/mmclassification/v0/resnet/resnetv1d101_batch256_imagenet_20200708-9cb302ef.pth",
|
||||
"resnet152_v1d": "https://download.openmmlab.com/mmclassification/v0/resnet/resnetv1d152_batch256_imagenet_20200708-e79cb6a2.pth",
|
||||
"resnext50_32x4d": "https://download.openmmlab.com/mmclassification/v0/resnext/resnext50_32x4d_b32x8_imagenet_20210429-56066e27.pth",
|
||||
"resnext101_32x4d": "https://download.openmmlab.com/mmclassification/v0/resnext/resnext101_32x4d_b32x8_imagenet_20210506-e0fa3dd5.pth",
|
||||
"resnext101_32x8d": "https://download.openmmlab.com/mmclassification/v0/resnext/resnext101_32x8d_b32x8_imagenet_20210506-23a247d5.pth",
|
||||
"resnext152_32x4d": "https://download.openmmlab.com/mmclassification/v0/resnext/resnext152_32x4d_b32x8_imagenet_20210524-927787be.pth",
|
||||
"se-resnet50": "https://download.openmmlab.com/mmclassification/v0/se-resnet/se-resnet50_batch256_imagenet_20200804-ae206104.pth",
|
||||
"se-resnet101": "https://download.openmmlab.com/mmclassification/v0/se-resnet/se-resnet101_batch256_imagenet_20200804-ba5b51d4.pth",
|
||||
"resnest50": "https://download.openmmlab.com/mmclassification/v0/resnest/resnest50_imagenet_converted-1ebf0afe.pth",
|
||||
"resnest101": "https://download.openmmlab.com/mmclassification/v0/resnest/resnest101_imagenet_converted-032caa52.pth",
|
||||
"resnest200": "https://download.openmmlab.com/mmclassification/v0/resnest/resnest200_imagenet_converted-581a60f2.pth",
|
||||
"resnest269": "https://download.openmmlab.com/mmclassification/v0/resnest/resnest269_imagenet_converted-59930960.pth",
|
||||
"shufflenet_v1": "https://download.openmmlab.com/mmclassification/v0/shufflenet_v1/shufflenet_v1_batch1024_imagenet_20200804-5d6cec73.pth",
|
||||
"shufflenet_v2": "https://download.openmmlab.com/mmclassification/v0/shufflenet_v2/shufflenet_v2_batch1024_imagenet_20200812-5bf4721e.pth",
|
||||
"mobilenet_v2": "https://download.openmmlab.com/mmclassification/v0/mobilenet_v2/mobilenet_v2_batch256_imagenet_20200708-3b2dc3af.pth"
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
{
|
||||
"vgg16_caffe": "https://download.openmmlab.com/pretrain/third_party/vgg16_caffe-292e1171.pth",
|
||||
"detectron/resnet50_caffe": "https://download.openmmlab.com/pretrain/third_party/resnet50_caffe-788b5fa3.pth",
|
||||
"detectron2/resnet50_caffe": "https://download.openmmlab.com/pretrain/third_party/resnet50_msra-5891d200.pth",
|
||||
"detectron/resnet101_caffe": "https://download.openmmlab.com/pretrain/third_party/resnet101_caffe-3ad79236.pth",
|
||||
"detectron2/resnet101_caffe": "https://download.openmmlab.com/pretrain/third_party/resnet101_msra-6cc46731.pth",
|
||||
"detectron2/resnext101_32x8d": "https://download.openmmlab.com/pretrain/third_party/resnext101_32x8d-1516f1aa.pth",
|
||||
"resnext50_32x4d": "https://download.openmmlab.com/pretrain/third_party/resnext50-32x4d-0ab1a123.pth",
|
||||
"resnext101_32x4d": "https://download.openmmlab.com/pretrain/third_party/resnext101_32x4d-a5af3160.pth",
|
||||
"resnext101_64x4d": "https://download.openmmlab.com/pretrain/third_party/resnext101_64x4d-ee2c6f71.pth",
|
||||
"contrib/resnet50_gn": "https://download.openmmlab.com/pretrain/third_party/resnet50_gn_thangvubk-ad1730dd.pth",
|
||||
"detectron/resnet50_gn": "https://download.openmmlab.com/pretrain/third_party/resnet50_gn-9186a21c.pth",
|
||||
"detectron/resnet101_gn": "https://download.openmmlab.com/pretrain/third_party/resnet101_gn-cac0ab98.pth",
|
||||
"jhu/resnet50_gn_ws": "https://download.openmmlab.com/pretrain/third_party/resnet50_gn_ws-15beedd8.pth",
|
||||
"jhu/resnet101_gn_ws": "https://download.openmmlab.com/pretrain/third_party/resnet101_gn_ws-3e3c308c.pth",
|
||||
"jhu/resnext50_32x4d_gn_ws": "https://download.openmmlab.com/pretrain/third_party/resnext50_32x4d_gn_ws-0d87ac85.pth",
|
||||
"jhu/resnext101_32x4d_gn_ws": "https://download.openmmlab.com/pretrain/third_party/resnext101_32x4d_gn_ws-34ac1a9e.pth",
|
||||
"jhu/resnext50_32x4d_gn": "https://download.openmmlab.com/pretrain/third_party/resnext50_32x4d_gn-c7e8b754.pth",
|
||||
"jhu/resnext101_32x4d_gn": "https://download.openmmlab.com/pretrain/third_party/resnext101_32x4d_gn-ac3bb84e.pth",
|
||||
"msra/hrnetv2_w18_small": "https://download.openmmlab.com/pretrain/third_party/hrnetv2_w18_small-b5a04e21.pth",
|
||||
"msra/hrnetv2_w18": "https://download.openmmlab.com/pretrain/third_party/hrnetv2_w18-00eb2006.pth",
|
||||
"msra/hrnetv2_w32": "https://download.openmmlab.com/pretrain/third_party/hrnetv2_w32-dc9eeb4f.pth",
|
||||
"msra/hrnetv2_w40": "https://download.openmmlab.com/pretrain/third_party/hrnetv2_w40-ed0b031c.pth",
|
||||
"msra/hrnetv2_w48": "https://download.openmmlab.com/pretrain/third_party/hrnetv2_w48-d2186c55.pth",
|
||||
"bninception_caffe": "https://download.openmmlab.com/pretrain/third_party/bn_inception_caffe-ed2e8665.pth",
|
||||
"kin400/i3d_r50_f32s2_k400": "https://download.openmmlab.com/pretrain/third_party/i3d_r50_f32s2_k400-2c57e077.pth",
|
||||
"kin400/nl3d_r50_f32s2_k400": "https://download.openmmlab.com/pretrain/third_party/nl3d_r50_f32s2_k400-fa7e7caa.pth",
|
||||
"res2net101_v1d_26w_4s": "https://download.openmmlab.com/pretrain/third_party/res2net101_v1d_26w_4s_mmdetv2-f0a600f9.pth",
|
||||
"regnetx_400mf": "https://download.openmmlab.com/pretrain/third_party/regnetx_400mf-a5b10d96.pth",
|
||||
"regnetx_800mf": "https://download.openmmlab.com/pretrain/third_party/regnetx_800mf-1f4be4c7.pth",
|
||||
"regnetx_1.6gf": "https://download.openmmlab.com/pretrain/third_party/regnetx_1.6gf-5791c176.pth",
|
||||
"regnetx_3.2gf": "https://download.openmmlab.com/pretrain/third_party/regnetx_3.2gf-c2599b0f.pth",
|
||||
"regnetx_4.0gf": "https://download.openmmlab.com/pretrain/third_party/regnetx_4.0gf-a88f671e.pth",
|
||||
"regnetx_6.4gf": "https://download.openmmlab.com/pretrain/third_party/regnetx_6.4gf-006af45d.pth",
|
||||
"regnetx_8.0gf": "https://download.openmmlab.com/pretrain/third_party/regnetx_8.0gf-3c68abe7.pth",
|
||||
"regnetx_12gf": "https://download.openmmlab.com/pretrain/third_party/regnetx_12gf-4c2a3350.pth",
|
||||
"resnet18_v1c": "https://download.openmmlab.com/pretrain/third_party/resnet18_v1c-b5776b93.pth",
|
||||
"resnet50_v1c": "https://download.openmmlab.com/pretrain/third_party/resnet50_v1c-2cccc1ad.pth",
|
||||
"resnet101_v1c": "https://download.openmmlab.com/pretrain/third_party/resnet101_v1c-e67eebb6.pth",
|
||||
"mmedit/vgg16": "https://download.openmmlab.com/mmediting/third_party/vgg_state_dict.pth",
|
||||
"mmedit/res34_en_nomixup": "https://download.openmmlab.com/mmediting/third_party/model_best_resnet34_En_nomixup.pth",
|
||||
"mmedit/mobilenet_v2": "https://download.openmmlab.com/mmediting/third_party/mobilenet_v2.pth",
|
||||
"contrib/mobilenet_v3_large": "https://download.openmmlab.com/pretrain/third_party/mobilenet_v3_large-bc2c3fd3.pth",
|
||||
"contrib/mobilenet_v3_small": "https://download.openmmlab.com/pretrain/third_party/mobilenet_v3_small-47085aa1.pth",
|
||||
"resnest50": "https://download.openmmlab.com/pretrain/third_party/resnest50_d2-7497a55b.pth",
|
||||
"resnest101": "https://download.openmmlab.com/pretrain/third_party/resnest101_d2-f3b931b2.pth",
|
||||
"resnest200": "https://download.openmmlab.com/pretrain/third_party/resnest200_d2-ca88e41f.pth",
|
||||
"darknet53": "https://download.openmmlab.com/pretrain/third_party/darknet53-a628ea1b.pth",
|
||||
"mmdet/mobilenet_v2": "https://download.openmmlab.com/mmdetection/v2.0/third_party/mobilenet_v2_batch256_imagenet-ff34753d.pth"
|
||||
}
|
||||
@@ -0,0 +1,81 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from .assign_score_withk import assign_score_withk
|
||||
from .ball_query import ball_query
|
||||
from .bbox import bbox_overlaps
|
||||
from .border_align import BorderAlign, border_align
|
||||
from .box_iou_rotated import box_iou_rotated
|
||||
from .carafe import CARAFE, CARAFENaive, CARAFEPack, carafe, carafe_naive
|
||||
from .cc_attention import CrissCrossAttention
|
||||
from .contour_expand import contour_expand
|
||||
from .corner_pool import CornerPool
|
||||
from .correlation import Correlation
|
||||
from .deform_conv import DeformConv2d, DeformConv2dPack, deform_conv2d
|
||||
from .deform_roi_pool import (DeformRoIPool, DeformRoIPoolPack,
|
||||
ModulatedDeformRoIPoolPack, deform_roi_pool)
|
||||
from .deprecated_wrappers import Conv2d_deprecated as Conv2d
|
||||
from .deprecated_wrappers import ConvTranspose2d_deprecated as ConvTranspose2d
|
||||
from .deprecated_wrappers import Linear_deprecated as Linear
|
||||
from .deprecated_wrappers import MaxPool2d_deprecated as MaxPool2d
|
||||
from .focal_loss import (SigmoidFocalLoss, SoftmaxFocalLoss,
|
||||
sigmoid_focal_loss, softmax_focal_loss)
|
||||
from .furthest_point_sample import (furthest_point_sample,
|
||||
furthest_point_sample_with_dist)
|
||||
from .fused_bias_leakyrelu import FusedBiasLeakyReLU, fused_bias_leakyrelu
|
||||
from .gather_points import gather_points
|
||||
from .group_points import GroupAll, QueryAndGroup, grouping_operation
|
||||
from .info import (get_compiler_version, get_compiling_cuda_version,
|
||||
get_onnxruntime_op_path)
|
||||
from .iou3d import boxes_iou_bev, nms_bev, nms_normal_bev
|
||||
from .knn import knn
|
||||
from .masked_conv import MaskedConv2d, masked_conv2d
|
||||
from .modulated_deform_conv import (ModulatedDeformConv2d,
|
||||
ModulatedDeformConv2dPack,
|
||||
modulated_deform_conv2d)
|
||||
from .multi_scale_deform_attn import MultiScaleDeformableAttention
|
||||
from .nms import batched_nms, nms, nms_match, nms_rotated, soft_nms
|
||||
from .pixel_group import pixel_group
|
||||
from .point_sample import (SimpleRoIAlign, point_sample,
|
||||
rel_roi_point_to_rel_img_point)
|
||||
from .points_in_boxes import (points_in_boxes_all, points_in_boxes_cpu,
|
||||
points_in_boxes_part)
|
||||
from .points_sampler import PointsSampler
|
||||
from .psa_mask import PSAMask
|
||||
from .roi_align import RoIAlign, roi_align
|
||||
from .roi_align_rotated import RoIAlignRotated, roi_align_rotated
|
||||
from .roi_pool import RoIPool, roi_pool
|
||||
from .roiaware_pool3d import RoIAwarePool3d
|
||||
from .roipoint_pool3d import RoIPointPool3d
|
||||
from .saconv import SAConv2d
|
||||
from .scatter_points import DynamicScatter, dynamic_scatter
|
||||
from .sync_bn import SyncBatchNorm
|
||||
from .three_interpolate import three_interpolate
|
||||
from .three_nn import three_nn
|
||||
from .tin_shift import TINShift, tin_shift
|
||||
from .upfirdn2d import upfirdn2d
|
||||
from .voxelize import Voxelization, voxelization
|
||||
|
||||
__all__ = [
|
||||
'bbox_overlaps', 'CARAFE', 'CARAFENaive', 'CARAFEPack', 'carafe',
|
||||
'carafe_naive', 'CornerPool', 'DeformConv2d', 'DeformConv2dPack',
|
||||
'deform_conv2d', 'DeformRoIPool', 'DeformRoIPoolPack',
|
||||
'ModulatedDeformRoIPoolPack', 'deform_roi_pool', 'SigmoidFocalLoss',
|
||||
'SoftmaxFocalLoss', 'sigmoid_focal_loss', 'softmax_focal_loss',
|
||||
'get_compiler_version', 'get_compiling_cuda_version',
|
||||
'get_onnxruntime_op_path', 'MaskedConv2d', 'masked_conv2d',
|
||||
'ModulatedDeformConv2d', 'ModulatedDeformConv2dPack',
|
||||
'modulated_deform_conv2d', 'batched_nms', 'nms', 'soft_nms', 'nms_match',
|
||||
'RoIAlign', 'roi_align', 'RoIPool', 'roi_pool', 'SyncBatchNorm', 'Conv2d',
|
||||
'ConvTranspose2d', 'Linear', 'MaxPool2d', 'CrissCrossAttention', 'PSAMask',
|
||||
'point_sample', 'rel_roi_point_to_rel_img_point', 'SimpleRoIAlign',
|
||||
'SAConv2d', 'TINShift', 'tin_shift', 'assign_score_withk',
|
||||
'box_iou_rotated', 'RoIPointPool3d', 'nms_rotated', 'knn', 'ball_query',
|
||||
'upfirdn2d', 'FusedBiasLeakyReLU', 'fused_bias_leakyrelu',
|
||||
'RoIAlignRotated', 'roi_align_rotated', 'pixel_group', 'QueryAndGroup',
|
||||
'GroupAll', 'grouping_operation', 'contour_expand', 'three_nn',
|
||||
'three_interpolate', 'MultiScaleDeformableAttention', 'BorderAlign',
|
||||
'border_align', 'gather_points', 'furthest_point_sample',
|
||||
'furthest_point_sample_with_dist', 'PointsSampler', 'Correlation',
|
||||
'boxes_iou_bev', 'nms_bev', 'nms_normal_bev', 'Voxelization',
|
||||
'voxelization', 'dynamic_scatter', 'DynamicScatter', 'RoIAwarePool3d',
|
||||
'points_in_boxes_part', 'points_in_boxes_cpu', 'points_in_boxes_all'
|
||||
]
|
||||
@@ -0,0 +1,123 @@
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['assign_score_withk_forward', 'assign_score_withk_backward'])
|
||||
|
||||
|
||||
class AssignScoreWithK(Function):
|
||||
r"""Perform weighted sum to generate output features according to scores.
|
||||
Modified from `PAConv <https://github.com/CVMI-Lab/PAConv/tree/main/
|
||||
scene_seg/lib/paconv_lib/src/gpu>`_.
|
||||
|
||||
This is a memory-efficient CUDA implementation of assign_scores operation,
|
||||
which first transform all point features with weight bank, then assemble
|
||||
neighbor features with ``knn_idx`` and perform weighted sum of ``scores``.
|
||||
|
||||
See the `paper <https://arxiv.org/pdf/2103.14635.pdf>`_ appendix Sec. D for
|
||||
more detailed descriptions.
|
||||
|
||||
Note:
|
||||
This implementation assumes using ``neighbor`` kernel input, which is
|
||||
(point_features - center_features, point_features).
|
||||
See https://github.com/CVMI-Lab/PAConv/blob/main/scene_seg/model/
|
||||
pointnet2/paconv.py#L128 for more details.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
scores,
|
||||
point_features,
|
||||
center_features,
|
||||
knn_idx,
|
||||
aggregate='sum'):
|
||||
"""
|
||||
Args:
|
||||
scores (torch.Tensor): (B, npoint, K, M), predicted scores to
|
||||
aggregate weight matrices in the weight bank.
|
||||
``npoint`` is the number of sampled centers.
|
||||
``K`` is the number of queried neighbors.
|
||||
``M`` is the number of weight matrices in the weight bank.
|
||||
point_features (torch.Tensor): (B, N, M, out_dim)
|
||||
Pre-computed point features to be aggregated.
|
||||
center_features (torch.Tensor): (B, N, M, out_dim)
|
||||
Pre-computed center features to be aggregated.
|
||||
knn_idx (torch.Tensor): (B, npoint, K), index of sampled kNN.
|
||||
We assume the first idx in each row is the idx of the center.
|
||||
aggregate (str, optional): Aggregation method.
|
||||
Can be 'sum', 'avg' or 'max'. Defaults: 'sum'.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: (B, out_dim, npoint, K), the aggregated features.
|
||||
"""
|
||||
agg = {'sum': 0, 'avg': 1, 'max': 2}
|
||||
|
||||
B, N, M, out_dim = point_features.size()
|
||||
_, npoint, K, _ = scores.size()
|
||||
|
||||
output = point_features.new_zeros((B, out_dim, npoint, K))
|
||||
ext_module.assign_score_withk_forward(
|
||||
point_features.contiguous(),
|
||||
center_features.contiguous(),
|
||||
scores.contiguous(),
|
||||
knn_idx.contiguous(),
|
||||
output,
|
||||
B=B,
|
||||
N0=N,
|
||||
N1=npoint,
|
||||
M=M,
|
||||
K=K,
|
||||
O=out_dim,
|
||||
aggregate=agg[aggregate])
|
||||
|
||||
ctx.save_for_backward(output, point_features, center_features, scores,
|
||||
knn_idx)
|
||||
ctx.agg = agg[aggregate]
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_out):
|
||||
"""
|
||||
Args:
|
||||
grad_out (torch.Tensor): (B, out_dim, npoint, K)
|
||||
|
||||
Returns:
|
||||
grad_scores (torch.Tensor): (B, npoint, K, M)
|
||||
grad_point_features (torch.Tensor): (B, N, M, out_dim)
|
||||
grad_center_features (torch.Tensor): (B, N, M, out_dim)
|
||||
"""
|
||||
_, point_features, center_features, scores, knn_idx = ctx.saved_tensors
|
||||
|
||||
agg = ctx.agg
|
||||
|
||||
B, N, M, out_dim = point_features.size()
|
||||
_, npoint, K, _ = scores.size()
|
||||
|
||||
grad_point_features = point_features.new_zeros(point_features.shape)
|
||||
grad_center_features = center_features.new_zeros(center_features.shape)
|
||||
grad_scores = scores.new_zeros(scores.shape)
|
||||
|
||||
ext_module.assign_score_withk_backward(
|
||||
grad_out.contiguous(),
|
||||
point_features.contiguous(),
|
||||
center_features.contiguous(),
|
||||
scores.contiguous(),
|
||||
knn_idx.contiguous(),
|
||||
grad_point_features,
|
||||
grad_center_features,
|
||||
grad_scores,
|
||||
B=B,
|
||||
N0=N,
|
||||
N1=npoint,
|
||||
M=M,
|
||||
K=K,
|
||||
O=out_dim,
|
||||
aggregate=agg)
|
||||
|
||||
return grad_scores, grad_point_features, \
|
||||
grad_center_features, None, None
|
||||
|
||||
|
||||
assign_score_withk = AssignScoreWithK.apply
|
||||
@@ -0,0 +1,55 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['ball_query_forward'])
|
||||
|
||||
|
||||
class BallQuery(Function):
|
||||
"""Find nearby points in spherical space."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, min_radius: float, max_radius: float, sample_num: int,
|
||||
xyz: torch.Tensor, center_xyz: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
min_radius (float): minimum radius of the balls.
|
||||
max_radius (float): maximum radius of the balls.
|
||||
sample_num (int): maximum number of features in the balls.
|
||||
xyz (Tensor): (B, N, 3) xyz coordinates of the features.
|
||||
center_xyz (Tensor): (B, npoint, 3) centers of the ball query.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, npoint, nsample) tensor with the indices of
|
||||
the features that form the query balls.
|
||||
"""
|
||||
assert center_xyz.is_contiguous()
|
||||
assert xyz.is_contiguous()
|
||||
assert min_radius < max_radius
|
||||
|
||||
B, N, _ = xyz.size()
|
||||
npoint = center_xyz.size(1)
|
||||
idx = xyz.new_zeros(B, npoint, sample_num, dtype=torch.int)
|
||||
|
||||
ext_module.ball_query_forward(
|
||||
center_xyz,
|
||||
xyz,
|
||||
idx,
|
||||
b=B,
|
||||
n=N,
|
||||
m=npoint,
|
||||
min_radius=min_radius,
|
||||
max_radius=max_radius,
|
||||
nsample=sample_num)
|
||||
if torch.__version__ != 'parrots':
|
||||
ctx.mark_non_differentiable(idx)
|
||||
return idx
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, a=None):
|
||||
return None, None, None, None
|
||||
|
||||
|
||||
ball_query = BallQuery.apply
|
||||
@@ -0,0 +1,72 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['bbox_overlaps'])
|
||||
|
||||
|
||||
def bbox_overlaps(bboxes1, bboxes2, mode='iou', aligned=False, offset=0):
|
||||
"""Calculate overlap between two set of bboxes.
|
||||
|
||||
If ``aligned`` is ``False``, then calculate the ious between each bbox
|
||||
of bboxes1 and bboxes2, otherwise the ious between each aligned pair of
|
||||
bboxes1 and bboxes2.
|
||||
|
||||
Args:
|
||||
bboxes1 (Tensor): shape (m, 4) in <x1, y1, x2, y2> format or empty.
|
||||
bboxes2 (Tensor): shape (n, 4) in <x1, y1, x2, y2> format or empty.
|
||||
If aligned is ``True``, then m and n must be equal.
|
||||
mode (str): "iou" (intersection over union) or iof (intersection over
|
||||
foreground).
|
||||
|
||||
Returns:
|
||||
ious(Tensor): shape (m, n) if aligned == False else shape (m, 1)
|
||||
|
||||
Example:
|
||||
>>> bboxes1 = torch.FloatTensor([
|
||||
>>> [0, 0, 10, 10],
|
||||
>>> [10, 10, 20, 20],
|
||||
>>> [32, 32, 38, 42],
|
||||
>>> ])
|
||||
>>> bboxes2 = torch.FloatTensor([
|
||||
>>> [0, 0, 10, 20],
|
||||
>>> [0, 10, 10, 19],
|
||||
>>> [10, 10, 20, 20],
|
||||
>>> ])
|
||||
>>> bbox_overlaps(bboxes1, bboxes2)
|
||||
tensor([[0.5000, 0.0000, 0.0000],
|
||||
[0.0000, 0.0000, 1.0000],
|
||||
[0.0000, 0.0000, 0.0000]])
|
||||
|
||||
Example:
|
||||
>>> empty = torch.FloatTensor([])
|
||||
>>> nonempty = torch.FloatTensor([
|
||||
>>> [0, 0, 10, 9],
|
||||
>>> ])
|
||||
>>> assert tuple(bbox_overlaps(empty, nonempty).shape) == (0, 1)
|
||||
>>> assert tuple(bbox_overlaps(nonempty, empty).shape) == (1, 0)
|
||||
>>> assert tuple(bbox_overlaps(empty, empty).shape) == (0, 0)
|
||||
"""
|
||||
|
||||
mode_dict = {'iou': 0, 'iof': 1}
|
||||
assert mode in mode_dict.keys()
|
||||
mode_flag = mode_dict[mode]
|
||||
# Either the boxes are empty or the length of boxes' last dimension is 4
|
||||
assert (bboxes1.size(-1) == 4 or bboxes1.size(0) == 0)
|
||||
assert (bboxes2.size(-1) == 4 or bboxes2.size(0) == 0)
|
||||
assert offset == 1 or offset == 0
|
||||
|
||||
rows = bboxes1.size(0)
|
||||
cols = bboxes2.size(0)
|
||||
if aligned:
|
||||
assert rows == cols
|
||||
|
||||
if rows * cols == 0:
|
||||
return bboxes1.new(rows, 1) if aligned else bboxes1.new(rows, cols)
|
||||
|
||||
if aligned:
|
||||
ious = bboxes1.new_zeros(rows)
|
||||
else:
|
||||
ious = bboxes1.new_zeros((rows, cols))
|
||||
ext_module.bbox_overlaps(
|
||||
bboxes1, bboxes2, ious, mode=mode_flag, aligned=aligned, offset=offset)
|
||||
return ious
|
||||
@@ -0,0 +1,109 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
# modified from
|
||||
# https://github.com/Megvii-BaseDetection/cvpods/blob/master/cvpods/layers/border_align.py
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['border_align_forward', 'border_align_backward'])
|
||||
|
||||
|
||||
class BorderAlignFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input, boxes, pool_size):
|
||||
return g.op(
|
||||
'mmcv::MMCVBorderAlign', input, boxes, pool_size_i=pool_size)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input, boxes, pool_size):
|
||||
ctx.pool_size = pool_size
|
||||
ctx.input_shape = input.size()
|
||||
|
||||
assert boxes.ndim == 3, 'boxes must be with shape [B, H*W, 4]'
|
||||
assert boxes.size(2) == 4, \
|
||||
'the last dimension of boxes must be (x1, y1, x2, y2)'
|
||||
assert input.size(1) % 4 == 0, \
|
||||
'the channel for input feature must be divisible by factor 4'
|
||||
|
||||
# [B, C//4, H*W, 4]
|
||||
output_shape = (input.size(0), input.size(1) // 4, boxes.size(1), 4)
|
||||
output = input.new_zeros(output_shape)
|
||||
# `argmax_idx` only used for backward
|
||||
argmax_idx = input.new_zeros(output_shape).to(torch.int)
|
||||
|
||||
ext_module.border_align_forward(
|
||||
input, boxes, output, argmax_idx, pool_size=ctx.pool_size)
|
||||
|
||||
ctx.save_for_backward(boxes, argmax_idx)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
boxes, argmax_idx = ctx.saved_tensors
|
||||
grad_input = grad_output.new_zeros(ctx.input_shape)
|
||||
# complex head architecture may cause grad_output uncontiguous
|
||||
grad_output = grad_output.contiguous()
|
||||
ext_module.border_align_backward(
|
||||
grad_output,
|
||||
boxes,
|
||||
argmax_idx,
|
||||
grad_input,
|
||||
pool_size=ctx.pool_size)
|
||||
return grad_input, None, None
|
||||
|
||||
|
||||
border_align = BorderAlignFunction.apply
|
||||
|
||||
|
||||
class BorderAlign(nn.Module):
|
||||
r"""Border align pooling layer.
|
||||
|
||||
Applies border_align over the input feature based on predicted bboxes.
|
||||
The details were described in the paper
|
||||
`BorderDet: Border Feature for Dense Object Detection
|
||||
<https://arxiv.org/abs/2007.11056>`_.
|
||||
|
||||
For each border line (e.g. top, left, bottom or right) of each box,
|
||||
border_align does the following:
|
||||
1. uniformly samples `pool_size`+1 positions on this line, involving \
|
||||
the start and end points.
|
||||
2. the corresponding features on these points are computed by \
|
||||
bilinear interpolation.
|
||||
3. max pooling over all the `pool_size`+1 positions are used for \
|
||||
computing pooled feature.
|
||||
|
||||
Args:
|
||||
pool_size (int): number of positions sampled over the boxes' borders
|
||||
(e.g. top, bottom, left, right).
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, pool_size):
|
||||
super(BorderAlign, self).__init__()
|
||||
self.pool_size = pool_size
|
||||
|
||||
def forward(self, input, boxes):
|
||||
"""
|
||||
Args:
|
||||
input: Features with shape [N,4C,H,W]. Channels ranged in [0,C),
|
||||
[C,2C), [2C,3C), [3C,4C) represent the top, left, bottom,
|
||||
right features respectively.
|
||||
boxes: Boxes with shape [N,H*W,4]. Coordinate format (x1,y1,x2,y2).
|
||||
|
||||
Returns:
|
||||
Tensor: Pooled features with shape [N,C,H*W,4]. The order is
|
||||
(top,left,bottom,right) for the last dimension.
|
||||
"""
|
||||
return border_align(input, boxes, self.pool_size)
|
||||
|
||||
def __repr__(self):
|
||||
s = self.__class__.__name__
|
||||
s += f'(pool_size={self.pool_size})'
|
||||
return s
|
||||
@@ -0,0 +1,45 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['box_iou_rotated'])
|
||||
|
||||
|
||||
def box_iou_rotated(bboxes1, bboxes2, mode='iou', aligned=False):
|
||||
"""Return intersection-over-union (Jaccard index) of boxes.
|
||||
|
||||
Both sets of boxes are expected to be in
|
||||
(x_center, y_center, width, height, angle) format.
|
||||
|
||||
If ``aligned`` is ``False``, then calculate the ious between each bbox
|
||||
of bboxes1 and bboxes2, otherwise the ious between each aligned pair of
|
||||
bboxes1 and bboxes2.
|
||||
|
||||
Arguments:
|
||||
boxes1 (Tensor): rotated bboxes 1. \
|
||||
It has shape (N, 5), indicating (x, y, w, h, theta) for each row.
|
||||
Note that theta is in radian.
|
||||
boxes2 (Tensor): rotated bboxes 2. \
|
||||
It has shape (M, 5), indicating (x, y, w, h, theta) for each row.
|
||||
Note that theta is in radian.
|
||||
mode (str): "iou" (intersection over union) or iof (intersection over
|
||||
foreground).
|
||||
|
||||
Returns:
|
||||
ious(Tensor): shape (N, M) if aligned == False else shape (N,)
|
||||
"""
|
||||
assert mode in ['iou', 'iof']
|
||||
mode_dict = {'iou': 0, 'iof': 1}
|
||||
mode_flag = mode_dict[mode]
|
||||
rows = bboxes1.size(0)
|
||||
cols = bboxes2.size(0)
|
||||
if aligned:
|
||||
ious = bboxes1.new_zeros(rows)
|
||||
else:
|
||||
ious = bboxes1.new_zeros((rows * cols))
|
||||
bboxes1 = bboxes1.contiguous()
|
||||
bboxes2 = bboxes2.contiguous()
|
||||
ext_module.box_iou_rotated(
|
||||
bboxes1, bboxes2, ious, mode_flag=mode_flag, aligned=aligned)
|
||||
if not aligned:
|
||||
ious = ious.view(rows, cols)
|
||||
return ious
|
||||
@@ -0,0 +1,287 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.autograd import Function
|
||||
from torch.nn.modules.module import Module
|
||||
|
||||
from ..cnn import UPSAMPLE_LAYERS, normal_init, xavier_init
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'carafe_naive_forward', 'carafe_naive_backward', 'carafe_forward',
|
||||
'carafe_backward'
|
||||
])
|
||||
|
||||
|
||||
class CARAFENaiveFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, features, masks, kernel_size, group_size, scale_factor):
|
||||
return g.op(
|
||||
'mmcv::MMCVCARAFENaive',
|
||||
features,
|
||||
masks,
|
||||
kernel_size_i=kernel_size,
|
||||
group_size_i=group_size,
|
||||
scale_factor_f=scale_factor)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, features, masks, kernel_size, group_size, scale_factor):
|
||||
assert scale_factor >= 1
|
||||
assert masks.size(1) == kernel_size * kernel_size * group_size
|
||||
assert masks.size(-1) == features.size(-1) * scale_factor
|
||||
assert masks.size(-2) == features.size(-2) * scale_factor
|
||||
assert features.size(1) % group_size == 0
|
||||
assert (kernel_size - 1) % 2 == 0 and kernel_size >= 1
|
||||
ctx.kernel_size = kernel_size
|
||||
ctx.group_size = group_size
|
||||
ctx.scale_factor = scale_factor
|
||||
ctx.feature_size = features.size()
|
||||
ctx.mask_size = masks.size()
|
||||
|
||||
n, c, h, w = features.size()
|
||||
output = features.new_zeros((n, c, h * scale_factor, w * scale_factor))
|
||||
ext_module.carafe_naive_forward(
|
||||
features,
|
||||
masks,
|
||||
output,
|
||||
kernel_size=kernel_size,
|
||||
group_size=group_size,
|
||||
scale_factor=scale_factor)
|
||||
|
||||
if features.requires_grad or masks.requires_grad:
|
||||
ctx.save_for_backward(features, masks)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
assert grad_output.is_cuda
|
||||
|
||||
features, masks = ctx.saved_tensors
|
||||
kernel_size = ctx.kernel_size
|
||||
group_size = ctx.group_size
|
||||
scale_factor = ctx.scale_factor
|
||||
|
||||
grad_input = torch.zeros_like(features)
|
||||
grad_masks = torch.zeros_like(masks)
|
||||
ext_module.carafe_naive_backward(
|
||||
grad_output.contiguous(),
|
||||
features,
|
||||
masks,
|
||||
grad_input,
|
||||
grad_masks,
|
||||
kernel_size=kernel_size,
|
||||
group_size=group_size,
|
||||
scale_factor=scale_factor)
|
||||
|
||||
return grad_input, grad_masks, None, None, None
|
||||
|
||||
|
||||
carafe_naive = CARAFENaiveFunction.apply
|
||||
|
||||
|
||||
class CARAFENaive(Module):
|
||||
|
||||
def __init__(self, kernel_size, group_size, scale_factor):
|
||||
super(CARAFENaive, self).__init__()
|
||||
|
||||
assert isinstance(kernel_size, int) and isinstance(
|
||||
group_size, int) and isinstance(scale_factor, int)
|
||||
self.kernel_size = kernel_size
|
||||
self.group_size = group_size
|
||||
self.scale_factor = scale_factor
|
||||
|
||||
def forward(self, features, masks):
|
||||
return carafe_naive(features, masks, self.kernel_size, self.group_size,
|
||||
self.scale_factor)
|
||||
|
||||
|
||||
class CARAFEFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, features, masks, kernel_size, group_size, scale_factor):
|
||||
return g.op(
|
||||
'mmcv::MMCVCARAFE',
|
||||
features,
|
||||
masks,
|
||||
kernel_size_i=kernel_size,
|
||||
group_size_i=group_size,
|
||||
scale_factor_f=scale_factor)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, features, masks, kernel_size, group_size, scale_factor):
|
||||
assert scale_factor >= 1
|
||||
assert masks.size(1) == kernel_size * kernel_size * group_size
|
||||
assert masks.size(-1) == features.size(-1) * scale_factor
|
||||
assert masks.size(-2) == features.size(-2) * scale_factor
|
||||
assert features.size(1) % group_size == 0
|
||||
assert (kernel_size - 1) % 2 == 0 and kernel_size >= 1
|
||||
ctx.kernel_size = kernel_size
|
||||
ctx.group_size = group_size
|
||||
ctx.scale_factor = scale_factor
|
||||
ctx.feature_size = features.size()
|
||||
ctx.mask_size = masks.size()
|
||||
|
||||
n, c, h, w = features.size()
|
||||
output = features.new_zeros((n, c, h * scale_factor, w * scale_factor))
|
||||
routput = features.new_zeros(output.size(), requires_grad=False)
|
||||
rfeatures = features.new_zeros(features.size(), requires_grad=False)
|
||||
rmasks = masks.new_zeros(masks.size(), requires_grad=False)
|
||||
ext_module.carafe_forward(
|
||||
features,
|
||||
masks,
|
||||
rfeatures,
|
||||
routput,
|
||||
rmasks,
|
||||
output,
|
||||
kernel_size=kernel_size,
|
||||
group_size=group_size,
|
||||
scale_factor=scale_factor)
|
||||
|
||||
if features.requires_grad or masks.requires_grad:
|
||||
ctx.save_for_backward(features, masks, rfeatures)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
assert grad_output.is_cuda
|
||||
|
||||
features, masks, rfeatures = ctx.saved_tensors
|
||||
kernel_size = ctx.kernel_size
|
||||
group_size = ctx.group_size
|
||||
scale_factor = ctx.scale_factor
|
||||
|
||||
rgrad_output = torch.zeros_like(grad_output, requires_grad=False)
|
||||
rgrad_input_hs = torch.zeros_like(grad_output, requires_grad=False)
|
||||
rgrad_input = torch.zeros_like(features, requires_grad=False)
|
||||
rgrad_masks = torch.zeros_like(masks, requires_grad=False)
|
||||
grad_input = torch.zeros_like(features, requires_grad=False)
|
||||
grad_masks = torch.zeros_like(masks, requires_grad=False)
|
||||
ext_module.carafe_backward(
|
||||
grad_output.contiguous(),
|
||||
rfeatures,
|
||||
masks,
|
||||
rgrad_output,
|
||||
rgrad_input_hs,
|
||||
rgrad_input,
|
||||
rgrad_masks,
|
||||
grad_input,
|
||||
grad_masks,
|
||||
kernel_size=kernel_size,
|
||||
group_size=group_size,
|
||||
scale_factor=scale_factor)
|
||||
return grad_input, grad_masks, None, None, None
|
||||
|
||||
|
||||
carafe = CARAFEFunction.apply
|
||||
|
||||
|
||||
class CARAFE(Module):
|
||||
""" CARAFE: Content-Aware ReAssembly of FEatures
|
||||
|
||||
Please refer to https://arxiv.org/abs/1905.02188 for more details.
|
||||
|
||||
Args:
|
||||
kernel_size (int): reassemble kernel size
|
||||
group_size (int): reassemble group size
|
||||
scale_factor (int): upsample ratio
|
||||
|
||||
Returns:
|
||||
upsampled feature map
|
||||
"""
|
||||
|
||||
def __init__(self, kernel_size, group_size, scale_factor):
|
||||
super(CARAFE, self).__init__()
|
||||
|
||||
assert isinstance(kernel_size, int) and isinstance(
|
||||
group_size, int) and isinstance(scale_factor, int)
|
||||
self.kernel_size = kernel_size
|
||||
self.group_size = group_size
|
||||
self.scale_factor = scale_factor
|
||||
|
||||
def forward(self, features, masks):
|
||||
return carafe(features, masks, self.kernel_size, self.group_size,
|
||||
self.scale_factor)
|
||||
|
||||
|
||||
@UPSAMPLE_LAYERS.register_module(name='carafe')
|
||||
class CARAFEPack(nn.Module):
|
||||
"""A unified package of CARAFE upsampler that contains: 1) channel
|
||||
compressor 2) content encoder 3) CARAFE op.
|
||||
|
||||
Official implementation of ICCV 2019 paper
|
||||
CARAFE: Content-Aware ReAssembly of FEatures
|
||||
Please refer to https://arxiv.org/abs/1905.02188 for more details.
|
||||
|
||||
Args:
|
||||
channels (int): input feature channels
|
||||
scale_factor (int): upsample ratio
|
||||
up_kernel (int): kernel size of CARAFE op
|
||||
up_group (int): group size of CARAFE op
|
||||
encoder_kernel (int): kernel size of content encoder
|
||||
encoder_dilation (int): dilation of content encoder
|
||||
compressed_channels (int): output channels of channels compressor
|
||||
|
||||
Returns:
|
||||
upsampled feature map
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
channels,
|
||||
scale_factor,
|
||||
up_kernel=5,
|
||||
up_group=1,
|
||||
encoder_kernel=3,
|
||||
encoder_dilation=1,
|
||||
compressed_channels=64):
|
||||
super(CARAFEPack, self).__init__()
|
||||
self.channels = channels
|
||||
self.scale_factor = scale_factor
|
||||
self.up_kernel = up_kernel
|
||||
self.up_group = up_group
|
||||
self.encoder_kernel = encoder_kernel
|
||||
self.encoder_dilation = encoder_dilation
|
||||
self.compressed_channels = compressed_channels
|
||||
self.channel_compressor = nn.Conv2d(channels, self.compressed_channels,
|
||||
1)
|
||||
self.content_encoder = nn.Conv2d(
|
||||
self.compressed_channels,
|
||||
self.up_kernel * self.up_kernel * self.up_group *
|
||||
self.scale_factor * self.scale_factor,
|
||||
self.encoder_kernel,
|
||||
padding=int((self.encoder_kernel - 1) * self.encoder_dilation / 2),
|
||||
dilation=self.encoder_dilation,
|
||||
groups=1)
|
||||
self.init_weights()
|
||||
|
||||
def init_weights(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.Conv2d):
|
||||
xavier_init(m, distribution='uniform')
|
||||
normal_init(self.content_encoder, std=0.001)
|
||||
|
||||
def kernel_normalizer(self, mask):
|
||||
mask = F.pixel_shuffle(mask, self.scale_factor)
|
||||
n, mask_c, h, w = mask.size()
|
||||
# use float division explicitly,
|
||||
# to void inconsistency while exporting to onnx
|
||||
mask_channel = int(mask_c / float(self.up_kernel**2))
|
||||
mask = mask.view(n, mask_channel, -1, h, w)
|
||||
|
||||
mask = F.softmax(mask, dim=2, dtype=mask.dtype)
|
||||
mask = mask.view(n, mask_c, h, w).contiguous()
|
||||
|
||||
return mask
|
||||
|
||||
def feature_reassemble(self, x, mask):
|
||||
x = carafe(x, mask, self.up_kernel, self.up_group, self.scale_factor)
|
||||
return x
|
||||
|
||||
def forward(self, x):
|
||||
compressed_x = self.channel_compressor(x)
|
||||
mask = self.content_encoder(compressed_x)
|
||||
mask = self.kernel_normalizer(mask)
|
||||
|
||||
x = self.feature_reassemble(x, mask)
|
||||
return x
|
||||
@@ -0,0 +1,83 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from custom_mmpkg.custom_mmcv.cnn import PLUGIN_LAYERS, Scale
|
||||
|
||||
|
||||
def NEG_INF_DIAG(n, device):
|
||||
"""Returns a diagonal matrix of size [n, n].
|
||||
|
||||
The diagonal are all "-inf". This is for avoiding calculating the
|
||||
overlapped element in the Criss-Cross twice.
|
||||
"""
|
||||
return torch.diag(torch.tensor(float('-inf')).to(device).repeat(n), 0)
|
||||
|
||||
|
||||
@PLUGIN_LAYERS.register_module()
|
||||
class CrissCrossAttention(nn.Module):
|
||||
"""Criss-Cross Attention Module.
|
||||
|
||||
.. note::
|
||||
Before v1.3.13, we use a CUDA op. Since v1.3.13, we switch
|
||||
to a pure PyTorch and equivalent implementation. For more
|
||||
details, please refer to https://github.com/open-mmlab/mmcv/pull/1201.
|
||||
|
||||
Speed comparison for one forward pass
|
||||
|
||||
- Input size: [2,512,97,97]
|
||||
- Device: 1 NVIDIA GeForce RTX 2080 Ti
|
||||
|
||||
+-----------------------+---------------+------------+---------------+
|
||||
| |PyTorch version|CUDA version|Relative speed |
|
||||
+=======================+===============+============+===============+
|
||||
|with torch.no_grad() |0.00554402 s |0.0299619 s |5.4x |
|
||||
+-----------------------+---------------+------------+---------------+
|
||||
|no with torch.no_grad()|0.00562803 s |0.0301349 s |5.4x |
|
||||
+-----------------------+---------------+------------+---------------+
|
||||
|
||||
Args:
|
||||
in_channels (int): Channels of the input feature map.
|
||||
"""
|
||||
|
||||
def __init__(self, in_channels):
|
||||
super().__init__()
|
||||
self.query_conv = nn.Conv2d(in_channels, in_channels // 8, 1)
|
||||
self.key_conv = nn.Conv2d(in_channels, in_channels // 8, 1)
|
||||
self.value_conv = nn.Conv2d(in_channels, in_channels, 1)
|
||||
self.gamma = Scale(0.)
|
||||
self.in_channels = in_channels
|
||||
|
||||
def forward(self, x):
|
||||
"""forward function of Criss-Cross Attention.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input feature. \
|
||||
shape (batch_size, in_channels, height, width)
|
||||
Returns:
|
||||
Tensor: Output of the layer, with shape of \
|
||||
(batch_size, in_channels, height, width)
|
||||
"""
|
||||
B, C, H, W = x.size()
|
||||
query = self.query_conv(x)
|
||||
key = self.key_conv(x)
|
||||
value = self.value_conv(x)
|
||||
energy_H = torch.einsum('bchw,bciw->bwhi', query, key) + NEG_INF_DIAG(
|
||||
H, query.device)
|
||||
energy_H = energy_H.transpose(1, 2)
|
||||
energy_W = torch.einsum('bchw,bchj->bhwj', query, key)
|
||||
attn = F.softmax(
|
||||
torch.cat([energy_H, energy_W], dim=-1), dim=-1) # [B,H,W,(H+W)]
|
||||
out = torch.einsum('bciw,bhwi->bchw', value, attn[..., :H])
|
||||
out += torch.einsum('bchj,bhwj->bchw', value, attn[..., H:])
|
||||
|
||||
out = self.gamma(out) + x
|
||||
out = out.contiguous()
|
||||
|
||||
return out
|
||||
|
||||
def __repr__(self):
|
||||
s = self.__class__.__name__
|
||||
s += f'(in_channels={self.in_channels})'
|
||||
return s
|
||||
@@ -0,0 +1,49 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['contour_expand'])
|
||||
|
||||
|
||||
def contour_expand(kernel_mask, internal_kernel_label, min_kernel_area,
|
||||
kernel_num):
|
||||
"""Expand kernel contours so that foreground pixels are assigned into
|
||||
instances.
|
||||
|
||||
Arguments:
|
||||
kernel_mask (np.array or Tensor): The instance kernel mask with
|
||||
size hxw.
|
||||
internal_kernel_label (np.array or Tensor): The instance internal
|
||||
kernel label with size hxw.
|
||||
min_kernel_area (int): The minimum kernel area.
|
||||
kernel_num (int): The instance kernel number.
|
||||
|
||||
Returns:
|
||||
label (list): The instance index map with size hxw.
|
||||
"""
|
||||
assert isinstance(kernel_mask, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(internal_kernel_label, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(min_kernel_area, int)
|
||||
assert isinstance(kernel_num, int)
|
||||
|
||||
if isinstance(kernel_mask, np.ndarray):
|
||||
kernel_mask = torch.from_numpy(kernel_mask)
|
||||
if isinstance(internal_kernel_label, np.ndarray):
|
||||
internal_kernel_label = torch.from_numpy(internal_kernel_label)
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
if kernel_mask.shape[0] == 0 or internal_kernel_label.shape[0] == 0:
|
||||
label = []
|
||||
else:
|
||||
label = ext_module.contour_expand(
|
||||
kernel_mask,
|
||||
internal_kernel_label,
|
||||
min_kernel_area=min_kernel_area,
|
||||
kernel_num=kernel_num)
|
||||
label = label.tolist()
|
||||
else:
|
||||
label = ext_module.contour_expand(kernel_mask, internal_kernel_label,
|
||||
min_kernel_area, kernel_num)
|
||||
return label
|
||||
@@ -0,0 +1,161 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'top_pool_forward', 'top_pool_backward', 'bottom_pool_forward',
|
||||
'bottom_pool_backward', 'left_pool_forward', 'left_pool_backward',
|
||||
'right_pool_forward', 'right_pool_backward'
|
||||
])
|
||||
|
||||
_mode_dict = {'top': 0, 'bottom': 1, 'left': 2, 'right': 3}
|
||||
|
||||
|
||||
class TopPoolFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input):
|
||||
output = g.op(
|
||||
'mmcv::MMCVCornerPool', input, mode_i=int(_mode_dict['top']))
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
output = ext_module.top_pool_forward(input)
|
||||
ctx.save_for_backward(input)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input, = ctx.saved_tensors
|
||||
output = ext_module.top_pool_backward(input, grad_output)
|
||||
return output
|
||||
|
||||
|
||||
class BottomPoolFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input):
|
||||
output = g.op(
|
||||
'mmcv::MMCVCornerPool', input, mode_i=int(_mode_dict['bottom']))
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
output = ext_module.bottom_pool_forward(input)
|
||||
ctx.save_for_backward(input)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input, = ctx.saved_tensors
|
||||
output = ext_module.bottom_pool_backward(input, grad_output)
|
||||
return output
|
||||
|
||||
|
||||
class LeftPoolFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input):
|
||||
output = g.op(
|
||||
'mmcv::MMCVCornerPool', input, mode_i=int(_mode_dict['left']))
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
output = ext_module.left_pool_forward(input)
|
||||
ctx.save_for_backward(input)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input, = ctx.saved_tensors
|
||||
output = ext_module.left_pool_backward(input, grad_output)
|
||||
return output
|
||||
|
||||
|
||||
class RightPoolFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input):
|
||||
output = g.op(
|
||||
'mmcv::MMCVCornerPool', input, mode_i=int(_mode_dict['right']))
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input):
|
||||
output = ext_module.right_pool_forward(input)
|
||||
ctx.save_for_backward(input)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input, = ctx.saved_tensors
|
||||
output = ext_module.right_pool_backward(input, grad_output)
|
||||
return output
|
||||
|
||||
|
||||
class CornerPool(nn.Module):
|
||||
"""Corner Pooling.
|
||||
|
||||
Corner Pooling is a new type of pooling layer that helps a
|
||||
convolutional network better localize corners of bounding boxes.
|
||||
|
||||
Please refer to https://arxiv.org/abs/1808.01244 for more details.
|
||||
Code is modified from https://github.com/princeton-vl/CornerNet-Lite.
|
||||
|
||||
Args:
|
||||
mode(str): Pooling orientation for the pooling layer
|
||||
|
||||
- 'bottom': Bottom Pooling
|
||||
- 'left': Left Pooling
|
||||
- 'right': Right Pooling
|
||||
- 'top': Top Pooling
|
||||
|
||||
Returns:
|
||||
Feature map after pooling.
|
||||
"""
|
||||
|
||||
pool_functions = {
|
||||
'bottom': BottomPoolFunction,
|
||||
'left': LeftPoolFunction,
|
||||
'right': RightPoolFunction,
|
||||
'top': TopPoolFunction,
|
||||
}
|
||||
|
||||
cummax_dim_flip = {
|
||||
'bottom': (2, False),
|
||||
'left': (3, True),
|
||||
'right': (3, False),
|
||||
'top': (2, True),
|
||||
}
|
||||
|
||||
def __init__(self, mode):
|
||||
super(CornerPool, self).__init__()
|
||||
assert mode in self.pool_functions
|
||||
self.mode = mode
|
||||
self.corner_pool = self.pool_functions[mode]
|
||||
|
||||
def forward(self, x):
|
||||
if torch.__version__ != 'parrots' and torch.__version__ >= '1.5.0':
|
||||
if torch.onnx.is_in_onnx_export():
|
||||
assert torch.__version__ >= '1.7.0', \
|
||||
'When `cummax` serves as an intermediate component whose '\
|
||||
'outputs is used as inputs for another modules, it\'s '\
|
||||
'expected that pytorch version must be >= 1.7.0, '\
|
||||
'otherwise Error appears like: `RuntimeError: tuple '\
|
||||
'appears in op that does not forward tuples, unsupported '\
|
||||
'kind: prim::PythonOp`.'
|
||||
|
||||
dim, flip = self.cummax_dim_flip[self.mode]
|
||||
if flip:
|
||||
x = x.flip(dim)
|
||||
pool_tensor, _ = torch.cummax(x, dim=dim)
|
||||
if flip:
|
||||
pool_tensor = pool_tensor.flip(dim)
|
||||
return pool_tensor
|
||||
else:
|
||||
return self.corner_pool.apply(x)
|
||||
@@ -0,0 +1,196 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
from torch import Tensor, nn
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
from torch.nn.modules.utils import _pair
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['correlation_forward', 'correlation_backward'])
|
||||
|
||||
|
||||
class CorrelationFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
input1,
|
||||
input2,
|
||||
kernel_size=1,
|
||||
max_displacement=1,
|
||||
stride=1,
|
||||
padding=1,
|
||||
dilation=1,
|
||||
dilation_patch=1):
|
||||
|
||||
ctx.save_for_backward(input1, input2)
|
||||
|
||||
kH, kW = ctx.kernel_size = _pair(kernel_size)
|
||||
patch_size = max_displacement * 2 + 1
|
||||
ctx.patch_size = patch_size
|
||||
dH, dW = ctx.stride = _pair(stride)
|
||||
padH, padW = ctx.padding = _pair(padding)
|
||||
dilationH, dilationW = ctx.dilation = _pair(dilation)
|
||||
dilation_patchH, dilation_patchW = ctx.dilation_patch = _pair(
|
||||
dilation_patch)
|
||||
|
||||
output_size = CorrelationFunction._output_size(ctx, input1)
|
||||
|
||||
output = input1.new_zeros(output_size)
|
||||
|
||||
ext_module.correlation_forward(
|
||||
input1,
|
||||
input2,
|
||||
output,
|
||||
kH=kH,
|
||||
kW=kW,
|
||||
patchH=patch_size,
|
||||
patchW=patch_size,
|
||||
padH=padH,
|
||||
padW=padW,
|
||||
dilationH=dilationH,
|
||||
dilationW=dilationW,
|
||||
dilation_patchH=dilation_patchH,
|
||||
dilation_patchW=dilation_patchW,
|
||||
dH=dH,
|
||||
dW=dW)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
input1, input2 = ctx.saved_tensors
|
||||
|
||||
kH, kW = ctx.kernel_size
|
||||
patch_size = ctx.patch_size
|
||||
padH, padW = ctx.padding
|
||||
dilationH, dilationW = ctx.dilation
|
||||
dilation_patchH, dilation_patchW = ctx.dilation_patch
|
||||
dH, dW = ctx.stride
|
||||
grad_input1 = torch.zeros_like(input1)
|
||||
grad_input2 = torch.zeros_like(input2)
|
||||
|
||||
ext_module.correlation_backward(
|
||||
grad_output,
|
||||
input1,
|
||||
input2,
|
||||
grad_input1,
|
||||
grad_input2,
|
||||
kH=kH,
|
||||
kW=kW,
|
||||
patchH=patch_size,
|
||||
patchW=patch_size,
|
||||
padH=padH,
|
||||
padW=padW,
|
||||
dilationH=dilationH,
|
||||
dilationW=dilationW,
|
||||
dilation_patchH=dilation_patchH,
|
||||
dilation_patchW=dilation_patchW,
|
||||
dH=dH,
|
||||
dW=dW)
|
||||
return grad_input1, grad_input2, None, None, None, None, None, None
|
||||
|
||||
@staticmethod
|
||||
def _output_size(ctx, input1):
|
||||
iH, iW = input1.size(2), input1.size(3)
|
||||
batch_size = input1.size(0)
|
||||
kH, kW = ctx.kernel_size
|
||||
patch_size = ctx.patch_size
|
||||
dH, dW = ctx.stride
|
||||
padH, padW = ctx.padding
|
||||
dilationH, dilationW = ctx.dilation
|
||||
dilatedKH = (kH - 1) * dilationH + 1
|
||||
dilatedKW = (kW - 1) * dilationW + 1
|
||||
|
||||
oH = int((iH + 2 * padH - dilatedKH) / dH + 1)
|
||||
oW = int((iW + 2 * padW - dilatedKW) / dW + 1)
|
||||
|
||||
output_size = (batch_size, patch_size, patch_size, oH, oW)
|
||||
return output_size
|
||||
|
||||
|
||||
class Correlation(nn.Module):
|
||||
r"""Correlation operator
|
||||
|
||||
This correlation operator works for optical flow correlation computation.
|
||||
|
||||
There are two batched tensors with shape :math:`(N, C, H, W)`,
|
||||
and the correlation output's shape is :math:`(N, max\_displacement \times
|
||||
2 + 1, max\_displacement * 2 + 1, H_{out}, W_{out})`
|
||||
|
||||
where
|
||||
|
||||
.. math::
|
||||
H_{out} = \left\lfloor\frac{H_{in} + 2 \times padding -
|
||||
dilation \times (kernel\_size - 1) - 1}
|
||||
{stride} + 1\right\rfloor
|
||||
|
||||
.. math::
|
||||
W_{out} = \left\lfloor\frac{W_{in} + 2 \times padding - dilation
|
||||
\times (kernel\_size - 1) - 1}
|
||||
{stride} + 1\right\rfloor
|
||||
|
||||
the correlation item :math:`(N_i, dy, dx)` is formed by taking the sliding
|
||||
window convolution between input1 and shifted input2,
|
||||
|
||||
.. math::
|
||||
Corr(N_i, dx, dy) =
|
||||
\sum_{c=0}^{C-1}
|
||||
input1(N_i, c) \star
|
||||
\mathcal{S}(input2(N_i, c), dy, dx)
|
||||
|
||||
where :math:`\star` is the valid 2d sliding window convolution operator,
|
||||
and :math:`\mathcal{S}` means shifting the input features (auto-complete
|
||||
zero marginal), and :math:`dx, dy` are shifting distance, :math:`dx, dy \in
|
||||
[-max\_displacement \times dilation\_patch, max\_displacement \times
|
||||
dilation\_patch]`.
|
||||
|
||||
Args:
|
||||
kernel_size (int): The size of sliding window i.e. local neighborhood
|
||||
representing the center points and involved in correlation
|
||||
computation. Defaults to 1.
|
||||
max_displacement (int): The radius for computing correlation volume,
|
||||
but the actual working space can be dilated by dilation_patch.
|
||||
Defaults to 1.
|
||||
stride (int): The stride of the sliding blocks in the input spatial
|
||||
dimensions. Defaults to 1.
|
||||
padding (int): Zero padding added to all four sides of the input1.
|
||||
Defaults to 0.
|
||||
dilation (int): The spacing of local neighborhood that will involved
|
||||
in correlation. Defaults to 1.
|
||||
dilation_patch (int): The spacing between position need to compute
|
||||
correlation. Defaults to 1.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
kernel_size: int = 1,
|
||||
max_displacement: int = 1,
|
||||
stride: int = 1,
|
||||
padding: int = 0,
|
||||
dilation: int = 1,
|
||||
dilation_patch: int = 1) -> None:
|
||||
super().__init__()
|
||||
self.kernel_size = kernel_size
|
||||
self.max_displacement = max_displacement
|
||||
self.stride = stride
|
||||
self.padding = padding
|
||||
self.dilation = dilation
|
||||
self.dilation_patch = dilation_patch
|
||||
|
||||
def forward(self, input1: Tensor, input2: Tensor) -> Tensor:
|
||||
return CorrelationFunction.apply(input1, input2, self.kernel_size,
|
||||
self.max_displacement, self.stride,
|
||||
self.padding, self.dilation,
|
||||
self.dilation_patch)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
s = self.__class__.__name__
|
||||
s += f'(kernel_size={self.kernel_size}, '
|
||||
s += f'max_displacement={self.max_displacement}, '
|
||||
s += f'stride={self.stride}, '
|
||||
s += f'padding={self.padding}, '
|
||||
s += f'dilation={self.dilation}, '
|
||||
s += f'dilation_patch={self.dilation_patch})'
|
||||
return s
|
||||
@@ -0,0 +1,405 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from typing import Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch import Tensor
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
from torch.nn.modules.utils import _pair, _single
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import deprecated_api_warning
|
||||
from ..cnn import CONV_LAYERS
|
||||
from ..utils import ext_loader, print_log
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'deform_conv_forward', 'deform_conv_backward_input',
|
||||
'deform_conv_backward_parameters'
|
||||
])
|
||||
|
||||
|
||||
class DeformConv2dFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g,
|
||||
input,
|
||||
offset,
|
||||
weight,
|
||||
stride,
|
||||
padding,
|
||||
dilation,
|
||||
groups,
|
||||
deform_groups,
|
||||
bias=False,
|
||||
im2col_step=32):
|
||||
return g.op(
|
||||
'mmcv::MMCVDeformConv2d',
|
||||
input,
|
||||
offset,
|
||||
weight,
|
||||
stride_i=stride,
|
||||
padding_i=padding,
|
||||
dilation_i=dilation,
|
||||
groups_i=groups,
|
||||
deform_groups_i=deform_groups,
|
||||
bias_i=bias,
|
||||
im2col_step_i=im2col_step)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
input,
|
||||
offset,
|
||||
weight,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
deform_groups=1,
|
||||
bias=False,
|
||||
im2col_step=32):
|
||||
if input is not None and input.dim() != 4:
|
||||
raise ValueError(
|
||||
f'Expected 4D tensor as input, got {input.dim()}D tensor \
|
||||
instead.')
|
||||
assert bias is False, 'Only support bias is False.'
|
||||
ctx.stride = _pair(stride)
|
||||
ctx.padding = _pair(padding)
|
||||
ctx.dilation = _pair(dilation)
|
||||
ctx.groups = groups
|
||||
ctx.deform_groups = deform_groups
|
||||
ctx.im2col_step = im2col_step
|
||||
|
||||
# When pytorch version >= 1.6.0, amp is adopted for fp16 mode;
|
||||
# amp won't cast the type of model (float32), but "offset" is cast
|
||||
# to float16 by nn.Conv2d automatically, leading to the type
|
||||
# mismatch with input (when it is float32) or weight.
|
||||
# The flag for whether to use fp16 or amp is the type of "offset",
|
||||
# we cast weight and input to temporarily support fp16 and amp
|
||||
# whatever the pytorch version is.
|
||||
input = input.type_as(offset)
|
||||
weight = weight.type_as(input)
|
||||
ctx.save_for_backward(input, offset, weight)
|
||||
|
||||
output = input.new_empty(
|
||||
DeformConv2dFunction._output_size(ctx, input, weight))
|
||||
|
||||
ctx.bufs_ = [input.new_empty(0), input.new_empty(0)] # columns, ones
|
||||
|
||||
cur_im2col_step = min(ctx.im2col_step, input.size(0))
|
||||
assert (input.size(0) %
|
||||
cur_im2col_step) == 0, 'im2col step must divide batchsize'
|
||||
ext_module.deform_conv_forward(
|
||||
input,
|
||||
weight,
|
||||
offset,
|
||||
output,
|
||||
ctx.bufs_[0],
|
||||
ctx.bufs_[1],
|
||||
kW=weight.size(3),
|
||||
kH=weight.size(2),
|
||||
dW=ctx.stride[1],
|
||||
dH=ctx.stride[0],
|
||||
padW=ctx.padding[1],
|
||||
padH=ctx.padding[0],
|
||||
dilationW=ctx.dilation[1],
|
||||
dilationH=ctx.dilation[0],
|
||||
group=ctx.groups,
|
||||
deformable_group=ctx.deform_groups,
|
||||
im2col_step=cur_im2col_step)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
input, offset, weight = ctx.saved_tensors
|
||||
|
||||
grad_input = grad_offset = grad_weight = None
|
||||
|
||||
cur_im2col_step = min(ctx.im2col_step, input.size(0))
|
||||
assert (input.size(0) % cur_im2col_step
|
||||
) == 0, 'batch size must be divisible by im2col_step'
|
||||
|
||||
grad_output = grad_output.contiguous()
|
||||
if ctx.needs_input_grad[0] or ctx.needs_input_grad[1]:
|
||||
grad_input = torch.zeros_like(input)
|
||||
grad_offset = torch.zeros_like(offset)
|
||||
ext_module.deform_conv_backward_input(
|
||||
input,
|
||||
offset,
|
||||
grad_output,
|
||||
grad_input,
|
||||
grad_offset,
|
||||
weight,
|
||||
ctx.bufs_[0],
|
||||
kW=weight.size(3),
|
||||
kH=weight.size(2),
|
||||
dW=ctx.stride[1],
|
||||
dH=ctx.stride[0],
|
||||
padW=ctx.padding[1],
|
||||
padH=ctx.padding[0],
|
||||
dilationW=ctx.dilation[1],
|
||||
dilationH=ctx.dilation[0],
|
||||
group=ctx.groups,
|
||||
deformable_group=ctx.deform_groups,
|
||||
im2col_step=cur_im2col_step)
|
||||
|
||||
if ctx.needs_input_grad[2]:
|
||||
grad_weight = torch.zeros_like(weight)
|
||||
ext_module.deform_conv_backward_parameters(
|
||||
input,
|
||||
offset,
|
||||
grad_output,
|
||||
grad_weight,
|
||||
ctx.bufs_[0],
|
||||
ctx.bufs_[1],
|
||||
kW=weight.size(3),
|
||||
kH=weight.size(2),
|
||||
dW=ctx.stride[1],
|
||||
dH=ctx.stride[0],
|
||||
padW=ctx.padding[1],
|
||||
padH=ctx.padding[0],
|
||||
dilationW=ctx.dilation[1],
|
||||
dilationH=ctx.dilation[0],
|
||||
group=ctx.groups,
|
||||
deformable_group=ctx.deform_groups,
|
||||
scale=1,
|
||||
im2col_step=cur_im2col_step)
|
||||
|
||||
return grad_input, grad_offset, grad_weight, \
|
||||
None, None, None, None, None, None, None
|
||||
|
||||
@staticmethod
|
||||
def _output_size(ctx, input, weight):
|
||||
channels = weight.size(0)
|
||||
output_size = (input.size(0), channels)
|
||||
for d in range(input.dim() - 2):
|
||||
in_size = input.size(d + 2)
|
||||
pad = ctx.padding[d]
|
||||
kernel = ctx.dilation[d] * (weight.size(d + 2) - 1) + 1
|
||||
stride_ = ctx.stride[d]
|
||||
output_size += ((in_size + (2 * pad) - kernel) // stride_ + 1, )
|
||||
if not all(map(lambda s: s > 0, output_size)):
|
||||
raise ValueError(
|
||||
'convolution input is too small (output would be ' +
|
||||
'x'.join(map(str, output_size)) + ')')
|
||||
return output_size
|
||||
|
||||
|
||||
deform_conv2d = DeformConv2dFunction.apply
|
||||
|
||||
|
||||
class DeformConv2d(nn.Module):
|
||||
r"""Deformable 2D convolution.
|
||||
|
||||
Applies a deformable 2D convolution over an input signal composed of
|
||||
several input planes. DeformConv2d was described in the paper
|
||||
`Deformable Convolutional Networks
|
||||
<https://arxiv.org/pdf/1703.06211.pdf>`_
|
||||
|
||||
Note:
|
||||
The argument ``im2col_step`` was added in version 1.3.17, which means
|
||||
number of samples processed by the ``im2col_cuda_kernel`` per call.
|
||||
It enables users to define ``batch_size`` and ``im2col_step`` more
|
||||
flexibly and solved `issue mmcv#1440
|
||||
<https://github.com/open-mmlab/mmcv/issues/1440>`_.
|
||||
|
||||
Args:
|
||||
in_channels (int): Number of channels in the input image.
|
||||
out_channels (int): Number of channels produced by the convolution.
|
||||
kernel_size(int, tuple): Size of the convolving kernel.
|
||||
stride(int, tuple): Stride of the convolution. Default: 1.
|
||||
padding (int or tuple): Zero-padding added to both sides of the input.
|
||||
Default: 0.
|
||||
dilation (int or tuple): Spacing between kernel elements. Default: 1.
|
||||
groups (int): Number of blocked connections from input.
|
||||
channels to output channels. Default: 1.
|
||||
deform_groups (int): Number of deformable group partitions.
|
||||
bias (bool): If True, adds a learnable bias to the output.
|
||||
Default: False.
|
||||
im2col_step (int): Number of samples processed by im2col_cuda_kernel
|
||||
per call. It will work when ``batch_size`` > ``im2col_step``, but
|
||||
``batch_size`` must be divisible by ``im2col_step``. Default: 32.
|
||||
`New in version 1.3.17.`
|
||||
"""
|
||||
|
||||
@deprecated_api_warning({'deformable_groups': 'deform_groups'},
|
||||
cls_name='DeformConv2d')
|
||||
def __init__(self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size: Union[int, Tuple[int, ...]],
|
||||
stride: Union[int, Tuple[int, ...]] = 1,
|
||||
padding: Union[int, Tuple[int, ...]] = 0,
|
||||
dilation: Union[int, Tuple[int, ...]] = 1,
|
||||
groups: int = 1,
|
||||
deform_groups: int = 1,
|
||||
bias: bool = False,
|
||||
im2col_step: int = 32) -> None:
|
||||
super(DeformConv2d, self).__init__()
|
||||
|
||||
assert not bias, \
|
||||
f'bias={bias} is not supported in DeformConv2d.'
|
||||
assert in_channels % groups == 0, \
|
||||
f'in_channels {in_channels} cannot be divisible by groups {groups}'
|
||||
assert out_channels % groups == 0, \
|
||||
f'out_channels {out_channels} cannot be divisible by groups \
|
||||
{groups}'
|
||||
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.kernel_size = _pair(kernel_size)
|
||||
self.stride = _pair(stride)
|
||||
self.padding = _pair(padding)
|
||||
self.dilation = _pair(dilation)
|
||||
self.groups = groups
|
||||
self.deform_groups = deform_groups
|
||||
self.im2col_step = im2col_step
|
||||
# enable compatibility with nn.Conv2d
|
||||
self.transposed = False
|
||||
self.output_padding = _single(0)
|
||||
|
||||
# only weight, no bias
|
||||
self.weight = nn.Parameter(
|
||||
torch.Tensor(out_channels, in_channels // self.groups,
|
||||
*self.kernel_size))
|
||||
|
||||
self.reset_parameters()
|
||||
|
||||
def reset_parameters(self):
|
||||
# switch the initialization of `self.weight` to the standard kaiming
|
||||
# method described in `Delving deep into rectifiers: Surpassing
|
||||
# human-level performance on ImageNet classification` - He, K. et al.
|
||||
# (2015), using a uniform distribution
|
||||
nn.init.kaiming_uniform_(self.weight, nonlinearity='relu')
|
||||
|
||||
def forward(self, x: Tensor, offset: Tensor) -> Tensor:
|
||||
"""Deformable Convolutional forward function.
|
||||
|
||||
Args:
|
||||
x (Tensor): Input feature, shape (B, C_in, H_in, W_in)
|
||||
offset (Tensor): Offset for deformable convolution, shape
|
||||
(B, deform_groups*kernel_size[0]*kernel_size[1]*2,
|
||||
H_out, W_out), H_out, W_out are equal to the output's.
|
||||
|
||||
An offset is like `[y0, x0, y1, x1, y2, x2, ..., y8, x8]`.
|
||||
The spatial arrangement is like:
|
||||
|
||||
.. code:: text
|
||||
|
||||
(x0, y0) (x1, y1) (x2, y2)
|
||||
(x3, y3) (x4, y4) (x5, y5)
|
||||
(x6, y6) (x7, y7) (x8, y8)
|
||||
|
||||
Returns:
|
||||
Tensor: Output of the layer.
|
||||
"""
|
||||
# To fix an assert error in deform_conv_cuda.cpp:128
|
||||
# input image is smaller than kernel
|
||||
input_pad = (x.size(2) < self.kernel_size[0]) or (x.size(3) <
|
||||
self.kernel_size[1])
|
||||
if input_pad:
|
||||
pad_h = max(self.kernel_size[0] - x.size(2), 0)
|
||||
pad_w = max(self.kernel_size[1] - x.size(3), 0)
|
||||
x = F.pad(x, (0, pad_w, 0, pad_h), 'constant', 0).contiguous()
|
||||
offset = F.pad(offset, (0, pad_w, 0, pad_h), 'constant', 0)
|
||||
offset = offset.contiguous()
|
||||
out = deform_conv2d(x, offset, self.weight, self.stride, self.padding,
|
||||
self.dilation, self.groups, self.deform_groups,
|
||||
False, self.im2col_step)
|
||||
if input_pad:
|
||||
out = out[:, :, :out.size(2) - pad_h, :out.size(3) -
|
||||
pad_w].contiguous()
|
||||
return out
|
||||
|
||||
def __repr__(self):
|
||||
s = self.__class__.__name__
|
||||
s += f'(in_channels={self.in_channels},\n'
|
||||
s += f'out_channels={self.out_channels},\n'
|
||||
s += f'kernel_size={self.kernel_size},\n'
|
||||
s += f'stride={self.stride},\n'
|
||||
s += f'padding={self.padding},\n'
|
||||
s += f'dilation={self.dilation},\n'
|
||||
s += f'groups={self.groups},\n'
|
||||
s += f'deform_groups={self.deform_groups},\n'
|
||||
# bias is not supported in DeformConv2d.
|
||||
s += 'bias=False)'
|
||||
return s
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module('DCN')
|
||||
class DeformConv2dPack(DeformConv2d):
|
||||
"""A Deformable Conv Encapsulation that acts as normal Conv layers.
|
||||
|
||||
The offset tensor is like `[y0, x0, y1, x1, y2, x2, ..., y8, x8]`.
|
||||
The spatial arrangement is like:
|
||||
|
||||
.. code:: text
|
||||
|
||||
(x0, y0) (x1, y1) (x2, y2)
|
||||
(x3, y3) (x4, y4) (x5, y5)
|
||||
(x6, y6) (x7, y7) (x8, y8)
|
||||
|
||||
Args:
|
||||
in_channels (int): Same as nn.Conv2d.
|
||||
out_channels (int): Same as nn.Conv2d.
|
||||
kernel_size (int or tuple[int]): Same as nn.Conv2d.
|
||||
stride (int or tuple[int]): Same as nn.Conv2d.
|
||||
padding (int or tuple[int]): Same as nn.Conv2d.
|
||||
dilation (int or tuple[int]): Same as nn.Conv2d.
|
||||
groups (int): Same as nn.Conv2d.
|
||||
bias (bool or str): If specified as `auto`, it will be decided by the
|
||||
norm_cfg. Bias will be set as True if norm_cfg is None, otherwise
|
||||
False.
|
||||
"""
|
||||
|
||||
_version = 2
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(DeformConv2dPack, self).__init__(*args, **kwargs)
|
||||
self.conv_offset = nn.Conv2d(
|
||||
self.in_channels,
|
||||
self.deform_groups * 2 * self.kernel_size[0] * self.kernel_size[1],
|
||||
kernel_size=self.kernel_size,
|
||||
stride=_pair(self.stride),
|
||||
padding=_pair(self.padding),
|
||||
dilation=_pair(self.dilation),
|
||||
bias=True)
|
||||
self.init_offset()
|
||||
|
||||
def init_offset(self):
|
||||
self.conv_offset.weight.data.zero_()
|
||||
self.conv_offset.bias.data.zero_()
|
||||
|
||||
def forward(self, x):
|
||||
offset = self.conv_offset(x)
|
||||
return deform_conv2d(x, offset, self.weight, self.stride, self.padding,
|
||||
self.dilation, self.groups, self.deform_groups,
|
||||
False, self.im2col_step)
|
||||
|
||||
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict,
|
||||
missing_keys, unexpected_keys, error_msgs):
|
||||
version = local_metadata.get('version', None)
|
||||
|
||||
if version is None or version < 2:
|
||||
# the key is different in early versions
|
||||
# In version < 2, DeformConvPack loads previous benchmark models.
|
||||
if (prefix + 'conv_offset.weight' not in state_dict
|
||||
and prefix[:-1] + '_offset.weight' in state_dict):
|
||||
state_dict[prefix + 'conv_offset.weight'] = state_dict.pop(
|
||||
prefix[:-1] + '_offset.weight')
|
||||
if (prefix + 'conv_offset.bias' not in state_dict
|
||||
and prefix[:-1] + '_offset.bias' in state_dict):
|
||||
state_dict[prefix +
|
||||
'conv_offset.bias'] = state_dict.pop(prefix[:-1] +
|
||||
'_offset.bias')
|
||||
|
||||
if version is not None and version > 1:
|
||||
print_log(
|
||||
f'DeformConv2dPack {prefix.rstrip(".")} is upgraded to '
|
||||
'version 2.',
|
||||
logger='root')
|
||||
|
||||
super()._load_from_state_dict(state_dict, prefix, local_metadata,
|
||||
strict, missing_keys, unexpected_keys,
|
||||
error_msgs)
|
||||
@@ -0,0 +1,204 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from torch import nn
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
from torch.nn.modules.utils import _pair
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['deform_roi_pool_forward', 'deform_roi_pool_backward'])
|
||||
|
||||
|
||||
class DeformRoIPoolFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input, rois, offset, output_size, spatial_scale,
|
||||
sampling_ratio, gamma):
|
||||
return g.op(
|
||||
'mmcv::MMCVDeformRoIPool',
|
||||
input,
|
||||
rois,
|
||||
offset,
|
||||
pooled_height_i=output_size[0],
|
||||
pooled_width_i=output_size[1],
|
||||
spatial_scale_f=spatial_scale,
|
||||
sampling_ratio_f=sampling_ratio,
|
||||
gamma_f=gamma)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
input,
|
||||
rois,
|
||||
offset,
|
||||
output_size,
|
||||
spatial_scale=1.0,
|
||||
sampling_ratio=0,
|
||||
gamma=0.1):
|
||||
if offset is None:
|
||||
offset = input.new_zeros(0)
|
||||
ctx.output_size = _pair(output_size)
|
||||
ctx.spatial_scale = float(spatial_scale)
|
||||
ctx.sampling_ratio = int(sampling_ratio)
|
||||
ctx.gamma = float(gamma)
|
||||
|
||||
assert rois.size(1) == 5, 'RoI must be (idx, x1, y1, x2, y2)!'
|
||||
|
||||
output_shape = (rois.size(0), input.size(1), ctx.output_size[0],
|
||||
ctx.output_size[1])
|
||||
output = input.new_zeros(output_shape)
|
||||
|
||||
ext_module.deform_roi_pool_forward(
|
||||
input,
|
||||
rois,
|
||||
offset,
|
||||
output,
|
||||
pooled_height=ctx.output_size[0],
|
||||
pooled_width=ctx.output_size[1],
|
||||
spatial_scale=ctx.spatial_scale,
|
||||
sampling_ratio=ctx.sampling_ratio,
|
||||
gamma=ctx.gamma)
|
||||
|
||||
ctx.save_for_backward(input, rois, offset)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
input, rois, offset = ctx.saved_tensors
|
||||
grad_input = grad_output.new_zeros(input.shape)
|
||||
grad_offset = grad_output.new_zeros(offset.shape)
|
||||
|
||||
ext_module.deform_roi_pool_backward(
|
||||
grad_output,
|
||||
input,
|
||||
rois,
|
||||
offset,
|
||||
grad_input,
|
||||
grad_offset,
|
||||
pooled_height=ctx.output_size[0],
|
||||
pooled_width=ctx.output_size[1],
|
||||
spatial_scale=ctx.spatial_scale,
|
||||
sampling_ratio=ctx.sampling_ratio,
|
||||
gamma=ctx.gamma)
|
||||
if grad_offset.numel() == 0:
|
||||
grad_offset = None
|
||||
return grad_input, None, grad_offset, None, None, None, None
|
||||
|
||||
|
||||
deform_roi_pool = DeformRoIPoolFunction.apply
|
||||
|
||||
|
||||
class DeformRoIPool(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
output_size,
|
||||
spatial_scale=1.0,
|
||||
sampling_ratio=0,
|
||||
gamma=0.1):
|
||||
super(DeformRoIPool, self).__init__()
|
||||
self.output_size = _pair(output_size)
|
||||
self.spatial_scale = float(spatial_scale)
|
||||
self.sampling_ratio = int(sampling_ratio)
|
||||
self.gamma = float(gamma)
|
||||
|
||||
def forward(self, input, rois, offset=None):
|
||||
return deform_roi_pool(input, rois, offset, self.output_size,
|
||||
self.spatial_scale, self.sampling_ratio,
|
||||
self.gamma)
|
||||
|
||||
|
||||
class DeformRoIPoolPack(DeformRoIPool):
|
||||
|
||||
def __init__(self,
|
||||
output_size,
|
||||
output_channels,
|
||||
deform_fc_channels=1024,
|
||||
spatial_scale=1.0,
|
||||
sampling_ratio=0,
|
||||
gamma=0.1):
|
||||
super(DeformRoIPoolPack, self).__init__(output_size, spatial_scale,
|
||||
sampling_ratio, gamma)
|
||||
|
||||
self.output_channels = output_channels
|
||||
self.deform_fc_channels = deform_fc_channels
|
||||
|
||||
self.offset_fc = nn.Sequential(
|
||||
nn.Linear(
|
||||
self.output_size[0] * self.output_size[1] *
|
||||
self.output_channels, self.deform_fc_channels),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(self.deform_fc_channels, self.deform_fc_channels),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(self.deform_fc_channels,
|
||||
self.output_size[0] * self.output_size[1] * 2))
|
||||
self.offset_fc[-1].weight.data.zero_()
|
||||
self.offset_fc[-1].bias.data.zero_()
|
||||
|
||||
def forward(self, input, rois):
|
||||
assert input.size(1) == self.output_channels
|
||||
x = deform_roi_pool(input, rois, None, self.output_size,
|
||||
self.spatial_scale, self.sampling_ratio,
|
||||
self.gamma)
|
||||
rois_num = rois.size(0)
|
||||
offset = self.offset_fc(x.view(rois_num, -1))
|
||||
offset = offset.view(rois_num, 2, self.output_size[0],
|
||||
self.output_size[1])
|
||||
return deform_roi_pool(input, rois, offset, self.output_size,
|
||||
self.spatial_scale, self.sampling_ratio,
|
||||
self.gamma)
|
||||
|
||||
|
||||
class ModulatedDeformRoIPoolPack(DeformRoIPool):
|
||||
|
||||
def __init__(self,
|
||||
output_size,
|
||||
output_channels,
|
||||
deform_fc_channels=1024,
|
||||
spatial_scale=1.0,
|
||||
sampling_ratio=0,
|
||||
gamma=0.1):
|
||||
super(ModulatedDeformRoIPoolPack,
|
||||
self).__init__(output_size, spatial_scale, sampling_ratio, gamma)
|
||||
|
||||
self.output_channels = output_channels
|
||||
self.deform_fc_channels = deform_fc_channels
|
||||
|
||||
self.offset_fc = nn.Sequential(
|
||||
nn.Linear(
|
||||
self.output_size[0] * self.output_size[1] *
|
||||
self.output_channels, self.deform_fc_channels),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(self.deform_fc_channels, self.deform_fc_channels),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(self.deform_fc_channels,
|
||||
self.output_size[0] * self.output_size[1] * 2))
|
||||
self.offset_fc[-1].weight.data.zero_()
|
||||
self.offset_fc[-1].bias.data.zero_()
|
||||
|
||||
self.mask_fc = nn.Sequential(
|
||||
nn.Linear(
|
||||
self.output_size[0] * self.output_size[1] *
|
||||
self.output_channels, self.deform_fc_channels),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(self.deform_fc_channels,
|
||||
self.output_size[0] * self.output_size[1] * 1),
|
||||
nn.Sigmoid())
|
||||
self.mask_fc[2].weight.data.zero_()
|
||||
self.mask_fc[2].bias.data.zero_()
|
||||
|
||||
def forward(self, input, rois):
|
||||
assert input.size(1) == self.output_channels
|
||||
x = deform_roi_pool(input, rois, None, self.output_size,
|
||||
self.spatial_scale, self.sampling_ratio,
|
||||
self.gamma)
|
||||
rois_num = rois.size(0)
|
||||
offset = self.offset_fc(x.view(rois_num, -1))
|
||||
offset = offset.view(rois_num, 2, self.output_size[0],
|
||||
self.output_size[1])
|
||||
mask = self.mask_fc(x.view(rois_num, -1))
|
||||
mask = mask.view(rois_num, 1, self.output_size[0], self.output_size[1])
|
||||
d = deform_roi_pool(input, rois, offset, self.output_size,
|
||||
self.spatial_scale, self.sampling_ratio,
|
||||
self.gamma)
|
||||
return d * mask
|
||||
@@ -0,0 +1,43 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
# This file is for backward compatibility.
|
||||
# Module wrappers for empty tensor have been moved to mmcv.cnn.bricks.
|
||||
import warnings
|
||||
|
||||
from ..cnn.bricks.wrappers import Conv2d, ConvTranspose2d, Linear, MaxPool2d
|
||||
|
||||
|
||||
class Conv2d_deprecated(Conv2d):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
warnings.warn(
|
||||
'Importing Conv2d wrapper from "mmcv.ops" will be deprecated in'
|
||||
' the future. Please import them from "mmcv.cnn" instead')
|
||||
|
||||
|
||||
class ConvTranspose2d_deprecated(ConvTranspose2d):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
warnings.warn(
|
||||
'Importing ConvTranspose2d wrapper from "mmcv.ops" will be '
|
||||
'deprecated in the future. Please import them from "mmcv.cnn" '
|
||||
'instead')
|
||||
|
||||
|
||||
class MaxPool2d_deprecated(MaxPool2d):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
warnings.warn(
|
||||
'Importing MaxPool2d wrapper from "mmcv.ops" will be deprecated in'
|
||||
' the future. Please import them from "mmcv.cnn" instead')
|
||||
|
||||
|
||||
class Linear_deprecated(Linear):
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
warnings.warn(
|
||||
'Importing Linear wrapper from "mmcv.ops" will be deprecated in'
|
||||
' the future. Please import them from "mmcv.cnn" instead')
|
||||
@@ -0,0 +1,212 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'sigmoid_focal_loss_forward', 'sigmoid_focal_loss_backward',
|
||||
'softmax_focal_loss_forward', 'softmax_focal_loss_backward'
|
||||
])
|
||||
|
||||
|
||||
class SigmoidFocalLossFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input, target, gamma, alpha, weight, reduction):
|
||||
return g.op(
|
||||
'mmcv::MMCVSigmoidFocalLoss',
|
||||
input,
|
||||
target,
|
||||
gamma_f=gamma,
|
||||
alpha_f=alpha,
|
||||
weight_f=weight,
|
||||
reduction_s=reduction)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
input,
|
||||
target,
|
||||
gamma=2.0,
|
||||
alpha=0.25,
|
||||
weight=None,
|
||||
reduction='mean'):
|
||||
|
||||
assert isinstance(target, (torch.LongTensor, torch.cuda.LongTensor))
|
||||
assert input.dim() == 2
|
||||
assert target.dim() == 1
|
||||
assert input.size(0) == target.size(0)
|
||||
if weight is None:
|
||||
weight = input.new_empty(0)
|
||||
else:
|
||||
assert weight.dim() == 1
|
||||
assert input.size(1) == weight.size(0)
|
||||
ctx.reduction_dict = {'none': 0, 'mean': 1, 'sum': 2}
|
||||
assert reduction in ctx.reduction_dict.keys()
|
||||
|
||||
ctx.gamma = float(gamma)
|
||||
ctx.alpha = float(alpha)
|
||||
ctx.reduction = ctx.reduction_dict[reduction]
|
||||
|
||||
output = input.new_zeros(input.size())
|
||||
|
||||
ext_module.sigmoid_focal_loss_forward(
|
||||
input, target, weight, output, gamma=ctx.gamma, alpha=ctx.alpha)
|
||||
if ctx.reduction == ctx.reduction_dict['mean']:
|
||||
output = output.sum() / input.size(0)
|
||||
elif ctx.reduction == ctx.reduction_dict['sum']:
|
||||
output = output.sum()
|
||||
ctx.save_for_backward(input, target, weight)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
input, target, weight = ctx.saved_tensors
|
||||
|
||||
grad_input = input.new_zeros(input.size())
|
||||
|
||||
ext_module.sigmoid_focal_loss_backward(
|
||||
input,
|
||||
target,
|
||||
weight,
|
||||
grad_input,
|
||||
gamma=ctx.gamma,
|
||||
alpha=ctx.alpha)
|
||||
|
||||
grad_input *= grad_output
|
||||
if ctx.reduction == ctx.reduction_dict['mean']:
|
||||
grad_input /= input.size(0)
|
||||
return grad_input, None, None, None, None, None
|
||||
|
||||
|
||||
sigmoid_focal_loss = SigmoidFocalLossFunction.apply
|
||||
|
||||
|
||||
class SigmoidFocalLoss(nn.Module):
|
||||
|
||||
def __init__(self, gamma, alpha, weight=None, reduction='mean'):
|
||||
super(SigmoidFocalLoss, self).__init__()
|
||||
self.gamma = gamma
|
||||
self.alpha = alpha
|
||||
self.register_buffer('weight', weight)
|
||||
self.reduction = reduction
|
||||
|
||||
def forward(self, input, target):
|
||||
return sigmoid_focal_loss(input, target, self.gamma, self.alpha,
|
||||
self.weight, self.reduction)
|
||||
|
||||
def __repr__(self):
|
||||
s = self.__class__.__name__
|
||||
s += f'(gamma={self.gamma}, '
|
||||
s += f'alpha={self.alpha}, '
|
||||
s += f'reduction={self.reduction})'
|
||||
return s
|
||||
|
||||
|
||||
class SoftmaxFocalLossFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input, target, gamma, alpha, weight, reduction):
|
||||
return g.op(
|
||||
'mmcv::MMCVSoftmaxFocalLoss',
|
||||
input,
|
||||
target,
|
||||
gamma_f=gamma,
|
||||
alpha_f=alpha,
|
||||
weight_f=weight,
|
||||
reduction_s=reduction)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
input,
|
||||
target,
|
||||
gamma=2.0,
|
||||
alpha=0.25,
|
||||
weight=None,
|
||||
reduction='mean'):
|
||||
|
||||
assert isinstance(target, (torch.LongTensor, torch.cuda.LongTensor))
|
||||
assert input.dim() == 2
|
||||
assert target.dim() == 1
|
||||
assert input.size(0) == target.size(0)
|
||||
if weight is None:
|
||||
weight = input.new_empty(0)
|
||||
else:
|
||||
assert weight.dim() == 1
|
||||
assert input.size(1) == weight.size(0)
|
||||
ctx.reduction_dict = {'none': 0, 'mean': 1, 'sum': 2}
|
||||
assert reduction in ctx.reduction_dict.keys()
|
||||
|
||||
ctx.gamma = float(gamma)
|
||||
ctx.alpha = float(alpha)
|
||||
ctx.reduction = ctx.reduction_dict[reduction]
|
||||
|
||||
channel_stats, _ = torch.max(input, dim=1)
|
||||
input_softmax = input - channel_stats.unsqueeze(1).expand_as(input)
|
||||
input_softmax.exp_()
|
||||
|
||||
channel_stats = input_softmax.sum(dim=1)
|
||||
input_softmax /= channel_stats.unsqueeze(1).expand_as(input)
|
||||
|
||||
output = input.new_zeros(input.size(0))
|
||||
ext_module.softmax_focal_loss_forward(
|
||||
input_softmax,
|
||||
target,
|
||||
weight,
|
||||
output,
|
||||
gamma=ctx.gamma,
|
||||
alpha=ctx.alpha)
|
||||
|
||||
if ctx.reduction == ctx.reduction_dict['mean']:
|
||||
output = output.sum() / input.size(0)
|
||||
elif ctx.reduction == ctx.reduction_dict['sum']:
|
||||
output = output.sum()
|
||||
ctx.save_for_backward(input_softmax, target, weight)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input_softmax, target, weight = ctx.saved_tensors
|
||||
buff = input_softmax.new_zeros(input_softmax.size(0))
|
||||
grad_input = input_softmax.new_zeros(input_softmax.size())
|
||||
|
||||
ext_module.softmax_focal_loss_backward(
|
||||
input_softmax,
|
||||
target,
|
||||
weight,
|
||||
buff,
|
||||
grad_input,
|
||||
gamma=ctx.gamma,
|
||||
alpha=ctx.alpha)
|
||||
|
||||
grad_input *= grad_output
|
||||
if ctx.reduction == ctx.reduction_dict['mean']:
|
||||
grad_input /= input_softmax.size(0)
|
||||
return grad_input, None, None, None, None, None
|
||||
|
||||
|
||||
softmax_focal_loss = SoftmaxFocalLossFunction.apply
|
||||
|
||||
|
||||
class SoftmaxFocalLoss(nn.Module):
|
||||
|
||||
def __init__(self, gamma, alpha, weight=None, reduction='mean'):
|
||||
super(SoftmaxFocalLoss, self).__init__()
|
||||
self.gamma = gamma
|
||||
self.alpha = alpha
|
||||
self.register_buffer('weight', weight)
|
||||
self.reduction = reduction
|
||||
|
||||
def forward(self, input, target):
|
||||
return softmax_focal_loss(input, target, self.gamma, self.alpha,
|
||||
self.weight, self.reduction)
|
||||
|
||||
def __repr__(self):
|
||||
s = self.__class__.__name__
|
||||
s += f'(gamma={self.gamma}, '
|
||||
s += f'alpha={self.alpha}, '
|
||||
s += f'reduction={self.reduction})'
|
||||
return s
|
||||
@@ -0,0 +1,83 @@
|
||||
import torch
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'furthest_point_sampling_forward',
|
||||
'furthest_point_sampling_with_dist_forward'
|
||||
])
|
||||
|
||||
|
||||
class FurthestPointSampling(Function):
|
||||
"""Uses iterative furthest point sampling to select a set of features whose
|
||||
corresponding points have the furthest distance."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, points_xyz: torch.Tensor,
|
||||
num_points: int) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
points_xyz (Tensor): (B, N, 3) where N > num_points.
|
||||
num_points (int): Number of points in the sampled set.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, num_points) indices of the sampled points.
|
||||
"""
|
||||
assert points_xyz.is_contiguous()
|
||||
|
||||
B, N = points_xyz.size()[:2]
|
||||
output = torch.cuda.IntTensor(B, num_points)
|
||||
temp = torch.cuda.FloatTensor(B, N).fill_(1e10)
|
||||
|
||||
ext_module.furthest_point_sampling_forward(
|
||||
points_xyz,
|
||||
temp,
|
||||
output,
|
||||
b=B,
|
||||
n=N,
|
||||
m=num_points,
|
||||
)
|
||||
if torch.__version__ != 'parrots':
|
||||
ctx.mark_non_differentiable(output)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(xyz, a=None):
|
||||
return None, None
|
||||
|
||||
|
||||
class FurthestPointSamplingWithDist(Function):
|
||||
"""Uses iterative furthest point sampling to select a set of features whose
|
||||
corresponding points have the furthest distance."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, points_dist: torch.Tensor,
|
||||
num_points: int) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
points_dist (Tensor): (B, N, N) Distance between each point pair.
|
||||
num_points (int): Number of points in the sampled set.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, num_points) indices of the sampled points.
|
||||
"""
|
||||
assert points_dist.is_contiguous()
|
||||
|
||||
B, N, _ = points_dist.size()
|
||||
output = points_dist.new_zeros([B, num_points], dtype=torch.int32)
|
||||
temp = points_dist.new_zeros([B, N]).fill_(1e10)
|
||||
|
||||
ext_module.furthest_point_sampling_with_dist_forward(
|
||||
points_dist, temp, output, b=B, n=N, m=num_points)
|
||||
if torch.__version__ != 'parrots':
|
||||
ctx.mark_non_differentiable(output)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(xyz, a=None):
|
||||
return None, None
|
||||
|
||||
|
||||
furthest_point_sample = FurthestPointSampling.apply
|
||||
furthest_point_sample_with_dist = FurthestPointSamplingWithDist.apply
|
||||
@@ -0,0 +1,268 @@
|
||||
# modified from https://github.com/rosinality/stylegan2-pytorch/blob/master/op/fused_act.py # noqa:E501
|
||||
|
||||
# Copyright (c) 2021, NVIDIA Corporation. All rights reserved.
|
||||
# NVIDIA Source Code License for StyleGAN2 with Adaptive Discriminator
|
||||
# Augmentation (ADA)
|
||||
# =======================================================================
|
||||
|
||||
# 1. Definitions
|
||||
|
||||
# "Licensor" means any person or entity that distributes its Work.
|
||||
|
||||
# "Software" means the original work of authorship made available under
|
||||
# this License.
|
||||
|
||||
# "Work" means the Software and any additions to or derivative works of
|
||||
# the Software that are made available under this License.
|
||||
|
||||
# The terms "reproduce," "reproduction," "derivative works," and
|
||||
# "distribution" have the meaning as provided under U.S. copyright law;
|
||||
# provided, however, that for the purposes of this License, derivative
|
||||
# works shall not include works that remain separable from, or merely
|
||||
# link (or bind by name) to the interfaces of, the Work.
|
||||
|
||||
# Works, including the Software, are "made available" under this License
|
||||
# by including in or with the Work either (a) a copyright notice
|
||||
# referencing the applicability of this License to the Work, or (b) a
|
||||
# copy of this License.
|
||||
|
||||
# 2. License Grants
|
||||
|
||||
# 2.1 Copyright Grant. Subject to the terms and conditions of this
|
||||
# License, each Licensor grants to you a perpetual, worldwide,
|
||||
# non-exclusive, royalty-free, copyright license to reproduce,
|
||||
# prepare derivative works of, publicly display, publicly perform,
|
||||
# sublicense and distribute its Work and any resulting derivative
|
||||
# works in any form.
|
||||
|
||||
# 3. Limitations
|
||||
|
||||
# 3.1 Redistribution. You may reproduce or distribute the Work only
|
||||
# if (a) you do so under this License, (b) you include a complete
|
||||
# copy of this License with your distribution, and (c) you retain
|
||||
# without modification any copyright, patent, trademark, or
|
||||
# attribution notices that are present in the Work.
|
||||
|
||||
# 3.2 Derivative Works. You may specify that additional or different
|
||||
# terms apply to the use, reproduction, and distribution of your
|
||||
# derivative works of the Work ("Your Terms") only if (a) Your Terms
|
||||
# provide that the use limitation in Section 3.3 applies to your
|
||||
# derivative works, and (b) you identify the specific derivative
|
||||
# works that are subject to Your Terms. Notwithstanding Your Terms,
|
||||
# this License (including the redistribution requirements in Section
|
||||
# 3.1) will continue to apply to the Work itself.
|
||||
|
||||
# 3.3 Use Limitation. The Work and any derivative works thereof only
|
||||
# may be used or intended for use non-commercially. Notwithstanding
|
||||
# the foregoing, NVIDIA and its affiliates may use the Work and any
|
||||
# derivative works commercially. As used herein, "non-commercially"
|
||||
# means for research or evaluation purposes only.
|
||||
|
||||
# 3.4 Patent Claims. If you bring or threaten to bring a patent claim
|
||||
# against any Licensor (including any claim, cross-claim or
|
||||
# counterclaim in a lawsuit) to enforce any patents that you allege
|
||||
# are infringed by any Work, then your rights under this License from
|
||||
# such Licensor (including the grant in Section 2.1) will terminate
|
||||
# immediately.
|
||||
|
||||
# 3.5 Trademarks. This License does not grant any rights to use any
|
||||
# Licensor’s or its affiliates’ names, logos, or trademarks, except
|
||||
# as necessary to reproduce the notices described in this License.
|
||||
|
||||
# 3.6 Termination. If you violate any term of this License, then your
|
||||
# rights under this License (including the grant in Section 2.1) will
|
||||
# terminate immediately.
|
||||
|
||||
# 4. Disclaimer of Warranty.
|
||||
|
||||
# THE WORK IS PROVIDED "AS IS" WITHOUT WARRANTIES OR CONDITIONS OF ANY
|
||||
# KIND, EITHER EXPRESS OR IMPLIED, INCLUDING WARRANTIES OR CONDITIONS OF
|
||||
# MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE, TITLE OR
|
||||
# NON-INFRINGEMENT. YOU BEAR THE RISK OF UNDERTAKING ANY ACTIVITIES UNDER
|
||||
# THIS LICENSE.
|
||||
|
||||
# 5. Limitation of Liability.
|
||||
|
||||
# EXCEPT AS PROHIBITED BY APPLICABLE LAW, IN NO EVENT AND UNDER NO LEGAL
|
||||
# THEORY, WHETHER IN TORT (INCLUDING NEGLIGENCE), CONTRACT, OR OTHERWISE
|
||||
# SHALL ANY LICENSOR BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY DIRECT,
|
||||
# INDIRECT, SPECIAL, INCIDENTAL, OR CONSEQUENTIAL DAMAGES ARISING OUT OF
|
||||
# OR RELATED TO THIS LICENSE, THE USE OR INABILITY TO USE THE WORK
|
||||
# (INCLUDING BUT NOT LIMITED TO LOSS OF GOODWILL, BUSINESS INTERRUPTION,
|
||||
# LOST PROFITS OR DATA, COMPUTER FAILURE OR MALFUNCTION, OR ANY OTHER
|
||||
# COMMERCIAL DAMAGES OR LOSSES), EVEN IF THE LICENSOR HAS BEEN ADVISED OF
|
||||
# THE POSSIBILITY OF SUCH DAMAGES.
|
||||
|
||||
# =======================================================================
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['fused_bias_leakyrelu'])
|
||||
|
||||
|
||||
class FusedBiasLeakyReLUFunctionBackward(Function):
|
||||
"""Calculate second order deviation.
|
||||
|
||||
This function is to compute the second order deviation for the fused leaky
|
||||
relu operation.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, grad_output, out, negative_slope, scale):
|
||||
ctx.save_for_backward(out)
|
||||
ctx.negative_slope = negative_slope
|
||||
ctx.scale = scale
|
||||
|
||||
empty = grad_output.new_empty(0)
|
||||
|
||||
grad_input = ext_module.fused_bias_leakyrelu(
|
||||
grad_output,
|
||||
empty,
|
||||
out,
|
||||
act=3,
|
||||
grad=1,
|
||||
alpha=negative_slope,
|
||||
scale=scale)
|
||||
|
||||
dim = [0]
|
||||
|
||||
if grad_input.ndim > 2:
|
||||
dim += list(range(2, grad_input.ndim))
|
||||
|
||||
grad_bias = grad_input.sum(dim).detach()
|
||||
|
||||
return grad_input, grad_bias
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, gradgrad_input, gradgrad_bias):
|
||||
out, = ctx.saved_tensors
|
||||
|
||||
# The second order deviation, in fact, contains two parts, while the
|
||||
# the first part is zero. Thus, we direct consider the second part
|
||||
# which is similar with the first order deviation in implementation.
|
||||
gradgrad_out = ext_module.fused_bias_leakyrelu(
|
||||
gradgrad_input,
|
||||
gradgrad_bias.to(out.dtype),
|
||||
out,
|
||||
act=3,
|
||||
grad=1,
|
||||
alpha=ctx.negative_slope,
|
||||
scale=ctx.scale)
|
||||
|
||||
return gradgrad_out, None, None, None
|
||||
|
||||
|
||||
class FusedBiasLeakyReLUFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input, bias, negative_slope, scale):
|
||||
empty = input.new_empty(0)
|
||||
|
||||
out = ext_module.fused_bias_leakyrelu(
|
||||
input,
|
||||
bias,
|
||||
empty,
|
||||
act=3,
|
||||
grad=0,
|
||||
alpha=negative_slope,
|
||||
scale=scale)
|
||||
ctx.save_for_backward(out)
|
||||
ctx.negative_slope = negative_slope
|
||||
ctx.scale = scale
|
||||
|
||||
return out
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
out, = ctx.saved_tensors
|
||||
|
||||
grad_input, grad_bias = FusedBiasLeakyReLUFunctionBackward.apply(
|
||||
grad_output, out, ctx.negative_slope, ctx.scale)
|
||||
|
||||
return grad_input, grad_bias, None, None
|
||||
|
||||
|
||||
class FusedBiasLeakyReLU(nn.Module):
|
||||
"""Fused bias leaky ReLU.
|
||||
|
||||
This function is introduced in the StyleGAN2:
|
||||
http://arxiv.org/abs/1912.04958
|
||||
|
||||
The bias term comes from the convolution operation. In addition, to keep
|
||||
the variance of the feature map or gradients unchanged, they also adopt a
|
||||
scale similarly with Kaiming initialization. However, since the
|
||||
:math:`1+{alpha}^2` : is too small, we can just ignore it. Therefore, the
|
||||
final scale is just :math:`\sqrt{2}`:. Of course, you may change it with # noqa: W605, E501
|
||||
your own scale.
|
||||
|
||||
TODO: Implement the CPU version.
|
||||
|
||||
Args:
|
||||
channel (int): The channel number of the feature map.
|
||||
negative_slope (float, optional): Same as nn.LeakyRelu.
|
||||
Defaults to 0.2.
|
||||
scale (float, optional): A scalar to adjust the variance of the feature
|
||||
map. Defaults to 2**0.5.
|
||||
"""
|
||||
|
||||
def __init__(self, num_channels, negative_slope=0.2, scale=2**0.5):
|
||||
super(FusedBiasLeakyReLU, self).__init__()
|
||||
|
||||
self.bias = nn.Parameter(torch.zeros(num_channels))
|
||||
self.negative_slope = negative_slope
|
||||
self.scale = scale
|
||||
|
||||
def forward(self, input):
|
||||
return fused_bias_leakyrelu(input, self.bias, self.negative_slope,
|
||||
self.scale)
|
||||
|
||||
|
||||
def fused_bias_leakyrelu(input, bias, negative_slope=0.2, scale=2**0.5):
|
||||
"""Fused bias leaky ReLU function.
|
||||
|
||||
This function is introduced in the StyleGAN2:
|
||||
http://arxiv.org/abs/1912.04958
|
||||
|
||||
The bias term comes from the convolution operation. In addition, to keep
|
||||
the variance of the feature map or gradients unchanged, they also adopt a
|
||||
scale similarly with Kaiming initialization. However, since the
|
||||
:math:`1+{alpha}^2` : is too small, we can just ignore it. Therefore, the
|
||||
final scale is just :math:`\sqrt{2}`:. Of course, you may change it with # noqa: W605, E501
|
||||
your own scale.
|
||||
|
||||
Args:
|
||||
input (torch.Tensor): Input feature map.
|
||||
bias (nn.Parameter): The bias from convolution operation.
|
||||
negative_slope (float, optional): Same as nn.LeakyRelu.
|
||||
Defaults to 0.2.
|
||||
scale (float, optional): A scalar to adjust the variance of the feature
|
||||
map. Defaults to 2**0.5.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Feature map after non-linear activation.
|
||||
"""
|
||||
|
||||
if not input.is_cuda:
|
||||
return bias_leakyrelu_ref(input, bias, negative_slope, scale)
|
||||
|
||||
return FusedBiasLeakyReLUFunction.apply(input, bias.to(input.dtype),
|
||||
negative_slope, scale)
|
||||
|
||||
|
||||
def bias_leakyrelu_ref(x, bias, negative_slope=0.2, scale=2**0.5):
|
||||
|
||||
if bias is not None:
|
||||
assert bias.ndim == 1
|
||||
assert bias.shape[0] == x.shape[1]
|
||||
x = x + bias.reshape([-1 if i == 1 else 1 for i in range(x.ndim)])
|
||||
|
||||
x = F.leaky_relu(x, negative_slope)
|
||||
if scale != 1:
|
||||
x = x * scale
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,57 @@
|
||||
import torch
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['gather_points_forward', 'gather_points_backward'])
|
||||
|
||||
|
||||
class GatherPoints(Function):
|
||||
"""Gather points with given index."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, features: torch.Tensor,
|
||||
indices: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
features (Tensor): (B, C, N) features to gather.
|
||||
indices (Tensor): (B, M) where M is the number of points.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, C, M) where M is the number of points.
|
||||
"""
|
||||
assert features.is_contiguous()
|
||||
assert indices.is_contiguous()
|
||||
|
||||
B, npoint = indices.size()
|
||||
_, C, N = features.size()
|
||||
output = torch.cuda.FloatTensor(B, C, npoint)
|
||||
|
||||
ext_module.gather_points_forward(
|
||||
features, indices, output, b=B, c=C, n=N, npoints=npoint)
|
||||
|
||||
ctx.for_backwards = (indices, C, N)
|
||||
if torch.__version__ != 'parrots':
|
||||
ctx.mark_non_differentiable(indices)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_out):
|
||||
idx, C, N = ctx.for_backwards
|
||||
B, npoint = idx.size()
|
||||
|
||||
grad_features = torch.cuda.FloatTensor(B, C, N).zero_()
|
||||
grad_out_data = grad_out.data.contiguous()
|
||||
ext_module.gather_points_backward(
|
||||
grad_out_data,
|
||||
idx,
|
||||
grad_features.data,
|
||||
b=B,
|
||||
c=C,
|
||||
n=N,
|
||||
npoints=npoint)
|
||||
return grad_features, None
|
||||
|
||||
|
||||
gather_points = GatherPoints.apply
|
||||
@@ -0,0 +1,224 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn as nn
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
from .ball_query import ball_query
|
||||
from .knn import knn
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['group_points_forward', 'group_points_backward'])
|
||||
|
||||
|
||||
class QueryAndGroup(nn.Module):
|
||||
"""Groups points with a ball query of radius.
|
||||
|
||||
Args:
|
||||
max_radius (float): The maximum radius of the balls.
|
||||
If None is given, we will use kNN sampling instead of ball query.
|
||||
sample_num (int): Maximum number of features to gather in the ball.
|
||||
min_radius (float, optional): The minimum radius of the balls.
|
||||
Default: 0.
|
||||
use_xyz (bool, optional): Whether to use xyz.
|
||||
Default: True.
|
||||
return_grouped_xyz (bool, optional): Whether to return grouped xyz.
|
||||
Default: False.
|
||||
normalize_xyz (bool, optional): Whether to normalize xyz.
|
||||
Default: False.
|
||||
uniform_sample (bool, optional): Whether to sample uniformly.
|
||||
Default: False
|
||||
return_unique_cnt (bool, optional): Whether to return the count of
|
||||
unique samples. Default: False.
|
||||
return_grouped_idx (bool, optional): Whether to return grouped idx.
|
||||
Default: False.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
max_radius,
|
||||
sample_num,
|
||||
min_radius=0,
|
||||
use_xyz=True,
|
||||
return_grouped_xyz=False,
|
||||
normalize_xyz=False,
|
||||
uniform_sample=False,
|
||||
return_unique_cnt=False,
|
||||
return_grouped_idx=False):
|
||||
super().__init__()
|
||||
self.max_radius = max_radius
|
||||
self.min_radius = min_radius
|
||||
self.sample_num = sample_num
|
||||
self.use_xyz = use_xyz
|
||||
self.return_grouped_xyz = return_grouped_xyz
|
||||
self.normalize_xyz = normalize_xyz
|
||||
self.uniform_sample = uniform_sample
|
||||
self.return_unique_cnt = return_unique_cnt
|
||||
self.return_grouped_idx = return_grouped_idx
|
||||
if self.return_unique_cnt:
|
||||
assert self.uniform_sample, \
|
||||
'uniform_sample should be True when ' \
|
||||
'returning the count of unique samples'
|
||||
if self.max_radius is None:
|
||||
assert not self.normalize_xyz, \
|
||||
'can not normalize grouped xyz when max_radius is None'
|
||||
|
||||
def forward(self, points_xyz, center_xyz, features=None):
|
||||
"""
|
||||
Args:
|
||||
points_xyz (Tensor): (B, N, 3) xyz coordinates of the features.
|
||||
center_xyz (Tensor): (B, npoint, 3) coordinates of the centriods.
|
||||
features (Tensor): (B, C, N) Descriptors of the features.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, 3 + C, npoint, sample_num) Grouped feature.
|
||||
"""
|
||||
# if self.max_radius is None, we will perform kNN instead of ball query
|
||||
# idx is of shape [B, npoint, sample_num]
|
||||
if self.max_radius is None:
|
||||
idx = knn(self.sample_num, points_xyz, center_xyz, False)
|
||||
idx = idx.transpose(1, 2).contiguous()
|
||||
else:
|
||||
idx = ball_query(self.min_radius, self.max_radius, self.sample_num,
|
||||
points_xyz, center_xyz)
|
||||
|
||||
if self.uniform_sample:
|
||||
unique_cnt = torch.zeros((idx.shape[0], idx.shape[1]))
|
||||
for i_batch in range(idx.shape[0]):
|
||||
for i_region in range(idx.shape[1]):
|
||||
unique_ind = torch.unique(idx[i_batch, i_region, :])
|
||||
num_unique = unique_ind.shape[0]
|
||||
unique_cnt[i_batch, i_region] = num_unique
|
||||
sample_ind = torch.randint(
|
||||
0,
|
||||
num_unique, (self.sample_num - num_unique, ),
|
||||
dtype=torch.long)
|
||||
all_ind = torch.cat((unique_ind, unique_ind[sample_ind]))
|
||||
idx[i_batch, i_region, :] = all_ind
|
||||
|
||||
xyz_trans = points_xyz.transpose(1, 2).contiguous()
|
||||
# (B, 3, npoint, sample_num)
|
||||
grouped_xyz = grouping_operation(xyz_trans, idx)
|
||||
grouped_xyz_diff = grouped_xyz - \
|
||||
center_xyz.transpose(1, 2).unsqueeze(-1) # relative offsets
|
||||
if self.normalize_xyz:
|
||||
grouped_xyz_diff /= self.max_radius
|
||||
|
||||
if features is not None:
|
||||
grouped_features = grouping_operation(features, idx)
|
||||
if self.use_xyz:
|
||||
# (B, C + 3, npoint, sample_num)
|
||||
new_features = torch.cat([grouped_xyz_diff, grouped_features],
|
||||
dim=1)
|
||||
else:
|
||||
new_features = grouped_features
|
||||
else:
|
||||
assert (self.use_xyz
|
||||
), 'Cannot have not features and not use xyz as a feature!'
|
||||
new_features = grouped_xyz_diff
|
||||
|
||||
ret = [new_features]
|
||||
if self.return_grouped_xyz:
|
||||
ret.append(grouped_xyz)
|
||||
if self.return_unique_cnt:
|
||||
ret.append(unique_cnt)
|
||||
if self.return_grouped_idx:
|
||||
ret.append(idx)
|
||||
if len(ret) == 1:
|
||||
return ret[0]
|
||||
else:
|
||||
return tuple(ret)
|
||||
|
||||
|
||||
class GroupAll(nn.Module):
|
||||
"""Group xyz with feature.
|
||||
|
||||
Args:
|
||||
use_xyz (bool): Whether to use xyz.
|
||||
"""
|
||||
|
||||
def __init__(self, use_xyz: bool = True):
|
||||
super().__init__()
|
||||
self.use_xyz = use_xyz
|
||||
|
||||
def forward(self,
|
||||
xyz: torch.Tensor,
|
||||
new_xyz: torch.Tensor,
|
||||
features: torch.Tensor = None):
|
||||
"""
|
||||
Args:
|
||||
xyz (Tensor): (B, N, 3) xyz coordinates of the features.
|
||||
new_xyz (Tensor): new xyz coordinates of the features.
|
||||
features (Tensor): (B, C, N) features to group.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, C + 3, 1, N) Grouped feature.
|
||||
"""
|
||||
grouped_xyz = xyz.transpose(1, 2).unsqueeze(2)
|
||||
if features is not None:
|
||||
grouped_features = features.unsqueeze(2)
|
||||
if self.use_xyz:
|
||||
# (B, 3 + C, 1, N)
|
||||
new_features = torch.cat([grouped_xyz, grouped_features],
|
||||
dim=1)
|
||||
else:
|
||||
new_features = grouped_features
|
||||
else:
|
||||
new_features = grouped_xyz
|
||||
|
||||
return new_features
|
||||
|
||||
|
||||
class GroupingOperation(Function):
|
||||
"""Group feature with given index."""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, features: torch.Tensor,
|
||||
indices: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
features (Tensor): (B, C, N) tensor of features to group.
|
||||
indices (Tensor): (B, npoint, nsample) the indices of
|
||||
features to group with.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, C, npoint, nsample) Grouped features.
|
||||
"""
|
||||
features = features.contiguous()
|
||||
indices = indices.contiguous()
|
||||
|
||||
B, nfeatures, nsample = indices.size()
|
||||
_, C, N = features.size()
|
||||
output = torch.cuda.FloatTensor(B, C, nfeatures, nsample)
|
||||
|
||||
ext_module.group_points_forward(B, C, N, nfeatures, nsample, features,
|
||||
indices, output)
|
||||
|
||||
ctx.for_backwards = (indices, N)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx,
|
||||
grad_out: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Args:
|
||||
grad_out (Tensor): (B, C, npoint, nsample) tensor of the gradients
|
||||
of the output from forward.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, C, N) gradient of the features.
|
||||
"""
|
||||
idx, N = ctx.for_backwards
|
||||
|
||||
B, C, npoint, nsample = grad_out.size()
|
||||
grad_features = torch.cuda.FloatTensor(B, C, N).zero_()
|
||||
|
||||
grad_out_data = grad_out.data.contiguous()
|
||||
ext_module.group_points_backward(B, C, N, npoint, nsample,
|
||||
grad_out_data, idx,
|
||||
grad_features.data)
|
||||
return grad_features, None
|
||||
|
||||
|
||||
grouping_operation = GroupingOperation.apply
|
||||
@@ -0,0 +1,36 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import glob
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
import parrots
|
||||
|
||||
def get_compiler_version():
|
||||
return 'GCC ' + parrots.version.compiler
|
||||
|
||||
def get_compiling_cuda_version():
|
||||
return parrots.version.cuda
|
||||
else:
|
||||
from ..utils import ext_loader
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['get_compiler_version', 'get_compiling_cuda_version'])
|
||||
|
||||
def get_compiler_version():
|
||||
return ext_module.get_compiler_version()
|
||||
|
||||
def get_compiling_cuda_version():
|
||||
return ext_module.get_compiling_cuda_version()
|
||||
|
||||
|
||||
def get_onnxruntime_op_path():
|
||||
wildcard = os.path.join(
|
||||
os.path.abspath(os.path.dirname(os.path.dirname(__file__))),
|
||||
'_ext_ort.*.so')
|
||||
|
||||
paths = glob.glob(wildcard)
|
||||
if len(paths) > 0:
|
||||
return paths[0]
|
||||
else:
|
||||
return ''
|
||||
@@ -0,0 +1,85 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import torch
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'iou3d_boxes_iou_bev_forward', 'iou3d_nms_forward',
|
||||
'iou3d_nms_normal_forward'
|
||||
])
|
||||
|
||||
|
||||
def boxes_iou_bev(boxes_a, boxes_b):
|
||||
"""Calculate boxes IoU in the Bird's Eye View.
|
||||
|
||||
Args:
|
||||
boxes_a (torch.Tensor): Input boxes a with shape (M, 5).
|
||||
boxes_b (torch.Tensor): Input boxes b with shape (N, 5).
|
||||
|
||||
Returns:
|
||||
ans_iou (torch.Tensor): IoU result with shape (M, N).
|
||||
"""
|
||||
ans_iou = boxes_a.new_zeros(
|
||||
torch.Size((boxes_a.shape[0], boxes_b.shape[0])))
|
||||
|
||||
ext_module.iou3d_boxes_iou_bev_forward(boxes_a.contiguous(),
|
||||
boxes_b.contiguous(), ans_iou)
|
||||
|
||||
return ans_iou
|
||||
|
||||
|
||||
def nms_bev(boxes, scores, thresh, pre_max_size=None, post_max_size=None):
|
||||
"""NMS function GPU implementation (for BEV boxes). The overlap of two
|
||||
boxes for IoU calculation is defined as the exact overlapping area of the
|
||||
two boxes. In this function, one can also set ``pre_max_size`` and
|
||||
``post_max_size``.
|
||||
|
||||
Args:
|
||||
boxes (torch.Tensor): Input boxes with the shape of [N, 5]
|
||||
([x1, y1, x2, y2, ry]).
|
||||
scores (torch.Tensor): Scores of boxes with the shape of [N].
|
||||
thresh (float): Overlap threshold of NMS.
|
||||
pre_max_size (int, optional): Max size of boxes before NMS.
|
||||
Default: None.
|
||||
post_max_size (int, optional): Max size of boxes after NMS.
|
||||
Default: None.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Indexes after NMS.
|
||||
"""
|
||||
assert boxes.size(1) == 5, 'Input boxes shape should be [N, 5]'
|
||||
order = scores.sort(0, descending=True)[1]
|
||||
|
||||
if pre_max_size is not None:
|
||||
order = order[:pre_max_size]
|
||||
boxes = boxes[order].contiguous()
|
||||
|
||||
keep = torch.zeros(boxes.size(0), dtype=torch.long)
|
||||
num_out = ext_module.iou3d_nms_forward(boxes, keep, thresh)
|
||||
keep = order[keep[:num_out].cuda(boxes.device)].contiguous()
|
||||
if post_max_size is not None:
|
||||
keep = keep[:post_max_size]
|
||||
return keep
|
||||
|
||||
|
||||
def nms_normal_bev(boxes, scores, thresh):
|
||||
"""Normal NMS function GPU implementation (for BEV boxes). The overlap of
|
||||
two boxes for IoU calculation is defined as the exact overlapping area of
|
||||
the two boxes WITH their yaw angle set to 0.
|
||||
|
||||
Args:
|
||||
boxes (torch.Tensor): Input boxes with shape (N, 5).
|
||||
scores (torch.Tensor): Scores of predicted boxes with shape (N).
|
||||
thresh (float): Overlap threshold of NMS.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Remaining indices with scores in descending order.
|
||||
"""
|
||||
assert boxes.shape[1] == 5, 'Input boxes shape should be [N, 5]'
|
||||
order = scores.sort(0, descending=True)[1]
|
||||
|
||||
boxes = boxes[order].contiguous()
|
||||
|
||||
keep = torch.zeros(boxes.size(0), dtype=torch.long)
|
||||
num_out = ext_module.iou3d_nms_normal_forward(boxes, keep, thresh)
|
||||
return order[keep[:num_out].cuda(boxes.device)].contiguous()
|
||||
@@ -0,0 +1,77 @@
|
||||
import torch
|
||||
from torch.autograd import Function
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['knn_forward'])
|
||||
|
||||
|
||||
class KNN(Function):
|
||||
r"""KNN (CUDA) based on heap data structure.
|
||||
Modified from `PAConv <https://github.com/CVMI-Lab/PAConv/tree/main/
|
||||
scene_seg/lib/pointops/src/knnquery_heap>`_.
|
||||
|
||||
Find k-nearest points.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
k: int,
|
||||
xyz: torch.Tensor,
|
||||
center_xyz: torch.Tensor = None,
|
||||
transposed: bool = False) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
k (int): number of nearest neighbors.
|
||||
xyz (Tensor): (B, N, 3) if transposed == False, else (B, 3, N).
|
||||
xyz coordinates of the features.
|
||||
center_xyz (Tensor, optional): (B, npoint, 3) if transposed ==
|
||||
False, else (B, 3, npoint). centers of the knn query.
|
||||
Default: None.
|
||||
transposed (bool, optional): whether the input tensors are
|
||||
transposed. Should not explicitly use this keyword when
|
||||
calling knn (=KNN.apply), just add the fourth param.
|
||||
Default: False.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, k, npoint) tensor with the indices of
|
||||
the features that form k-nearest neighbours.
|
||||
"""
|
||||
assert (k > 0) & (k < 100), 'k should be in range(0, 100)'
|
||||
|
||||
if center_xyz is None:
|
||||
center_xyz = xyz
|
||||
|
||||
if transposed:
|
||||
xyz = xyz.transpose(2, 1).contiguous()
|
||||
center_xyz = center_xyz.transpose(2, 1).contiguous()
|
||||
|
||||
assert xyz.is_contiguous() # [B, N, 3]
|
||||
assert center_xyz.is_contiguous() # [B, npoint, 3]
|
||||
|
||||
center_xyz_device = center_xyz.get_device()
|
||||
assert center_xyz_device == xyz.get_device(), \
|
||||
'center_xyz and xyz should be put on the same device'
|
||||
if torch.cuda.current_device() != center_xyz_device:
|
||||
torch.cuda.set_device(center_xyz_device)
|
||||
|
||||
B, npoint, _ = center_xyz.shape
|
||||
N = xyz.shape[1]
|
||||
|
||||
idx = center_xyz.new_zeros((B, npoint, k)).int()
|
||||
dist2 = center_xyz.new_zeros((B, npoint, k)).float()
|
||||
|
||||
ext_module.knn_forward(
|
||||
xyz, center_xyz, idx, dist2, b=B, n=N, m=npoint, nsample=k)
|
||||
# idx shape to [B, k, npoint]
|
||||
idx = idx.transpose(2, 1).contiguous()
|
||||
if torch.__version__ != 'parrots':
|
||||
ctx.mark_non_differentiable(idx)
|
||||
return idx
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, a=None):
|
||||
return None, None, None
|
||||
|
||||
|
||||
knn = KNN.apply
|
||||
@@ -0,0 +1,111 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
from torch.nn.modules.utils import _pair
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['masked_im2col_forward', 'masked_col2im_forward'])
|
||||
|
||||
|
||||
class MaskedConv2dFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, features, mask, weight, bias, padding, stride):
|
||||
return g.op(
|
||||
'mmcv::MMCVMaskedConv2d',
|
||||
features,
|
||||
mask,
|
||||
weight,
|
||||
bias,
|
||||
padding_i=padding,
|
||||
stride_i=stride)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, features, mask, weight, bias, padding=0, stride=1):
|
||||
assert mask.dim() == 3 and mask.size(0) == 1
|
||||
assert features.dim() == 4 and features.size(0) == 1
|
||||
assert features.size()[2:] == mask.size()[1:]
|
||||
pad_h, pad_w = _pair(padding)
|
||||
stride_h, stride_w = _pair(stride)
|
||||
if stride_h != 1 or stride_w != 1:
|
||||
raise ValueError(
|
||||
'Stride could not only be 1 in masked_conv2d currently.')
|
||||
out_channel, in_channel, kernel_h, kernel_w = weight.size()
|
||||
|
||||
batch_size = features.size(0)
|
||||
out_h = int(
|
||||
math.floor((features.size(2) + 2 * pad_h -
|
||||
(kernel_h - 1) - 1) / stride_h + 1))
|
||||
out_w = int(
|
||||
math.floor((features.size(3) + 2 * pad_w -
|
||||
(kernel_h - 1) - 1) / stride_w + 1))
|
||||
mask_inds = torch.nonzero(mask[0] > 0, as_tuple=False)
|
||||
output = features.new_zeros(batch_size, out_channel, out_h, out_w)
|
||||
if mask_inds.numel() > 0:
|
||||
mask_h_idx = mask_inds[:, 0].contiguous()
|
||||
mask_w_idx = mask_inds[:, 1].contiguous()
|
||||
data_col = features.new_zeros(in_channel * kernel_h * kernel_w,
|
||||
mask_inds.size(0))
|
||||
ext_module.masked_im2col_forward(
|
||||
features,
|
||||
mask_h_idx,
|
||||
mask_w_idx,
|
||||
data_col,
|
||||
kernel_h=kernel_h,
|
||||
kernel_w=kernel_w,
|
||||
pad_h=pad_h,
|
||||
pad_w=pad_w)
|
||||
|
||||
masked_output = torch.addmm(1, bias[:, None], 1,
|
||||
weight.view(out_channel, -1), data_col)
|
||||
ext_module.masked_col2im_forward(
|
||||
masked_output,
|
||||
mask_h_idx,
|
||||
mask_w_idx,
|
||||
output,
|
||||
height=out_h,
|
||||
width=out_w,
|
||||
channels=out_channel)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
return (None, ) * 5
|
||||
|
||||
|
||||
masked_conv2d = MaskedConv2dFunction.apply
|
||||
|
||||
|
||||
class MaskedConv2d(nn.Conv2d):
|
||||
"""A MaskedConv2d which inherits the official Conv2d.
|
||||
|
||||
The masked forward doesn't implement the backward function and only
|
||||
supports the stride parameter to be 1 currently.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
bias=True):
|
||||
super(MaskedConv2d,
|
||||
self).__init__(in_channels, out_channels, kernel_size, stride,
|
||||
padding, dilation, groups, bias)
|
||||
|
||||
def forward(self, input, mask=None):
|
||||
if mask is None: # fallback to the normal Conv2d
|
||||
return super(MaskedConv2d, self).forward(input)
|
||||
else:
|
||||
return masked_conv2d(input, mask, self.weight, self.bias,
|
||||
self.padding)
|
||||
@@ -0,0 +1,149 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
from abc import abstractmethod
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from ..cnn import ConvModule
|
||||
|
||||
|
||||
class BaseMergeCell(nn.Module):
|
||||
"""The basic class for cells used in NAS-FPN and NAS-FCOS.
|
||||
|
||||
BaseMergeCell takes 2 inputs. After applying convolution
|
||||
on them, they are resized to the target size. Then,
|
||||
they go through binary_op, which depends on the type of cell.
|
||||
If with_out_conv is True, the result of output will go through
|
||||
another convolution layer.
|
||||
|
||||
Args:
|
||||
in_channels (int): number of input channels in out_conv layer.
|
||||
out_channels (int): number of output channels in out_conv layer.
|
||||
with_out_conv (bool): Whether to use out_conv layer
|
||||
out_conv_cfg (dict): Config dict for convolution layer, which should
|
||||
contain "groups", "kernel_size", "padding", "bias" to build
|
||||
out_conv layer.
|
||||
out_norm_cfg (dict): Config dict for normalization layer in out_conv.
|
||||
out_conv_order (tuple): The order of conv/norm/activation layers in
|
||||
out_conv.
|
||||
with_input1_conv (bool): Whether to use convolution on input1.
|
||||
with_input2_conv (bool): Whether to use convolution on input2.
|
||||
input_conv_cfg (dict): Config dict for building input1_conv layer and
|
||||
input2_conv layer, which is expected to contain the type of
|
||||
convolution.
|
||||
Default: None, which means using conv2d.
|
||||
input_norm_cfg (dict): Config dict for normalization layer in
|
||||
input1_conv and input2_conv layer. Default: None.
|
||||
upsample_mode (str): Interpolation method used to resize the output
|
||||
of input1_conv and input2_conv to target size. Currently, we
|
||||
support ['nearest', 'bilinear']. Default: 'nearest'.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
fused_channels=256,
|
||||
out_channels=256,
|
||||
with_out_conv=True,
|
||||
out_conv_cfg=dict(
|
||||
groups=1, kernel_size=3, padding=1, bias=True),
|
||||
out_norm_cfg=None,
|
||||
out_conv_order=('act', 'conv', 'norm'),
|
||||
with_input1_conv=False,
|
||||
with_input2_conv=False,
|
||||
input_conv_cfg=None,
|
||||
input_norm_cfg=None,
|
||||
upsample_mode='nearest'):
|
||||
super(BaseMergeCell, self).__init__()
|
||||
assert upsample_mode in ['nearest', 'bilinear']
|
||||
self.with_out_conv = with_out_conv
|
||||
self.with_input1_conv = with_input1_conv
|
||||
self.with_input2_conv = with_input2_conv
|
||||
self.upsample_mode = upsample_mode
|
||||
|
||||
if self.with_out_conv:
|
||||
self.out_conv = ConvModule(
|
||||
fused_channels,
|
||||
out_channels,
|
||||
**out_conv_cfg,
|
||||
norm_cfg=out_norm_cfg,
|
||||
order=out_conv_order)
|
||||
|
||||
self.input1_conv = self._build_input_conv(
|
||||
out_channels, input_conv_cfg,
|
||||
input_norm_cfg) if with_input1_conv else nn.Sequential()
|
||||
self.input2_conv = self._build_input_conv(
|
||||
out_channels, input_conv_cfg,
|
||||
input_norm_cfg) if with_input2_conv else nn.Sequential()
|
||||
|
||||
def _build_input_conv(self, channel, conv_cfg, norm_cfg):
|
||||
return ConvModule(
|
||||
channel,
|
||||
channel,
|
||||
3,
|
||||
padding=1,
|
||||
conv_cfg=conv_cfg,
|
||||
norm_cfg=norm_cfg,
|
||||
bias=True)
|
||||
|
||||
@abstractmethod
|
||||
def _binary_op(self, x1, x2):
|
||||
pass
|
||||
|
||||
def _resize(self, x, size):
|
||||
if x.shape[-2:] == size:
|
||||
return x
|
||||
elif x.shape[-2:] < size:
|
||||
return F.interpolate(x, size=size, mode=self.upsample_mode)
|
||||
else:
|
||||
assert x.shape[-2] % size[-2] == 0 and x.shape[-1] % size[-1] == 0
|
||||
kernel_size = x.shape[-1] // size[-1]
|
||||
x = F.max_pool2d(x, kernel_size=kernel_size, stride=kernel_size)
|
||||
return x
|
||||
|
||||
def forward(self, x1, x2, out_size=None):
|
||||
assert x1.shape[:2] == x2.shape[:2]
|
||||
assert out_size is None or len(out_size) == 2
|
||||
if out_size is None: # resize to larger one
|
||||
out_size = max(x1.size()[2:], x2.size()[2:])
|
||||
|
||||
x1 = self.input1_conv(x1)
|
||||
x2 = self.input2_conv(x2)
|
||||
|
||||
x1 = self._resize(x1, out_size)
|
||||
x2 = self._resize(x2, out_size)
|
||||
|
||||
x = self._binary_op(x1, x2)
|
||||
if self.with_out_conv:
|
||||
x = self.out_conv(x)
|
||||
return x
|
||||
|
||||
|
||||
class SumCell(BaseMergeCell):
|
||||
|
||||
def __init__(self, in_channels, out_channels, **kwargs):
|
||||
super(SumCell, self).__init__(in_channels, out_channels, **kwargs)
|
||||
|
||||
def _binary_op(self, x1, x2):
|
||||
return x1 + x2
|
||||
|
||||
|
||||
class ConcatCell(BaseMergeCell):
|
||||
|
||||
def __init__(self, in_channels, out_channels, **kwargs):
|
||||
super(ConcatCell, self).__init__(in_channels * 2, out_channels,
|
||||
**kwargs)
|
||||
|
||||
def _binary_op(self, x1, x2):
|
||||
ret = torch.cat([x1, x2], dim=1)
|
||||
return ret
|
||||
|
||||
|
||||
class GlobalPoolingCell(BaseMergeCell):
|
||||
|
||||
def __init__(self, in_channels=None, out_channels=None, **kwargs):
|
||||
super().__init__(in_channels, out_channels, **kwargs)
|
||||
self.global_pool = nn.AdaptiveAvgPool2d((1, 1))
|
||||
|
||||
def _binary_op(self, x1, x2):
|
||||
x2_att = self.global_pool(x2).sigmoid()
|
||||
return x2 + x2_att * x1
|
||||
@@ -0,0 +1,282 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.autograd import Function
|
||||
from torch.autograd.function import once_differentiable
|
||||
from torch.nn.modules.utils import _pair, _single
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import deprecated_api_warning
|
||||
from ..cnn import CONV_LAYERS
|
||||
from ..utils import ext_loader, print_log
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext',
|
||||
['modulated_deform_conv_forward', 'modulated_deform_conv_backward'])
|
||||
|
||||
|
||||
class ModulatedDeformConv2dFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input, offset, mask, weight, bias, stride, padding,
|
||||
dilation, groups, deform_groups):
|
||||
input_tensors = [input, offset, mask, weight]
|
||||
if bias is not None:
|
||||
input_tensors.append(bias)
|
||||
return g.op(
|
||||
'mmcv::MMCVModulatedDeformConv2d',
|
||||
*input_tensors,
|
||||
stride_i=stride,
|
||||
padding_i=padding,
|
||||
dilation_i=dilation,
|
||||
groups_i=groups,
|
||||
deform_groups_i=deform_groups)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx,
|
||||
input,
|
||||
offset,
|
||||
mask,
|
||||
weight,
|
||||
bias=None,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
deform_groups=1):
|
||||
if input is not None and input.dim() != 4:
|
||||
raise ValueError(
|
||||
f'Expected 4D tensor as input, got {input.dim()}D tensor \
|
||||
instead.')
|
||||
ctx.stride = _pair(stride)
|
||||
ctx.padding = _pair(padding)
|
||||
ctx.dilation = _pair(dilation)
|
||||
ctx.groups = groups
|
||||
ctx.deform_groups = deform_groups
|
||||
ctx.with_bias = bias is not None
|
||||
if not ctx.with_bias:
|
||||
bias = input.new_empty(0) # fake tensor
|
||||
# When pytorch version >= 1.6.0, amp is adopted for fp16 mode;
|
||||
# amp won't cast the type of model (float32), but "offset" is cast
|
||||
# to float16 by nn.Conv2d automatically, leading to the type
|
||||
# mismatch with input (when it is float32) or weight.
|
||||
# The flag for whether to use fp16 or amp is the type of "offset",
|
||||
# we cast weight and input to temporarily support fp16 and amp
|
||||
# whatever the pytorch version is.
|
||||
input = input.type_as(offset)
|
||||
weight = weight.type_as(input)
|
||||
ctx.save_for_backward(input, offset, mask, weight, bias)
|
||||
output = input.new_empty(
|
||||
ModulatedDeformConv2dFunction._output_size(ctx, input, weight))
|
||||
ctx._bufs = [input.new_empty(0), input.new_empty(0)]
|
||||
ext_module.modulated_deform_conv_forward(
|
||||
input,
|
||||
weight,
|
||||
bias,
|
||||
ctx._bufs[0],
|
||||
offset,
|
||||
mask,
|
||||
output,
|
||||
ctx._bufs[1],
|
||||
kernel_h=weight.size(2),
|
||||
kernel_w=weight.size(3),
|
||||
stride_h=ctx.stride[0],
|
||||
stride_w=ctx.stride[1],
|
||||
pad_h=ctx.padding[0],
|
||||
pad_w=ctx.padding[1],
|
||||
dilation_h=ctx.dilation[0],
|
||||
dilation_w=ctx.dilation[1],
|
||||
group=ctx.groups,
|
||||
deformable_group=ctx.deform_groups,
|
||||
with_bias=ctx.with_bias)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
input, offset, mask, weight, bias = ctx.saved_tensors
|
||||
grad_input = torch.zeros_like(input)
|
||||
grad_offset = torch.zeros_like(offset)
|
||||
grad_mask = torch.zeros_like(mask)
|
||||
grad_weight = torch.zeros_like(weight)
|
||||
grad_bias = torch.zeros_like(bias)
|
||||
grad_output = grad_output.contiguous()
|
||||
ext_module.modulated_deform_conv_backward(
|
||||
input,
|
||||
weight,
|
||||
bias,
|
||||
ctx._bufs[0],
|
||||
offset,
|
||||
mask,
|
||||
ctx._bufs[1],
|
||||
grad_input,
|
||||
grad_weight,
|
||||
grad_bias,
|
||||
grad_offset,
|
||||
grad_mask,
|
||||
grad_output,
|
||||
kernel_h=weight.size(2),
|
||||
kernel_w=weight.size(3),
|
||||
stride_h=ctx.stride[0],
|
||||
stride_w=ctx.stride[1],
|
||||
pad_h=ctx.padding[0],
|
||||
pad_w=ctx.padding[1],
|
||||
dilation_h=ctx.dilation[0],
|
||||
dilation_w=ctx.dilation[1],
|
||||
group=ctx.groups,
|
||||
deformable_group=ctx.deform_groups,
|
||||
with_bias=ctx.with_bias)
|
||||
if not ctx.with_bias:
|
||||
grad_bias = None
|
||||
|
||||
return (grad_input, grad_offset, grad_mask, grad_weight, grad_bias,
|
||||
None, None, None, None, None)
|
||||
|
||||
@staticmethod
|
||||
def _output_size(ctx, input, weight):
|
||||
channels = weight.size(0)
|
||||
output_size = (input.size(0), channels)
|
||||
for d in range(input.dim() - 2):
|
||||
in_size = input.size(d + 2)
|
||||
pad = ctx.padding[d]
|
||||
kernel = ctx.dilation[d] * (weight.size(d + 2) - 1) + 1
|
||||
stride_ = ctx.stride[d]
|
||||
output_size += ((in_size + (2 * pad) - kernel) // stride_ + 1, )
|
||||
if not all(map(lambda s: s > 0, output_size)):
|
||||
raise ValueError(
|
||||
'convolution input is too small (output would be ' +
|
||||
'x'.join(map(str, output_size)) + ')')
|
||||
return output_size
|
||||
|
||||
|
||||
modulated_deform_conv2d = ModulatedDeformConv2dFunction.apply
|
||||
|
||||
|
||||
class ModulatedDeformConv2d(nn.Module):
|
||||
|
||||
@deprecated_api_warning({'deformable_groups': 'deform_groups'},
|
||||
cls_name='ModulatedDeformConv2d')
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
out_channels,
|
||||
kernel_size,
|
||||
stride=1,
|
||||
padding=0,
|
||||
dilation=1,
|
||||
groups=1,
|
||||
deform_groups=1,
|
||||
bias=True):
|
||||
super(ModulatedDeformConv2d, self).__init__()
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels
|
||||
self.kernel_size = _pair(kernel_size)
|
||||
self.stride = _pair(stride)
|
||||
self.padding = _pair(padding)
|
||||
self.dilation = _pair(dilation)
|
||||
self.groups = groups
|
||||
self.deform_groups = deform_groups
|
||||
# enable compatibility with nn.Conv2d
|
||||
self.transposed = False
|
||||
self.output_padding = _single(0)
|
||||
|
||||
self.weight = nn.Parameter(
|
||||
torch.Tensor(out_channels, in_channels // groups,
|
||||
*self.kernel_size))
|
||||
if bias:
|
||||
self.bias = nn.Parameter(torch.Tensor(out_channels))
|
||||
else:
|
||||
self.register_parameter('bias', None)
|
||||
self.init_weights()
|
||||
|
||||
def init_weights(self):
|
||||
n = self.in_channels
|
||||
for k in self.kernel_size:
|
||||
n *= k
|
||||
stdv = 1. / math.sqrt(n)
|
||||
self.weight.data.uniform_(-stdv, stdv)
|
||||
if self.bias is not None:
|
||||
self.bias.data.zero_()
|
||||
|
||||
def forward(self, x, offset, mask):
|
||||
return modulated_deform_conv2d(x, offset, mask, self.weight, self.bias,
|
||||
self.stride, self.padding,
|
||||
self.dilation, self.groups,
|
||||
self.deform_groups)
|
||||
|
||||
|
||||
@CONV_LAYERS.register_module('DCNv2')
|
||||
class ModulatedDeformConv2dPack(ModulatedDeformConv2d):
|
||||
"""A ModulatedDeformable Conv Encapsulation that acts as normal Conv
|
||||
layers.
|
||||
|
||||
Args:
|
||||
in_channels (int): Same as nn.Conv2d.
|
||||
out_channels (int): Same as nn.Conv2d.
|
||||
kernel_size (int or tuple[int]): Same as nn.Conv2d.
|
||||
stride (int): Same as nn.Conv2d, while tuple is not supported.
|
||||
padding (int): Same as nn.Conv2d, while tuple is not supported.
|
||||
dilation (int): Same as nn.Conv2d, while tuple is not supported.
|
||||
groups (int): Same as nn.Conv2d.
|
||||
bias (bool or str): If specified as `auto`, it will be decided by the
|
||||
norm_cfg. Bias will be set as True if norm_cfg is None, otherwise
|
||||
False.
|
||||
"""
|
||||
|
||||
_version = 2
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super(ModulatedDeformConv2dPack, self).__init__(*args, **kwargs)
|
||||
self.conv_offset = nn.Conv2d(
|
||||
self.in_channels,
|
||||
self.deform_groups * 3 * self.kernel_size[0] * self.kernel_size[1],
|
||||
kernel_size=self.kernel_size,
|
||||
stride=self.stride,
|
||||
padding=self.padding,
|
||||
dilation=self.dilation,
|
||||
bias=True)
|
||||
self.init_weights()
|
||||
|
||||
def init_weights(self):
|
||||
super(ModulatedDeformConv2dPack, self).init_weights()
|
||||
if hasattr(self, 'conv_offset'):
|
||||
self.conv_offset.weight.data.zero_()
|
||||
self.conv_offset.bias.data.zero_()
|
||||
|
||||
def forward(self, x):
|
||||
out = self.conv_offset(x)
|
||||
o1, o2, mask = torch.chunk(out, 3, dim=1)
|
||||
offset = torch.cat((o1, o2), dim=1)
|
||||
mask = torch.sigmoid(mask)
|
||||
return modulated_deform_conv2d(x, offset, mask, self.weight, self.bias,
|
||||
self.stride, self.padding,
|
||||
self.dilation, self.groups,
|
||||
self.deform_groups)
|
||||
|
||||
def _load_from_state_dict(self, state_dict, prefix, local_metadata, strict,
|
||||
missing_keys, unexpected_keys, error_msgs):
|
||||
version = local_metadata.get('version', None)
|
||||
|
||||
if version is None or version < 2:
|
||||
# the key is different in early versions
|
||||
# In version < 2, ModulatedDeformConvPack
|
||||
# loads previous benchmark models.
|
||||
if (prefix + 'conv_offset.weight' not in state_dict
|
||||
and prefix[:-1] + '_offset.weight' in state_dict):
|
||||
state_dict[prefix + 'conv_offset.weight'] = state_dict.pop(
|
||||
prefix[:-1] + '_offset.weight')
|
||||
if (prefix + 'conv_offset.bias' not in state_dict
|
||||
and prefix[:-1] + '_offset.bias' in state_dict):
|
||||
state_dict[prefix +
|
||||
'conv_offset.bias'] = state_dict.pop(prefix[:-1] +
|
||||
'_offset.bias')
|
||||
|
||||
if version is not None and version > 1:
|
||||
print_log(
|
||||
f'ModulatedDeformConvPack {prefix.rstrip(".")} is upgraded to '
|
||||
'version 2.',
|
||||
logger='root')
|
||||
|
||||
super()._load_from_state_dict(state_dict, prefix, local_metadata,
|
||||
strict, missing_keys, unexpected_keys,
|
||||
error_msgs)
|
||||
@@ -0,0 +1,358 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import math
|
||||
import warnings
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.autograd.function import Function, once_differentiable
|
||||
|
||||
from custom_mmpkg.custom_mmcv import deprecated_api_warning
|
||||
from custom_mmpkg.custom_mmcv.cnn import constant_init, xavier_init
|
||||
from custom_mmpkg.custom_mmcv.cnn.bricks.registry import ATTENTION
|
||||
from custom_mmpkg.custom_mmcv.runner import BaseModule
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['ms_deform_attn_backward', 'ms_deform_attn_forward'])
|
||||
|
||||
|
||||
class MultiScaleDeformableAttnFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, value, value_spatial_shapes, value_level_start_index,
|
||||
sampling_locations, attention_weights, im2col_step):
|
||||
"""GPU version of multi-scale deformable attention.
|
||||
|
||||
Args:
|
||||
value (Tensor): The value has shape
|
||||
(bs, num_keys, mum_heads, embed_dims//num_heads)
|
||||
value_spatial_shapes (Tensor): Spatial shape of
|
||||
each feature map, has shape (num_levels, 2),
|
||||
last dimension 2 represent (h, w)
|
||||
sampling_locations (Tensor): The location of sampling points,
|
||||
has shape
|
||||
(bs ,num_queries, num_heads, num_levels, num_points, 2),
|
||||
the last dimension 2 represent (x, y).
|
||||
attention_weights (Tensor): The weight of sampling points used
|
||||
when calculate the attention, has shape
|
||||
(bs ,num_queries, num_heads, num_levels, num_points),
|
||||
im2col_step (Tensor): The step used in image to column.
|
||||
|
||||
Returns:
|
||||
Tensor: has shape (bs, num_queries, embed_dims)
|
||||
"""
|
||||
|
||||
ctx.im2col_step = im2col_step
|
||||
output = ext_module.ms_deform_attn_forward(
|
||||
value,
|
||||
value_spatial_shapes,
|
||||
value_level_start_index,
|
||||
sampling_locations,
|
||||
attention_weights,
|
||||
im2col_step=ctx.im2col_step)
|
||||
ctx.save_for_backward(value, value_spatial_shapes,
|
||||
value_level_start_index, sampling_locations,
|
||||
attention_weights)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
@once_differentiable
|
||||
def backward(ctx, grad_output):
|
||||
"""GPU version of backward function.
|
||||
|
||||
Args:
|
||||
grad_output (Tensor): Gradient
|
||||
of output tensor of forward.
|
||||
|
||||
Returns:
|
||||
Tuple[Tensor]: Gradient
|
||||
of input tensors in forward.
|
||||
"""
|
||||
value, value_spatial_shapes, value_level_start_index,\
|
||||
sampling_locations, attention_weights = ctx.saved_tensors
|
||||
grad_value = torch.zeros_like(value)
|
||||
grad_sampling_loc = torch.zeros_like(sampling_locations)
|
||||
grad_attn_weight = torch.zeros_like(attention_weights)
|
||||
|
||||
ext_module.ms_deform_attn_backward(
|
||||
value,
|
||||
value_spatial_shapes,
|
||||
value_level_start_index,
|
||||
sampling_locations,
|
||||
attention_weights,
|
||||
grad_output.contiguous(),
|
||||
grad_value,
|
||||
grad_sampling_loc,
|
||||
grad_attn_weight,
|
||||
im2col_step=ctx.im2col_step)
|
||||
|
||||
return grad_value, None, None, \
|
||||
grad_sampling_loc, grad_attn_weight, None
|
||||
|
||||
|
||||
def multi_scale_deformable_attn_pytorch(value, value_spatial_shapes,
|
||||
sampling_locations, attention_weights):
|
||||
"""CPU version of multi-scale deformable attention.
|
||||
|
||||
Args:
|
||||
value (Tensor): The value has shape
|
||||
(bs, num_keys, mum_heads, embed_dims//num_heads)
|
||||
value_spatial_shapes (Tensor): Spatial shape of
|
||||
each feature map, has shape (num_levels, 2),
|
||||
last dimension 2 represent (h, w)
|
||||
sampling_locations (Tensor): The location of sampling points,
|
||||
has shape
|
||||
(bs ,num_queries, num_heads, num_levels, num_points, 2),
|
||||
the last dimension 2 represent (x, y).
|
||||
attention_weights (Tensor): The weight of sampling points used
|
||||
when calculate the attention, has shape
|
||||
(bs ,num_queries, num_heads, num_levels, num_points),
|
||||
|
||||
Returns:
|
||||
Tensor: has shape (bs, num_queries, embed_dims)
|
||||
"""
|
||||
|
||||
bs, _, num_heads, embed_dims = value.shape
|
||||
_, num_queries, num_heads, num_levels, num_points, _ =\
|
||||
sampling_locations.shape
|
||||
value_list = value.split([H_ * W_ for H_, W_ in value_spatial_shapes],
|
||||
dim=1)
|
||||
sampling_grids = 2 * sampling_locations - 1
|
||||
sampling_value_list = []
|
||||
for level, (H_, W_) in enumerate(value_spatial_shapes):
|
||||
# bs, H_*W_, num_heads, embed_dims ->
|
||||
# bs, H_*W_, num_heads*embed_dims ->
|
||||
# bs, num_heads*embed_dims, H_*W_ ->
|
||||
# bs*num_heads, embed_dims, H_, W_
|
||||
value_l_ = value_list[level].flatten(2).transpose(1, 2).reshape(
|
||||
bs * num_heads, embed_dims, H_, W_)
|
||||
# bs, num_queries, num_heads, num_points, 2 ->
|
||||
# bs, num_heads, num_queries, num_points, 2 ->
|
||||
# bs*num_heads, num_queries, num_points, 2
|
||||
sampling_grid_l_ = sampling_grids[:, :, :,
|
||||
level].transpose(1, 2).flatten(0, 1)
|
||||
# bs*num_heads, embed_dims, num_queries, num_points
|
||||
sampling_value_l_ = F.grid_sample(
|
||||
value_l_,
|
||||
sampling_grid_l_,
|
||||
mode='bilinear',
|
||||
padding_mode='zeros',
|
||||
align_corners=False)
|
||||
sampling_value_list.append(sampling_value_l_)
|
||||
# (bs, num_queries, num_heads, num_levels, num_points) ->
|
||||
# (bs, num_heads, num_queries, num_levels, num_points) ->
|
||||
# (bs, num_heads, 1, num_queries, num_levels*num_points)
|
||||
attention_weights = attention_weights.transpose(1, 2).reshape(
|
||||
bs * num_heads, 1, num_queries, num_levels * num_points)
|
||||
output = (torch.stack(sampling_value_list, dim=-2).flatten(-2) *
|
||||
attention_weights).sum(-1).view(bs, num_heads * embed_dims,
|
||||
num_queries)
|
||||
return output.transpose(1, 2).contiguous()
|
||||
|
||||
|
||||
@ATTENTION.register_module()
|
||||
class MultiScaleDeformableAttention(BaseModule):
|
||||
"""An attention module used in Deformable-Detr.
|
||||
|
||||
`Deformable DETR: Deformable Transformers for End-to-End Object Detection.
|
||||
<https://arxiv.org/pdf/2010.04159.pdf>`_.
|
||||
|
||||
Args:
|
||||
embed_dims (int): The embedding dimension of Attention.
|
||||
Default: 256.
|
||||
num_heads (int): Parallel attention heads. Default: 64.
|
||||
num_levels (int): The number of feature map used in
|
||||
Attention. Default: 4.
|
||||
num_points (int): The number of sampling points for
|
||||
each query in each head. Default: 4.
|
||||
im2col_step (int): The step used in image_to_column.
|
||||
Default: 64.
|
||||
dropout (float): A Dropout layer on `inp_identity`.
|
||||
Default: 0.1.
|
||||
batch_first (bool): Key, Query and Value are shape of
|
||||
(batch, n, embed_dim)
|
||||
or (n, batch, embed_dim). Default to False.
|
||||
norm_cfg (dict): Config dict for normalization layer.
|
||||
Default: None.
|
||||
init_cfg (obj:`mmcv.ConfigDict`): The Config for initialization.
|
||||
Default: None.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
embed_dims=256,
|
||||
num_heads=8,
|
||||
num_levels=4,
|
||||
num_points=4,
|
||||
im2col_step=64,
|
||||
dropout=0.1,
|
||||
batch_first=False,
|
||||
norm_cfg=None,
|
||||
init_cfg=None):
|
||||
super().__init__(init_cfg)
|
||||
if embed_dims % num_heads != 0:
|
||||
raise ValueError(f'embed_dims must be divisible by num_heads, '
|
||||
f'but got {embed_dims} and {num_heads}')
|
||||
dim_per_head = embed_dims // num_heads
|
||||
self.norm_cfg = norm_cfg
|
||||
self.dropout = nn.Dropout(dropout)
|
||||
self.batch_first = batch_first
|
||||
|
||||
# you'd better set dim_per_head to a power of 2
|
||||
# which is more efficient in the CUDA implementation
|
||||
def _is_power_of_2(n):
|
||||
if (not isinstance(n, int)) or (n < 0):
|
||||
raise ValueError(
|
||||
'invalid input for _is_power_of_2: {} (type: {})'.format(
|
||||
n, type(n)))
|
||||
return (n & (n - 1) == 0) and n != 0
|
||||
|
||||
if not _is_power_of_2(dim_per_head):
|
||||
warnings.warn(
|
||||
"You'd better set embed_dims in "
|
||||
'MultiScaleDeformAttention to make '
|
||||
'the dimension of each attention head a power of 2 '
|
||||
'which is more efficient in our CUDA implementation.')
|
||||
|
||||
self.im2col_step = im2col_step
|
||||
self.embed_dims = embed_dims
|
||||
self.num_levels = num_levels
|
||||
self.num_heads = num_heads
|
||||
self.num_points = num_points
|
||||
self.sampling_offsets = nn.Linear(
|
||||
embed_dims, num_heads * num_levels * num_points * 2)
|
||||
self.attention_weights = nn.Linear(embed_dims,
|
||||
num_heads * num_levels * num_points)
|
||||
self.value_proj = nn.Linear(embed_dims, embed_dims)
|
||||
self.output_proj = nn.Linear(embed_dims, embed_dims)
|
||||
self.init_weights()
|
||||
|
||||
def init_weights(self):
|
||||
"""Default initialization for Parameters of Module."""
|
||||
constant_init(self.sampling_offsets, 0.)
|
||||
thetas = torch.arange(
|
||||
self.num_heads,
|
||||
dtype=torch.float32) * (2.0 * math.pi / self.num_heads)
|
||||
grid_init = torch.stack([thetas.cos(), thetas.sin()], -1)
|
||||
grid_init = (grid_init /
|
||||
grid_init.abs().max(-1, keepdim=True)[0]).view(
|
||||
self.num_heads, 1, 1,
|
||||
2).repeat(1, self.num_levels, self.num_points, 1)
|
||||
for i in range(self.num_points):
|
||||
grid_init[:, :, i, :] *= i + 1
|
||||
|
||||
self.sampling_offsets.bias.data = grid_init.view(-1)
|
||||
constant_init(self.attention_weights, val=0., bias=0.)
|
||||
xavier_init(self.value_proj, distribution='uniform', bias=0.)
|
||||
xavier_init(self.output_proj, distribution='uniform', bias=0.)
|
||||
self._is_init = True
|
||||
|
||||
@deprecated_api_warning({'residual': 'identity'},
|
||||
cls_name='MultiScaleDeformableAttention')
|
||||
def forward(self,
|
||||
query,
|
||||
key=None,
|
||||
value=None,
|
||||
identity=None,
|
||||
query_pos=None,
|
||||
key_padding_mask=None,
|
||||
reference_points=None,
|
||||
spatial_shapes=None,
|
||||
level_start_index=None,
|
||||
**kwargs):
|
||||
"""Forward Function of MultiScaleDeformAttention.
|
||||
|
||||
Args:
|
||||
query (Tensor): Query of Transformer with shape
|
||||
(num_query, bs, embed_dims).
|
||||
key (Tensor): The key tensor with shape
|
||||
`(num_key, bs, embed_dims)`.
|
||||
value (Tensor): The value tensor with shape
|
||||
`(num_key, bs, embed_dims)`.
|
||||
identity (Tensor): The tensor used for addition, with the
|
||||
same shape as `query`. Default None. If None,
|
||||
`query` will be used.
|
||||
query_pos (Tensor): The positional encoding for `query`.
|
||||
Default: None.
|
||||
key_pos (Tensor): The positional encoding for `key`. Default
|
||||
None.
|
||||
reference_points (Tensor): The normalized reference
|
||||
points with shape (bs, num_query, num_levels, 2),
|
||||
all elements is range in [0, 1], top-left (0,0),
|
||||
bottom-right (1, 1), including padding area.
|
||||
or (N, Length_{query}, num_levels, 4), add
|
||||
additional two dimensions is (w, h) to
|
||||
form reference boxes.
|
||||
key_padding_mask (Tensor): ByteTensor for `query`, with
|
||||
shape [bs, num_key].
|
||||
spatial_shapes (Tensor): Spatial shape of features in
|
||||
different levels. With shape (num_levels, 2),
|
||||
last dimension represents (h, w).
|
||||
level_start_index (Tensor): The start index of each level.
|
||||
A tensor has shape ``(num_levels, )`` and can be represented
|
||||
as [0, h_0*w_0, h_0*w_0+h_1*w_1, ...].
|
||||
|
||||
Returns:
|
||||
Tensor: forwarded results with shape [num_query, bs, embed_dims].
|
||||
"""
|
||||
|
||||
if value is None:
|
||||
value = query
|
||||
|
||||
if identity is None:
|
||||
identity = query
|
||||
if query_pos is not None:
|
||||
query = query + query_pos
|
||||
if not self.batch_first:
|
||||
# change to (bs, num_query ,embed_dims)
|
||||
query = query.permute(1, 0, 2)
|
||||
value = value.permute(1, 0, 2)
|
||||
|
||||
bs, num_query, _ = query.shape
|
||||
bs, num_value, _ = value.shape
|
||||
assert (spatial_shapes[:, 0] * spatial_shapes[:, 1]).sum() == num_value
|
||||
|
||||
value = self.value_proj(value)
|
||||
if key_padding_mask is not None:
|
||||
value = value.masked_fill(key_padding_mask[..., None], 0.0)
|
||||
value = value.view(bs, num_value, self.num_heads, -1)
|
||||
sampling_offsets = self.sampling_offsets(query).view(
|
||||
bs, num_query, self.num_heads, self.num_levels, self.num_points, 2)
|
||||
attention_weights = self.attention_weights(query).view(
|
||||
bs, num_query, self.num_heads, self.num_levels * self.num_points)
|
||||
attention_weights = attention_weights.softmax(-1)
|
||||
|
||||
attention_weights = attention_weights.view(bs, num_query,
|
||||
self.num_heads,
|
||||
self.num_levels,
|
||||
self.num_points)
|
||||
if reference_points.shape[-1] == 2:
|
||||
offset_normalizer = torch.stack(
|
||||
[spatial_shapes[..., 1], spatial_shapes[..., 0]], -1)
|
||||
sampling_locations = reference_points[:, :, None, :, None, :] \
|
||||
+ sampling_offsets \
|
||||
/ offset_normalizer[None, None, None, :, None, :]
|
||||
elif reference_points.shape[-1] == 4:
|
||||
sampling_locations = reference_points[:, :, None, :, None, :2] \
|
||||
+ sampling_offsets / self.num_points \
|
||||
* reference_points[:, :, None, :, None, 2:] \
|
||||
* 0.5
|
||||
else:
|
||||
raise ValueError(
|
||||
f'Last dim of reference_points must be'
|
||||
f' 2 or 4, but get {reference_points.shape[-1]} instead.')
|
||||
if torch.cuda.is_available() and value.is_cuda:
|
||||
output = MultiScaleDeformableAttnFunction.apply(
|
||||
value, spatial_shapes, level_start_index, sampling_locations,
|
||||
attention_weights, self.im2col_step)
|
||||
else:
|
||||
output = multi_scale_deformable_attn_pytorch(
|
||||
value, spatial_shapes, sampling_locations, attention_weights)
|
||||
|
||||
output = self.output_proj(output)
|
||||
|
||||
if not self.batch_first:
|
||||
# (num_query, bs ,embed_dims)
|
||||
output = output.permute(1, 0, 2)
|
||||
|
||||
return self.dropout(output) + identity
|
||||
@@ -0,0 +1,417 @@
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from custom_mmpkg.custom_mmcv.utils import deprecated_api_warning
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext(
|
||||
'_ext', ['nms', 'softnms', 'nms_match', 'nms_rotated'])
|
||||
|
||||
|
||||
# This function is modified from: https://github.com/pytorch/vision/
|
||||
class NMSop(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, bboxes, scores, iou_threshold, offset, score_threshold,
|
||||
max_num):
|
||||
is_filtering_by_score = score_threshold > 0
|
||||
if is_filtering_by_score:
|
||||
valid_mask = scores > score_threshold
|
||||
bboxes, scores = bboxes[valid_mask], scores[valid_mask]
|
||||
valid_inds = torch.nonzero(
|
||||
valid_mask, as_tuple=False).squeeze(dim=1)
|
||||
|
||||
inds = ext_module.nms(
|
||||
bboxes, scores, iou_threshold=float(iou_threshold), offset=offset)
|
||||
|
||||
if max_num > 0:
|
||||
inds = inds[:max_num]
|
||||
if is_filtering_by_score:
|
||||
inds = valid_inds[inds]
|
||||
return inds
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, bboxes, scores, iou_threshold, offset, score_threshold,
|
||||
max_num):
|
||||
from ..onnx import is_custom_op_loaded
|
||||
has_custom_op = is_custom_op_loaded()
|
||||
# TensorRT nms plugin is aligned with original nms in ONNXRuntime
|
||||
is_trt_backend = os.environ.get('ONNX_BACKEND') == 'MMCVTensorRT'
|
||||
if has_custom_op and (not is_trt_backend):
|
||||
return g.op(
|
||||
'mmcv::NonMaxSuppression',
|
||||
bboxes,
|
||||
scores,
|
||||
iou_threshold_f=float(iou_threshold),
|
||||
offset_i=int(offset))
|
||||
else:
|
||||
from torch.onnx.symbolic_opset9 import select, squeeze, unsqueeze
|
||||
from ..onnx.onnx_utils.symbolic_helper import _size_helper
|
||||
|
||||
boxes = unsqueeze(g, bboxes, 0)
|
||||
scores = unsqueeze(g, unsqueeze(g, scores, 0), 0)
|
||||
|
||||
if max_num > 0:
|
||||
max_num = g.op(
|
||||
'Constant',
|
||||
value_t=torch.tensor(max_num, dtype=torch.long))
|
||||
else:
|
||||
dim = g.op('Constant', value_t=torch.tensor(0))
|
||||
max_num = _size_helper(g, bboxes, dim)
|
||||
max_output_per_class = max_num
|
||||
iou_threshold = g.op(
|
||||
'Constant',
|
||||
value_t=torch.tensor([iou_threshold], dtype=torch.float))
|
||||
score_threshold = g.op(
|
||||
'Constant',
|
||||
value_t=torch.tensor([score_threshold], dtype=torch.float))
|
||||
nms_out = g.op('NonMaxSuppression', boxes, scores,
|
||||
max_output_per_class, iou_threshold,
|
||||
score_threshold)
|
||||
return squeeze(
|
||||
g,
|
||||
select(
|
||||
g, nms_out, 1,
|
||||
g.op(
|
||||
'Constant',
|
||||
value_t=torch.tensor([2], dtype=torch.long))), 1)
|
||||
|
||||
|
||||
class SoftNMSop(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, boxes, scores, iou_threshold, sigma, min_score, method,
|
||||
offset):
|
||||
dets = boxes.new_empty((boxes.size(0), 5), device='cpu')
|
||||
inds = ext_module.softnms(
|
||||
boxes.cpu(),
|
||||
scores.cpu(),
|
||||
dets.cpu(),
|
||||
iou_threshold=float(iou_threshold),
|
||||
sigma=float(sigma),
|
||||
min_score=float(min_score),
|
||||
method=int(method),
|
||||
offset=int(offset))
|
||||
return dets, inds
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, boxes, scores, iou_threshold, sigma, min_score, method,
|
||||
offset):
|
||||
from packaging import version
|
||||
assert version.parse(torch.__version__) >= version.parse('1.7.0')
|
||||
nms_out = g.op(
|
||||
'mmcv::SoftNonMaxSuppression',
|
||||
boxes,
|
||||
scores,
|
||||
iou_threshold_f=float(iou_threshold),
|
||||
sigma_f=float(sigma),
|
||||
min_score_f=float(min_score),
|
||||
method_i=int(method),
|
||||
offset_i=int(offset),
|
||||
outputs=2)
|
||||
return nms_out
|
||||
|
||||
|
||||
@deprecated_api_warning({'iou_thr': 'iou_threshold'})
|
||||
def nms(boxes, scores, iou_threshold, offset=0, score_threshold=0, max_num=-1):
|
||||
"""Dispatch to either CPU or GPU NMS implementations.
|
||||
|
||||
The input can be either torch tensor or numpy array. GPU NMS will be used
|
||||
if the input is gpu tensor, otherwise CPU NMS
|
||||
will be used. The returned type will always be the same as inputs.
|
||||
|
||||
Arguments:
|
||||
boxes (torch.Tensor or np.ndarray): boxes in shape (N, 4).
|
||||
scores (torch.Tensor or np.ndarray): scores in shape (N, ).
|
||||
iou_threshold (float): IoU threshold for NMS.
|
||||
offset (int, 0 or 1): boxes' width or height is (x2 - x1 + offset).
|
||||
score_threshold (float): score threshold for NMS.
|
||||
max_num (int): maximum number of boxes after NMS.
|
||||
|
||||
Returns:
|
||||
tuple: kept dets(boxes and scores) and indice, which is always the \
|
||||
same data type as the input.
|
||||
|
||||
Example:
|
||||
>>> boxes = np.array([[49.1, 32.4, 51.0, 35.9],
|
||||
>>> [49.3, 32.9, 51.0, 35.3],
|
||||
>>> [49.2, 31.8, 51.0, 35.4],
|
||||
>>> [35.1, 11.5, 39.1, 15.7],
|
||||
>>> [35.6, 11.8, 39.3, 14.2],
|
||||
>>> [35.3, 11.5, 39.9, 14.5],
|
||||
>>> [35.2, 11.7, 39.7, 15.7]], dtype=np.float32)
|
||||
>>> scores = np.array([0.9, 0.9, 0.5, 0.5, 0.5, 0.4, 0.3],\
|
||||
dtype=np.float32)
|
||||
>>> iou_threshold = 0.6
|
||||
>>> dets, inds = nms(boxes, scores, iou_threshold)
|
||||
>>> assert len(inds) == len(dets) == 3
|
||||
"""
|
||||
assert isinstance(boxes, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(scores, (torch.Tensor, np.ndarray))
|
||||
is_numpy = False
|
||||
if isinstance(boxes, np.ndarray):
|
||||
is_numpy = True
|
||||
boxes = torch.from_numpy(boxes)
|
||||
if isinstance(scores, np.ndarray):
|
||||
scores = torch.from_numpy(scores)
|
||||
assert boxes.size(1) == 4
|
||||
assert boxes.size(0) == scores.size(0)
|
||||
assert offset in (0, 1)
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
indata_list = [boxes, scores]
|
||||
indata_dict = {
|
||||
'iou_threshold': float(iou_threshold),
|
||||
'offset': int(offset)
|
||||
}
|
||||
inds = ext_module.nms(*indata_list, **indata_dict)
|
||||
else:
|
||||
inds = NMSop.apply(boxes, scores, iou_threshold, offset,
|
||||
score_threshold, max_num)
|
||||
dets = torch.cat((boxes[inds], scores[inds].reshape(-1, 1)), dim=1)
|
||||
if is_numpy:
|
||||
dets = dets.cpu().numpy()
|
||||
inds = inds.cpu().numpy()
|
||||
return dets, inds
|
||||
|
||||
|
||||
@deprecated_api_warning({'iou_thr': 'iou_threshold'})
|
||||
def soft_nms(boxes,
|
||||
scores,
|
||||
iou_threshold=0.3,
|
||||
sigma=0.5,
|
||||
min_score=1e-3,
|
||||
method='linear',
|
||||
offset=0):
|
||||
"""Dispatch to only CPU Soft NMS implementations.
|
||||
|
||||
The input can be either a torch tensor or numpy array.
|
||||
The returned type will always be the same as inputs.
|
||||
|
||||
Arguments:
|
||||
boxes (torch.Tensor or np.ndarray): boxes in shape (N, 4).
|
||||
scores (torch.Tensor or np.ndarray): scores in shape (N, ).
|
||||
iou_threshold (float): IoU threshold for NMS.
|
||||
sigma (float): hyperparameter for gaussian method
|
||||
min_score (float): score filter threshold
|
||||
method (str): either 'linear' or 'gaussian'
|
||||
offset (int, 0 or 1): boxes' width or height is (x2 - x1 + offset).
|
||||
|
||||
Returns:
|
||||
tuple: kept dets(boxes and scores) and indice, which is always the \
|
||||
same data type as the input.
|
||||
|
||||
Example:
|
||||
>>> boxes = np.array([[4., 3., 5., 3.],
|
||||
>>> [4., 3., 5., 4.],
|
||||
>>> [3., 1., 3., 1.],
|
||||
>>> [3., 1., 3., 1.],
|
||||
>>> [3., 1., 3., 1.],
|
||||
>>> [3., 1., 3., 1.]], dtype=np.float32)
|
||||
>>> scores = np.array([0.9, 0.9, 0.5, 0.5, 0.4, 0.0], dtype=np.float32)
|
||||
>>> iou_threshold = 0.6
|
||||
>>> dets, inds = soft_nms(boxes, scores, iou_threshold, sigma=0.5)
|
||||
>>> assert len(inds) == len(dets) == 5
|
||||
"""
|
||||
|
||||
assert isinstance(boxes, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(scores, (torch.Tensor, np.ndarray))
|
||||
is_numpy = False
|
||||
if isinstance(boxes, np.ndarray):
|
||||
is_numpy = True
|
||||
boxes = torch.from_numpy(boxes)
|
||||
if isinstance(scores, np.ndarray):
|
||||
scores = torch.from_numpy(scores)
|
||||
assert boxes.size(1) == 4
|
||||
assert boxes.size(0) == scores.size(0)
|
||||
assert offset in (0, 1)
|
||||
method_dict = {'naive': 0, 'linear': 1, 'gaussian': 2}
|
||||
assert method in method_dict.keys()
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
dets = boxes.new_empty((boxes.size(0), 5), device='cpu')
|
||||
indata_list = [boxes.cpu(), scores.cpu(), dets.cpu()]
|
||||
indata_dict = {
|
||||
'iou_threshold': float(iou_threshold),
|
||||
'sigma': float(sigma),
|
||||
'min_score': min_score,
|
||||
'method': method_dict[method],
|
||||
'offset': int(offset)
|
||||
}
|
||||
inds = ext_module.softnms(*indata_list, **indata_dict)
|
||||
else:
|
||||
dets, inds = SoftNMSop.apply(boxes.cpu(), scores.cpu(),
|
||||
float(iou_threshold), float(sigma),
|
||||
float(min_score), method_dict[method],
|
||||
int(offset))
|
||||
|
||||
dets = dets[:inds.size(0)]
|
||||
|
||||
if is_numpy:
|
||||
dets = dets.cpu().numpy()
|
||||
inds = inds.cpu().numpy()
|
||||
return dets, inds
|
||||
else:
|
||||
return dets.to(device=boxes.device), inds.to(device=boxes.device)
|
||||
|
||||
|
||||
def batched_nms(boxes, scores, idxs, nms_cfg, class_agnostic=False):
|
||||
"""Performs non-maximum suppression in a batched fashion.
|
||||
|
||||
Modified from https://github.com/pytorch/vision/blob
|
||||
/505cd6957711af790211896d32b40291bea1bc21/torchvision/ops/boxes.py#L39.
|
||||
In order to perform NMS independently per class, we add an offset to all
|
||||
the boxes. The offset is dependent only on the class idx, and is large
|
||||
enough so that boxes from different classes do not overlap.
|
||||
|
||||
Arguments:
|
||||
boxes (torch.Tensor): boxes in shape (N, 4).
|
||||
scores (torch.Tensor): scores in shape (N, ).
|
||||
idxs (torch.Tensor): each index value correspond to a bbox cluster,
|
||||
and NMS will not be applied between elements of different idxs,
|
||||
shape (N, ).
|
||||
nms_cfg (dict): specify nms type and other parameters like iou_thr.
|
||||
Possible keys includes the following.
|
||||
|
||||
- iou_thr (float): IoU threshold used for NMS.
|
||||
- split_thr (float): threshold number of boxes. In some cases the
|
||||
number of boxes is large (e.g., 200k). To avoid OOM during
|
||||
training, the users could set `split_thr` to a small value.
|
||||
If the number of boxes is greater than the threshold, it will
|
||||
perform NMS on each group of boxes separately and sequentially.
|
||||
Defaults to 10000.
|
||||
class_agnostic (bool): if true, nms is class agnostic,
|
||||
i.e. IoU thresholding happens over all boxes,
|
||||
regardless of the predicted class.
|
||||
|
||||
Returns:
|
||||
tuple: kept dets and indice.
|
||||
"""
|
||||
nms_cfg_ = nms_cfg.copy()
|
||||
class_agnostic = nms_cfg_.pop('class_agnostic', class_agnostic)
|
||||
if class_agnostic:
|
||||
boxes_for_nms = boxes
|
||||
else:
|
||||
max_coordinate = boxes.max()
|
||||
offsets = idxs.to(boxes) * (max_coordinate + torch.tensor(1).to(boxes))
|
||||
boxes_for_nms = boxes + offsets[:, None]
|
||||
|
||||
nms_type = nms_cfg_.pop('type', 'nms')
|
||||
nms_op = eval(nms_type)
|
||||
|
||||
split_thr = nms_cfg_.pop('split_thr', 10000)
|
||||
# Won't split to multiple nms nodes when exporting to onnx
|
||||
if boxes_for_nms.shape[0] < split_thr or torch.onnx.is_in_onnx_export():
|
||||
dets, keep = nms_op(boxes_for_nms, scores, **nms_cfg_)
|
||||
boxes = boxes[keep]
|
||||
# -1 indexing works abnormal in TensorRT
|
||||
# This assumes `dets` has 5 dimensions where
|
||||
# the last dimension is score.
|
||||
# TODO: more elegant way to handle the dimension issue.
|
||||
# Some type of nms would reweight the score, such as SoftNMS
|
||||
scores = dets[:, 4]
|
||||
else:
|
||||
max_num = nms_cfg_.pop('max_num', -1)
|
||||
total_mask = scores.new_zeros(scores.size(), dtype=torch.bool)
|
||||
# Some type of nms would reweight the score, such as SoftNMS
|
||||
scores_after_nms = scores.new_zeros(scores.size())
|
||||
for id in torch.unique(idxs):
|
||||
mask = (idxs == id).nonzero(as_tuple=False).view(-1)
|
||||
dets, keep = nms_op(boxes_for_nms[mask], scores[mask], **nms_cfg_)
|
||||
total_mask[mask[keep]] = True
|
||||
scores_after_nms[mask[keep]] = dets[:, -1]
|
||||
keep = total_mask.nonzero(as_tuple=False).view(-1)
|
||||
|
||||
scores, inds = scores_after_nms[keep].sort(descending=True)
|
||||
keep = keep[inds]
|
||||
boxes = boxes[keep]
|
||||
|
||||
if max_num > 0:
|
||||
keep = keep[:max_num]
|
||||
boxes = boxes[:max_num]
|
||||
scores = scores[:max_num]
|
||||
|
||||
return torch.cat([boxes, scores[:, None]], -1), keep
|
||||
|
||||
|
||||
def nms_match(dets, iou_threshold):
|
||||
"""Matched dets into different groups by NMS.
|
||||
|
||||
NMS match is Similar to NMS but when a bbox is suppressed, nms match will
|
||||
record the indice of suppressed bbox and form a group with the indice of
|
||||
kept bbox. In each group, indice is sorted as score order.
|
||||
|
||||
Arguments:
|
||||
dets (torch.Tensor | np.ndarray): Det boxes with scores, shape (N, 5).
|
||||
iou_thr (float): IoU thresh for NMS.
|
||||
|
||||
Returns:
|
||||
List[torch.Tensor | np.ndarray]: The outer list corresponds different
|
||||
matched group, the inner Tensor corresponds the indices for a group
|
||||
in score order.
|
||||
"""
|
||||
if dets.shape[0] == 0:
|
||||
matched = []
|
||||
else:
|
||||
assert dets.shape[-1] == 5, 'inputs dets.shape should be (N, 5), ' \
|
||||
f'but get {dets.shape}'
|
||||
if isinstance(dets, torch.Tensor):
|
||||
dets_t = dets.detach().cpu()
|
||||
else:
|
||||
dets_t = torch.from_numpy(dets)
|
||||
indata_list = [dets_t]
|
||||
indata_dict = {'iou_threshold': float(iou_threshold)}
|
||||
matched = ext_module.nms_match(*indata_list, **indata_dict)
|
||||
if torch.__version__ == 'parrots':
|
||||
matched = matched.tolist()
|
||||
|
||||
if isinstance(dets, torch.Tensor):
|
||||
return [dets.new_tensor(m, dtype=torch.long) for m in matched]
|
||||
else:
|
||||
return [np.array(m, dtype=np.int) for m in matched]
|
||||
|
||||
|
||||
def nms_rotated(dets, scores, iou_threshold, labels=None):
|
||||
"""Performs non-maximum suppression (NMS) on the rotated boxes according to
|
||||
their intersection-over-union (IoU).
|
||||
|
||||
Rotated NMS iteratively removes lower scoring rotated boxes which have an
|
||||
IoU greater than iou_threshold with another (higher scoring) rotated box.
|
||||
|
||||
Args:
|
||||
boxes (Tensor): Rotated boxes in shape (N, 5). They are expected to \
|
||||
be in (x_ctr, y_ctr, width, height, angle_radian) format.
|
||||
scores (Tensor): scores in shape (N, ).
|
||||
iou_threshold (float): IoU thresh for NMS.
|
||||
labels (Tensor): boxes' label in shape (N,).
|
||||
|
||||
Returns:
|
||||
tuple: kept dets(boxes and scores) and indice, which is always the \
|
||||
same data type as the input.
|
||||
"""
|
||||
if dets.shape[0] == 0:
|
||||
return dets, None
|
||||
multi_label = labels is not None
|
||||
if multi_label:
|
||||
dets_wl = torch.cat((dets, labels.unsqueeze(1)), 1)
|
||||
else:
|
||||
dets_wl = dets
|
||||
_, order = scores.sort(0, descending=True)
|
||||
dets_sorted = dets_wl.index_select(0, order)
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
keep_inds = ext_module.nms_rotated(
|
||||
dets_wl,
|
||||
scores,
|
||||
order,
|
||||
dets_sorted,
|
||||
iou_threshold=iou_threshold,
|
||||
multi_label=multi_label)
|
||||
else:
|
||||
keep_inds = ext_module.nms_rotated(dets_wl, scores, order, dets_sorted,
|
||||
iou_threshold, multi_label)
|
||||
dets = torch.cat((dets[keep_inds], scores[keep_inds].reshape(-1, 1)),
|
||||
dim=1)
|
||||
return dets, keep_inds
|
||||
@@ -0,0 +1,75 @@
|
||||
# Copyright (c) OpenMMLab. All rights reserved.
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', ['pixel_group'])
|
||||
|
||||
|
||||
def pixel_group(score, mask, embedding, kernel_label, kernel_contour,
|
||||
kernel_region_num, distance_threshold):
|
||||
"""Group pixels into text instances, which is widely used text detection
|
||||
methods.
|
||||
|
||||
Arguments:
|
||||
score (np.array or Tensor): The foreground score with size hxw.
|
||||
mask (np.array or Tensor): The foreground mask with size hxw.
|
||||
embedding (np.array or Tensor): The embedding with size hxwxc to
|
||||
distinguish instances.
|
||||
kernel_label (np.array or Tensor): The instance kernel index with
|
||||
size hxw.
|
||||
kernel_contour (np.array or Tensor): The kernel contour with size hxw.
|
||||
kernel_region_num (int): The instance kernel region number.
|
||||
distance_threshold (float): The embedding distance threshold between
|
||||
kernel and pixel in one instance.
|
||||
|
||||
Returns:
|
||||
pixel_assignment (List[List[float]]): The instance coordinate list.
|
||||
Each element consists of averaged confidence, pixel number, and
|
||||
coordinates (x_i, y_i for all pixels) in order.
|
||||
"""
|
||||
assert isinstance(score, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(mask, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(embedding, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(kernel_label, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(kernel_contour, (torch.Tensor, np.ndarray))
|
||||
assert isinstance(kernel_region_num, int)
|
||||
assert isinstance(distance_threshold, float)
|
||||
|
||||
if isinstance(score, np.ndarray):
|
||||
score = torch.from_numpy(score)
|
||||
if isinstance(mask, np.ndarray):
|
||||
mask = torch.from_numpy(mask)
|
||||
if isinstance(embedding, np.ndarray):
|
||||
embedding = torch.from_numpy(embedding)
|
||||
if isinstance(kernel_label, np.ndarray):
|
||||
kernel_label = torch.from_numpy(kernel_label)
|
||||
if isinstance(kernel_contour, np.ndarray):
|
||||
kernel_contour = torch.from_numpy(kernel_contour)
|
||||
|
||||
if torch.__version__ == 'parrots':
|
||||
label = ext_module.pixel_group(
|
||||
score,
|
||||
mask,
|
||||
embedding,
|
||||
kernel_label,
|
||||
kernel_contour,
|
||||
kernel_region_num=kernel_region_num,
|
||||
distance_threshold=distance_threshold)
|
||||
label = label.tolist()
|
||||
label = label[0]
|
||||
list_index = kernel_region_num
|
||||
pixel_assignment = []
|
||||
for x in range(kernel_region_num):
|
||||
pixel_assignment.append(
|
||||
np.array(
|
||||
label[list_index:list_index + int(label[x])],
|
||||
dtype=np.float))
|
||||
list_index = list_index + int(label[x])
|
||||
else:
|
||||
pixel_assignment = ext_module.pixel_group(score, mask, embedding,
|
||||
kernel_label, kernel_contour,
|
||||
kernel_region_num,
|
||||
distance_threshold)
|
||||
return pixel_assignment
|
||||
@@ -0,0 +1,336 @@
|
||||
# Modified from https://github.com/facebookresearch/detectron2/tree/master/projects/PointRend # noqa
|
||||
|
||||
from os import path as osp
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.modules.utils import _pair
|
||||
from torch.onnx.operators import shape_as_tensor
|
||||
|
||||
|
||||
def bilinear_grid_sample(im, grid, align_corners=False):
|
||||
"""Given an input and a flow-field grid, computes the output using input
|
||||
values and pixel locations from grid. Supported only bilinear interpolation
|
||||
method to sample the input pixels.
|
||||
|
||||
Args:
|
||||
im (torch.Tensor): Input feature map, shape (N, C, H, W)
|
||||
grid (torch.Tensor): Point coordinates, shape (N, Hg, Wg, 2)
|
||||
align_corners {bool}: If set to True, the extrema (-1 and 1) are
|
||||
considered as referring to the center points of the input’s
|
||||
corner pixels. If set to False, they are instead considered as
|
||||
referring to the corner points of the input’s corner pixels,
|
||||
making the sampling more resolution agnostic.
|
||||
Returns:
|
||||
torch.Tensor: A tensor with sampled points, shape (N, C, Hg, Wg)
|
||||
"""
|
||||
n, c, h, w = im.shape
|
||||
gn, gh, gw, _ = grid.shape
|
||||
assert n == gn
|
||||
|
||||
x = grid[:, :, :, 0]
|
||||
y = grid[:, :, :, 1]
|
||||
|
||||
if align_corners:
|
||||
x = ((x + 1) / 2) * (w - 1)
|
||||
y = ((y + 1) / 2) * (h - 1)
|
||||
else:
|
||||
x = ((x + 1) * w - 1) / 2
|
||||
y = ((y + 1) * h - 1) / 2
|
||||
|
||||
x = x.view(n, -1)
|
||||
y = y.view(n, -1)
|
||||
|
||||
x0 = torch.floor(x).long()
|
||||
y0 = torch.floor(y).long()
|
||||
x1 = x0 + 1
|
||||
y1 = y0 + 1
|
||||
|
||||
wa = ((x1 - x) * (y1 - y)).unsqueeze(1)
|
||||
wb = ((x1 - x) * (y - y0)).unsqueeze(1)
|
||||
wc = ((x - x0) * (y1 - y)).unsqueeze(1)
|
||||
wd = ((x - x0) * (y - y0)).unsqueeze(1)
|
||||
|
||||
# Apply default for grid_sample function zero padding
|
||||
im_padded = F.pad(im, pad=[1, 1, 1, 1], mode='constant', value=0)
|
||||
padded_h = h + 2
|
||||
padded_w = w + 2
|
||||
# save points positions after padding
|
||||
x0, x1, y0, y1 = x0 + 1, x1 + 1, y0 + 1, y1 + 1
|
||||
|
||||
# Clip coordinates to padded image size
|
||||
x0 = torch.where(x0 < 0, torch.tensor(0), x0)
|
||||
x0 = torch.where(x0 > padded_w - 1, torch.tensor(padded_w - 1), x0)
|
||||
x1 = torch.where(x1 < 0, torch.tensor(0), x1)
|
||||
x1 = torch.where(x1 > padded_w - 1, torch.tensor(padded_w - 1), x1)
|
||||
y0 = torch.where(y0 < 0, torch.tensor(0), y0)
|
||||
y0 = torch.where(y0 > padded_h - 1, torch.tensor(padded_h - 1), y0)
|
||||
y1 = torch.where(y1 < 0, torch.tensor(0), y1)
|
||||
y1 = torch.where(y1 > padded_h - 1, torch.tensor(padded_h - 1), y1)
|
||||
|
||||
im_padded = im_padded.view(n, c, -1)
|
||||
|
||||
x0_y0 = (x0 + y0 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x0_y1 = (x0 + y1 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x1_y0 = (x1 + y0 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x1_y1 = (x1 + y1 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
|
||||
Ia = torch.gather(im_padded, 2, x0_y0)
|
||||
Ib = torch.gather(im_padded, 2, x0_y1)
|
||||
Ic = torch.gather(im_padded, 2, x1_y0)
|
||||
Id = torch.gather(im_padded, 2, x1_y1)
|
||||
|
||||
return (Ia * wa + Ib * wb + Ic * wc + Id * wd).reshape(n, c, gh, gw)
|
||||
|
||||
|
||||
def is_in_onnx_export_without_custom_ops():
|
||||
from custom_mmpkg.custom_mmcv.ops import get_onnxruntime_op_path
|
||||
ort_custom_op_path = get_onnxruntime_op_path()
|
||||
return torch.onnx.is_in_onnx_export(
|
||||
) and not osp.exists(ort_custom_op_path)
|
||||
|
||||
|
||||
def normalize(grid):
|
||||
"""Normalize input grid from [-1, 1] to [0, 1]
|
||||
Args:
|
||||
grid (Tensor): The grid to be normalize, range [-1, 1].
|
||||
Returns:
|
||||
Tensor: Normalized grid, range [0, 1].
|
||||
"""
|
||||
|
||||
return (grid + 1.0) / 2.0
|
||||
|
||||
|
||||
def denormalize(grid):
|
||||
"""Denormalize input grid from range [0, 1] to [-1, 1]
|
||||
Args:
|
||||
grid (Tensor): The grid to be denormalize, range [0, 1].
|
||||
Returns:
|
||||
Tensor: Denormalized grid, range [-1, 1].
|
||||
"""
|
||||
|
||||
return grid * 2.0 - 1.0
|
||||
|
||||
|
||||
def generate_grid(num_grid, size, device):
|
||||
"""Generate regular square grid of points in [0, 1] x [0, 1] coordinate
|
||||
space.
|
||||
|
||||
Args:
|
||||
num_grid (int): The number of grids to sample, one for each region.
|
||||
size (tuple(int, int)): The side size of the regular grid.
|
||||
device (torch.device): Desired device of returned tensor.
|
||||
|
||||
Returns:
|
||||
(torch.Tensor): A tensor of shape (num_grid, size[0]*size[1], 2) that
|
||||
contains coordinates for the regular grids.
|
||||
"""
|
||||
|
||||
affine_trans = torch.tensor([[[1., 0., 0.], [0., 1., 0.]]], device=device)
|
||||
grid = F.affine_grid(
|
||||
affine_trans, torch.Size((1, 1, *size)), align_corners=False)
|
||||
grid = normalize(grid)
|
||||
return grid.view(1, -1, 2).expand(num_grid, -1, -1)
|
||||
|
||||
|
||||
def rel_roi_point_to_abs_img_point(rois, rel_roi_points):
|
||||
"""Convert roi based relative point coordinates to image based absolute
|
||||
point coordinates.
|
||||
|
||||
Args:
|
||||
rois (Tensor): RoIs or BBoxes, shape (N, 4) or (N, 5)
|
||||
rel_roi_points (Tensor): Point coordinates inside RoI, relative to
|
||||
RoI, location, range (0, 1), shape (N, P, 2)
|
||||
Returns:
|
||||
Tensor: Image based absolute point coordinates, shape (N, P, 2)
|
||||
"""
|
||||
|
||||
with torch.no_grad():
|
||||
assert rel_roi_points.size(0) == rois.size(0)
|
||||
assert rois.dim() == 2
|
||||
assert rel_roi_points.dim() == 3
|
||||
assert rel_roi_points.size(2) == 2
|
||||
# remove batch idx
|
||||
if rois.size(1) == 5:
|
||||
rois = rois[:, 1:]
|
||||
abs_img_points = rel_roi_points.clone()
|
||||
# To avoid an error during exporting to onnx use independent
|
||||
# variables instead inplace computation
|
||||
xs = abs_img_points[:, :, 0] * (rois[:, None, 2] - rois[:, None, 0])
|
||||
ys = abs_img_points[:, :, 1] * (rois[:, None, 3] - rois[:, None, 1])
|
||||
xs += rois[:, None, 0]
|
||||
ys += rois[:, None, 1]
|
||||
abs_img_points = torch.stack([xs, ys], dim=2)
|
||||
return abs_img_points
|
||||
|
||||
|
||||
def get_shape_from_feature_map(x):
|
||||
"""Get spatial resolution of input feature map considering exporting to
|
||||
onnx mode.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor, shape (N, C, H, W)
|
||||
Returns:
|
||||
torch.Tensor: Spatial resolution (width, height), shape (1, 1, 2)
|
||||
"""
|
||||
if torch.onnx.is_in_onnx_export():
|
||||
img_shape = shape_as_tensor(x)[2:].flip(0).view(1, 1, 2).to(
|
||||
x.device).float()
|
||||
else:
|
||||
img_shape = torch.tensor(x.shape[2:]).flip(0).view(1, 1, 2).to(
|
||||
x.device).float()
|
||||
return img_shape
|
||||
|
||||
|
||||
def abs_img_point_to_rel_img_point(abs_img_points, img, spatial_scale=1.):
|
||||
"""Convert image based absolute point coordinates to image based relative
|
||||
coordinates for sampling.
|
||||
|
||||
Args:
|
||||
abs_img_points (Tensor): Image based absolute point coordinates,
|
||||
shape (N, P, 2)
|
||||
img (tuple/Tensor): (height, width) of image or feature map.
|
||||
spatial_scale (float): Scale points by this factor. Default: 1.
|
||||
|
||||
Returns:
|
||||
Tensor: Image based relative point coordinates for sampling,
|
||||
shape (N, P, 2)
|
||||
"""
|
||||
|
||||
assert (isinstance(img, tuple) and len(img) == 2) or \
|
||||
(isinstance(img, torch.Tensor) and len(img.shape) == 4)
|
||||
|
||||
if isinstance(img, tuple):
|
||||
h, w = img
|
||||
scale = torch.tensor([w, h],
|
||||
dtype=torch.float,
|
||||
device=abs_img_points.device)
|
||||
scale = scale.view(1, 1, 2)
|
||||
else:
|
||||
scale = get_shape_from_feature_map(img)
|
||||
|
||||
return abs_img_points / scale * spatial_scale
|
||||
|
||||
|
||||
def rel_roi_point_to_rel_img_point(rois,
|
||||
rel_roi_points,
|
||||
img,
|
||||
spatial_scale=1.):
|
||||
"""Convert roi based relative point coordinates to image based absolute
|
||||
point coordinates.
|
||||
|
||||
Args:
|
||||
rois (Tensor): RoIs or BBoxes, shape (N, 4) or (N, 5)
|
||||
rel_roi_points (Tensor): Point coordinates inside RoI, relative to
|
||||
RoI, location, range (0, 1), shape (N, P, 2)
|
||||
img (tuple/Tensor): (height, width) of image or feature map.
|
||||
spatial_scale (float): Scale points by this factor. Default: 1.
|
||||
|
||||
Returns:
|
||||
Tensor: Image based relative point coordinates for sampling,
|
||||
shape (N, P, 2)
|
||||
"""
|
||||
|
||||
abs_img_point = rel_roi_point_to_abs_img_point(rois, rel_roi_points)
|
||||
rel_img_point = abs_img_point_to_rel_img_point(abs_img_point, img,
|
||||
spatial_scale)
|
||||
|
||||
return rel_img_point
|
||||
|
||||
|
||||
def point_sample(input, points, align_corners=False, **kwargs):
|
||||
"""A wrapper around :func:`grid_sample` to support 3D point_coords tensors
|
||||
Unlike :func:`torch.nn.functional.grid_sample` it assumes point_coords to
|
||||
lie inside ``[0, 1] x [0, 1]`` square.
|
||||
|
||||
Args:
|
||||
input (Tensor): Feature map, shape (N, C, H, W).
|
||||
points (Tensor): Image based absolute point coordinates (normalized),
|
||||
range [0, 1] x [0, 1], shape (N, P, 2) or (N, Hgrid, Wgrid, 2).
|
||||
align_corners (bool): Whether align_corners. Default: False
|
||||
|
||||
Returns:
|
||||
Tensor: Features of `point` on `input`, shape (N, C, P) or
|
||||
(N, C, Hgrid, Wgrid).
|
||||
"""
|
||||
|
||||
add_dim = False
|
||||
if points.dim() == 3:
|
||||
add_dim = True
|
||||
points = points.unsqueeze(2)
|
||||
if is_in_onnx_export_without_custom_ops():
|
||||
# If custom ops for onnx runtime not compiled use python
|
||||
# implementation of grid_sample function to make onnx graph
|
||||
# with supported nodes
|
||||
output = bilinear_grid_sample(
|
||||
input, denormalize(points), align_corners=align_corners)
|
||||
else:
|
||||
output = F.grid_sample(
|
||||
input, denormalize(points), align_corners=align_corners, **kwargs)
|
||||
if add_dim:
|
||||
output = output.squeeze(3)
|
||||
return output
|
||||
|
||||
|
||||
class SimpleRoIAlign(nn.Module):
|
||||
|
||||
def __init__(self, output_size, spatial_scale, aligned=True):
|
||||
"""Simple RoI align in PointRend, faster than standard RoIAlign.
|
||||
|
||||
Args:
|
||||
output_size (tuple[int]): h, w
|
||||
spatial_scale (float): scale the input boxes by this number
|
||||
aligned (bool): if False, use the legacy implementation in
|
||||
MMDetection, align_corners=True will be used in F.grid_sample.
|
||||
If True, align the results more perfectly.
|
||||
"""
|
||||
|
||||
super(SimpleRoIAlign, self).__init__()
|
||||
self.output_size = _pair(output_size)
|
||||
self.spatial_scale = float(spatial_scale)
|
||||
# to be consistent with other RoI ops
|
||||
self.use_torchvision = False
|
||||
self.aligned = aligned
|
||||
|
||||
def forward(self, features, rois):
|
||||
num_imgs = features.size(0)
|
||||
num_rois = rois.size(0)
|
||||
rel_roi_points = generate_grid(
|
||||
num_rois, self.output_size, device=rois.device)
|
||||
|
||||
if torch.onnx.is_in_onnx_export():
|
||||
rel_img_points = rel_roi_point_to_rel_img_point(
|
||||
rois, rel_roi_points, features, self.spatial_scale)
|
||||
rel_img_points = rel_img_points.reshape(num_imgs, -1,
|
||||
*rel_img_points.shape[1:])
|
||||
point_feats = point_sample(
|
||||
features, rel_img_points, align_corners=not self.aligned)
|
||||
point_feats = point_feats.transpose(1, 2)
|
||||
else:
|
||||
point_feats = []
|
||||
for batch_ind in range(num_imgs):
|
||||
# unravel batch dim
|
||||
feat = features[batch_ind].unsqueeze(0)
|
||||
inds = (rois[:, 0].long() == batch_ind)
|
||||
if inds.any():
|
||||
rel_img_points = rel_roi_point_to_rel_img_point(
|
||||
rois[inds], rel_roi_points[inds], feat,
|
||||
self.spatial_scale).unsqueeze(0)
|
||||
point_feat = point_sample(
|
||||
feat, rel_img_points, align_corners=not self.aligned)
|
||||
point_feat = point_feat.squeeze(0).transpose(0, 1)
|
||||
point_feats.append(point_feat)
|
||||
|
||||
point_feats = torch.cat(point_feats, dim=0)
|
||||
|
||||
channels = features.size(1)
|
||||
roi_feats = point_feats.reshape(num_rois, channels, *self.output_size)
|
||||
|
||||
return roi_feats
|
||||
|
||||
def __repr__(self):
|
||||
format_str = self.__class__.__name__
|
||||
format_str += '(output_size={}, spatial_scale={}'.format(
|
||||
self.output_size, self.spatial_scale)
|
||||
return format_str
|
||||
@@ -0,0 +1,133 @@
|
||||
import torch
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext', [
|
||||
'points_in_boxes_part_forward', 'points_in_boxes_cpu_forward',
|
||||
'points_in_boxes_all_forward'
|
||||
])
|
||||
|
||||
|
||||
def points_in_boxes_part(points, boxes):
|
||||
"""Find the box in which each point is (CUDA).
|
||||
|
||||
Args:
|
||||
points (torch.Tensor): [B, M, 3], [x, y, z] in LiDAR/DEPTH coordinate
|
||||
boxes (torch.Tensor): [B, T, 7],
|
||||
num_valid_boxes <= T, [x, y, z, x_size, y_size, z_size, rz] in
|
||||
LiDAR/DEPTH coordinate, (x, y, z) is the bottom center
|
||||
|
||||
Returns:
|
||||
box_idxs_of_pts (torch.Tensor): (B, M), default background = -1
|
||||
"""
|
||||
assert points.shape[0] == boxes.shape[0], \
|
||||
'Points and boxes should have the same batch size, ' \
|
||||
f'but got {points.shape[0]} and {boxes.shape[0]}'
|
||||
assert boxes.shape[2] == 7, \
|
||||
'boxes dimension should be 7, ' \
|
||||
f'but got unexpected shape {boxes.shape[2]}'
|
||||
assert points.shape[2] == 3, \
|
||||
'points dimension should be 3, ' \
|
||||
f'but got unexpected shape {points.shape[2]}'
|
||||
batch_size, num_points, _ = points.shape
|
||||
|
||||
box_idxs_of_pts = points.new_zeros((batch_size, num_points),
|
||||
dtype=torch.int).fill_(-1)
|
||||
|
||||
# If manually put the tensor 'points' or 'boxes' on a device
|
||||
# which is not the current device, some temporary variables
|
||||
# will be created on the current device in the cuda op,
|
||||
# and the output will be incorrect.
|
||||
# Therefore, we force the current device to be the same
|
||||
# as the device of the tensors if it was not.
|
||||
# Please refer to https://github.com/open-mmlab/mmdetection3d/issues/305
|
||||
# for the incorrect output before the fix.
|
||||
points_device = points.get_device()
|
||||
assert points_device == boxes.get_device(), \
|
||||
'Points and boxes should be put on the same device'
|
||||
if torch.cuda.current_device() != points_device:
|
||||
torch.cuda.set_device(points_device)
|
||||
|
||||
ext_module.points_in_boxes_part_forward(boxes.contiguous(),
|
||||
points.contiguous(),
|
||||
box_idxs_of_pts)
|
||||
|
||||
return box_idxs_of_pts
|
||||
|
||||
|
||||
def points_in_boxes_cpu(points, boxes):
|
||||
"""Find all boxes in which each point is (CPU). The CPU version of
|
||||
:meth:`points_in_boxes_all`.
|
||||
|
||||
Args:
|
||||
points (torch.Tensor): [B, M, 3], [x, y, z] in
|
||||
LiDAR/DEPTH coordinate
|
||||
boxes (torch.Tensor): [B, T, 7],
|
||||
num_valid_boxes <= T, [x, y, z, x_size, y_size, z_size, rz],
|
||||
(x, y, z) is the bottom center.
|
||||
|
||||
Returns:
|
||||
box_idxs_of_pts (torch.Tensor): (B, M, T), default background = 0.
|
||||
"""
|
||||
assert points.shape[0] == boxes.shape[0], \
|
||||
'Points and boxes should have the same batch size, ' \
|
||||
f'but got {points.shape[0]} and {boxes.shape[0]}'
|
||||
assert boxes.shape[2] == 7, \
|
||||
'boxes dimension should be 7, ' \
|
||||
f'but got unexpected shape {boxes.shape[2]}'
|
||||
assert points.shape[2] == 3, \
|
||||
'points dimension should be 3, ' \
|
||||
f'but got unexpected shape {points.shape[2]}'
|
||||
batch_size, num_points, _ = points.shape
|
||||
num_boxes = boxes.shape[1]
|
||||
|
||||
point_indices = points.new_zeros((batch_size, num_boxes, num_points),
|
||||
dtype=torch.int)
|
||||
for b in range(batch_size):
|
||||
ext_module.points_in_boxes_cpu_forward(boxes[b].float().contiguous(),
|
||||
points[b].float().contiguous(),
|
||||
point_indices[b])
|
||||
point_indices = point_indices.transpose(1, 2)
|
||||
|
||||
return point_indices
|
||||
|
||||
|
||||
def points_in_boxes_all(points, boxes):
|
||||
"""Find all boxes in which each point is (CUDA).
|
||||
|
||||
Args:
|
||||
points (torch.Tensor): [B, M, 3], [x, y, z] in LiDAR/DEPTH coordinate
|
||||
boxes (torch.Tensor): [B, T, 7],
|
||||
num_valid_boxes <= T, [x, y, z, x_size, y_size, z_size, rz],
|
||||
(x, y, z) is the bottom center.
|
||||
|
||||
Returns:
|
||||
box_idxs_of_pts (torch.Tensor): (B, M, T), default background = 0.
|
||||
"""
|
||||
assert boxes.shape[0] == points.shape[0], \
|
||||
'Points and boxes should have the same batch size, ' \
|
||||
f'but got {boxes.shape[0]} and {boxes.shape[0]}'
|
||||
assert boxes.shape[2] == 7, \
|
||||
'boxes dimension should be 7, ' \
|
||||
f'but got unexpected shape {boxes.shape[2]}'
|
||||
assert points.shape[2] == 3, \
|
||||
'points dimension should be 3, ' \
|
||||
f'but got unexpected shape {points.shape[2]}'
|
||||
batch_size, num_points, _ = points.shape
|
||||
num_boxes = boxes.shape[1]
|
||||
|
||||
box_idxs_of_pts = points.new_zeros((batch_size, num_points, num_boxes),
|
||||
dtype=torch.int).fill_(0)
|
||||
|
||||
# Same reason as line 25-32
|
||||
points_device = points.get_device()
|
||||
assert points_device == boxes.get_device(), \
|
||||
'Points and boxes should be put on the same device'
|
||||
if torch.cuda.current_device() != points_device:
|
||||
torch.cuda.set_device(points_device)
|
||||
|
||||
ext_module.points_in_boxes_all_forward(boxes.contiguous(),
|
||||
points.contiguous(),
|
||||
box_idxs_of_pts)
|
||||
|
||||
return box_idxs_of_pts
|
||||
@@ -0,0 +1,177 @@
|
||||
from typing import List
|
||||
|
||||
import torch
|
||||
from torch import nn as nn
|
||||
|
||||
from custom_mmpkg.custom_mmcv.runner import force_fp32
|
||||
from .furthest_point_sample import (furthest_point_sample,
|
||||
furthest_point_sample_with_dist)
|
||||
|
||||
|
||||
def calc_square_dist(point_feat_a, point_feat_b, norm=True):
|
||||
"""Calculating square distance between a and b.
|
||||
|
||||
Args:
|
||||
point_feat_a (Tensor): (B, N, C) Feature vector of each point.
|
||||
point_feat_b (Tensor): (B, M, C) Feature vector of each point.
|
||||
norm (Bool, optional): Whether to normalize the distance.
|
||||
Default: True.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, N, M) Distance between each pair points.
|
||||
"""
|
||||
num_channel = point_feat_a.shape[-1]
|
||||
# [bs, n, 1]
|
||||
a_square = torch.sum(point_feat_a.unsqueeze(dim=2).pow(2), dim=-1)
|
||||
# [bs, 1, m]
|
||||
b_square = torch.sum(point_feat_b.unsqueeze(dim=1).pow(2), dim=-1)
|
||||
|
||||
corr_matrix = torch.matmul(point_feat_a, point_feat_b.transpose(1, 2))
|
||||
|
||||
dist = a_square + b_square - 2 * corr_matrix
|
||||
if norm:
|
||||
dist = torch.sqrt(dist) / num_channel
|
||||
return dist
|
||||
|
||||
|
||||
def get_sampler_cls(sampler_type):
|
||||
"""Get the type and mode of points sampler.
|
||||
|
||||
Args:
|
||||
sampler_type (str): The type of points sampler.
|
||||
The valid value are "D-FPS", "F-FPS", or "FS".
|
||||
|
||||
Returns:
|
||||
class: Points sampler type.
|
||||
"""
|
||||
sampler_mappings = {
|
||||
'D-FPS': DFPSSampler,
|
||||
'F-FPS': FFPSSampler,
|
||||
'FS': FSSampler,
|
||||
}
|
||||
try:
|
||||
return sampler_mappings[sampler_type]
|
||||
except KeyError:
|
||||
raise KeyError(
|
||||
f'Supported `sampler_type` are {sampler_mappings.keys()}, but got \
|
||||
{sampler_type}')
|
||||
|
||||
|
||||
class PointsSampler(nn.Module):
|
||||
"""Points sampling.
|
||||
|
||||
Args:
|
||||
num_point (list[int]): Number of sample points.
|
||||
fps_mod_list (list[str], optional): Type of FPS method, valid mod
|
||||
['F-FPS', 'D-FPS', 'FS'], Default: ['D-FPS'].
|
||||
F-FPS: using feature distances for FPS.
|
||||
D-FPS: using Euclidean distances of points for FPS.
|
||||
FS: using F-FPS and D-FPS simultaneously.
|
||||
fps_sample_range_list (list[int], optional):
|
||||
Range of points to apply FPS. Default: [-1].
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
num_point: List[int],
|
||||
fps_mod_list: List[str] = ['D-FPS'],
|
||||
fps_sample_range_list: List[int] = [-1]):
|
||||
super().__init__()
|
||||
# FPS would be applied to different fps_mod in the list,
|
||||
# so the length of the num_point should be equal to
|
||||
# fps_mod_list and fps_sample_range_list.
|
||||
assert len(num_point) == len(fps_mod_list) == len(
|
||||
fps_sample_range_list)
|
||||
self.num_point = num_point
|
||||
self.fps_sample_range_list = fps_sample_range_list
|
||||
self.samplers = nn.ModuleList()
|
||||
for fps_mod in fps_mod_list:
|
||||
self.samplers.append(get_sampler_cls(fps_mod)())
|
||||
self.fp16_enabled = False
|
||||
|
||||
@force_fp32()
|
||||
def forward(self, points_xyz, features):
|
||||
"""
|
||||
Args:
|
||||
points_xyz (Tensor): (B, N, 3) xyz coordinates of the features.
|
||||
features (Tensor): (B, C, N) Descriptors of the features.
|
||||
|
||||
Returns:
|
||||
Tensor: (B, npoint, sample_num) Indices of sampled points.
|
||||
"""
|
||||
indices = []
|
||||
last_fps_end_index = 0
|
||||
|
||||
for fps_sample_range, sampler, npoint in zip(
|
||||
self.fps_sample_range_list, self.samplers, self.num_point):
|
||||
assert fps_sample_range < points_xyz.shape[1]
|
||||
|
||||
if fps_sample_range == -1:
|
||||
sample_points_xyz = points_xyz[:, last_fps_end_index:]
|
||||
if features is not None:
|
||||
sample_features = features[:, :, last_fps_end_index:]
|
||||
else:
|
||||
sample_features = None
|
||||
else:
|
||||
sample_points_xyz = \
|
||||
points_xyz[:, last_fps_end_index:fps_sample_range]
|
||||
if features is not None:
|
||||
sample_features = features[:, :, last_fps_end_index:
|
||||
fps_sample_range]
|
||||
else:
|
||||
sample_features = None
|
||||
|
||||
fps_idx = sampler(sample_points_xyz.contiguous(), sample_features,
|
||||
npoint)
|
||||
|
||||
indices.append(fps_idx + last_fps_end_index)
|
||||
last_fps_end_index += fps_sample_range
|
||||
indices = torch.cat(indices, dim=1)
|
||||
|
||||
return indices
|
||||
|
||||
|
||||
class DFPSSampler(nn.Module):
|
||||
"""Using Euclidean distances of points for FPS."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, points, features, npoint):
|
||||
"""Sampling points with D-FPS."""
|
||||
fps_idx = furthest_point_sample(points.contiguous(), npoint)
|
||||
return fps_idx
|
||||
|
||||
|
||||
class FFPSSampler(nn.Module):
|
||||
"""Using feature distances for FPS."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, points, features, npoint):
|
||||
"""Sampling points with F-FPS."""
|
||||
assert features is not None, \
|
||||
'feature input to FFPS_Sampler should not be None'
|
||||
features_for_fps = torch.cat([points, features.transpose(1, 2)], dim=2)
|
||||
features_dist = calc_square_dist(
|
||||
features_for_fps, features_for_fps, norm=False)
|
||||
fps_idx = furthest_point_sample_with_dist(features_dist, npoint)
|
||||
return fps_idx
|
||||
|
||||
|
||||
class FSSampler(nn.Module):
|
||||
"""Using F-FPS and D-FPS simultaneously."""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, points, features, npoint):
|
||||
"""Sampling points with FS_Sampling."""
|
||||
assert features is not None, \
|
||||
'feature input to FS_Sampler should not be None'
|
||||
ffps_sampler = FFPSSampler()
|
||||
dfps_sampler = DFPSSampler()
|
||||
fps_idx_ffps = ffps_sampler(points, features, npoint)
|
||||
fps_idx_dfps = dfps_sampler(points, features, npoint)
|
||||
fps_idx = torch.cat([fps_idx_ffps, fps_idx_dfps], dim=1)
|
||||
return fps_idx
|
||||
@@ -0,0 +1,92 @@
|
||||
# Modified from https://github.com/hszhao/semseg/blob/master/lib/psa
|
||||
from torch import nn
|
||||
from torch.autograd import Function
|
||||
from torch.nn.modules.utils import _pair
|
||||
|
||||
from ..utils import ext_loader
|
||||
|
||||
ext_module = ext_loader.load_ext('_ext',
|
||||
['psamask_forward', 'psamask_backward'])
|
||||
|
||||
|
||||
class PSAMaskFunction(Function):
|
||||
|
||||
@staticmethod
|
||||
def symbolic(g, input, psa_type, mask_size):
|
||||
return g.op(
|
||||
'mmcv::MMCVPSAMask',
|
||||
input,
|
||||
psa_type_i=psa_type,
|
||||
mask_size_i=mask_size)
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, input, psa_type, mask_size):
|
||||
ctx.psa_type = psa_type
|
||||
ctx.mask_size = _pair(mask_size)
|
||||
ctx.save_for_backward(input)
|
||||
|
||||
h_mask, w_mask = ctx.mask_size
|
||||
batch_size, channels, h_feature, w_feature = input.size()
|
||||
assert channels == h_mask * w_mask
|
||||
output = input.new_zeros(
|
||||
(batch_size, h_feature * w_feature, h_feature, w_feature))
|
||||
|
||||
ext_module.psamask_forward(
|
||||
input,
|
||||
output,
|
||||
psa_type=psa_type,
|
||||
num_=batch_size,
|
||||
h_feature=h_feature,
|
||||
w_feature=w_feature,
|
||||
h_mask=h_mask,
|
||||
w_mask=w_mask,
|
||||
half_h_mask=(h_mask - 1) // 2,
|
||||
half_w_mask=(w_mask - 1) // 2)
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
input = ctx.saved_tensors[0]
|
||||
psa_type = ctx.psa_type
|
||||
h_mask, w_mask = ctx.mask_size
|
||||
batch_size, channels, h_feature, w_feature = input.size()
|
||||
grad_input = grad_output.new_zeros(
|
||||
(batch_size, channels, h_feature, w_feature))
|
||||
ext_module.psamask_backward(
|
||||
grad_output,
|
||||
grad_input,
|
||||
psa_type=psa_type,
|
||||
num_=batch_size,
|
||||
h_feature=h_feature,
|
||||
w_feature=w_feature,
|
||||
h_mask=h_mask,
|
||||
w_mask=w_mask,
|
||||
half_h_mask=(h_mask - 1) // 2,
|
||||
half_w_mask=(w_mask - 1) // 2)
|
||||
return grad_input, None, None, None
|
||||
|
||||
|
||||
psa_mask = PSAMaskFunction.apply
|
||||
|
||||
|
||||
class PSAMask(nn.Module):
|
||||
|
||||
def __init__(self, psa_type, mask_size=None):
|
||||
super(PSAMask, self).__init__()
|
||||
assert psa_type in ['collect', 'distribute']
|
||||
if psa_type == 'collect':
|
||||
psa_type_enum = 0
|
||||
else:
|
||||
psa_type_enum = 1
|
||||
self.psa_type_enum = psa_type_enum
|
||||
self.mask_size = mask_size
|
||||
self.psa_type = psa_type
|
||||
|
||||
def forward(self, input):
|
||||
return psa_mask(input, self.psa_type_enum, self.mask_size)
|
||||
|
||||
def __repr__(self):
|
||||
s = self.__class__.__name__
|
||||
s += f'(psa_type={self.psa_type}, '
|
||||
s += f'mask_size={self.mask_size})'
|
||||
return s
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user