initial commit

This commit is contained in:
kijai
2024-11-07 21:25:03 +02:00
parent 766c95fc91
commit 97f33da323
89 changed files with 38037 additions and 0 deletions
+9
View File
@@ -0,0 +1,9 @@
output/
*__pycache__/
samples*/
runs/
checkpoints/
master_ip
logs/
*.DS_Store
.idea
+13
View File
@@ -0,0 +1,13 @@
S-Lab License 1.0
Copyright 2024 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.
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.
+7
View File
@@ -0,0 +1,7 @@
# ComfyUI nodes to use GIMM-VFI frame interpolation
Requires cupy, currently tested only with `cupy-cuda12==13.3.0`
Original repository:
https://github.com/GSeanCDAT/GIMM-VFI
+3
View File
@@ -0,0 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+61
View File
@@ -0,0 +1,61 @@
trainer: stage_inr
dataset:
type: fast_vimeo_flow
path: ./data/vimeo90k/vimeo_triplet
add_objects: false
expansion: false
random_t: false
aug: true
t_scale: 10
pair: false
arch: # needs to add encoder, modulation type
type: gimm
ema: null
modulated_layer_idxs: [1]
coord_range: [-1., 1.]
hyponet:
type: mlp
n_layer: 5 # including the output layer
hidden_dim: [128] # list, assert len(hidden_dim) in [1, n_layers-1]
use_bias: true
input_dim: 3
output_dim: 2
output_bias: 0.5
activation:
type: siren
siren_w0: 1.0
initialization:
weight_init_type: siren
bias_init_type: siren
loss:
type: mse #now unnecessary
optimizer:
type: adam
init_lr: 0.0001
weight_decay: 0.0
betas: [0.9, 0.999] #[0.9, 0.95]
ft: false
warmup:
epoch: 0
multiplier: 1
buffer_epoch: 0
min_lr: 0.0001
mode: fix
start_from_zero: True
max_gn: null
experiment:
amp: True
batch_size: 32
total_batch_size: 64
epochs: 400
save_ckpt_freq: 20
test_freq: 10
test_imlog_freq: 10
+57
View File
@@ -0,0 +1,57 @@
trainer: stage_inr
dataset:
type: vimeo_arb
path: ./data/vimeo90k/vimeo_septuplet
aug: true
arch:
type: gimmvfi_f
ema: true
modulated_layer_idxs: [1]
coord_range: [-1., 1.]
hyponet:
type: mlp
n_layer: 5 # including the output layer
hidden_dim: [128] # list, assert len(hidden_dim) in [1, n_layers-1]
use_bias: true
input_dim: 3
output_dim: 2
output_bias: 0.5
activation:
type: siren
siren_w0: 1.0
initialization:
weight_init_type: siren
bias_init_type: siren
loss:
subsample:
type: random
ratio: 0.1
optimizer:
type: adamw
init_lr: 0.00008
weight_decay: 0.00004
betas: [0.9, 0.999]
ft: true
warmup:
epoch: 1
multiplier: 1
buffer_epoch: 0
min_lr: 0.000008
mode: fix
start_from_zero: True
max_gn: null
experiment:
amp: True
batch_size: 4
total_batch_size: 32
epochs: 60
save_ckpt_freq: 10
test_freq: 10
test_imlog_freq: 10
+57
View File
@@ -0,0 +1,57 @@
trainer: stage_inr
dataset:
type: vimeo_arb
path: ./data/vimeo90k/vimeo_septuplet
aug: true
arch:
type: gimmvfi_r
ema: true
modulated_layer_idxs: [1]
coord_range: [-1., 1.]
hyponet:
type: mlp
n_layer: 5 # including the output layer
hidden_dim: [128] # list, assert len(hidden_dim) in [1, n_layers-1]
use_bias: true
input_dim: 3
output_dim: 2
output_bias: 0.5
activation:
type: siren
siren_w0: 1.0
initialization:
weight_init_type: siren
bias_init_type: siren
loss:
subsample:
type: random
ratio: 0.1
optimizer:
type: adamw
init_lr: 0.00008
weight_decay: 0.00004
betas: [0.9, 0.999]
ft: true
warmup:
epoch: 1
multiplier: 1
buffer_epoch: 0
min_lr: 0.000008
mode: fix
start_from_zero: True
max_gn: null
experiment:
amp: True
batch_size: 4
total_batch_size: 32
epochs: 60
save_ckpt_freq: 10
test_freq: 10
test_imlog_freq: 10
+26
View File
@@ -0,0 +1,26 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
# from .gimm import GIMM
# from .gimmvfi_f import GIMMVFI_F
# from .gimmvfi_r import GIMMVFI_R
# def gimm(config):
# return GIMM(config)
# def gimmvfi_f(config):
# return GIMMVFI_F(config)
# def gimmvfi_r(config):
# return GIMMVFI_R(config)
+57
View File
@@ -0,0 +1,57 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
from typing import List, Optional
from dataclasses import dataclass, field
from omegaconf import OmegaConf, MISSING
from .modules.module_config import HypoNetConfig
@dataclass
class GIMMConfig:
type: str = "gimm"
ema: Optional[bool] = None
ema_value: Optional[float] = None
fwarp_type: str = "linear"
hyponet: HypoNetConfig = field(default_factory=HypoNetConfig)
coord_range: List[float] = MISSING
modulated_layer_idxs: Optional[List[int]] = None
@classmethod
def create(cls, config):
# We need to specify the type of the default DataEncoderConfig.
# Otherwise, data_encoder will be initialized & structured as "unfold" type (which is default value)
# hence merging with the config with other type would cause config error.
defaults = OmegaConf.structured(cls(ema=False))
config = OmegaConf.merge(defaults, config)
return config
@dataclass
class GIMMVFIConfig:
type: str = "gimmvfi"
ema: Optional[bool] = None
ema_value: Optional[float] = None
fwarp_type: str = "linear"
rec_weight: float = 0.1
hyponet: HypoNetConfig = field(default_factory=HypoNetConfig)
raft_iter: int = 20
coord_range: List[float] = MISSING
modulated_layer_idxs: Optional[List[int]] = None
@classmethod
def create(cls, config):
# We need to specify the type of the default DataEncoderConfig.
# Otherwise, data_encoder will be initialized & structured as "unfold" type (which is default value)
# hence merging with the config with other type would cause config error.
defaults = OmegaConf.structured(cls(ema=False))
config = OmegaConf.merge(defaults, config)
return config
@@ -0,0 +1,201 @@
Apache License
Version 2.0, January 2004
http://www.apache.org/licenses/
TERMS AND CONDITIONS FOR USE, REPRODUCTION, AND DISTRIBUTION
1. Definitions.
"License" shall mean the terms and conditions for use, reproduction,
and distribution as defined by Sections 1 through 9 of this document.
"Licensor" shall mean the copyright owner or entity authorized by
the copyright owner that is granting the License.
"Legal Entity" shall mean the union of the acting entity and all
other entities that control, are controlled by, or are under common
control with that entity. For the purposes of this definition,
"control" means (i) the power, direct or indirect, to cause the
direction or management of such entity, whether by contract or
otherwise, or (ii) ownership of fifty percent (50%) or more of the
outstanding shares, or (iii) beneficial ownership of such entity.
"You" (or "Your") shall mean an individual or Legal Entity
exercising permissions granted by this License.
"Source" form shall mean the preferred form for making modifications,
including but not limited to software source code, documentation
source, and configuration files.
"Object" form shall mean any form resulting from mechanical
transformation or translation of a Source form, including but
not limited to compiled object code, generated documentation,
and conversions to other media types.
"Work" shall mean the work of authorship, whether in Source or
Object form, made available under the License, as indicated by a
copyright notice that is included in or attached to the work
(an example is provided in the Appendix below).
"Derivative Works" shall mean any work, whether in Source or Object
form, that is based on (or derived from) the Work and for which the
editorial revisions, annotations, elaborations, or other modifications
represent, as a whole, an original work of authorship. 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 and Derivative Works thereof.
"Contribution" shall mean any work of authorship, including
the original version of the Work and any modifications or additions
to that Work or Derivative Works thereof, that is intentionally
submitted to Licensor for inclusion in the Work by the copyright owner
or by an individual or Legal Entity authorized to submit on behalf of
the copyright owner. For the purposes of this definition, "submitted"
means any form of electronic, verbal, or written communication sent
to the Licensor or its representatives, including but not limited to
communication on electronic mailing lists, source code control systems,
and issue tracking systems that are managed by, or on behalf of, the
Licensor for the purpose of discussing and improving the Work, but
excluding communication that is conspicuously marked or otherwise
designated in writing by the copyright owner as "Not a Contribution."
"Contributor" shall mean Licensor and any individual or Legal Entity
on behalf of whom a Contribution has been received by Licensor and
subsequently incorporated within the Work.
2. Grant of Copyright License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
copyright license to reproduce, prepare Derivative Works of,
publicly display, publicly perform, sublicense, and distribute the
Work and such Derivative Works in Source or Object form.
3. Grant of Patent License. Subject to the terms and conditions of
this License, each Contributor hereby grants to You a perpetual,
worldwide, non-exclusive, no-charge, royalty-free, irrevocable
(except as stated in this section) patent license to make, have made,
use, offer to sell, sell, import, and otherwise transfer the Work,
where such license applies only to those patent claims licensable
by such Contributor that are necessarily infringed by their
Contribution(s) alone or by combination of their Contribution(s)
with the Work to which such Contribution(s) was submitted. If You
institute patent litigation against any entity (including a
cross-claim or counterclaim in a lawsuit) alleging that the Work
or a Contribution incorporated within the Work constitutes direct
or contributory patent infringement, then any patent licenses
granted to You under this License for that Work shall terminate
as of the date such litigation is filed.
4. Redistribution. You may reproduce and distribute copies of the
Work or Derivative Works thereof in any medium, with or without
modifications, and in Source or Object form, provided that You
meet the following conditions:
(a) You must give any other recipients of the Work or
Derivative Works a copy of this License; and
(b) You must cause any modified files to carry prominent notices
stating that You changed the files; and
(c) You must retain, in the Source form of any Derivative Works
that You distribute, all copyright, patent, trademark, and
attribution notices from the Source form of the Work,
excluding those notices that do not pertain to any part of
the Derivative Works; and
(d) If the Work includes a "NOTICE" text file as part of its
distribution, then any Derivative Works that You distribute must
include a readable copy of the attribution notices contained
within such NOTICE file, excluding those notices that do not
pertain to any part of the Derivative Works, in at least one
of the following places: within a NOTICE text file distributed
as part of the Derivative Works; within the Source form or
documentation, if provided along with the Derivative Works; or,
within a display generated by the Derivative Works, if and
wherever such third-party notices normally appear. The contents
of the NOTICE file are for informational purposes only and
do not modify the License. You may add Your own attribution
notices within Derivative Works that You distribute, alongside
or as an addendum to the NOTICE text from the Work, provided
that such additional attribution notices cannot be construed
as modifying the License.
You may add Your own copyright statement to Your modifications and
may provide additional or different license terms and conditions
for use, reproduction, or distribution of Your modifications, or
for any such Derivative Works as a whole, provided Your use,
reproduction, and distribution of the Work otherwise complies with
the conditions stated in this License.
5. Submission of Contributions. Unless You explicitly state otherwise,
any Contribution intentionally submitted for inclusion in the Work
by You to the Licensor shall be under the terms and conditions of
this License, without any additional terms or conditions.
Notwithstanding the above, nothing herein shall supersede or modify
the terms of any separate license agreement you may have executed
with Licensor regarding such Contributions.
6. Trademarks. This License does not grant permission to use the trade
names, trademarks, service marks, or product names of the Licensor,
except as required for reasonable and customary use in describing the
origin of the Work and reproducing the content of the NOTICE file.
7. Disclaimer of Warranty. Unless required by applicable law or
agreed to in writing, Licensor provides the Work (and each
Contributor provides its Contributions) on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or
implied, including, without limitation, any warranties or conditions
of TITLE, NON-INFRINGEMENT, MERCHANTABILITY, or FITNESS FOR A
PARTICULAR PURPOSE. You are solely responsible for determining the
appropriateness of using or redistributing the Work and assume any
risks associated with Your exercise of permissions under this License.
8. Limitation of Liability. In no event and under no legal theory,
whether in tort (including negligence), contract, or otherwise,
unless required by applicable law (such as deliberate and grossly
negligent acts) or agreed to in writing, shall any Contributor be
liable to You for damages, including any direct, indirect, special,
incidental, or consequential damages of any character arising as a
result of this License or out of the use or inability to use the
Work (including but not limited to damages for loss of goodwill,
work stoppage, computer failure or malfunction, or any and all
other commercial damages or losses), even if such Contributor
has been advised of the possibility of such damages.
9. Accepting Warranty or Additional Liability. While redistributing
the Work or Derivative Works thereof, You may choose to offer,
and charge a fee for, acceptance of support, warranty, indemnity,
or other liability obligations and/or rights consistent with this
License. However, in accepting such obligations, You may act only
on Your own behalf and on Your sole responsibility, not on behalf
of any other Contributor, and only if You agree to indemnify,
defend, and hold each Contributor harmless for any liability
incurred by, or claims asserted against, such Contributor by reason
of your accepting any such warranty or additional liability.
END OF TERMS AND CONDITIONS
APPENDIX: How to apply the Apache License to your work.
To apply the Apache License to your work, attach the following
boilerplate notice, with the fields enclosed by brackets "[]"
replaced with your own identifying information. (Don't include
the brackets!) The text should be enclosed in the appropriate
comment syntax for the file format. We also recommend that a
file or class name and description of purpose be included on the
same "printed page" as the copyright notice for easier
identification within third-party archives.
Copyright [yyyy] [name of copyright owner]
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
@@ -0,0 +1,164 @@
# FlowFormer: A Transformer Architecture for Optical Flow
### [Project Page](https://drinkingcoder.github.io/publication/flowformer/)
> FlowFormer: A Transformer Architecture for Optical Flow
> [Zhaoyang Huang](https://drinkingcoder.github.io)<sup>\*</sup>, Xiaoyu Shi<sup>\*</sup>, Chao Zhang, Qiang Wang, Ka Chun Cheung, [Hongwei Qin](http://qinhongwei.com/academic/), [Jifeng Dai](https://jifengdai.org/), [Hongsheng Li](https://www.ee.cuhk.edu.hk/~hsli/)
> ECCV 2022
<img src="assets/teaser.png">
## News
Our FlowFormer++ and VideoFlow are accepted by CVPR and ICCV, which ranks 2nd and 1st on the Sintel benchmark!
Please also refer to our [FlowFormer++](https://github.com/XiaoyuShi97/FlowFormerPlusPlus) and [VideoFlow](https://github.com/XiaoyuShi97/VideoFlow).
## TODO List
- [x] Code release (2022-8-1)
- [x] Models release (2022-8-1)
## Data Preparation
Similar to RAFT, to evaluate/train FlowFormer, you will need to download the required datasets.
* [FlyingChairs](https://lmb.informatik.uni-freiburg.de/resources/datasets/FlyingChairs.en.html#flyingchairs)
* [FlyingThings3D](https://lmb.informatik.uni-freiburg.de/resources/datasets/SceneFlowDatasets.en.html)
* [Sintel](http://sintel.is.tue.mpg.de/)
* [KITTI](http://www.cvlibs.net/datasets/kitti/eval_scene_flow.php?benchmark=flow)
* [HD1K](http://hci-benchmark.iwr.uni-heidelberg.de/) (optional)
By default `datasets.py` will search for the datasets in these locations. You can create symbolic links to wherever the datasets were downloaded in the `datasets` folder
```Shell
├── datasets
├── Sintel
├── test
├── training
├── KITTI
├── testing
├── training
├── devkit
├── FlyingChairs_release
├── data
├── FlyingThings3D
├── frames_cleanpass
├── frames_finalpass
├── optical_flow
```
## Requirements
```shell
conda create --name flowformer
conda activate flowformer
conda install pytorch=1.6.0 torchvision=0.7.0 cudatoolkit=10.1 matplotlib tensorboard scipy opencv -c pytorch
pip install yacs loguru einops timm==0.4.12 imageio
```
## Training
The script will load the config according to the training stage. The trained model will be saved in a directory in `logs` and `checkpoints`. For example, the following script will load the config `configs/default.py`. The trained model will be saved as `logs/xxxx/final` and `checkpoints/chairs.pth`.
```shell
python -u train_FlowFormer.py --name chairs --stage chairs --validation chairs
```
To finish the entire training schedule, you can run:
```shell
./run_train.sh
```
## Models
We provide [models](https://drive.google.com/drive/folders/1K2dcWxaqOLiQ3PoqRdokrgWsGIf3yBA_?usp=sharing) trained in the four stages. The default path of the models for evaluation is:
```Shell
├── checkpoints
├── chairs.pth
├── things.pth
├── sintel.pth
├── kitti.pth
├── flowformer-small.pth
├── things_kitti.pth
```
flowformer-small.pth is a small version of our flowformer. things_kitti.pth is the FlowFormer# introduced in our [supplementary](https://drinkingcoder.github.io/publication/flowformer/images/FlowFormer-supp.pdf), used for KITTI training set evaluation.
## Evaluation
The model to be evaluated is assigned by the `_CN.model` in the config file.
Evaluating the model on the Sintel training set and the KITTI training set. The corresponding config file is `configs/things_eval.py`.
```Shell
# with tiling technique
python evaluate_FlowFormer_tile.py --eval sintel_validation
python evaluate_FlowFormer_tile.py --eval kitti_validation --model checkpoints/things_kitti.pth
# without tiling technique
python evaluate_FlowFormer.py --dataset sintel
```
||with tile|w/o tile|
|----|-----|--------|
|clean|0.94|1.01|
|final|2.33|2.40|
Evaluating the small version model. The corresponding config file is `configs/small_things_eval.py`.
```Shell
# with tiling technique
python evaluate_FlowFormer_tile.py --eval sintel_validation --small
# without tiling technique
python evaluate_FlowFormer.py --dataset sintel --small
```
||with tile|w/o tile|
|----|-----|--------|
|clean|1.21|1.32|
|final|2.61|2.68|
Generating the submission for the Sintel and KITTI benchmarks. The corresponding config file is `configs/submission.py`.
```Shell
python evaluate_FlowFormer_tile.py --eval sintel_submission
python evaluate_FlowFormer_tile.py --eval kitti_submission
```
Visualizing the sintel dataset:
```Shell
python visualize_flow.py --eval_type sintel --keep_size
```
Visualizing an image sequence extracted from a video:
```Shell
python visualize_flow.py --eval_type seq
```
The default image sequence format is:
```Shell
├── demo_data
├── mihoyo
├── 000001.png
├── 000002.png
├── 000003.png
.
.
.
├── 001000.png
```
## License
FlowFormer is released under the Apache License
## Citation
```bibtex
@article{huang2022flowformer,
title={{FlowFormer}: A Transformer Architecture for Optical Flow},
author={Huang, Zhaoyang and Shi, Xiaoyu and Zhang, Chao and Wang, Qiang and Cheung, Ka Chun and Qin, Hongwei and Dai, Jifeng and Li, Hongsheng},
journal={{ECCV}},
year={2022}
}
@inproceedings{shi2023flowformer++,
title={Flowformer++: Masked cost volume autoencoding for pretraining optical flow estimation},
author={Shi, Xiaoyu and Huang, Zhaoyang and Li, Dasong and Zhang, Manyuan and Cheung, Ka Chun and See, Simon and Qin, Hongwei and Dai, Jifeng and Li, Hongsheng},
booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition},
pages={1599--1610},
year={2023}
}
@article{huang2023flowformer,
title={FlowFormer: A Transformer Architecture and Its Masked Cost Volume Autoencoding for Optical Flow},
author={Huang, Zhaoyang and Shi, Xiaoyu and Zhang, Chao and Wang, Qiang and Li, Yijin and Qin, Hongwei and Dai, Jifeng and Wang, Xiaogang and Li, Hongsheng},
journal={arXiv preprint arXiv:2306.05442},
year={2023}
}
```
## Acknowledgement
In this project, we use parts of codes in:
- [RAFT](https://github.com/princeton-vl/RAFT)
- [GMA](https://github.com/zacjiang/GMA)
- [timm](https://github.com/rwightman/pytorch-image-models)
@@ -0,0 +1,18 @@
import torch
from .configs.submission import get_cfg
from .core.FlowFormer import build_flowformer
def initialize_Flowformer():
cfg = get_cfg()
model = build_flowformer(cfg)
ckpt = torch.load(cfg.model, map_location="cpu")
def convert(param):
return {k.replace("module.", ""): v for k, v in param.items() if "module" in k}
ckpt = convert(ckpt)
model.load_state_dict(ckpt)
return model
@@ -0,0 +1,54 @@
#include <torch/extension.h>
#include <vector>
// CUDA forward declarations
std::vector<torch::Tensor> corr_cuda_forward(
torch::Tensor fmap1,
torch::Tensor fmap2,
torch::Tensor coords,
int radius);
std::vector<torch::Tensor> corr_cuda_backward(
torch::Tensor fmap1,
torch::Tensor fmap2,
torch::Tensor coords,
torch::Tensor corr_grad,
int radius);
// C++ interface
#define CHECK_CUDA(x) TORCH_CHECK(x.type().is_cuda(), #x " must be a CUDA tensor")
#define CHECK_CONTIGUOUS(x) TORCH_CHECK(x.is_contiguous(), #x " must be contiguous")
#define CHECK_INPUT(x) CHECK_CUDA(x); CHECK_CONTIGUOUS(x)
std::vector<torch::Tensor> corr_forward(
torch::Tensor fmap1,
torch::Tensor fmap2,
torch::Tensor coords,
int radius) {
CHECK_INPUT(fmap1);
CHECK_INPUT(fmap2);
CHECK_INPUT(coords);
return corr_cuda_forward(fmap1, fmap2, coords, radius);
}
std::vector<torch::Tensor> corr_backward(
torch::Tensor fmap1,
torch::Tensor fmap2,
torch::Tensor coords,
torch::Tensor corr_grad,
int radius) {
CHECK_INPUT(fmap1);
CHECK_INPUT(fmap2);
CHECK_INPUT(coords);
CHECK_INPUT(corr_grad);
return corr_cuda_backward(fmap1, fmap2, coords, corr_grad, radius);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("forward", &corr_forward, "CORR forward");
m.def("backward", &corr_backward, "CORR backward");
}
@@ -0,0 +1,324 @@
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>
#include <vector>
#define BLOCK_H 4
#define BLOCK_W 8
#define BLOCK_HW BLOCK_H * BLOCK_W
#define CHANNEL_STRIDE 32
__forceinline__ __device__
bool within_bounds(int h, int w, int H, int W) {
return h >= 0 && h < H && w >= 0 && w < W;
}
template <typename scalar_t>
__global__ void corr_forward_kernel(
const torch::PackedTensorAccessor32<scalar_t,4,torch::RestrictPtrTraits> fmap1,
const torch::PackedTensorAccessor32<scalar_t,4,torch::RestrictPtrTraits> fmap2,
const torch::PackedTensorAccessor32<scalar_t,5,torch::RestrictPtrTraits> coords,
torch::PackedTensorAccessor32<scalar_t,5,torch::RestrictPtrTraits> corr,
int r)
{
const int b = blockIdx.x;
const int h0 = blockIdx.y * blockDim.x;
const int w0 = blockIdx.z * blockDim.y;
const int tid = threadIdx.x * blockDim.y + threadIdx.y;
const int H1 = fmap1.size(1);
const int W1 = fmap1.size(2);
const int H2 = fmap2.size(1);
const int W2 = fmap2.size(2);
const int N = coords.size(1);
const int C = fmap1.size(3);
__shared__ scalar_t f1[CHANNEL_STRIDE][BLOCK_HW+1];
__shared__ scalar_t f2[CHANNEL_STRIDE][BLOCK_HW+1];
__shared__ scalar_t x2s[BLOCK_HW];
__shared__ scalar_t y2s[BLOCK_HW];
for (int c=0; c<C; c+=CHANNEL_STRIDE) {
for (int k=0; k<BLOCK_HW; k+=BLOCK_HW/CHANNEL_STRIDE) {
int k1 = k + tid / CHANNEL_STRIDE;
int h1 = h0 + k1 / BLOCK_W;
int w1 = w0 + k1 % BLOCK_W;
int c1 = tid % CHANNEL_STRIDE;
auto fptr = fmap1[b][h1][w1];
if (within_bounds(h1, w1, H1, W1))
f1[c1][k1] = fptr[c+c1];
else
f1[c1][k1] = 0.0;
}
__syncthreads();
for (int n=0; n<N; n++) {
int h1 = h0 + threadIdx.x;
int w1 = w0 + threadIdx.y;
if (within_bounds(h1, w1, H1, W1)) {
x2s[tid] = coords[b][n][h1][w1][0];
y2s[tid] = coords[b][n][h1][w1][1];
}
scalar_t dx = x2s[tid] - floor(x2s[tid]);
scalar_t dy = y2s[tid] - floor(y2s[tid]);
int rd = 2*r + 1;
for (int iy=0; iy<rd+1; iy++) {
for (int ix=0; ix<rd+1; ix++) {
for (int k=0; k<BLOCK_HW; k+=BLOCK_HW/CHANNEL_STRIDE) {
int k1 = k + tid / CHANNEL_STRIDE;
int h2 = static_cast<int>(floor(y2s[k1]))-r+iy;
int w2 = static_cast<int>(floor(x2s[k1]))-r+ix;
int c2 = tid % CHANNEL_STRIDE;
auto fptr = fmap2[b][h2][w2];
if (within_bounds(h2, w2, H2, W2))
f2[c2][k1] = fptr[c+c2];
else
f2[c2][k1] = 0.0;
}
__syncthreads();
scalar_t s = 0.0;
for (int k=0; k<CHANNEL_STRIDE; k++)
s += f1[k][tid] * f2[k][tid];
int ix_nw = H1*W1*((iy-1) + rd*(ix-1));
int ix_ne = H1*W1*((iy-1) + rd*ix);
int ix_sw = H1*W1*(iy + rd*(ix-1));
int ix_se = H1*W1*(iy + rd*ix);
scalar_t nw = s * (dy) * (dx);
scalar_t ne = s * (dy) * (1-dx);
scalar_t sw = s * (1-dy) * (dx);
scalar_t se = s * (1-dy) * (1-dx);
scalar_t* corr_ptr = &corr[b][n][0][h1][w1];
if (iy > 0 && ix > 0 && within_bounds(h1, w1, H1, W1))
*(corr_ptr + ix_nw) += nw;
if (iy > 0 && ix < rd && within_bounds(h1, w1, H1, W1))
*(corr_ptr + ix_ne) += ne;
if (iy < rd && ix > 0 && within_bounds(h1, w1, H1, W1))
*(corr_ptr + ix_sw) += sw;
if (iy < rd && ix < rd && within_bounds(h1, w1, H1, W1))
*(corr_ptr + ix_se) += se;
}
}
}
}
}
template <typename scalar_t>
__global__ void corr_backward_kernel(
const torch::PackedTensorAccessor32<scalar_t,4,torch::RestrictPtrTraits> fmap1,
const torch::PackedTensorAccessor32<scalar_t,4,torch::RestrictPtrTraits> fmap2,
const torch::PackedTensorAccessor32<scalar_t,5,torch::RestrictPtrTraits> coords,
const torch::PackedTensorAccessor32<scalar_t,5,torch::RestrictPtrTraits> corr_grad,
torch::PackedTensorAccessor32<scalar_t,4,torch::RestrictPtrTraits> fmap1_grad,
torch::PackedTensorAccessor32<scalar_t,4,torch::RestrictPtrTraits> fmap2_grad,
torch::PackedTensorAccessor32<scalar_t,5,torch::RestrictPtrTraits> coords_grad,
int r)
{
const int b = blockIdx.x;
const int h0 = blockIdx.y * blockDim.x;
const int w0 = blockIdx.z * blockDim.y;
const int tid = threadIdx.x * blockDim.y + threadIdx.y;
const int H1 = fmap1.size(1);
const int W1 = fmap1.size(2);
const int H2 = fmap2.size(1);
const int W2 = fmap2.size(2);
const int N = coords.size(1);
const int C = fmap1.size(3);
__shared__ scalar_t f1[CHANNEL_STRIDE][BLOCK_HW+1];
__shared__ scalar_t f2[CHANNEL_STRIDE][BLOCK_HW+1];
__shared__ scalar_t f1_grad[CHANNEL_STRIDE][BLOCK_HW+1];
__shared__ scalar_t f2_grad[CHANNEL_STRIDE][BLOCK_HW+1];
__shared__ scalar_t x2s[BLOCK_HW];
__shared__ scalar_t y2s[BLOCK_HW];
for (int c=0; c<C; c+=CHANNEL_STRIDE) {
for (int k=0; k<BLOCK_HW; k+=BLOCK_HW/CHANNEL_STRIDE) {
int k1 = k + tid / CHANNEL_STRIDE;
int h1 = h0 + k1 / BLOCK_W;
int w1 = w0 + k1 % BLOCK_W;
int c1 = tid % CHANNEL_STRIDE;
auto fptr = fmap1[b][h1][w1];
if (within_bounds(h1, w1, H1, W1))
f1[c1][k1] = fptr[c+c1];
else
f1[c1][k1] = 0.0;
f1_grad[c1][k1] = 0.0;
}
__syncthreads();
int h1 = h0 + threadIdx.x;
int w1 = w0 + threadIdx.y;
for (int n=0; n<N; n++) {
x2s[tid] = coords[b][n][h1][w1][0];
y2s[tid] = coords[b][n][h1][w1][1];
scalar_t dx = x2s[tid] - floor(x2s[tid]);
scalar_t dy = y2s[tid] - floor(y2s[tid]);
int rd = 2*r + 1;
for (int iy=0; iy<rd+1; iy++) {
for (int ix=0; ix<rd+1; ix++) {
for (int k=0; k<BLOCK_HW; k+=BLOCK_HW/CHANNEL_STRIDE) {
int k1 = k + tid / CHANNEL_STRIDE;
int h2 = static_cast<int>(floor(y2s[k1]))-r+iy;
int w2 = static_cast<int>(floor(x2s[k1]))-r+ix;
int c2 = tid % CHANNEL_STRIDE;
auto fptr = fmap2[b][h2][w2];
if (within_bounds(h2, w2, H2, W2))
f2[c2][k1] = fptr[c+c2];
else
f2[c2][k1] = 0.0;
f2_grad[c2][k1] = 0.0;
}
__syncthreads();
const scalar_t* grad_ptr = &corr_grad[b][n][0][h1][w1];
scalar_t g = 0.0;
int ix_nw = H1*W1*((iy-1) + rd*(ix-1));
int ix_ne = H1*W1*((iy-1) + rd*ix);
int ix_sw = H1*W1*(iy + rd*(ix-1));
int ix_se = H1*W1*(iy + rd*ix);
if (iy > 0 && ix > 0 && within_bounds(h1, w1, H1, W1))
g += *(grad_ptr + ix_nw) * dy * dx;
if (iy > 0 && ix < rd && within_bounds(h1, w1, H1, W1))
g += *(grad_ptr + ix_ne) * dy * (1-dx);
if (iy < rd && ix > 0 && within_bounds(h1, w1, H1, W1))
g += *(grad_ptr + ix_sw) * (1-dy) * dx;
if (iy < rd && ix < rd && within_bounds(h1, w1, H1, W1))
g += *(grad_ptr + ix_se) * (1-dy) * (1-dx);
for (int k=0; k<CHANNEL_STRIDE; k++) {
f1_grad[k][tid] += g * f2[k][tid];
f2_grad[k][tid] += g * f1[k][tid];
}
for (int k=0; k<BLOCK_HW; k+=BLOCK_HW/CHANNEL_STRIDE) {
int k1 = k + tid / CHANNEL_STRIDE;
int h2 = static_cast<int>(floor(y2s[k1]))-r+iy;
int w2 = static_cast<int>(floor(x2s[k1]))-r+ix;
int c2 = tid % CHANNEL_STRIDE;
scalar_t* fptr = &fmap2_grad[b][h2][w2][0];
if (within_bounds(h2, w2, H2, W2))
atomicAdd(fptr+c+c2, f2_grad[c2][k1]);
}
}
}
}
__syncthreads();
for (int k=0; k<BLOCK_HW; k+=BLOCK_HW/CHANNEL_STRIDE) {
int k1 = k + tid / CHANNEL_STRIDE;
int h1 = h0 + k1 / BLOCK_W;
int w1 = w0 + k1 % BLOCK_W;
int c1 = tid % CHANNEL_STRIDE;
scalar_t* fptr = &fmap1_grad[b][h1][w1][0];
if (within_bounds(h1, w1, H1, W1))
fptr[c+c1] += f1_grad[c1][k1];
}
}
}
std::vector<torch::Tensor> corr_cuda_forward(
torch::Tensor fmap1,
torch::Tensor fmap2,
torch::Tensor coords,
int radius)
{
const auto B = coords.size(0);
const auto N = coords.size(1);
const auto H = coords.size(2);
const auto W = coords.size(3);
const auto rd = 2 * radius + 1;
auto opts = fmap1.options();
auto corr = torch::zeros({B, N, rd*rd, H, W}, opts);
const dim3 blocks(B, (H+BLOCK_H-1)/BLOCK_H, (W+BLOCK_W-1)/BLOCK_W);
const dim3 threads(BLOCK_H, BLOCK_W);
corr_forward_kernel<float><<<blocks, threads>>>(
fmap1.packed_accessor32<float,4,torch::RestrictPtrTraits>(),
fmap2.packed_accessor32<float,4,torch::RestrictPtrTraits>(),
coords.packed_accessor32<float,5,torch::RestrictPtrTraits>(),
corr.packed_accessor32<float,5,torch::RestrictPtrTraits>(),
radius);
return {corr};
}
std::vector<torch::Tensor> corr_cuda_backward(
torch::Tensor fmap1,
torch::Tensor fmap2,
torch::Tensor coords,
torch::Tensor corr_grad,
int radius)
{
const auto B = coords.size(0);
const auto N = coords.size(1);
const auto H1 = fmap1.size(1);
const auto W1 = fmap1.size(2);
const auto H2 = fmap2.size(1);
const auto W2 = fmap2.size(2);
const auto C = fmap1.size(3);
auto opts = fmap1.options();
auto fmap1_grad = torch::zeros({B, H1, W1, C}, opts);
auto fmap2_grad = torch::zeros({B, H2, W2, C}, opts);
auto coords_grad = torch::zeros({B, N, H1, W1, 2}, opts);
const dim3 blocks(B, (H1+BLOCK_H-1)/BLOCK_H, (W1+BLOCK_W-1)/BLOCK_W);
const dim3 threads(BLOCK_H, BLOCK_W);
corr_backward_kernel<float><<<blocks, threads>>>(
fmap1.packed_accessor32<float,4,torch::RestrictPtrTraits>(),
fmap2.packed_accessor32<float,4,torch::RestrictPtrTraits>(),
coords.packed_accessor32<float,5,torch::RestrictPtrTraits>(),
corr_grad.packed_accessor32<float,5,torch::RestrictPtrTraits>(),
fmap1_grad.packed_accessor32<float,4,torch::RestrictPtrTraits>(),
fmap2_grad.packed_accessor32<float,4,torch::RestrictPtrTraits>(),
coords_grad.packed_accessor32<float,5,torch::RestrictPtrTraits>(),
radius);
return {fmap1_grad, fmap2_grad, coords_grad};
}
@@ -0,0 +1,15 @@
from setuptools import setup
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
setup(
name="correlation",
ext_modules=[
CUDAExtension(
"alt_cuda_corr",
sources=["correlation.cpp", "correlation_kernel.cu"],
extra_compile_args={"cxx": [], "nvcc": ["-O3"]},
),
],
cmdclass={"build_ext": BuildExtension},
)
Binary file not shown.

After

Width:  |  Height:  |  Size: 845 KiB

File diff suppressed because it is too large Load Diff
@@ -0,0 +1,78 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = "default"
_CN.suffix = "arxiv2"
_CN.gamma = 0.8
_CN.max_flow = 400
_CN.batch_size = 8
_CN.sum_freq = 100
_CN.val_freq = 5000
_CN.image_size = [368, 496]
_CN.add_noise = True
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.restore_ckpt = None
###########################################
# latentcostformer
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.predictor_dim = 128
_CN.latentcostformer.motion_feature_dim = 209 # use concat, so double query_latent_dim
_CN.latentcostformer.arc_type = "transformer"
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.gma = True
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 12
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
##########################################
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 25e-5
_CN.trainer.adamw_decay = 1e-4
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,83 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = "kitti"
_CN.suffix = "kitti"
_CN.gamma = 0.85
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 499999999
_CN.image_size = [432, 960]
_CN.add_noise = True
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.model = None
_CN.restore_ckpt = "checkpoints/sintel.pth"
# latentcostformer
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.predictor_dim = 128
_CN.latentcostformer.motion_feature_dim = 209 # use concat, so double query_latent_dim
_CN.latentcostformer.arc_type = "transformer"
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.vertical_encoder_attn = "twins"
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.pwc_aug = False
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.no_sc = False
_CN.latentcostformer.use_rpe = False
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 12
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-5
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 50000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,77 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = "default"
_CN.suffix = "sintel"
_CN.gamma = 0.85
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 5000000
_CN.image_size = [432, 960]
_CN.add_noise = True
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.restore_ckpt = "checkpoints/things.pth"
# latentcostformer
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.arc_type = "transformer"
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.no_pe = False
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.kernel_size = 9
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.no_sc = False
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 12
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-5
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,77 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = ""
_CN.suffix = ""
_CN.gamma = 0.8
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 5000000
_CN.image_size = [432, 960]
_CN.add_noise = False
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.model = "checkpoints/flowformer-small/things.pth"
# latentcostformer
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 4
_CN.latentcostformer.cost_latent_dim = 32
_CN.latentcostformer.arc_type = "transformer"
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 1
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.no_pe = False
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.kernel_size = 9
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 0
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "basicencoder"
_CN.latentcostformer.fnet = "basicencoder"
_CN.latentcostformer.no_sc = False
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 32
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-4
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,77 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = ""
_CN.suffix = ""
_CN.gamma = 0.8
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 5000000
_CN.image_size = [432, 960]
_CN.add_noise = False
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.model = "pretrained_ckpt/flowformer_sintel.pth"
# latentcostformer
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.arc_type = "transformer"
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.no_pe = False
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.kernel_size = 9
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.no_sc = False
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 32
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-4
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,76 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = ""
_CN.suffix = ""
_CN.gamma = 0.8
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 5000000
_CN.image_size = [432, 960]
_CN.add_noise = True
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.restore_ckpt = "checkpoints/chairs.pth"
#######################################
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.nat_rep = "abs"
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.no_pe = False
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.kernel_size = 9
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 12
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-4
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,77 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = ""
_CN.suffix = ""
_CN.gamma = 0.8
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 5000000
_CN.image_size = [432, 960]
_CN.add_noise = False
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.model = "checkpoints/things.pth"
# latentcostformer
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.arc_type = "transformer"
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.no_pe = False
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.kernel_size = 9
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.no_sc = False
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 32
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-4
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,76 @@
from yacs.config import CfgNode as CN
_CN = CN()
_CN.name = ""
_CN.suffix = ""
_CN.gamma = 0.8
_CN.max_flow = 400
_CN.batch_size = 6
_CN.sum_freq = 100
_CN.val_freq = 5000000
_CN.image_size = [400, 720]
_CN.add_noise = True
_CN.critical_params = []
_CN.transformer = "latentcostformer"
_CN.restore_ckpt = "checkpoints/chairs.pth"
#######################################
_CN.latentcostformer = CN()
_CN.latentcostformer.pe = "linear"
_CN.latentcostformer.dropout = 0.0
_CN.latentcostformer.encoder_latent_dim = 256 # in twins, this is 256
_CN.latentcostformer.query_latent_dim = 64
_CN.latentcostformer.cost_latent_input_dim = 64
_CN.latentcostformer.cost_latent_token_num = 8
_CN.latentcostformer.cost_latent_dim = 128
_CN.latentcostformer.cost_heads_num = 1
# encoder
_CN.latentcostformer.pretrain = True
_CN.latentcostformer.context_concat = False
_CN.latentcostformer.encoder_depth = 3
_CN.latentcostformer.feat_cross_attn = False
_CN.latentcostformer.nat_rep = "abs"
_CN.latentcostformer.patch_size = 8
_CN.latentcostformer.patch_embed = "single"
_CN.latentcostformer.no_pe = False
_CN.latentcostformer.gma = "GMA"
_CN.latentcostformer.kernel_size = 9
_CN.latentcostformer.rm_res = True
_CN.latentcostformer.vert_c_dim = 64
_CN.latentcostformer.cost_encoder_res = True
_CN.latentcostformer.cnet = "twins"
_CN.latentcostformer.fnet = "twins"
_CN.latentcostformer.only_global = False
_CN.latentcostformer.add_flow_token = True
_CN.latentcostformer.use_mlp = False
_CN.latentcostformer.vertical_conv = False
# decoder
_CN.latentcostformer.decoder_depth = 12
_CN.latentcostformer.critical_params = [
"cost_heads_num",
"vert_c_dim",
"cnet",
"pretrain",
"add_flow_token",
"encoder_depth",
"gma",
"cost_encoder_res",
]
### TRAINER
_CN.trainer = CN()
_CN.trainer.scheduler = "OneCycleLR"
_CN.trainer.optimizer = "adamw"
_CN.trainer.canonical_lr = 12.5e-5
_CN.trainer.adamw_decay = 1e-4
_CN.trainer.clip = 1.0
_CN.trainer.num_steps = 120000
_CN.trainer.epsilon = 1e-8
_CN.trainer.anneal_strategy = "linear"
def get_cfg():
return _CN.clone()
@@ -0,0 +1,197 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import einsum
from einops.layers.torch import Rearrange
from einops import rearrange
class BroadMultiHeadAttention(nn.Module):
def __init__(self, dim, heads):
super(BroadMultiHeadAttention, self).__init__()
self.dim = dim
self.heads = heads
self.scale = (dim / heads) ** -0.5
self.attend = nn.Softmax(dim=-1)
def attend_with_rpe(self, Q, K):
Q = rearrange(Q.squeeze(), "i (heads d) -> heads i d", heads=self.heads)
K = rearrange(K, "b j (heads d) -> b heads j d", heads=self.heads)
dots = einsum("hid, bhjd -> bhij", Q, K) * self.scale # (b hw) heads 1 pointnum
return self.attend(dots)
def forward(self, Q, K, V):
attn = self.attend_with_rpe(Q, K)
B, _, _ = K.shape
_, N, _ = Q.shape
V = rearrange(V, "b j (heads d) -> b heads j d", heads=self.heads)
out = einsum("bhij, bhjd -> bhid", attn, V)
out = rearrange(out, "b heads n d -> b n (heads d)", b=B, n=N)
return out
class MultiHeadAttention(nn.Module):
def __init__(self, dim, heads):
super(MultiHeadAttention, self).__init__()
self.dim = dim
self.heads = heads
self.scale = (dim / heads) ** -0.5
self.attend = nn.Softmax(dim=-1)
def attend_with_rpe(self, Q, K):
Q = rearrange(Q, "b i (heads d) -> b heads i d", heads=self.heads)
K = rearrange(K, "b j (heads d) -> b heads j d", heads=self.heads)
dots = (
einsum("bhid, bhjd -> bhij", Q, K) * self.scale
) # (b hw) heads 1 pointnum
return self.attend(dots)
def forward(self, Q, K, V):
attn = self.attend_with_rpe(Q, K)
B, HW, _ = Q.shape
V = rearrange(V, "b j (heads d) -> b heads j d", heads=self.heads)
out = einsum("bhij, bhjd -> bhid", attn, V)
out = rearrange(out, "b heads hw d -> b hw (heads d)", b=B, hw=HW)
return out
# class MultiHeadAttentionRelative_encoder(nn.Module):
# def __init__(self, dim, heads):
# super(MultiHeadAttentionRelative, self).__init__()
# self.dim = dim
# self.heads = heads
# self.scale = (dim/heads) ** -0.5
# self.attend = nn.Softmax(dim=-1)
# def attend_with_rpe(self, Q, K, Q_r, K_r):
# """
# Q: [BH1W1, H3W3, dim]
# K: [BH1W1, H3W3, dim]
# Q_r: [BH1W1, H3W3, H3W3, dim]
# K_r: [BH1W1, H3W3, H3W3, dim]
# """
# Q = rearrange(Q, 'b i (heads d) -> b heads i d', heads=self.heads) # [BH1W1, heads, H3W3, dim]
# K = rearrange(K, 'b j (heads d) -> b heads j d', heads=self.heads) # [BH1W1, heads, H3W3, dim]
# K_r = rearrange(K_r, 'b j (heads d) -> b heads j d', heads=self.heads) # [BH1W1, heads, H3W3, dim]
# Q_r = rearrange(Q_r, 'b j (heads d) -> b heads j d', heads=self.heads) # [BH1W1, heads, H3W3, dim]
# # context-context similarity
# c_c = einsum('bhid, bhjd -> bhij', Q, K) * self.scale # [(B H1W1) heads H3W3 H3W3]
# # context-position similarity
# c_p = einsum('bhid, bhjd -> bhij', Q, K_r) * self.scale # [(B H1W1) heads 1 H3W3]
# # position-context similarity
# p_c = einsum('bhijd, bhikd -> bhijk', Q_r[:,:,:,None,:], K[:,:,:,None,:])
# p_c = torch.squeeze(p_c, dim=4)
# p_c = p_c.permute(0, 1, 3, 2)
# dots = c_c + c_p + p_c
# return self.attend(dots)
# def forward(self, Q, K, V, Q_r, K_r):
# attn = self.attend_with_rpe(Q, K, Q_r, K_r)
# B, HW, _ = Q.shape
# V = rearrange(V, 'b j (heads d) -> b heads j d', heads=self.heads)
# out = einsum('bhij, bhjd -> bhid', attn, V)
# out = rearrange(out, 'b heads hw d -> b hw (heads d)', b=B, hw=HW)
# return out
class MultiHeadAttentionRelative(nn.Module):
def __init__(self, dim, heads):
super(MultiHeadAttentionRelative, self).__init__()
self.dim = dim
self.heads = heads
self.scale = (dim / heads) ** -0.5
self.attend = nn.Softmax(dim=-1)
def attend_with_rpe(self, Q, K, Q_r, K_r):
"""
Q: [BH1W1, 1, dim]
K: [BH1W1, H3W3, dim]
Q_r: [BH1W1, H3W3, dim]
K_r: [BH1W1, H3W3, dim]
"""
Q = rearrange(
Q, "b i (heads d) -> b heads i d", heads=self.heads
) # [BH1W1, heads, 1, dim]
K = rearrange(
K, "b j (heads d) -> b heads j d", heads=self.heads
) # [BH1W1, heads, H3W3, dim]
K_r = rearrange(
K_r, "b j (heads d) -> b heads j d", heads=self.heads
) # [BH1W1, heads, H3W3, dim]
Q_r = rearrange(
Q_r, "b j (heads d) -> b heads j d", heads=self.heads
) # [BH1W1, heads, H3W3, dim]
# context-context similarity
c_c = einsum("bhid, bhjd -> bhij", Q, K) * self.scale # [(B H1W1) heads 1 H3W3]
# context-position similarity
c_p = (
einsum("bhid, bhjd -> bhij", Q, K_r) * self.scale
) # [(B H1W1) heads 1 H3W3]
# position-context similarity
p_c = (
einsum("bhijd, bhikd -> bhijk", Q_r[:, :, :, None, :], K[:, :, :, None, :])
* self.scale
)
p_c = torch.squeeze(p_c, dim=4)
p_c = p_c.permute(0, 1, 3, 2)
dots = c_c + c_p + p_c
return self.attend(dots)
def forward(self, Q, K, V, Q_r, K_r):
attn = self.attend_with_rpe(Q, K, Q_r, K_r)
B, HW, _ = Q.shape
V = rearrange(V, "b j (heads d) -> b heads j d", heads=self.heads)
out = einsum("bhij, bhjd -> bhid", attn, V)
out = rearrange(out, "b heads hw d -> b hw (heads d)", b=B, hw=HW)
return out
def LinearPositionEmbeddingSine(x, dim=128, NORMALIZE_FACOR=1 / 200):
# 200 should be enough for a 8x downsampled image
# assume x to be [_, _, 2]
freq_bands = torch.linspace(0, dim // 4 - 1, dim // 4).to(x.device)
return torch.cat(
[
torch.sin(3.14 * x[..., -2:-1] * freq_bands * NORMALIZE_FACOR),
torch.cos(3.14 * x[..., -2:-1] * freq_bands * NORMALIZE_FACOR),
torch.sin(3.14 * x[..., -1:] * freq_bands * NORMALIZE_FACOR),
torch.cos(3.14 * x[..., -1:] * freq_bands * NORMALIZE_FACOR),
],
dim=-1,
)
def ExpPositionEmbeddingSine(x, dim=128, NORMALIZE_FACOR=1 / 200):
# 200 should be enough for a 8x downsampled image
# assume x to be [_, _, 2]
freq_bands = torch.linspace(0, dim // 4 - 1, dim // 4).to(x.device)
return torch.cat(
[
torch.sin(x[..., -2:-1] * (NORMALIZE_FACOR * 2**freq_bands)),
torch.cos(x[..., -2:-1] * (NORMALIZE_FACOR * 2**freq_bands)),
torch.sin(x[..., -1:] * (NORMALIZE_FACOR * 2**freq_bands)),
torch.cos(x[..., -1:] * (NORMALIZE_FACOR * 2**freq_bands)),
],
dim=-1,
)
@@ -0,0 +1,649 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from timm.models.layers import Mlp, DropPath, to_2tuple, trunc_normal_
import math
import numpy as np
class ResidualBlock(nn.Module):
def __init__(self, in_planes, planes, norm_fn="group", stride=1):
super(ResidualBlock, self).__init__()
self.conv1 = nn.Conv2d(
in_planes, planes, kernel_size=3, padding=1, stride=stride
)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1)
self.relu = nn.ReLU(inplace=True)
num_groups = planes // 8
if norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
if not stride == 1:
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
elif norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(planes)
self.norm2 = nn.BatchNorm2d(planes)
if not stride == 1:
self.norm3 = nn.BatchNorm2d(planes)
elif norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(planes)
self.norm2 = nn.InstanceNorm2d(planes)
if not stride == 1:
self.norm3 = nn.InstanceNorm2d(planes)
elif norm_fn == "none":
self.norm1 = nn.Sequential()
self.norm2 = nn.Sequential()
if not stride == 1:
self.norm3 = nn.Sequential()
if stride == 1:
self.downsample = None
else:
self.downsample = nn.Sequential(
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm3
)
def forward(self, x):
y = x
y = self.relu(self.norm1(self.conv1(y)))
y = self.relu(self.norm2(self.conv2(y)))
if self.downsample is not None:
x = self.downsample(x)
return self.relu(x + y)
class BottleneckBlock(nn.Module):
def __init__(self, in_planes, planes, norm_fn="group", stride=1):
super(BottleneckBlock, self).__init__()
self.conv1 = nn.Conv2d(in_planes, planes // 4, kernel_size=1, padding=0)
self.conv2 = nn.Conv2d(
planes // 4, planes // 4, kernel_size=3, padding=1, stride=stride
)
self.conv3 = nn.Conv2d(planes // 4, planes, kernel_size=1, padding=0)
self.relu = nn.ReLU(inplace=True)
num_groups = planes // 8
if norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // 4)
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // 4)
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
if not stride == 1:
self.norm4 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
elif norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(planes // 4)
self.norm2 = nn.BatchNorm2d(planes // 4)
self.norm3 = nn.BatchNorm2d(planes)
if not stride == 1:
self.norm4 = nn.BatchNorm2d(planes)
elif norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(planes // 4)
self.norm2 = nn.InstanceNorm2d(planes // 4)
self.norm3 = nn.InstanceNorm2d(planes)
if not stride == 1:
self.norm4 = nn.InstanceNorm2d(planes)
elif norm_fn == "none":
self.norm1 = nn.Sequential()
self.norm2 = nn.Sequential()
self.norm3 = nn.Sequential()
if not stride == 1:
self.norm4 = nn.Sequential()
if stride == 1:
self.downsample = None
else:
self.downsample = nn.Sequential(
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm4
)
def forward(self, x):
y = x
y = self.relu(self.norm1(self.conv1(y)))
y = self.relu(self.norm2(self.conv2(y)))
y = self.relu(self.norm3(self.conv3(y)))
if self.downsample is not None:
x = self.downsample(x)
return self.relu(x + y)
class BasicEncoder(nn.Module):
def __init__(self, input_dim=3, output_dim=128, norm_fn="batch", dropout=0.0):
super(BasicEncoder, self).__init__()
self.norm_fn = norm_fn
mul = input_dim // 3
if self.norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=64 * mul)
elif self.norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(64 * mul)
elif self.norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(64 * mul)
elif self.norm_fn == "none":
self.norm1 = nn.Sequential()
self.conv1 = nn.Conv2d(input_dim, 64 * mul, kernel_size=7, stride=2, padding=3)
self.relu1 = nn.ReLU(inplace=True)
self.in_planes = 64 * mul
self.layer1 = self._make_layer(64 * mul, stride=1)
self.layer2 = self._make_layer(96 * mul, stride=2)
self.layer3 = self._make_layer(128 * mul, stride=2)
# output convolution
self.conv2 = nn.Conv2d(128 * mul, output_dim, kernel_size=1)
self.dropout = None
if dropout > 0:
self.dropout = nn.Dropout2d(p=dropout)
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
if m.weight is not None:
nn.init.constant_(m.weight, 1)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def _make_layer(self, dim, stride=1):
layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride)
layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1)
layers = (layer1, layer2)
self.in_planes = dim
return nn.Sequential(*layers)
def compute_params(self):
num = 0
for param in self.parameters():
num += np.prod(param.size())
return num
def forward(self, x):
# if input is list, combine batch dimension
is_list = isinstance(x, tuple) or isinstance(x, list)
if is_list:
batch_dim = x[0].shape[0]
x = torch.cat(x, dim=0)
x = self.conv1(x)
x = self.norm1(x)
x = self.relu1(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.conv2(x)
if self.training and self.dropout is not None:
x = self.dropout(x)
if is_list:
x = torch.split(x, [batch_dim, batch_dim], dim=0)
return x
class SmallEncoder(nn.Module):
def __init__(self, output_dim=128, norm_fn="batch", dropout=0.0):
super(SmallEncoder, self).__init__()
self.norm_fn = norm_fn
if self.norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=32)
elif self.norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(32)
elif self.norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(32)
elif self.norm_fn == "none":
self.norm1 = nn.Sequential()
self.conv1 = nn.Conv2d(3, 32, kernel_size=7, stride=2, padding=3)
self.relu1 = nn.ReLU(inplace=True)
self.in_planes = 32
self.layer1 = self._make_layer(32, stride=1)
self.layer2 = self._make_layer(64, stride=2)
self.layer3 = self._make_layer(96, stride=2)
self.dropout = None
if dropout > 0:
self.dropout = nn.Dropout2d(p=dropout)
self.conv2 = nn.Conv2d(96, output_dim, kernel_size=1)
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
if m.weight is not None:
nn.init.constant_(m.weight, 1)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def _make_layer(self, dim, stride=1):
layer1 = BottleneckBlock(self.in_planes, dim, self.norm_fn, stride=stride)
layer2 = BottleneckBlock(dim, dim, self.norm_fn, stride=1)
layers = (layer1, layer2)
self.in_planes = dim
return nn.Sequential(*layers)
def forward(self, x):
# if input is list, combine batch dimension
is_list = isinstance(x, tuple) or isinstance(x, list)
if is_list:
batch_dim = x[0].shape[0]
x = torch.cat(x, dim=0)
x = self.conv1(x)
x = self.norm1(x)
x = self.relu1(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.conv2(x)
if self.training and self.dropout is not None:
x = self.dropout(x)
if is_list:
x = torch.split(x, [batch_dim, batch_dim], dim=0)
return x
class ConvNets(nn.Module):
def __init__(self, in_dim, out_dim, inter_dim, depth, stride=1):
super(ConvNets, self).__init__()
self.conv_first = nn.Conv2d(
in_dim, inter_dim, kernel_size=3, padding=1, stride=stride
)
self.conv_last = nn.Conv2d(
inter_dim, out_dim, kernel_size=3, padding=1, stride=stride
)
self.relu = nn.ReLU(inplace=True)
self.inter_convs = nn.ModuleList(
[
ResidualBlock(inter_dim, inter_dim, norm_fn="none", stride=1)
for i in range(depth)
]
)
def forward(self, x):
x = self.relu(self.conv_first(x))
for inter_conv in self.inter_convs:
x = inter_conv(x)
x = self.conv_last(x)
return x
class FlowHead(nn.Module):
def __init__(self, input_dim=128, hidden_dim=256):
super(FlowHead, self).__init__()
self.conv1 = nn.Conv2d(input_dim, hidden_dim, 3, padding=1)
self.conv2 = nn.Conv2d(hidden_dim, 2, 3, padding=1)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
return self.conv2(self.relu(self.conv1(x)))
class ConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(ConvGRU, self).__init__()
self.convz = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convr = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convq = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
def forward(self, h, x):
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz(hx))
r = torch.sigmoid(self.convr(hx))
q = torch.tanh(self.convq(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class SepConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(SepConvGRU, self).__init__()
self.convz1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convr1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convq1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convz2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convr2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convq2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
def forward(self, h, x):
# horizontal
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz1(hx))
r = torch.sigmoid(self.convr1(hx))
q = torch.tanh(self.convq1(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
# vertical
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz2(hx))
r = torch.sigmoid(self.convr2(hx))
q = torch.tanh(self.convq2(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class BasicMotionEncoder(nn.Module):
def __init__(self, args):
super(BasicMotionEncoder, self).__init__()
cor_planes = args.motion_feature_dim
self.convc1 = nn.Conv2d(cor_planes, 256, 1, padding=0)
self.convc2 = nn.Conv2d(256, 192, 3, padding=1)
self.convf1 = nn.Conv2d(2, 128, 7, padding=3)
self.convf2 = nn.Conv2d(128, 64, 3, padding=1)
self.conv = nn.Conv2d(64 + 192, 128 - 2, 3, padding=1)
def forward(self, flow, corr):
cor = F.relu(self.convc1(corr))
cor = F.relu(self.convc2(cor))
flo = F.relu(self.convf1(flow))
flo = F.relu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
out = F.relu(self.conv(cor_flo))
return torch.cat([out, flow], dim=1)
class BasicFuseMotion(nn.Module):
def __init__(self, args):
super(BasicFuseMotion, self).__init__()
cor_planes = args.motion_feature_dim
out_planes = args.query_latent_dim
self.normf1 = nn.InstanceNorm2d(128)
self.normf2 = nn.InstanceNorm2d(128)
self.convf1 = nn.Conv2d(2, 128, 3, padding=1)
self.convf2 = nn.Conv2d(128, 128, 3, padding=1)
self.convf3 = nn.Conv2d(128, 64, 3, padding=1)
s = 1
self.normc1 = nn.InstanceNorm2d(256 * s)
self.normc2 = nn.InstanceNorm2d(256 * s)
self.normc3 = nn.InstanceNorm2d(256 * s)
self.convc1 = nn.Conv2d(cor_planes + 128, 256 * s, 1, padding=0)
self.convc2 = nn.Conv2d(256 * s, 256 * s, 3, padding=1)
self.convc3 = nn.Conv2d(256 * s, 256 * s, 3, padding=1)
self.convc4 = nn.Conv2d(256 * s, 256 * s, 3, padding=1)
self.conv = nn.Conv2d(256 * s + 64, out_planes, 1, padding=0)
def forward(self, flow, feat, context1=None):
flo = F.relu(self.normf1(self.convf1(flow)))
flo = F.relu(self.normf2(self.convf2(flo)))
flo = self.convf3(flo)
feat = torch.cat([feat, context1], dim=1)
feat = F.relu(self.normc1(self.convc1(feat)))
feat = F.relu(self.normc2(self.convc2(feat)))
feat = F.relu(self.normc3(self.convc3(feat)))
feat = self.convc4(feat)
feat = torch.cat([flo, feat], dim=1)
feat = F.relu(self.conv(feat))
return feat
class BasicUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=128, input_dim=128):
super(BasicUpdateBlock, self).__init__()
self.args = args
self.encoder = BasicMotionEncoder(args)
self.gru = SepConvGRU(hidden_dim=hidden_dim, input_dim=128 + hidden_dim)
self.flow_head = FlowHead(hidden_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
def forward(self, net, inp, corr, flow, upsample=True):
motion_features = self.encoder(flow, corr)
inp = torch.cat([inp, motion_features], dim=1)
net = self.gru(net, inp)
delta_flow = self.flow_head(net)
# scale mask to balence gradients
mask = 0.25 * self.mask(net)
return net, mask, delta_flow
class DirectMeanMaskPredictor(nn.Module):
def __init__(self, args):
super(DirectMeanMaskPredictor, self).__init__()
self.flow_head = FlowHead(args.predictor_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(args.predictor_dim, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
def forward(self, motion_features):
delta_flow = self.flow_head(motion_features)
mask = 0.25 * self.mask(motion_features)
return mask, delta_flow
class BaiscMeanPredictor(nn.Module):
def __init__(self, args, hidden_dim=128):
super(BaiscMeanPredictor, self).__init__()
self.args = args
self.encoder = BasicMotionEncoder(args)
self.flow_head = FlowHead(hidden_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
def forward(self, latent, flow):
motion_features = self.encoder(flow, latent)
delta_flow = self.flow_head(motion_features)
mask = 0.25 * self.mask(motion_features)
return mask, delta_flow
class BasicRPEEncoder(nn.Module):
def __init__(self, args):
super(BasicRPEEncoder, self).__init__()
self.args = args
dim = args.query_latent_dim
self.encoder = nn.Sequential(
nn.Linear(2, dim // 2),
nn.ReLU(inplace=True),
nn.Linear(dim // 2, dim),
nn.ReLU(inplace=True),
nn.Linear(dim, dim),
)
def forward(self, rpe_tokens):
return self.encoder(rpe_tokens)
from .twins import Block, CrossBlock
class TwinsSelfAttentionLayer(nn.Module):
def __init__(self, args):
super(TwinsSelfAttentionLayer, self).__init__()
self.args = args
embed_dim = 256
num_heads = 8
mlp_ratio = 4
ws = 7
sr_ratio = 4
dpr = 0.0
drop_rate = 0.0
attn_drop_rate = 0.0
self.local_block = Block(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=ws,
with_rpe=True,
)
self.global_block = Block(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=1,
with_rpe=True,
)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=0.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
elif isinstance(m, nn.Conv2d):
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
fan_out //= m.groups
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
if m.bias is not None:
m.bias.data.zero_()
elif isinstance(m, nn.BatchNorm2d):
m.weight.data.fill_(1.0)
m.bias.data.zero_()
def forward(self, x, tgt, size):
x = self.local_block(x, size)
x = self.global_block(x, size)
tgt = self.local_block(tgt, size)
tgt = self.global_block(tgt, size)
return x, tgt
class TwinsCrossAttentionLayer(nn.Module):
def __init__(self, args):
super(TwinsCrossAttentionLayer, self).__init__()
self.args = args
embed_dim = 256
num_heads = 8
mlp_ratio = 4
ws = 7
sr_ratio = 4
dpr = 0.0
drop_rate = 0.0
attn_drop_rate = 0.0
self.local_block = Block(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=ws,
with_rpe=True,
)
self.global_block = CrossBlock(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=1,
with_rpe=True,
)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
trunc_normal_(m.weight, std=0.02)
if isinstance(m, nn.Linear) and m.bias is not None:
nn.init.constant_(m.bias, 0)
elif isinstance(m, nn.LayerNorm):
nn.init.constant_(m.bias, 0)
nn.init.constant_(m.weight, 1.0)
elif isinstance(m, nn.Conv2d):
fan_out = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
fan_out //= m.groups
m.weight.data.normal_(0, math.sqrt(2.0 / fan_out))
if m.bias is not None:
m.bias.data.zero_()
elif isinstance(m, nn.BatchNorm2d):
m.weight.data.fill_(1.0)
m.bias.data.zero_()
def forward(self, x, tgt, size):
x = self.local_block(x, size)
tgt = self.local_block(tgt, size)
x, tgt = self.global_block(x, tgt, size)
return x, tgt
@@ -0,0 +1,98 @@
from turtle import forward
import torch
from torch import nn
import torch.nn.functional as F
import numpy as np
class ConvNextLayer(nn.Module):
def __init__(self, dim, depth=4):
super().__init__()
self.net = nn.Sequential(*[ConvNextBlock(dim=dim) for j in range(depth)])
def forward(self, x):
return self.net(x)
def compute_params(self):
num = 0
for param in self.parameters():
num += np.prod(param.size())
return num
class ConvNextBlock(nn.Module):
r"""ConvNeXt Block. There are two equivalent implementations:
(1) DwConv -> LayerNorm (channels_first) -> 1x1 Conv -> GELU -> 1x1 Conv; all in (N, C, H, W)
(2) DwConv -> Permute to (N, H, W, C); LayerNorm (channels_last) -> Linear -> GELU -> Linear; Permute back
We use (2) as we find it slightly faster in PyTorch
Args:
dim (int): Number of input channels.
drop_path (float): Stochastic depth rate. Default: 0.0
layer_scale_init_value (float): Init value for Layer Scale. Default: 1e-6.
"""
def __init__(self, dim, layer_scale_init_value=1e-6):
super().__init__()
self.dwconv = nn.Conv2d(
dim, dim, kernel_size=7, padding=3, groups=dim
) # depthwise conv
self.norm = LayerNorm(dim, eps=1e-6)
self.pwconv1 = nn.Linear(
dim, 4 * dim
) # pointwise/1x1 convs, implemented with linear layers
self.act = nn.GELU()
self.pwconv2 = nn.Linear(4 * dim, dim)
self.gamma = (
nn.Parameter(layer_scale_init_value * torch.ones((dim)), requires_grad=True)
if layer_scale_init_value > 0
else None
)
# self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity()
# print(f"conv next layer")
def forward(self, x):
input = x
x = self.dwconv(x)
x = x.permute(0, 2, 3, 1) # (N, C, H, W) -> (N, H, W, C)
x = self.norm(x)
x = self.pwconv1(x)
x = self.act(x)
x = self.pwconv2(x)
if self.gamma is not None:
x = self.gamma * x
x = x.permute(0, 3, 1, 2) # (N, H, W, C) -> (N, C, H, W)
x = input + x
return x
class LayerNorm(nn.Module):
r"""LayerNorm that supports two data formats: channels_last (default) or channels_first.
The ordering of the dimensions in the inputs. channels_last corresponds to inputs with
shape (batch_size, height, width, channels) while channels_first corresponds to inputs
with shape (batch_size, channels, height, width).
"""
def __init__(self, normalized_shape, eps=1e-6, data_format="channels_last"):
super().__init__()
self.weight = nn.Parameter(torch.ones(normalized_shape))
self.bias = nn.Parameter(torch.zeros(normalized_shape))
self.eps = eps
self.data_format = data_format
if self.data_format not in ["channels_last", "channels_first"]:
raise NotImplementedError
self.normalized_shape = (normalized_shape,)
def forward(self, x):
if self.data_format == "channels_last":
return F.layer_norm(
x, self.normalized_shape, self.weight, self.bias, self.eps
)
elif self.data_format == "channels_first":
u = x.mean(1, keepdim=True)
s = (x - u).pow(2).mean(1, keepdim=True)
x = (x - u) / torch.sqrt(s + self.eps)
x = self.weight[:, None, None] * x + self.bias[:, None, None]
return x
@@ -0,0 +1,321 @@
import loguru
import torch
import math
import torch.nn as nn
import torch.nn.functional as F
from torch import einsum
from einops.layers.torch import Rearrange
from einops import rearrange
from ...utils.utils import coords_grid, bilinear_sampler, upflow8
from .attention import (
MultiHeadAttention,
LinearPositionEmbeddingSine,
ExpPositionEmbeddingSine,
)
from typing import Optional, Tuple
from timm.models.layers import DropPath, to_2tuple, trunc_normal_
from .gru import BasicUpdateBlock, GMAUpdateBlock
from .gma import Attention
def initialize_flow(img):
"""Flow is represented as difference between two means flow = mean1 - mean0"""
N, C, H, W = img.shape
mean = coords_grid(N, H, W).to(img.device)
mean_init = coords_grid(N, H, W).to(img.device)
# optical flow computed as difference: flow = mean1 - mean0
return mean, mean_init
class CrossAttentionLayer(nn.Module):
# def __init__(self, dim, cfg, num_heads=8, attn_drop=0., proj_drop=0., drop_path=0., dropout=0.):
def __init__(
self,
qk_dim,
v_dim,
query_token_dim,
tgt_token_dim,
add_flow_token=True,
num_heads=8,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.0,
dropout=0.0,
pe="linear",
):
super(CrossAttentionLayer, self).__init__()
head_dim = qk_dim // num_heads
self.scale = head_dim**-0.5
self.query_token_dim = query_token_dim
self.pe = pe
self.norm1 = nn.LayerNorm(query_token_dim)
self.norm2 = nn.LayerNorm(query_token_dim)
self.multi_head_attn = MultiHeadAttention(qk_dim, num_heads)
self.q, self.k, self.v = (
nn.Linear(query_token_dim, qk_dim, bias=True),
nn.Linear(tgt_token_dim, qk_dim, bias=True),
nn.Linear(tgt_token_dim, v_dim, bias=True),
)
self.proj = nn.Linear(v_dim * 2, query_token_dim)
self.proj_drop = nn.Dropout(proj_drop)
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
self.ffn = nn.Sequential(
nn.Linear(query_token_dim, query_token_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(query_token_dim, query_token_dim),
nn.Dropout(dropout),
)
self.add_flow_token = add_flow_token
self.dim = qk_dim
def forward(self, query, key, value, memory, query_coord, patch_size, size_h3w3):
"""
query_coord [B, 2, H1, W1]
"""
B, _, H1, W1 = query_coord.shape
if key is None and value is None:
key = self.k(memory)
value = self.v(memory)
# [B, 2, H1, W1] -> [BH1W1, 1, 2]
query_coord = query_coord.contiguous()
query_coord = (
query_coord.view(B, 2, -1)
.permute(0, 2, 1)[:, :, None, :]
.contiguous()
.view(B * H1 * W1, 1, 2)
)
if self.pe == "linear":
query_coord_enc = LinearPositionEmbeddingSine(query_coord, dim=self.dim)
elif self.pe == "exp":
query_coord_enc = ExpPositionEmbeddingSine(query_coord, dim=self.dim)
short_cut = query
query = self.norm1(query)
if self.add_flow_token:
q = self.q(query + query_coord_enc)
else:
q = self.q(query_coord_enc)
k, v = key, value
x = self.multi_head_attn(q, k, v)
x = self.proj(torch.cat([x, short_cut], dim=2))
x = short_cut + self.proj_drop(x)
x = x + self.drop_path(self.ffn(self.norm2(x)))
return x, k, v
class MemoryDecoderLayer(nn.Module):
def __init__(self, dim, cfg):
super(MemoryDecoderLayer, self).__init__()
self.cfg = cfg
self.patch_size = cfg.patch_size # for converting coords into H2', W2' space
query_token_dim, tgt_token_dim = cfg.query_latent_dim, cfg.cost_latent_dim
qk_dim, v_dim = query_token_dim, query_token_dim
self.cross_attend = CrossAttentionLayer(
qk_dim,
v_dim,
query_token_dim,
tgt_token_dim,
add_flow_token=cfg.add_flow_token,
dropout=cfg.dropout,
)
def forward(self, query, key, value, memory, coords1, size, size_h3w3):
"""
x: [B*H1*W1, 1, C]
memory: [B*H1*W1, H2'*W2', C]
coords1 [B, 2, H2, W2]
size: B, C, H1, W1
1. Note that here coords0 and coords1 are in H2, W2 space.
Should first convert it into H2', W2' space.
2. We assume the upper-left point to be [0, 0], instead of letting center of upper-left patch to be [0, 0]
"""
x_global, k, v = self.cross_attend(
query, key, value, memory, coords1, self.patch_size, size_h3w3
)
B, C, H1, W1 = size
C = self.cfg.query_latent_dim
x_global = x_global.view(B, H1, W1, C).permute(0, 3, 1, 2)
return x_global, k, v
class ReverseCostExtractor(nn.Module):
def __init__(self, cfg):
super(ReverseCostExtractor, self).__init__()
self.cfg = cfg
def forward(self, cost_maps, coords0, coords1):
"""
cost_maps - B*H1*W1, cost_heads_num, H2, W2
coords - B, 2, H1, W1
"""
BH1W1, heads, H2, W2 = cost_maps.shape
B, _, H1, W1 = coords1.shape
assert (H1 == H2) and (W1 == W2)
assert BH1W1 == B * H1 * W1
cost_maps = cost_maps.reshape(B, H1 * W1 * heads, H2, W2)
coords = coords1.permute(0, 2, 3, 1)
corr = bilinear_sampler(cost_maps, coords) # [B, H1*W1*heads, H2, W2]
corr = rearrange(
corr,
"b (h1 w1 heads) h2 w2 -> (b h2 w2) heads h1 w1",
b=B,
heads=heads,
h1=H1,
w1=W1,
h2=H2,
w2=W2,
)
r = 4
dx = torch.linspace(-r, r, 2 * r + 1)
dy = torch.linspace(-r, r, 2 * r + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(coords0.device)
centroid = coords0.permute(0, 2, 3, 1).reshape(BH1W1, 1, 1, 2)
delta = delta.view(1, 2 * r + 1, 2 * r + 1, 2)
coords = centroid + delta
corr = bilinear_sampler(corr, coords)
corr = corr.view(B, H1, W1, -1).permute(0, 3, 1, 2)
return corr
class MemoryDecoder(nn.Module):
def __init__(self, cfg):
super(MemoryDecoder, self).__init__()
dim = self.dim = cfg.query_latent_dim
self.cfg = cfg
self.flow_token_encoder = nn.Sequential(
nn.Conv2d(81 * cfg.cost_heads_num, dim, 1, 1),
nn.GELU(),
nn.Conv2d(dim, dim, 1, 1),
)
self.proj = nn.Conv2d(256, 256, 1)
self.depth = cfg.decoder_depth
self.decoder_layer = MemoryDecoderLayer(dim, cfg)
if self.cfg.gma:
self.update_block = GMAUpdateBlock(self.cfg, hidden_dim=128)
self.att = Attention(
args=self.cfg, dim=128, heads=1, max_pos_size=160, dim_head=128
)
else:
self.update_block = BasicUpdateBlock(self.cfg, hidden_dim=128)
def upsample_flow(self, flow, mask):
"""Upsample flow field [H/8, W/8, 2] -> [H, W, 2] using convex combination"""
N, _, H, W = flow.shape
mask = mask.view(N, 1, 9, 8, 8, H, W)
mask = torch.softmax(mask, dim=2)
up_flow = F.unfold(8 * flow, [3, 3], padding=1)
up_flow = up_flow.view(N, 2, 9, 1, 1, H, W)
up_flow = torch.sum(mask * up_flow, dim=2)
up_flow = up_flow.permute(0, 1, 4, 2, 5, 3)
return up_flow.reshape(N, 2, 8 * H, 8 * W)
def encode_flow_token(self, cost_maps, coords):
"""
cost_maps - B*H1*W1, cost_heads_num, H2, W2
coords - B, 2, H1, W1
"""
coords = coords.permute(0, 2, 3, 1)
batch, h1, w1, _ = coords.shape
r = 4
dx = torch.linspace(-r, r, 2 * r + 1)
dy = torch.linspace(-r, r, 2 * r + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(coords.device)
centroid = coords.reshape(batch * h1 * w1, 1, 1, 2)
delta = delta.view(1, 2 * r + 1, 2 * r + 1, 2)
coords = centroid + delta
corr = bilinear_sampler(cost_maps, coords)
corr = corr.view(batch, h1, w1, -1).permute(0, 3, 1, 2)
return corr
def forward(self, cost_memory, context, data={}, flow_init=None, iters=None):
"""
memory: [B*H1*W1, H2'*W2', C]
context: [B, D, H1, W1]
"""
cost_maps = data["cost_maps"]
coords0, coords1 = initialize_flow(context)
if flow_init is not None:
# print("[Using warm start]")
coords1 = coords1 + flow_init
# flow = coords1
flow_predictions = []
context = self.proj(context)
net, inp = torch.split(context, [128, 128], dim=1)
net = torch.tanh(net)
inp = torch.relu(inp)
if self.cfg.gma:
attention = self.att(inp)
size = net.shape
key, value = None, None
if iters is None:
iters = self.depth
for idx in range(iters):
coords1 = coords1.detach()
cost_forward = self.encode_flow_token(cost_maps, coords1)
# cost_backward = self.reverse_cost_extractor(cost_maps, coords0, coords1)
query = self.flow_token_encoder(cost_forward)
query = (
query.permute(0, 2, 3, 1)
.contiguous()
.view(size[0] * size[2] * size[3], 1, self.dim)
)
cost_global, key, value = self.decoder_layer(
query, key, value, cost_memory, coords1, size, data["H3W3"]
)
if self.cfg.only_global:
corr = cost_global
else:
corr = torch.cat([cost_global, cost_forward], dim=1)
flow = coords1 - coords0
if self.cfg.gma:
net, up_mask, delta_flow = self.update_block(
net, inp, corr, flow, attention
)
else:
net, up_mask, delta_flow = self.update_block(net, inp, corr, flow)
# flow = delta_flow
coords1 = coords1 + delta_flow
flow_up = self.upsample_flow(coords1 - coords0, up_mask)
flow_predictions.append(flow_up)
# if self.training:
# return flow_predictions
# else:
return flow_predictions[-1], coords1 - coords0
@@ -0,0 +1,539 @@
import loguru
import torch
import math
import torch.nn as nn
import torch.nn.functional as F
from torch import einsum
import numpy as np
from einops.layers.torch import Rearrange
from einops import rearrange
import sys
from ...utils.utils import coords_grid, bilinear_sampler, upflow8
from .attention import (
BroadMultiHeadAttention,
MultiHeadAttention,
LinearPositionEmbeddingSine,
ExpPositionEmbeddingSine,
)
from ..encoders import twins_svt_large
from typing import Optional, Tuple
from .twins import Size_, PosConv
from .cnn import TwinsSelfAttentionLayer, TwinsCrossAttentionLayer, BasicEncoder
from .mlpmixer import MLPMixerLayer
from .convnext import ConvNextLayer
import time
from timm.models.layers import Mlp, DropPath, to_2tuple, trunc_normal_
class PatchEmbed(nn.Module):
def __init__(self, patch_size=16, in_chans=1, embed_dim=64, pe="linear"):
super().__init__()
self.patch_size = patch_size
self.dim = embed_dim
self.pe = pe
# assert patch_size == 8
if patch_size == 8:
self.proj = nn.Sequential(
nn.Conv2d(in_chans, embed_dim // 4, kernel_size=6, stride=2, padding=2),
nn.ReLU(),
nn.Conv2d(
embed_dim // 4, embed_dim // 2, kernel_size=6, stride=2, padding=2
),
nn.ReLU(),
nn.Conv2d(
embed_dim // 2, embed_dim, kernel_size=6, stride=2, padding=2
),
)
elif patch_size == 4:
self.proj = nn.Sequential(
nn.Conv2d(in_chans, embed_dim // 4, kernel_size=6, stride=2, padding=2),
nn.ReLU(),
nn.Conv2d(
embed_dim // 4, embed_dim, kernel_size=6, stride=2, padding=2
),
)
else:
print(f"patch size = {patch_size} is unacceptable.")
self.ffn_with_coord = nn.Sequential(
nn.Conv2d(embed_dim * 2, embed_dim * 2, kernel_size=1),
nn.ReLU(),
nn.Conv2d(embed_dim * 2, embed_dim * 2, kernel_size=1),
)
self.norm = nn.LayerNorm(embed_dim * 2)
def forward(self, x) -> Tuple[torch.Tensor, Size_]:
B, C, H, W = x.shape # C == 1
pad_l = pad_t = 0
pad_r = (self.patch_size - W % self.patch_size) % self.patch_size
pad_b = (self.patch_size - H % self.patch_size) % self.patch_size
x = F.pad(x, (pad_l, pad_r, pad_t, pad_b))
x = self.proj(x)
out_size = x.shape[2:]
patch_coord = (
coords_grid(B, out_size[0], out_size[1]).to(x.device) * self.patch_size
+ self.patch_size / 2
) # in feature coordinate space
patch_coord = patch_coord.view(B, 2, -1).permute(0, 2, 1)
if self.pe == "linear":
patch_coord_enc = LinearPositionEmbeddingSine(patch_coord, dim=self.dim)
elif self.pe == "exp":
patch_coord_enc = ExpPositionEmbeddingSine(patch_coord, dim=self.dim)
patch_coord_enc = patch_coord_enc.permute(0, 2, 1).view(
B, -1, out_size[0], out_size[1]
)
x_pe = torch.cat([x, patch_coord_enc], dim=1)
x = self.ffn_with_coord(x_pe)
x = self.norm(x.flatten(2).transpose(1, 2))
return x, out_size
from .twins import Block, CrossBlock
class GroupVerticalSelfAttentionLayer(nn.Module):
def __init__(
self,
dim,
cfg,
num_heads=8,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.0,
dropout=0.0,
):
super(GroupVerticalSelfAttentionLayer, self).__init__()
self.cfg = cfg
self.dim = dim
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = head_dim**-0.5
embed_dim = dim
mlp_ratio = 4
ws = 7
sr_ratio = 4
dpr = 0.0
drop_rate = dropout
attn_drop_rate = 0.0
self.block = Block(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=ws,
with_rpe=True,
vert_c_dim=cfg.vert_c_dim,
groupattention=True,
cfg=self.cfg,
)
def forward(self, x, size, context=None):
x = self.block(x, size, context)
return x
class VerticalSelfAttentionLayer(nn.Module):
def __init__(
self,
dim,
cfg,
num_heads=8,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.0,
dropout=0.0,
):
super(VerticalSelfAttentionLayer, self).__init__()
self.cfg = cfg
self.dim = dim
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = head_dim**-0.5
embed_dim = dim
mlp_ratio = 4
ws = 7
sr_ratio = 4
dpr = 0.0
drop_rate = dropout
attn_drop_rate = 0.0
self.local_block = Block(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=ws,
with_rpe=True,
vert_c_dim=cfg.vert_c_dim,
)
self.global_block = Block(
dim=embed_dim,
num_heads=num_heads,
mlp_ratio=mlp_ratio,
drop=drop_rate,
attn_drop=attn_drop_rate,
drop_path=dpr,
sr_ratio=sr_ratio,
ws=1,
with_rpe=True,
vert_c_dim=cfg.vert_c_dim,
)
def forward(self, x, size, context=None):
x = self.local_block(x, size, context)
x = self.global_block(x, size, context)
return x
def compute_params(self):
num = 0
for param in self.parameters():
num += np.prod(param.size())
return num
class SelfAttentionLayer(nn.Module):
def __init__(
self,
dim,
cfg,
num_heads=8,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.0,
dropout=0.0,
):
super(SelfAttentionLayer, self).__init__()
assert (
dim % num_heads == 0
), f"dim {dim} should be divided by num_heads {num_heads}."
self.dim = dim
self.num_heads = num_heads
head_dim = dim // num_heads
self.scale = head_dim**-0.5
self.norm1 = nn.LayerNorm(dim)
self.norm2 = nn.LayerNorm(dim)
self.multi_head_attn = MultiHeadAttention(dim, num_heads)
self.q, self.k, self.v = (
nn.Linear(dim, dim, bias=True),
nn.Linear(dim, dim, bias=True),
nn.Linear(dim, dim, bias=True),
)
self.proj = nn.Linear(dim, dim)
self.proj_drop = nn.Dropout(proj_drop)
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
self.ffn = nn.Sequential(
nn.Linear(dim, dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(dim, dim),
nn.Dropout(dropout),
)
def forward(self, x):
"""
x: [BH1W1, H3W3, D]
"""
short_cut = x
x = self.norm1(x)
q, k, v = self.q(x), self.k(x), self.v(x)
x = self.multi_head_attn(q, k, v)
x = self.proj(x)
x = short_cut + self.proj_drop(x)
x = x + self.drop_path(self.ffn(self.norm2(x)))
return x
def compute_params(self):
num = 0
for param in self.parameters():
num += np.prod(param.size())
return num
class CrossAttentionLayer(nn.Module):
def __init__(
self,
qk_dim,
v_dim,
query_token_dim,
tgt_token_dim,
num_heads=8,
attn_drop=0.0,
proj_drop=0.0,
drop_path=0.0,
dropout=0.0,
):
super(CrossAttentionLayer, self).__init__()
assert (
qk_dim % num_heads == 0
), f"dim {qk_dim} should be divided by num_heads {num_heads}."
assert (
v_dim % num_heads == 0
), f"dim {v_dim} should be divided by num_heads {num_heads}."
"""
Query Token: [N, C] -> [N, qk_dim] (Q)
Target Token: [M, D] -> [M, qk_dim] (K), [M, v_dim] (V)
"""
self.num_heads = num_heads
head_dim = qk_dim // num_heads
self.scale = head_dim**-0.5
self.norm1 = nn.LayerNorm(query_token_dim)
self.norm2 = nn.LayerNorm(query_token_dim)
self.multi_head_attn = BroadMultiHeadAttention(qk_dim, num_heads)
self.q, self.k, self.v = (
nn.Linear(query_token_dim, qk_dim, bias=True),
nn.Linear(tgt_token_dim, qk_dim, bias=True),
nn.Linear(tgt_token_dim, v_dim, bias=True),
)
self.proj = nn.Linear(v_dim, query_token_dim)
self.proj_drop = nn.Dropout(proj_drop)
self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()
self.ffn = nn.Sequential(
nn.Linear(query_token_dim, query_token_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(query_token_dim, query_token_dim),
nn.Dropout(dropout),
)
def forward(self, query, tgt_token):
"""
x: [BH1W1, H3W3, D]
"""
short_cut = query
query = self.norm1(query)
q, k, v = self.q(query), self.k(tgt_token), self.v(tgt_token)
x = self.multi_head_attn(q, k, v)
x = short_cut + self.proj_drop(self.proj(x))
x = x + self.drop_path(self.ffn(self.norm2(x)))
return x
class CostPerceiverEncoder(nn.Module):
def __init__(self, cfg):
super(CostPerceiverEncoder, self).__init__()
self.cfg = cfg
self.patch_size = cfg.patch_size
self.patch_embed = PatchEmbed(
in_chans=self.cfg.cost_heads_num,
patch_size=self.patch_size,
embed_dim=cfg.cost_latent_input_dim,
pe=cfg.pe,
)
self.depth = cfg.encoder_depth
self.latent_tokens = nn.Parameter(
torch.randn(1, cfg.cost_latent_token_num, cfg.cost_latent_dim)
)
query_token_dim, tgt_token_dim = (
cfg.cost_latent_dim,
cfg.cost_latent_input_dim * 2,
)
qk_dim, v_dim = query_token_dim, query_token_dim
self.input_layer = CrossAttentionLayer(
qk_dim, v_dim, query_token_dim, tgt_token_dim, dropout=cfg.dropout
)
if cfg.use_mlp:
self.encoder_layers = nn.ModuleList(
[
MLPMixerLayer(cfg.cost_latent_dim, cfg, dropout=cfg.dropout)
for idx in range(self.depth)
]
)
else:
self.encoder_layers = nn.ModuleList(
[
SelfAttentionLayer(cfg.cost_latent_dim, cfg, dropout=cfg.dropout)
for idx in range(self.depth)
]
)
if self.cfg.vertical_conv:
self.vertical_encoder_layers = nn.ModuleList(
[ConvNextLayer(cfg.cost_latent_dim) for idx in range(self.depth)]
)
else:
self.vertical_encoder_layers = nn.ModuleList(
[
VerticalSelfAttentionLayer(
cfg.cost_latent_dim, cfg, dropout=cfg.dropout
)
for idx in range(self.depth)
]
)
self.cost_scale_aug = None
if "cost_scale_aug" in cfg.keys():
self.cost_scale_aug = cfg.cost_scale_aug
print("[Using cost_scale_aug: {}]".format(self.cost_scale_aug))
def forward(self, cost_volume, data, context=None):
B, heads, H1, W1, H2, W2 = cost_volume.shape
cost_maps = (
cost_volume.permute(0, 2, 3, 1, 4, 5)
.contiguous()
.view(B * H1 * W1, self.cfg.cost_heads_num, H2, W2)
)
data["cost_maps"] = cost_maps
if self.cost_scale_aug is not None:
scale_factor = (
torch.FloatTensor(B * H1 * W1, self.cfg.cost_heads_num, H2, W2)
.uniform_(self.cost_scale_aug[0], self.cost_scale_aug[1])
.cuda()
)
cost_maps = cost_maps * scale_factor
x, size = self.patch_embed(cost_maps) # B*H1*W1, size[0]*size[1], C
data["H3W3"] = size
H3, W3 = size
x = self.input_layer(self.latent_tokens, x)
short_cut = x
for idx, layer in enumerate(self.encoder_layers):
x = layer(x)
if self.cfg.vertical_conv:
# B, H1*W1, K, D -> B, K, D, H1*W1 -> B*K, D, H1, W1
x = (
x.view(B, H1 * W1, self.cfg.cost_latent_token_num, -1)
.permute(0, 3, 1, 2)
.reshape(B * self.cfg.cost_latent_token_num, -1, H1, W1)
)
x = self.vertical_encoder_layers[idx](x)
# B*K, D, H1, W1 -> B, K, D, H1*W1 -> B, H1*W1, K, D
x = (
x.view(B, self.cfg.cost_latent_token_num, -1, H1 * W1)
.permute(0, 2, 3, 1)
.reshape(B * H1 * W1, self.cfg.cost_latent_token_num, -1)
)
else:
x = (
x.view(B, H1 * W1, self.cfg.cost_latent_token_num, -1)
.permute(0, 2, 1, 3)
.reshape(B * self.cfg.cost_latent_token_num, H1 * W1, -1)
)
x = self.vertical_encoder_layers[idx](x, (H1, W1), context)
x = (
x.view(B, self.cfg.cost_latent_token_num, H1 * W1, -1)
.permute(0, 2, 1, 3)
.reshape(B * H1 * W1, self.cfg.cost_latent_token_num, -1)
)
if self.cfg.cost_encoder_res is True:
x = x + short_cut
# print("~~~~")
return x
class MemoryEncoder(nn.Module):
def __init__(self, cfg):
super(MemoryEncoder, self).__init__()
self.cfg = cfg
if cfg.fnet == "twins":
self.feat_encoder = twins_svt_large(pretrained=self.cfg.pretrain)
elif cfg.fnet == "basicencoder":
self.feat_encoder = BasicEncoder(output_dim=256, norm_fn="instance")
else:
exit()
self.channel_convertor = nn.Conv2d(
cfg.encoder_latent_dim, cfg.encoder_latent_dim, 1, padding=0, bias=False
)
self.cost_perceiver_encoder = CostPerceiverEncoder(cfg)
def corr(self, fmap1, fmap2):
batch, dim, ht, wd = fmap1.shape
fmap1 = rearrange(
fmap1, "b (heads d) h w -> b heads (h w) d", heads=self.cfg.cost_heads_num
)
fmap2 = rearrange(
fmap2, "b (heads d) h w -> b heads (h w) d", heads=self.cfg.cost_heads_num
)
corr = einsum("bhid, bhjd -> bhij", fmap1, fmap2)
corr = corr.permute(0, 2, 1, 3).view(
batch * ht * wd, self.cfg.cost_heads_num, ht, wd
)
# corr = self.norm(self.relu(corr))
corr = corr.view(batch, ht * wd, self.cfg.cost_heads_num, ht * wd).permute(
0, 2, 1, 3
)
corr = corr.view(batch, self.cfg.cost_heads_num, ht, wd, ht, wd)
return corr
def forward(self, img1, img2, data, context=None, return_feat=False):
# The original implementation
# feat_s = self.feat_encoder(img1)
# feat_t = self.feat_encoder(img2)
# feat_s = self.channel_convertor(feat_s)
# feat_t = self.channel_convertor(feat_t)
imgs = torch.cat([img1, img2], dim=0)
feats = self.feat_encoder(imgs)
feats = self.channel_convertor(feats)
B = feats.shape[0] // 2
feat_s = feats[:B]
if return_feat:
ffeat = feats[:B]
feat_t = feats[B:]
B, C, H, W = feat_s.shape
size = (H, W)
if self.cfg.feat_cross_attn:
feat_s = feat_s.flatten(2).transpose(1, 2)
feat_t = feat_t.flatten(2).transpose(1, 2)
for layer in self.layers:
feat_s, feat_t = layer(feat_s, feat_t, size)
feat_s = feat_s.reshape(B, *size, -1).permute(0, 3, 1, 2).contiguous()
feat_t = feat_t.reshape(B, *size, -1).permute(0, 3, 1, 2).contiguous()
cost_volume = self.corr(feat_s, feat_t)
x = self.cost_perceiver_encoder(cost_volume, data, context)
if return_feat:
return x, ffeat
return x
@@ -0,0 +1,123 @@
import torch
from torch import nn, einsum
from einops import rearrange
class RelPosEmb(nn.Module):
def __init__(self, max_pos_size, dim_head):
super().__init__()
self.rel_height = nn.Embedding(2 * max_pos_size - 1, dim_head)
self.rel_width = nn.Embedding(2 * max_pos_size - 1, dim_head)
deltas = torch.arange(max_pos_size).view(1, -1) - torch.arange(
max_pos_size
).view(-1, 1)
rel_ind = deltas + max_pos_size - 1
self.register_buffer("rel_ind", rel_ind)
def forward(self, q):
batch, heads, h, w, c = q.shape
height_emb = self.rel_height(self.rel_ind[:h, :h].reshape(-1))
width_emb = self.rel_width(self.rel_ind[:w, :w].reshape(-1))
height_emb = rearrange(height_emb, "(x u) d -> x u () d", x=h)
width_emb = rearrange(width_emb, "(y v) d -> y () v d", y=w)
height_score = einsum("b h x y d, x u v d -> b h x y u v", q, height_emb)
width_score = einsum("b h x y d, y u v d -> b h x y u v", q, width_emb)
return height_score + width_score
class Attention(nn.Module):
def __init__(
self,
*,
args,
dim,
max_pos_size=100,
heads=4,
dim_head=128,
):
super().__init__()
self.args = args
self.heads = heads
self.scale = dim_head**-0.5
inner_dim = heads * dim_head
self.to_qk = nn.Conv2d(dim, inner_dim * 2, 1, bias=False)
self.pos_emb = RelPosEmb(max_pos_size, dim_head)
for param in self.pos_emb.parameters():
param.requires_grad = False
def forward(self, fmap):
heads, b, c, h, w = self.heads, *fmap.shape
q, k = self.to_qk(fmap).chunk(2, dim=1)
q, k = map(lambda t: rearrange(t, "b (h d) x y -> b h x y d", h=heads), (q, k))
q = self.scale * q
# if self.args.position_only:
# sim = self.pos_emb(q)
# elif self.args.position_and_content:
# sim_content = einsum('b h x y d, b h u v d -> b h x y u v', q, k)
# sim_pos = self.pos_emb(q)
# sim = sim_content + sim_pos
# else:
sim = einsum("b h x y d, b h u v d -> b h x y u v", q, k)
sim = rearrange(sim, "b h x y u v -> b h (x y) (u v)")
attn = sim.softmax(dim=-1)
return attn
class Aggregate(nn.Module):
def __init__(
self,
args,
dim,
heads=4,
dim_head=128,
):
super().__init__()
self.args = args
self.heads = heads
self.scale = dim_head**-0.5
inner_dim = heads * dim_head
self.to_v = nn.Conv2d(dim, inner_dim, 1, bias=False)
self.gamma = nn.Parameter(torch.zeros(1))
if dim != inner_dim:
self.project = nn.Conv2d(inner_dim, dim, 1, bias=False)
else:
self.project = None
def forward(self, attn, fmap):
heads, b, c, h, w = self.heads, *fmap.shape
v = self.to_v(fmap)
v = rearrange(v, "b (h d) x y -> b h (x y) d", h=heads)
out = einsum("b h i j, b h j d -> b h i d", attn, v)
out = rearrange(out, "b h (x y) d -> b (h d) x y", x=h, y=w)
if self.project is not None:
out = self.project(out)
out = fmap + self.gamma * out
return out
if __name__ == "__main__":
att = Attention(dim=128, heads=1)
fmap = torch.randn(2, 128, 40, 90)
out = att(fmap)
print(out.shape)
@@ -0,0 +1,160 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class FlowHead(nn.Module):
def __init__(self, input_dim=128, hidden_dim=256):
super(FlowHead, self).__init__()
self.conv1 = nn.Conv2d(input_dim, hidden_dim, 3, padding=1)
self.conv2 = nn.Conv2d(hidden_dim, 2, 3, padding=1)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
return self.conv2(self.relu(self.conv1(x)))
class ConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(ConvGRU, self).__init__()
self.convz = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convr = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convq = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
def forward(self, h, x):
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz(hx))
r = torch.sigmoid(self.convr(hx))
q = torch.tanh(self.convq(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class SepConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(SepConvGRU, self).__init__()
self.convz1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convr1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convq1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convz2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convr2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convq2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
def forward(self, h, x):
# horizontal
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz1(hx))
r = torch.sigmoid(self.convr1(hx))
q = torch.tanh(self.convq1(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
# vertical
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz2(hx))
r = torch.sigmoid(self.convr2(hx))
q = torch.tanh(self.convq2(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class BasicMotionEncoder(nn.Module):
def __init__(self, args):
super(BasicMotionEncoder, self).__init__()
if args.only_global:
print("[Decoding with only global cost]")
cor_planes = args.query_latent_dim
else:
cor_planes = 81 + args.query_latent_dim
self.convc1 = nn.Conv2d(cor_planes, 256, 1, padding=0)
self.convc2 = nn.Conv2d(256, 192, 3, padding=1)
self.convf1 = nn.Conv2d(2, 128, 7, padding=3)
self.convf2 = nn.Conv2d(128, 64, 3, padding=1)
self.conv = nn.Conv2d(64 + 192, 128 - 2, 3, padding=1)
def forward(self, flow, corr):
cor = F.relu(self.convc1(corr))
cor = F.relu(self.convc2(cor))
flo = F.relu(self.convf1(flow))
flo = F.relu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
out = F.relu(self.conv(cor_flo))
return torch.cat([out, flow], dim=1)
class BasicUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=128, input_dim=128):
super(BasicUpdateBlock, self).__init__()
self.args = args
self.encoder = BasicMotionEncoder(args)
self.gru = SepConvGRU(hidden_dim=hidden_dim, input_dim=128 + hidden_dim)
self.flow_head = FlowHead(hidden_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
def forward(self, net, inp, corr, flow, upsample=True):
motion_features = self.encoder(flow, corr)
inp = torch.cat([inp, motion_features], dim=1)
net = self.gru(net, inp)
delta_flow = self.flow_head(net)
# scale mask to balence gradients
mask = 0.25 * self.mask(net)
return net, mask, delta_flow
from .gma import Aggregate
class GMAUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=128):
super().__init__()
self.args = args
self.encoder = BasicMotionEncoder(args)
self.gru = SepConvGRU(
hidden_dim=hidden_dim, input_dim=128 + hidden_dim + hidden_dim
)
self.flow_head = FlowHead(hidden_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
self.aggregator = Aggregate(args=self.args, dim=128, dim_head=128, heads=1)
def forward(self, net, inp, corr, flow, attention):
motion_features = self.encoder(flow, corr)
motion_features_global = self.aggregator(attention, motion_features)
inp_cat = torch.cat([inp, motion_features, motion_features_global], dim=1)
# Attentional update
net = self.gru(net, inp_cat)
delta_flow = self.flow_head(net)
# scale mask to balence gradients
mask = 0.25 * self.mask(net)
return net, mask, delta_flow
@@ -0,0 +1,55 @@
from torch import nn
from einops.layers.torch import Rearrange, Reduce
from functools import partial
import numpy as np
class PreNormResidual(nn.Module):
def __init__(self, dim, fn):
super().__init__()
self.fn = fn
self.norm = nn.LayerNorm(dim)
def forward(self, x):
return self.fn(self.norm(x)) + x
def FeedForward(dim, expansion_factor=4, dropout=0.0, dense=nn.Linear):
return nn.Sequential(
dense(dim, dim * expansion_factor),
nn.GELU(),
nn.Dropout(dropout),
dense(dim * expansion_factor, dim),
nn.Dropout(dropout),
)
class MLPMixerLayer(nn.Module):
def __init__(self, dim, cfg, drop_path=0.0, dropout=0.0):
super(MLPMixerLayer, self).__init__()
# print(f"use mlp mixer layer")
K = cfg.cost_latent_token_num
expansion_factor = cfg.mlp_expansion_factor
chan_first, chan_last = partial(nn.Conv1d, kernel_size=1), nn.Linear
self.mlpmixer = nn.Sequential(
PreNormResidual(dim, FeedForward(K, expansion_factor, dropout, chan_first)),
PreNormResidual(
dim, FeedForward(dim, expansion_factor, dropout, chan_last)
),
)
def compute_params(self):
num = 0
for param in self.mlpmixer.parameters():
num += np.prod(param.size())
return num
def forward(self, x):
"""
x: [BH1W1, K, D]
"""
return self.mlpmixer(x)
@@ -0,0 +1,74 @@
import loguru
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import einsum
from einops.layers.torch import Rearrange
from einops import rearrange
from ...utils.utils import coords_grid, bilinear_sampler, upflow8
from ..common import (
FeedForward,
pyramid_retrieve_tokens,
sampler,
sampler_gaussian_fix,
retrieve_tokens,
MultiHeadAttention,
MLP,
)
from ..encoders import twins_svt_large_context, twins_svt_large
from ...position_encoding import PositionEncodingSine, LinearPositionEncoding
from .twins import PosConv
from .encoder import MemoryEncoder
from .decoder import MemoryDecoder
from .cnn import BasicEncoder
class FlowFormer(nn.Module):
def __init__(self, cfg):
super(FlowFormer, self).__init__()
self.cfg = cfg
self.memory_encoder = MemoryEncoder(cfg)
self.memory_decoder = MemoryDecoder(cfg)
if cfg.cnet == "twins":
self.context_encoder = twins_svt_large(pretrained=self.cfg.pretrain)
elif cfg.cnet == "basicencoder":
self.context_encoder = BasicEncoder(output_dim=256, norm_fn="instance")
def build_coord(self, img):
N, C, H, W = img.shape
coords = coords_grid(N, H // 8, W // 8)
return coords
def forward(
self, image1, image2, output=None, flow_init=None, return_feat=False, iters=None
):
# Following https://github.com/princeton-vl/RAFT/
image1 = 2 * (image1 / 255.0) - 1.0
image2 = 2 * (image2 / 255.0) - 1.0
data = {}
if self.cfg.context_concat:
context = self.context_encoder(torch.cat([image1, image2], dim=1))
else:
if return_feat:
context, cfeat = self.context_encoder(image1, return_feat=return_feat)
else:
context = self.context_encoder(image1)
if return_feat:
cost_memory, ffeat = self.memory_encoder(
image1, image2, data, context, return_feat=return_feat
)
else:
cost_memory = self.memory_encoder(image1, image2, data, context)
flow_predictions = self.memory_decoder(
cost_memory, context, data, flow_init=flow_init, iters=iters
)
if return_feat:
return flow_predictions, cfeat, ffeat
return flow_predictions
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,11 @@
import torch
def build_flowformer(cfg):
name = cfg.transformer
if name == "latentcostformer":
from .LatentCostFormer.transformer import FlowFormer
else:
raise ValueError(f"FlowFormer = {name} is not a valid architecture!")
return FlowFormer(cfg[name])
@@ -0,0 +1,566 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch import einsum
from einops.layers.torch import Rearrange
from einops import rearrange
from ..utils.utils import coords_grid, bilinear_sampler, indexing
from loguru import logger
import math
def nerf_encoding(x, L=6, NORMALIZE_FACOR=1 / 300):
"""
x is of shape [*, 2]. The last dimension are two coordinates (x and y).
"""
freq_bands = 2.0 ** torch.linspace(0, L, L - 1).to(x.device)
return torch.cat(
[
x * NORMALIZE_FACOR,
torch.sin(3.14 * x[..., -2:-1] * freq_bands * NORMALIZE_FACOR),
torch.cos(3.14 * x[..., -2:-1] * freq_bands * NORMALIZE_FACOR),
torch.sin(3.14 * x[..., -1:] * freq_bands * NORMALIZE_FACOR),
torch.cos(3.14 * x[..., -1:] * freq_bands * NORMALIZE_FACOR),
],
dim=-1,
)
def sampler_gaussian(latent, mean, std, image_size, point_num=25):
# latent [B, H*W, D]
# mean [B, 2, H, W]
# std [B, 1, H, W]
H, W = image_size
B, HW, D = latent.shape
STD_MAX = 20
latent = rearrange(
latent, "b (h w) c -> b c h w", h=H, w=W
) # latent = latent.view(B, H, W, D).permute(0, 3, 1, 2)
mean = mean.permute(0, 2, 3, 1) # [B, H, W, 2]
dx = torch.linspace(-1, 1, int(point_num**0.5))
dy = torch.linspace(-1, 1, int(point_num**0.5))
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
delta_3sigma = (
F.sigmoid(std.permute(0, 2, 3, 1).reshape(B * HW, 1, 1, 1))
* STD_MAX
* delta
* 3
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
centroid = mean.reshape(B * H * W, 1, 1, 2)
coords = centroid + delta_3sigma
coords = rearrange(coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W)
sampled_latents = bilinear_sampler(
latent, coords
) # [B*H*W, dim, point_num**0.5, point_num**0.5]
sampled_latents = sampled_latents.permute(0, 2, 3, 1)
sampled_weights = -(torch.sum(delta.pow(2), dim=-1))
return sampled_latents, sampled_weights
def sampler_gaussian_zy(
latent, mean, std, image_size, point_num=25, return_deltaXY=False, beta=1
):
# latent [B, H*W, D]
# mean [B, 2, H, W]
# std [B, 1, H, W]
H, W = image_size
B, HW, D = latent.shape
latent = rearrange(
latent, "b (h w) c -> b c h w", h=H, w=W
) # latent = latent.view(B, H, W, D).permute(0, 3, 1, 2)
mean = mean.permute(0, 2, 3, 1) # [B, H, W, 2]
dx = torch.linspace(-1, 1, int(point_num**0.5))
dy = torch.linspace(-1, 1, int(point_num**0.5))
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
delta_3sigma = (
std.permute(0, 2, 3, 1).reshape(B * HW, 1, 1, 1) * delta * 3
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
centroid = mean.reshape(B * H * W, 1, 1, 2)
coords = centroid + delta_3sigma
coords = rearrange(coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W)
sampled_latents = bilinear_sampler(
latent, coords
) # [B*H*W, dim, point_num**0.5, point_num**0.5]
sampled_latents = sampled_latents.permute(0, 2, 3, 1)
sampled_weights = -(torch.sum(delta.pow(2), dim=-1)) / beta
if return_deltaXY:
return sampled_latents, sampled_weights, delta_3sigma
else:
return sampled_latents, sampled_weights
def sampler_gaussian(latent, mean, std, image_size, point_num=25, return_deltaXY=False):
# latent [B, H*W, D]
# mean [B, 2, H, W]
# std [B, 1, H, W]
H, W = image_size
B, HW, D = latent.shape
STD_MAX = 20
latent = rearrange(
latent, "b (h w) c -> b c h w", h=H, w=W
) # latent = latent.view(B, H, W, D).permute(0, 3, 1, 2)
mean = mean.permute(0, 2, 3, 1) # [B, H, W, 2]
dx = torch.linspace(-1, 1, int(point_num**0.5))
dy = torch.linspace(-1, 1, int(point_num**0.5))
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
delta_3sigma = (
F.sigmoid(std.permute(0, 2, 3, 1).reshape(B * HW, 1, 1, 1))
* STD_MAX
* delta
* 3
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
centroid = mean.reshape(B * H * W, 1, 1, 2)
coords = centroid + delta_3sigma
coords = rearrange(coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W)
sampled_latents = bilinear_sampler(
latent, coords
) # [B*H*W, dim, point_num**0.5, point_num**0.5]
sampled_latents = sampled_latents.permute(0, 2, 3, 1)
sampled_weights = -(torch.sum(delta.pow(2), dim=-1))
if return_deltaXY:
return sampled_latents, sampled_weights, delta_3sigma
else:
return sampled_latents, sampled_weights
def sampler_gaussian_fix(latent, mean, image_size, point_num=49):
# latent [B, H*W, D]
# mean [B, 2, H, W]
H, W = image_size
B, HW, D = latent.shape
STD_MAX = 20
latent = rearrange(
latent, "b (h w) c -> b c h w", h=H, w=W
) # latent = latent.view(B, H, W, D).permute(0, 3, 1, 2)
mean = mean.permute(0, 2, 3, 1) # [B, H, W, 2]
radius = int((int(point_num**0.5) - 1) / 2)
dx = torch.linspace(-radius, radius, 2 * radius + 1)
dy = torch.linspace(-radius, radius, 2 * radius + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
centroid = mean.reshape(B * H * W, 1, 1, 2)
coords = centroid + delta
coords = rearrange(coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W)
sampled_latents = bilinear_sampler(
latent, coords
) # [B*H*W, dim, point_num**0.5, point_num**0.5]
sampled_latents = sampled_latents.permute(0, 2, 3, 1)
sampled_weights = -(torch.sum(delta.pow(2), dim=-1)) / point_num # smooth term
return sampled_latents, sampled_weights
def sampler_gaussian_fix_pyramid(
latent, feat_pyramid, scale_weight, mean, image_size, point_num=25
):
# latent [B, H*W, D]
# mean [B, 2, H, W]
# scale weight [B, H*W, layer_num]
H, W = image_size
B, HW, D = latent.shape
STD_MAX = 20
latent = rearrange(
latent, "b (h w) c -> b c h w", h=H, w=W
) # latent = latent.view(B, H, W, D).permute(0, 3, 1, 2)
mean = mean.permute(0, 2, 3, 1) # [B, H, W, 2]
radius = int((int(point_num**0.5) - 1) / 2)
dx = torch.linspace(-radius, radius, 2 * radius + 1)
dy = torch.linspace(-radius, radius, 2 * radius + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
sampled_latents = []
for i in range(len(feat_pyramid)):
centroid = mean.reshape(B * H * W, 1, 1, 2)
coords = (centroid + delta) / 2**i
coords = rearrange(
coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W
)
sampled_latents.append(bilinear_sampler(feat_pyramid[i], coords))
sampled_latents = torch.stack(
sampled_latents, dim=1
) # [B, layer_num, dim, H*W, point_num]
sampled_latents = sampled_latents.permute(
0, 3, 4, 2, 1
) # [B, H*W, point_num, dim, layer_num]
scale_weight = F.softmax(scale_weight, dim=2) # [B, H*W, layer_num]
vis_out = scale_weight
scale_weight = torch.unsqueeze(
torch.unsqueeze(scale_weight, dim=2), dim=2
) # [B, HW, 1, 1, layer_num]
weighted_latent = torch.sum(
sampled_latents * scale_weight, dim=-1
) # [B, H*W, point_num, dim]
sampled_weights = -(torch.sum(delta.pow(2), dim=-1)) / point_num # smooth term
return weighted_latent, sampled_weights, vis_out
def sampler_gaussian_pyramid(
latent, feat_pyramid, scale_weight, mean, std, image_size, point_num=25
):
# latent [B, H*W, D]
# mean [B, 2, H, W]
# scale weight [B, H*W, layer_num]
H, W = image_size
B, HW, D = latent.shape
STD_MAX = 20
latent = rearrange(
latent, "b (h w) c -> b c h w", h=H, w=W
) # latent = latent.view(B, H, W, D).permute(0, 3, 1, 2)
mean = mean.permute(0, 2, 3, 1) # [B, H, W, 2]
radius = int((int(point_num**0.5) - 1) / 2)
dx = torch.linspace(-1, 1, int(point_num**0.5))
dy = torch.linspace(-1, 1, int(point_num**0.5))
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
delta_3sigma = (
std.permute(0, 2, 3, 1).reshape(B * HW, 1, 1, 1) * delta * 3
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
sampled_latents = []
for i in range(len(feat_pyramid)):
centroid = mean.reshape(B * H * W, 1, 1, 2)
coords = (centroid + delta_3sigma) / 2**i
coords = rearrange(
coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W
)
sampled_latents.append(bilinear_sampler(feat_pyramid[i], coords))
sampled_latents = torch.stack(
sampled_latents, dim=1
) # [B, layer_num, dim, H*W, point_num]
sampled_latents = sampled_latents.permute(
0, 3, 4, 2, 1
) # [B, H*W, point_num, dim, layer_num]
scale_weight = F.softmax(scale_weight, dim=2) # [B, H*W, layer_num]
vis_out = scale_weight
scale_weight = torch.unsqueeze(
torch.unsqueeze(scale_weight, dim=2), dim=2
) # [B, HW, 1, 1, layer_num]
weighted_latent = torch.sum(
sampled_latents * scale_weight, dim=-1
) # [B, H*W, point_num, dim]
sampled_weights = -(torch.sum(delta.pow(2), dim=-1)) / point_num # smooth term
return weighted_latent, sampled_weights, vis_out
def sampler_gaussian_fix_MH(latent, mean, image_size, point_num=25):
"""different heads have different mean"""
# latent [B, H*W, D]
# mean [B, 2, H, W, heands]
H, W = image_size
B, HW, D = latent.shape
_, _, _, _, HEADS = mean.shape
STD_MAX = 20
latent = rearrange(latent, "b (h w) c -> b c h w", h=H, w=W)
mean = mean.permute(0, 2, 3, 4, 1) # [B, H, W, heads, 2]
radius = int((int(point_num**0.5) - 1) / 2)
dx = torch.linspace(-radius, radius, 2 * radius + 1)
dy = torch.linspace(-radius, radius, 2 * radius + 1)
delta = (
torch.stack(torch.meshgrid(dy, dx), axis=-1)
.to(mean.device)
.repeat(HEADS, 1, 1, 1)
) # [HEADS, point_num**0.5, point_num**0.5, 2]
centroid = mean.reshape(B * H * W, HEADS, 1, 1, 2)
coords = centroid + delta
coords = rearrange(
coords, "(b h w) H r1 r2 c -> b (h w H) (r1 r2) c", b=B, h=H, w=W, H=HEADS
)
sampled_latents = bilinear_sampler(latent, coords) # [B, dim, H*W*HEADS, pointnum]
sampled_latents = sampled_latents.permute(
0, 2, 3, 1
) # [B, H*W*HEADS, pointnum, dim]
sampled_weights = -(torch.sum(delta.pow(2), dim=-1)) / point_num # smooth term
return sampled_latents, sampled_weights
def sampler_gaussian_fix_pyramid_MH(
latent, feat_pyramid, scale_head_weight, mean, image_size, point_num=25
):
# latent [B, H*W, D]
# mean [B, 2, H, W, heands]
# scale_head weight [B, H*W, layer_num*heads]
H, W = image_size
B, HW, D = latent.shape
_, _, _, _, HEADS = mean.shape
latent = rearrange(latent, "b (h w) c -> b c h w", h=H, w=W)
mean = mean.permute(0, 2, 3, 4, 1) # [B, H, W, heads, 2]
radius = int((int(point_num**0.5) - 1) / 2)
dx = torch.linspace(-radius, radius, 2 * radius + 1)
dy = torch.linspace(-radius, radius, 2 * radius + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
mean.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
sampled_latents = []
centroid = mean.reshape(B * H * W, HEADS, 1, 1, 2)
for i in range(len(feat_pyramid)):
coords = (centroid) / 2**i + delta
coords = rearrange(
coords, "(b h w) H r1 r2 c -> b (h w H) (r1 r2) c", b=B, h=H, w=W, H=HEADS
)
sampled_latents.append(
bilinear_sampler(feat_pyramid[i], coords)
) # [B, dim, H*W*HEADS, point_num]
sampled_latents = torch.stack(
sampled_latents, dim=1
) # [B, layer_num, dim, H*W*HEADS, point_num]
sampled_latents = sampled_latents.permute(
0, 3, 4, 2, 1
) # [B, H*W*HEADS, point_num, dim, layer_num]
scale_head_weight = scale_head_weight.reshape(B, H * W * HEADS, -1)
scale_head_weight = F.softmax(scale_head_weight, dim=2) # [B, H*W*HEADS, layer_num]
scale_head_weight = torch.unsqueeze(
torch.unsqueeze(scale_head_weight, dim=2), dim=2
) # [B, H*W*HEADS, 1, 1, layer_num]
weighted_latent = torch.sum(
sampled_latents * scale_head_weight, dim=-1
) # [B, H*W*HEADS, point_num, dim]
sampled_weights = -(torch.sum(delta.pow(2), dim=-1)) / point_num # smooth term
return weighted_latent, sampled_weights
def sampler(feat, center, window_size):
# feat [B, C, H, W]
# center [B, 2, H, W]
center = center.permute(0, 2, 3, 1) # [B, H, W, 2]
B, H, W, C = center.shape
radius = window_size // 2
dx = torch.linspace(-radius, radius, 2 * radius + 1)
dy = torch.linspace(-radius, radius, 2 * radius + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
center.device
) # [B*H*W, window_size, point_num**0.5, 2]
center = center.reshape(B * H * W, 1, 1, 2)
coords = center + delta
coords = rearrange(coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W)
sampled_latents = bilinear_sampler(
feat, coords
) # [B*H*W, dim, window_size, window_size]
# sampled_latents = sampled_latents.permute(0, 2, 3, 1)
return sampled_latents
def retrieve_tokens(feat, center, window_size, sampler):
# feat [B, C, H, W]
# center [B, 2, H, W]
radius = window_size // 2
dx = torch.linspace(-radius, radius, 2 * radius + 1)
dy = torch.linspace(-radius, radius, 2 * radius + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(
center.device
) # [B*H*W, point_num**0.5, point_num**0.5, 2]
B, H, W, C = center.shape
centroid = center.reshape(B * H * W, 1, 1, 2)
coords = centroid + delta
coords = rearrange(coords, "(b h w) r1 r2 c -> b (h w) (r1 r2) c", b=B, h=H, w=W)
if sampler == "nn":
sampled_latents = indexing(feat, coords)
elif sampler == "bilinear":
sampled_latents = bilinear_sampler(feat, coords)
else:
raise ValueError("invalid sampler")
# [B, dim, H*W, point_num]
return sampled_latents
def pyramid_retrieve_tokens(
feat_pyramid, center, image_size, window_sizes, sampler="bilinear"
):
center = center.permute(0, 2, 3, 1) # [B, H, W, 2]
sampled_latents_pyramid = []
for idx in range(len(window_sizes)):
sampled_latents_pyramid.append(
retrieve_tokens(feat_pyramid[idx], center, window_sizes[idx], sampler)
)
center = center / 2
return torch.cat(sampled_latents_pyramid, dim=-1)
class FeedForward(nn.Module):
def __init__(self, dim, dropout=0.0):
super().__init__()
self.net = nn.Sequential(
nn.Linear(dim, dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(dim, dim),
nn.Dropout(dropout),
)
def forward(self, x):
x = self.net(x)
return x
class MLP(nn.Module):
def __init__(self, in_dim=22, out_dim=1, innter_dim=96, depth=5):
super().__init__()
self.FC1 = nn.Linear(in_dim, innter_dim)
self.FC_out = nn.Linear(innter_dim, out_dim)
self.relu = torch.nn.LeakyReLU(0.2)
self.FC_inter = nn.ModuleList(
[nn.Linear(innter_dim, innter_dim) for i in range(depth)]
)
def forward(self, x):
x = self.FC1(x)
x = self.relu(x)
for inter_fc in self.FC_inter:
x = inter_fc(x)
x = self.relu(x)
x = self.FC_out(x)
return x
class MultiHeadAttention(nn.Module):
def __init__(self, dim, heads, num_kv_tokens, cfg, rpe_bias=None, use_rpe=False):
super(MultiHeadAttention, self).__init__()
self.dim = dim
self.heads = heads
self.num_kv_tokens = num_kv_tokens
self.scale = (dim / heads) ** -0.5
self.rpe = cfg.rpe
self.attend = nn.Softmax(dim=-1)
self.use_rpe = use_rpe
if use_rpe:
if rpe_bias is None:
if self.rpe == "element-wise":
self.rpe_bias = nn.Parameter(
torch.zeros(heads, self.num_kv_tokens, dim // heads)
)
elif self.rpe == "head-wise":
self.rpe_bias = nn.Parameter(
torch.zeros(1, heads, 1, self.num_kv_tokens)
)
elif self.rpe == "token-wise":
self.rpe_bias = nn.Parameter(
torch.zeros(1, 1, 1, self.num_kv_tokens)
) # 81 is point_num
elif self.rpe == "implicit":
pass
# self.implicit_pe_fn = MLP(in_dim=22, out_dim=self.dim, innter_dim=int(self.dim//2.4), depth=2)
# raise ValueError('Implicit Encoding Not Implemented')
elif self.rpe == "element-wise-value":
self.rpe_bias = nn.Parameter(
torch.zeros(heads, self.num_kv_tokens, dim // heads)
)
self.rpe_value = nn.Parameter(torch.randn(self.num_kv_tokens, dim))
else:
raise ValueError("Not Implemented")
else:
self.rpe_bias = rpe_bias
def attend_with_rpe(self, Q, K, rpe_bias):
Q = rearrange(Q, "b i (heads d) -> b heads i d", heads=self.heads)
K = rearrange(K, "b j (heads d) -> b heads j d", heads=self.heads)
dots = (
einsum("bhid, bhjd -> bhij", Q, K) * self.scale
) # (b hw) heads 1 pointnum
if self.use_rpe:
if self.rpe == "element-wise":
rpe_bias_weight = (
einsum("bhid, hjd -> bhij", Q, rpe_bias) * self.scale
) # (b hw) heads 1 pointnum
dots = dots + rpe_bias_weight
elif self.rpe == "implicit":
pass
rpe_bias_weight = (
einsum("bhid, bhjd -> bhij", Q, rpe_bias) * self.scale
) # (b hw) heads 1 pointnum
dots = dots + rpe_bias_weight
elif self.rpe == "head-wise" or self.rpe == "token-wise":
dots = dots + rpe_bias
return self.attend(dots), dots
def forward(self, Q, K, V, rpe_bias=None):
if self.use_rpe:
if rpe_bias is None or self.rpe == "element-wise":
rpe_bias = self.rpe_bias
else:
rpe_bias = rearrange(
rpe_bias, "b hw pn (heads d) -> (b hw) heads pn d", heads=self.heads
)
attn, dots = self.attend_with_rpe(Q, K, rpe_bias)
else:
attn, dots = self.attend_with_rpe(Q, K, None)
B, HW, _ = Q.shape
if V is not None:
V = rearrange(V, "b j (heads d) -> b heads j d", heads=self.heads)
out = einsum("bhij, bhjd -> bhid", attn, V)
out = rearrange(out, "b heads hw d -> b hw (heads d)", b=B, hw=HW)
else:
out = None
# dots = torch.squeeze(dots, 2)
# dots = rearrange(dots, '(b hw) heads d -> b hw (heads d)', b=B, hw=HW)
return out, dots
@@ -0,0 +1,115 @@
import torch
import torch.nn as nn
import timm
import numpy as np
class twins_svt_large(nn.Module):
def __init__(self, pretrained=True):
super().__init__()
self.svt = timm.create_model("twins_svt_large", pretrained=pretrained)
del self.svt.head
del self.svt.patch_embeds[2]
del self.svt.patch_embeds[2]
del self.svt.blocks[2]
del self.svt.blocks[2]
del self.svt.pos_block[2]
del self.svt.pos_block[2]
self.svt.norm.weight.requires_grad = False
self.svt.norm.bias.requires_grad = False
def forward(self, x, data=None, layer=2, return_feat=False):
B = x.shape[0]
if return_feat:
feat = []
for i, (embed, drop, blocks, pos_blk) in enumerate(
zip(
self.svt.patch_embeds,
self.svt.pos_drops,
self.svt.blocks,
self.svt.pos_block,
)
):
x, size = embed(x)
x = drop(x)
for j, blk in enumerate(blocks):
x = blk(x, size)
if j == 0:
x = pos_blk(x, size)
if i < len(self.svt.depths) - 1:
x = x.reshape(B, *size, -1).permute(0, 3, 1, 2).contiguous()
if return_feat:
feat.append(x)
if i == layer - 1:
break
if return_feat:
return x, feat
return x
def compute_params(self, layer=2):
num = 0
for i, (embed, drop, blocks, pos_blk) in enumerate(
zip(
self.svt.patch_embeds,
self.svt.pos_drops,
self.svt.blocks,
self.svt.pos_block,
)
):
for param in embed.parameters():
num += np.prod(param.size())
for param in drop.parameters():
num += np.prod(param.size())
for param in blocks.parameters():
num += np.prod(param.size())
for param in pos_blk.parameters():
num += np.prod(param.size())
if i == layer - 1:
break
for param in self.svt.head.parameters():
num += np.prod(param.size())
return num
class twins_svt_large_context(nn.Module):
def __init__(self, pretrained=True):
super().__init__()
self.svt = timm.create_model("twins_svt_large_context", pretrained=pretrained)
def forward(self, x, data=None, layer=2):
B = x.shape[0]
for i, (embed, drop, blocks, pos_blk) in enumerate(
zip(
self.svt.patch_embeds,
self.svt.pos_drops,
self.svt.blocks,
self.svt.pos_block,
)
):
x, size = embed(x)
x = drop(x)
for j, blk in enumerate(blocks):
x = blk(x, size)
if j == 0:
x = pos_blk(x, size)
if i < len(self.svt.depths) - 1:
x = x.reshape(B, *size, -1).permute(0, 3, 1, 2).contiguous()
if i == layer - 1:
break
return x
if __name__ == "__main__":
m = twins_svt_large()
input = torch.randn(2, 3, 400, 800)
out = m.extract_feature(input)
print(out.shape)
@@ -0,0 +1,90 @@
import torch
import torch.nn.functional as F
from .utils.utils import bilinear_sampler, coords_grid
try:
import alt_cuda_corr
except:
# alt_cuda_corr is not compiled
pass
class CorrBlock:
def __init__(self, fmap1, fmap2, num_levels=4, radius=4):
self.num_levels = num_levels
self.radius = radius
self.corr_pyramid = []
# all pairs correlation
corr = CorrBlock.corr(fmap1, fmap2)
batch, h1, w1, dim, h2, w2 = corr.shape
corr = corr.reshape(batch * h1 * w1, dim, h2, w2)
self.corr_pyramid.append(corr)
for i in range(self.num_levels - 1):
corr = F.avg_pool2d(corr, 2, stride=2)
self.corr_pyramid.append(corr)
def __call__(self, coords):
r = self.radius
coords = coords.permute(0, 2, 3, 1)
batch, h1, w1, _ = coords.shape
out_pyramid = []
for i in range(self.num_levels):
corr = self.corr_pyramid[i]
dx = torch.linspace(-r, r, 2 * r + 1)
dy = torch.linspace(-r, r, 2 * r + 1)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1).to(coords.device)
centroid_lvl = coords.reshape(batch * h1 * w1, 1, 1, 2) / 2**i
delta_lvl = delta.view(1, 2 * r + 1, 2 * r + 1, 2)
coords_lvl = centroid_lvl + delta_lvl
corr = bilinear_sampler(corr, coords_lvl)
corr = corr.view(batch, h1, w1, -1)
out_pyramid.append(corr)
out = torch.cat(out_pyramid, dim=-1)
return out.permute(0, 3, 1, 2).contiguous().float()
@staticmethod
def corr(fmap1, fmap2):
batch, dim, ht, wd = fmap1.shape
fmap1 = fmap1.view(batch, dim, ht * wd)
fmap2 = fmap2.view(batch, dim, ht * wd)
corr = torch.matmul(fmap1.transpose(1, 2), fmap2)
corr = corr.view(batch, ht, wd, 1, ht, wd)
return corr / torch.sqrt(torch.tensor(dim).float())
class AlternateCorrBlock:
def __init__(self, fmap1, fmap2, num_levels=4, radius=4):
self.num_levels = num_levels
self.radius = radius
self.pyramid = [(fmap1, fmap2)]
for i in range(self.num_levels):
fmap1 = F.avg_pool2d(fmap1, 2, stride=2)
fmap2 = F.avg_pool2d(fmap2, 2, stride=2)
self.pyramid.append((fmap1, fmap2))
def __call__(self, coords):
coords = coords.permute(0, 2, 3, 1)
B, H, W, _ = coords.shape
dim = self.pyramid[0][0].shape[1]
corr_list = []
for i in range(self.num_levels):
r = self.radius
fmap1_i = self.pyramid[0][0].permute(0, 2, 3, 1).contiguous()
fmap2_i = self.pyramid[i][1].permute(0, 2, 3, 1).contiguous()
coords_i = (coords / 2**i).reshape(B, 1, H, W, 2).contiguous()
(corr,) = alt_cuda_corr.forward(fmap1_i, fmap2_i, coords_i, r)
corr_list.append(corr.squeeze(1))
corr = torch.stack(corr_list, dim=1)
corr = corr.reshape(B, -1, H, W)
return corr / torch.sqrt(torch.tensor(dim).float())
@@ -0,0 +1,297 @@
# Data loading based on https://github.com/NVIDIA/flownet2-pytorch
import numpy as np
import torch
import torch.utils.data as data
import torch.nn.functional as F
import os
import math
import random
from glob import glob
import os.path as osp
from .utils import frame_utils
from .utils.augmentor import FlowAugmentor, SparseFlowAugmentor
class FlowDataset(data.Dataset):
def __init__(self, aug_params=None, sparse=False):
self.augmentor = None
self.sparse = sparse
if aug_params is not None:
if sparse:
self.augmentor = SparseFlowAugmentor(**aug_params)
else:
self.augmentor = FlowAugmentor(**aug_params)
self.is_test = False
self.init_seed = False
self.flow_list = []
self.image_list = []
self.extra_info = []
def __getitem__(self, index):
if self.is_test:
img1 = frame_utils.read_gen(self.image_list[index][0])
img2 = frame_utils.read_gen(self.image_list[index][1])
img1 = np.array(img1).astype(np.uint8)[..., :3]
img2 = np.array(img2).astype(np.uint8)[..., :3]
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
return img1, img2, self.extra_info[index]
if not self.init_seed:
worker_info = torch.utils.data.get_worker_info()
if worker_info is not None:
torch.manual_seed(worker_info.id)
np.random.seed(worker_info.id)
random.seed(worker_info.id)
self.init_seed = True
index = index % len(self.image_list)
valid = None
if self.sparse:
flow, valid = frame_utils.readFlowKITTI(self.flow_list[index])
else:
flow = frame_utils.read_gen(self.flow_list[index])
img1 = frame_utils.read_gen(self.image_list[index][0])
img2 = frame_utils.read_gen(self.image_list[index][1])
flow = np.array(flow).astype(np.float32)
img1 = np.array(img1).astype(np.uint8)
img2 = np.array(img2).astype(np.uint8)
# grayscale images
if len(img1.shape) == 2:
img1 = np.tile(img1[..., None], (1, 1, 3))
img2 = np.tile(img2[..., None], (1, 1, 3))
else:
img1 = img1[..., :3]
img2 = img2[..., :3]
if self.augmentor is not None:
if self.sparse:
img1, img2, flow, valid = self.augmentor(img1, img2, flow, valid)
else:
img1, img2, flow = self.augmentor(img1, img2, flow)
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
flow = torch.from_numpy(flow).permute(2, 0, 1).float()
if valid is not None:
valid = torch.from_numpy(valid)
else:
valid = (flow[0].abs() < 1000) & (flow[1].abs() < 1000)
return img1, img2, flow, valid.float()
def __rmul__(self, v):
self.flow_list = v * self.flow_list
self.image_list = v * self.image_list
return self
def __len__(self):
return len(self.image_list)
class MpiSintel(FlowDataset):
def __init__(
self, aug_params=None, split="training", root="datasets/Sintel", dstype="clean"
):
super(MpiSintel, self).__init__(aug_params)
flow_root = osp.join(root, split, "flow")
image_root = osp.join(root, split, dstype)
if split == "test":
self.is_test = True
for scene in os.listdir(image_root):
image_list = sorted(glob(osp.join(image_root, scene, "*.png")))
for i in range(len(image_list) - 1):
self.image_list += [[image_list[i], image_list[i + 1]]]
self.extra_info += [(scene, i)] # scene and frame_id
if split != "test":
self.flow_list += sorted(glob(osp.join(flow_root, scene, "*.flo")))
class FlyingChairs(FlowDataset):
def __init__(
self, aug_params=None, split="train", root="datasets/FlyingChairs_release/data"
):
super(FlyingChairs, self).__init__(aug_params)
images = sorted(glob(osp.join(root, "*.ppm")))
flows = sorted(glob(osp.join(root, "*.flo")))
assert len(images) // 2 == len(flows)
split_list = np.loadtxt("chairs_split.txt", dtype=np.int32)
for i in range(len(flows)):
xid = split_list[i]
if (split == "training" and xid == 1) or (
split == "validation" and xid == 2
):
self.flow_list += [flows[i]]
self.image_list += [[images[2 * i], images[2 * i + 1]]]
class FlyingThings3D(FlowDataset):
def __init__(
self,
aug_params=None,
root="datasets/FlyingThings3D",
dstype="frames_cleanpass",
split="training",
):
super(FlyingThings3D, self).__init__(aug_params)
split_dir = "TRAIN" if split == "training" else "TEST"
for cam in ["left"]:
for direction in ["into_future", "into_past"]:
image_dirs = sorted(glob(osp.join(root, dstype, f"{split_dir}/*/*")))
image_dirs = sorted([osp.join(f, cam) for f in image_dirs])
flow_dirs = sorted(
glob(osp.join(root, f"optical_flow/{split_dir}/*/*"))
)
flow_dirs = sorted([osp.join(f, direction, cam) for f in flow_dirs])
for idir, fdir in zip(image_dirs, flow_dirs):
images = sorted(glob(osp.join(idir, "*.png")))
flows = sorted(glob(osp.join(fdir, "*.pfm")))
for i in range(len(flows) - 1):
if direction == "into_future":
self.image_list += [[images[i], images[i + 1]]]
self.flow_list += [flows[i]]
elif direction == "into_past":
self.image_list += [[images[i + 1], images[i]]]
self.flow_list += [flows[i + 1]]
class KITTI(FlowDataset):
def __init__(self, aug_params=None, split="training", root="datasets/KITTI"):
super(KITTI, self).__init__(aug_params, sparse=True)
if split == "testing":
self.is_test = True
root = osp.join(root, split)
images1 = sorted(glob(osp.join(root, "image_2/*_10.png")))
images2 = sorted(glob(osp.join(root, "image_2/*_11.png")))
for img1, img2 in zip(images1, images2):
frame_id = img1.split("/")[-1]
self.extra_info += [[frame_id]]
self.image_list += [[img1, img2]]
if split == "training":
self.flow_list = sorted(glob(osp.join(root, "flow_occ/*_10.png")))
class HD1K(FlowDataset):
def __init__(self, aug_params=None, root="datasets/HD1k"):
super(HD1K, self).__init__(aug_params, sparse=True)
seq_ix = 0
while 1:
flows = sorted(
glob(os.path.join(root, "hd1k_flow_gt", "flow_occ/%06d_*.png" % seq_ix))
)
images = sorted(
glob(os.path.join(root, "hd1k_input", "image_2/%06d_*.png" % seq_ix))
)
if len(flows) == 0:
break
for i in range(len(flows) - 1):
self.flow_list += [flows[i]]
self.image_list += [[images[i], images[i + 1]]]
seq_ix += 1
def fetch_dataloader(args, TRAIN_DS="C+T+K+S+H"):
"""Create the data loader for the corresponding trainign set"""
if args.stage == "chairs":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.1,
"max_scale": 1.0,
"do_flip": True,
}
train_dataset = FlyingChairs(aug_params, split="training")
elif args.stage == "things":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.4,
"max_scale": 0.8,
"do_flip": True,
}
clean_dataset = FlyingThings3D(aug_params, dstype="frames_cleanpass")
final_dataset = FlyingThings3D(aug_params, dstype="frames_finalpass")
train_dataset = clean_dataset + final_dataset
elif args.stage == "sintel":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.2,
"max_scale": 0.6,
"do_flip": True,
}
things = FlyingThings3D(aug_params, dstype="frames_cleanpass")
sintel_clean = MpiSintel(aug_params, split="training", dstype="clean")
sintel_final = MpiSintel(aug_params, split="training", dstype="final")
if TRAIN_DS == "C+T+K+S+H":
kitti = KITTI(
{
"crop_size": args.image_size,
"min_scale": -0.3,
"max_scale": 0.5,
"do_flip": True,
}
)
hd1k = HD1K(
{
"crop_size": args.image_size,
"min_scale": -0.5,
"max_scale": 0.2,
"do_flip": True,
}
)
train_dataset = (
100 * sintel_clean
+ 100 * sintel_final
+ 200 * kitti
+ 5 * hd1k
+ things
)
elif TRAIN_DS == "C+T+K/S":
train_dataset = 100 * sintel_clean + 100 * sintel_final + things
elif args.stage == "kitti":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.2,
"max_scale": 0.4,
"do_flip": False,
}
train_dataset = KITTI(aug_params, split="training")
train_loader = data.DataLoader(
train_dataset,
batch_size=args.batch_size,
pin_memory=False,
shuffle=True,
num_workers=128,
drop_last=True,
)
print("Training with %d image pairs" % len(train_dataset))
return train_loader
@@ -0,0 +1,267 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class ResidualBlock(nn.Module):
def __init__(self, in_planes, planes, norm_fn="group", stride=1):
super(ResidualBlock, self).__init__()
self.conv1 = nn.Conv2d(
in_planes, planes, kernel_size=3, padding=1, stride=stride
)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1)
self.relu = nn.ReLU(inplace=True)
num_groups = planes // 8
if norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
if not stride == 1:
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
elif norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(planes)
self.norm2 = nn.BatchNorm2d(planes)
if not stride == 1:
self.norm3 = nn.BatchNorm2d(planes)
elif norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(planes)
self.norm2 = nn.InstanceNorm2d(planes)
if not stride == 1:
self.norm3 = nn.InstanceNorm2d(planes)
elif norm_fn == "none":
self.norm1 = nn.Sequential()
self.norm2 = nn.Sequential()
if not stride == 1:
self.norm3 = nn.Sequential()
if stride == 1:
self.downsample = None
else:
self.downsample = nn.Sequential(
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm3
)
def forward(self, x):
y = x
y = self.relu(self.norm1(self.conv1(y)))
y = self.relu(self.norm2(self.conv2(y)))
if self.downsample is not None:
x = self.downsample(x)
return self.relu(x + y)
class BottleneckBlock(nn.Module):
def __init__(self, in_planes, planes, norm_fn="group", stride=1):
super(BottleneckBlock, self).__init__()
self.conv1 = nn.Conv2d(in_planes, planes // 4, kernel_size=1, padding=0)
self.conv2 = nn.Conv2d(
planes // 4, planes // 4, kernel_size=3, padding=1, stride=stride
)
self.conv3 = nn.Conv2d(planes // 4, planes, kernel_size=1, padding=0)
self.relu = nn.ReLU(inplace=True)
num_groups = planes // 8
if norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // 4)
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // 4)
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
if not stride == 1:
self.norm4 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
elif norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(planes // 4)
self.norm2 = nn.BatchNorm2d(planes // 4)
self.norm3 = nn.BatchNorm2d(planes)
if not stride == 1:
self.norm4 = nn.BatchNorm2d(planes)
elif norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(planes // 4)
self.norm2 = nn.InstanceNorm2d(planes // 4)
self.norm3 = nn.InstanceNorm2d(planes)
if not stride == 1:
self.norm4 = nn.InstanceNorm2d(planes)
elif norm_fn == "none":
self.norm1 = nn.Sequential()
self.norm2 = nn.Sequential()
self.norm3 = nn.Sequential()
if not stride == 1:
self.norm4 = nn.Sequential()
if stride == 1:
self.downsample = None
else:
self.downsample = nn.Sequential(
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm4
)
def forward(self, x):
y = x
y = self.relu(self.norm1(self.conv1(y)))
y = self.relu(self.norm2(self.conv2(y)))
y = self.relu(self.norm3(self.conv3(y)))
if self.downsample is not None:
x = self.downsample(x)
return self.relu(x + y)
class BasicEncoder(nn.Module):
def __init__(self, output_dim=128, norm_fn="batch", dropout=0.0):
super(BasicEncoder, self).__init__()
self.norm_fn = norm_fn
if self.norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=64)
elif self.norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(64)
elif self.norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(64)
elif self.norm_fn == "none":
self.norm1 = nn.Sequential()
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3)
self.relu1 = nn.ReLU(inplace=True)
self.in_planes = 64
self.layer1 = self._make_layer(64, stride=1)
self.layer2 = self._make_layer(96, stride=2)
self.layer3 = self._make_layer(128, stride=2)
# output convolution
self.conv2 = nn.Conv2d(128, output_dim, kernel_size=1)
self.dropout = None
if dropout > 0:
self.dropout = nn.Dropout2d(p=dropout)
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
if m.weight is not None:
nn.init.constant_(m.weight, 1)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def _make_layer(self, dim, stride=1):
layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride)
layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1)
layers = (layer1, layer2)
self.in_planes = dim
return nn.Sequential(*layers)
def forward(self, x):
# if input is list, combine batch dimension
is_list = isinstance(x, tuple) or isinstance(x, list)
if is_list:
batch_dim = x[0].shape[0]
x = torch.cat(x, dim=0)
x = self.conv1(x)
x = self.norm1(x)
x = self.relu1(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.conv2(x)
if self.training and self.dropout is not None:
x = self.dropout(x)
if is_list:
x = torch.split(x, [batch_dim, batch_dim], dim=0)
return x
class SmallEncoder(nn.Module):
def __init__(self, output_dim=128, norm_fn="batch", dropout=0.0):
super(SmallEncoder, self).__init__()
self.norm_fn = norm_fn
if self.norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=32)
elif self.norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(32)
elif self.norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(32)
elif self.norm_fn == "none":
self.norm1 = nn.Sequential()
self.conv1 = nn.Conv2d(3, 32, kernel_size=7, stride=2, padding=3)
self.relu1 = nn.ReLU(inplace=True)
self.in_planes = 32
self.layer1 = self._make_layer(32, stride=1)
self.layer2 = self._make_layer(64, stride=2)
self.layer3 = self._make_layer(96, stride=2)
self.dropout = None
if dropout > 0:
self.dropout = nn.Dropout2d(p=dropout)
self.conv2 = nn.Conv2d(96, output_dim, kernel_size=1)
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
if m.weight is not None:
nn.init.constant_(m.weight, 1)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def _make_layer(self, dim, stride=1):
layer1 = BottleneckBlock(self.in_planes, dim, self.norm_fn, stride=stride)
layer2 = BottleneckBlock(dim, dim, self.norm_fn, stride=1)
layers = (layer1, layer2)
self.in_planes = dim
return nn.Sequential(*layers)
def forward(self, x):
# if input is list, combine batch dimension
is_list = isinstance(x, tuple) or isinstance(x, list)
if is_list:
batch_dim = x[0].shape[0]
x = torch.cat(x, dim=0)
x = self.conv1(x)
x = self.norm1(x)
x = self.relu1(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.conv2(x)
if self.training and self.dropout is not None:
x = self.dropout(x)
if is_list:
x = torch.split(x, [batch_dim, batch_dim], dim=0)
return x
@@ -0,0 +1,40 @@
import torch
MAX_FLOW = 400
def sequence_loss(flow_preds, flow_gt, valid, cfg):
"""Loss function defined over sequence of flow predictions"""
gamma = cfg.gamma
max_flow = cfg.max_flow
n_predictions = len(flow_preds)
flow_loss = 0.0
flow_gt_thresholds = [5, 10, 20]
# exlude invalid pixels and extremely large diplacements
mag = torch.sum(flow_gt**2, dim=1).sqrt()
valid = (valid >= 0.5) & (mag < max_flow)
for i in range(n_predictions):
i_weight = gamma ** (n_predictions - i - 1)
i_loss = (flow_preds[i] - flow_gt).abs()
flow_loss += i_weight * (valid[:, None] * i_loss).mean()
epe = torch.sum((flow_preds[-1] - flow_gt) ** 2, dim=1).sqrt()
epe = epe.view(-1)[valid.view(-1)]
metrics = {
"epe": epe.mean().item(),
"1px": (epe < 1).float().mean().item(),
"3px": (epe < 3).float().mean().item(),
"5px": (epe < 5).float().mean().item(),
}
flow_gt_length = torch.sum(flow_gt**2, dim=1).sqrt()
flow_gt_length = flow_gt_length.view(-1)[valid.view(-1)]
for t in flow_gt_thresholds:
e = epe[flow_gt_length < t]
metrics.update({f"{t}-th-5px": (e < 5).float().mean().item()})
return flow_loss, metrics
@@ -0,0 +1,118 @@
import torch
from torch.optim.lr_scheduler import (
MultiStepLR,
CosineAnnealingLR,
ExponentialLR,
OneCycleLR,
)
def fetch_optimizer(model, cfg):
"""Create the optimizer and learning rate scheduler"""
# optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.wdecay, eps=args.epsilon)
# scheduler = optim.lr_scheduler.OneCycleLR(optimizer, args.lr, args.num_steps+100,
# pct_start=0.05, cycle_momentum=False, anneal_strategy='linear')
optimizer = build_optimizer(model, cfg)
scheduler = build_scheduler(cfg, optimizer)
return optimizer, scheduler
def build_optimizer(model, config):
name = config.optimizer
lr = config.canonical_lr
if name == "adam":
return torch.optim.Adam(
model.parameters(),
lr=lr,
weight_decay=config.adam_decay,
eps=config.epsilon,
)
elif name == "adamw":
if hasattr(config, "twins_lr_factor"):
factor = config.twins_lr_factor
print("[Decrease lr of pre-trained model by factor {}]".format(factor))
param_dicts = [
{
"params": [
p
for n, p in model.named_parameters()
if "feat_encoder" not in n
and "context_encoder" not in n
and p.requires_grad
]
},
{
"params": [
p
for n, p in model.named_parameters()
if ("feat_encoder" in n or "context_encoder" in n)
and p.requires_grad
],
"lr": lr * factor,
},
]
full = [n for n, _ in model.named_parameters()]
return torch.optim.AdamW(
param_dicts, lr=lr, weight_decay=config.adamw_decay, eps=config.epsilon
)
else:
return torch.optim.AdamW(
model.parameters(),
lr=lr,
weight_decay=config.adamw_decay,
eps=config.epsilon,
)
else:
raise ValueError(f"TRAINER.OPTIMIZER = {name} is not a valid optimizer!")
def build_scheduler(config, optimizer):
"""
Returns:
scheduler (dict):{
'scheduler': lr_scheduler,
'interval': 'step', # or 'epoch'
}
"""
# scheduler = {'interval': config.TRAINER.SCHEDULER_INTERVAL}
name = config.scheduler
lr = config.canonical_lr
if name == "OneCycleLR":
# scheduler = OneCycleLR(optimizer, )
if hasattr(config, "twins_lr_factor"):
factor = config.twins_lr_factor
scheduler = OneCycleLR(
optimizer,
[lr, lr * factor],
config.num_steps + 100,
pct_start=0.05,
cycle_momentum=False,
anneal_strategy=config.anneal_strategy,
)
else:
scheduler = OneCycleLR(
optimizer,
lr,
config.num_steps + 100,
pct_start=0.05,
cycle_momentum=False,
anneal_strategy=config.anneal_strategy,
)
# elif name == 'MultiStepLR':
# scheduler.update(
# {'scheduler': MultiStepLR(optimizer, config.TRAINER.MSLR_MILESTONES, gamma=config.TRAINER.MSLR_GAMMA)})
# elif name == 'CosineAnnealing':
# scheduler = CosineAnnealingLR(optimizer, config.num_steps+100)
# scheduler.update(
# {'scheduler': CosineAnnealingLR(optimizer, config.TRAINER.COSA_TMAX)})
# elif name == 'ExponentialLR':
# scheduler.update(
# {'scheduler': ExponentialLR(optimizer, config.TRAINER.ELR_GAMMA)})
else:
raise NotImplementedError()
return scheduler
@@ -0,0 +1,101 @@
from loguru import logger
import math
import torch
from torch import nn
class PositionEncodingSine(nn.Module):
"""
This is a sinusoidal position encoding that generalized to 2-dimensional images
"""
def __init__(self, d_model, max_shape=(256, 256)):
"""
Args:
max_shape (tuple): for 1/8 featmap, the max length of 256 corresponds to 2048 pixels
"""
super().__init__()
pe = torch.zeros((d_model, *max_shape))
y_position = torch.ones(max_shape).cumsum(0).float().unsqueeze(0)
x_position = torch.ones(max_shape).cumsum(1).float().unsqueeze(0)
div_term = torch.exp(
torch.arange(0, d_model // 2, 2).float()
* (-math.log(10000.0) / d_model // 2)
)
div_term = div_term[:, None, None] # [C//4, 1, 1]
pe[0::4, :, :] = torch.sin(x_position * div_term)
pe[1::4, :, :] = torch.cos(x_position * div_term)
pe[2::4, :, :] = torch.sin(y_position * div_term)
pe[3::4, :, :] = torch.cos(y_position * div_term)
self.register_buffer("pe", pe.unsqueeze(0)) # [1, C, H, W]
def forward(self, x):
"""
Args:
x: [N, C, H, W]
"""
return x + self.pe[:, :, : x.size(2), : x.size(3)]
class LinearPositionEncoding(nn.Module):
"""
This is a sinusoidal position encoding that generalized to 2-dimensional images
"""
def __init__(self, d_model, max_shape=(256, 256)):
"""
Args:
max_shape (tuple): for 1/8 featmap, the max length of 256 corresponds to 2048 pixels
"""
super().__init__()
pe = torch.zeros((d_model, *max_shape))
y_position = (
torch.ones(max_shape).cumsum(0).float().unsqueeze(0) - 1
) / max_shape[0]
x_position = (
torch.ones(max_shape).cumsum(1).float().unsqueeze(0) - 1
) / max_shape[1]
div_term = torch.arange(0, d_model // 2, 2).float()
div_term = div_term[:, None, None] # [C//4, 1, 1]
pe[0::4, :, :] = torch.sin(x_position * div_term * math.pi)
pe[1::4, :, :] = torch.cos(x_position * div_term * math.pi)
pe[2::4, :, :] = torch.sin(y_position * div_term * math.pi)
pe[3::4, :, :] = torch.cos(y_position * div_term * math.pi)
self.register_buffer("pe", pe.unsqueeze(0), persistent=False) # [1, C, H, W]
def forward(self, x):
"""
Args:
x: [N, C, H, W]
"""
# assert x.shape[2] == 80 and x.shape[3] == 80
return x + self.pe[:, :, : x.size(2), : x.size(3)]
class LearnedPositionEncoding(nn.Module):
"""
This is a sinusoidal position encoding that generalized to 2-dimensional images
"""
def __init__(self, d_model, max_shape=(80, 80)):
"""
Args:
max_shape (tuple): for 1/8 featmap, the max length of 256 corresponds to 2048 pixels
"""
super().__init__()
self.pe = nn.Parameter(torch.randn(1, max_shape[0], max_shape[1], d_model))
def forward(self, x):
"""
Args:
x: [N, C, H, W]
"""
# assert x.shape[2] == 80 and x.shape[3] == 80
return x + self.pe
@@ -0,0 +1,155 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from update import BasicUpdateBlock, SmallUpdateBlock
from extractor import BasicEncoder, SmallEncoder
from corr import CorrBlock, AlternateCorrBlock
from .utils.utils import bilinear_sampler, coords_grid, upflow8
try:
autocast = torch.cuda.amp.autocast
except:
# dummy autocast for PyTorch < 1.6
class autocast:
def __init__(self, enabled):
pass
def __enter__(self):
pass
def __exit__(self, *args):
pass
class RAFT(nn.Module):
def __init__(self, args):
super(RAFT, self).__init__()
self.args = args
if args.small:
self.hidden_dim = hdim = 96
self.context_dim = cdim = 64
args.corr_levels = 4
args.corr_radius = 3
else:
self.hidden_dim = hdim = 128
self.context_dim = cdim = 128
args.corr_levels = 4
args.corr_radius = 4
if "dropout" not in self.args:
self.args.dropout = 0
if "alternate_corr" not in self.args:
self.args.alternate_corr = False
# feature network, context network, and update block
if args.small:
self.fnet = SmallEncoder(
output_dim=128, norm_fn="instance", dropout=args.dropout
)
self.cnet = SmallEncoder(
output_dim=hdim + cdim, norm_fn="none", dropout=args.dropout
)
self.update_block = SmallUpdateBlock(self.args, hidden_dim=hdim)
else:
self.fnet = BasicEncoder(
output_dim=256, norm_fn="instance", dropout=args.dropout
)
self.cnet = BasicEncoder(
output_dim=hdim + cdim, norm_fn="batch", dropout=args.dropout
)
self.update_block = BasicUpdateBlock(self.args, hidden_dim=hdim)
def freeze_bn(self):
for m in self.modules():
if isinstance(m, nn.BatchNorm2d):
m.eval()
def initialize_flow(self, img):
"""Flow is represented as difference between two coordinate grids flow = coords1 - coords0"""
N, C, H, W = img.shape
coords0 = coords_grid(N, H // 8, W // 8).to(img.device)
coords1 = coords_grid(N, H // 8, W // 8).to(img.device)
# optical flow computed as difference: flow = coords1 - coords0
return coords0, coords1
def upsample_flow(self, flow, mask):
"""Upsample flow field [H/8, W/8, 2] -> [H, W, 2] using convex combination"""
N, _, H, W = flow.shape
mask = mask.view(N, 1, 9, 8, 8, H, W)
mask = torch.softmax(mask, dim=2)
up_flow = F.unfold(8 * flow, [3, 3], padding=1)
up_flow = up_flow.view(N, 2, 9, 1, 1, H, W)
up_flow = torch.sum(mask * up_flow, dim=2)
up_flow = up_flow.permute(0, 1, 4, 2, 5, 3)
return up_flow.reshape(N, 2, 8 * H, 8 * W)
def forward(
self, image1, image2, iters=12, flow_init=None, upsample=True, test_mode=False
):
"""Estimate optical flow between pair of frames"""
image1 = 2 * (image1 / 255.0) - 1.0
image2 = 2 * (image2 / 255.0) - 1.0
image1 = image1.contiguous()
image2 = image2.contiguous()
hdim = self.hidden_dim
cdim = self.context_dim
# run the feature network
with autocast(enabled=self.args.mixed_precision):
fmap1, fmap2 = self.fnet([image1, image2])
fmap1 = fmap1.float()
fmap2 = fmap2.float()
if self.args.alternate_corr:
corr_fn = AlternateCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
else:
corr_fn = CorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
# run the context network
with autocast(enabled=self.args.mixed_precision):
cnet = self.cnet(image1)
net, inp = torch.split(cnet, [hdim, cdim], dim=1)
net = torch.tanh(net)
inp = torch.relu(inp)
coords0, coords1 = self.initialize_flow(image1)
if flow_init is not None:
coords1 = coords1 + flow_init
flow_predictions = []
for itr in range(iters):
coords1 = coords1.detach()
corr = corr_fn(coords1) # index correlation volume
flow = coords1 - coords0
with autocast(enabled=self.args.mixed_precision):
net, up_mask, delta_flow = self.update_block(net, inp, corr, flow)
# F(t+1) = F(t) + \Delta(t)
coords1 = coords1 + delta_flow
# upsample predictions
if up_mask is None:
flow_up = upflow8(coords1 - coords0)
else:
flow_up = self.upsample_flow(coords1 - coords0, up_mask)
flow_predictions.append(flow_up)
if test_mode:
return coords1 - coords0, flow_up
return flow_predictions
@@ -0,0 +1,154 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class FlowHead(nn.Module):
def __init__(self, input_dim=128, hidden_dim=256):
super(FlowHead, self).__init__()
self.conv1 = nn.Conv2d(input_dim, hidden_dim, 3, padding=1)
self.conv2 = nn.Conv2d(hidden_dim, 2, 3, padding=1)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
return self.conv2(self.relu(self.conv1(x)))
class ConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(ConvGRU, self).__init__()
self.convz = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convr = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convq = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
def forward(self, h, x):
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz(hx))
r = torch.sigmoid(self.convr(hx))
q = torch.tanh(self.convq(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class SepConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(SepConvGRU, self).__init__()
self.convz1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convr1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convq1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convz2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convr2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convq2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
def forward(self, h, x):
# horizontal
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz1(hx))
r = torch.sigmoid(self.convr1(hx))
q = torch.tanh(self.convq1(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
# vertical
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz2(hx))
r = torch.sigmoid(self.convr2(hx))
q = torch.tanh(self.convq2(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class SmallMotionEncoder(nn.Module):
def __init__(self, args):
super(SmallMotionEncoder, self).__init__()
cor_planes = args.corr_levels * (2 * args.corr_radius + 1) ** 2
self.convc1 = nn.Conv2d(cor_planes, 96, 1, padding=0)
self.convf1 = nn.Conv2d(2, 64, 7, padding=3)
self.convf2 = nn.Conv2d(64, 32, 3, padding=1)
self.conv = nn.Conv2d(128, 80, 3, padding=1)
def forward(self, flow, corr):
cor = F.relu(self.convc1(corr))
flo = F.relu(self.convf1(flow))
flo = F.relu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
out = F.relu(self.conv(cor_flo))
return torch.cat([out, flow], dim=1)
class BasicMotionEncoder(nn.Module):
def __init__(self, args):
super(BasicMotionEncoder, self).__init__()
cor_planes = args.corr_levels * (2 * args.corr_radius + 1) ** 2
self.convc1 = nn.Conv2d(cor_planes, 256, 1, padding=0)
self.convc2 = nn.Conv2d(256, 192, 3, padding=1)
self.convf1 = nn.Conv2d(2, 128, 7, padding=3)
self.convf2 = nn.Conv2d(128, 64, 3, padding=1)
self.conv = nn.Conv2d(64 + 192, 128 - 2, 3, padding=1)
def forward(self, flow, corr):
cor = F.relu(self.convc1(corr))
cor = F.relu(self.convc2(cor))
flo = F.relu(self.convf1(flow))
flo = F.relu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
out = F.relu(self.conv(cor_flo))
return torch.cat([out, flow], dim=1)
class SmallUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=96):
super(SmallUpdateBlock, self).__init__()
self.encoder = SmallMotionEncoder(args)
self.gru = ConvGRU(hidden_dim=hidden_dim, input_dim=82 + 64)
self.flow_head = FlowHead(hidden_dim, hidden_dim=128)
def forward(self, net, inp, corr, flow):
motion_features = self.encoder(flow, corr)
inp = torch.cat([inp, motion_features], dim=1)
net = self.gru(net, inp)
delta_flow = self.flow_head(net)
return net, None, delta_flow
class BasicUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=128, input_dim=128):
super(BasicUpdateBlock, self).__init__()
self.args = args
self.encoder = BasicMotionEncoder(args)
self.gru = SepConvGRU(hidden_dim=hidden_dim, input_dim=128 + hidden_dim)
self.flow_head = FlowHead(hidden_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
def forward(self, net, inp, corr, flow, upsample=True):
motion_features = self.encoder(flow, corr)
inp = torch.cat([inp, motion_features], dim=1)
net = self.gru(net, inp)
delta_flow = self.flow_head(net)
# scale mask to balence gradients
mask = 0.25 * self.mask(net)
return net, mask, delta_flow
@@ -0,0 +1,336 @@
import numpy as np
import random
import math
from PIL import Image
import cv2
cv2.setNumThreads(0)
cv2.ocl.setUseOpenCL(False)
import torch
from torchvision.transforms import ColorJitter
import torch.nn.functional as F
from . import flow_transforms
class FlowAugmentor:
def __init__(
self, crop_size, min_scale=-0.2, max_scale=0.5, do_flip=True, pwc_aug=False
):
# spatial augmentation params
self.crop_size = crop_size
self.min_scale = min_scale
self.max_scale = max_scale
self.spatial_aug_prob = 0.8
self.stretch_prob = 0.8
self.max_stretch = 0.2
# flip augmentation params
self.do_flip = do_flip
self.h_flip_prob = 0.5
self.v_flip_prob = 0.1
# photometric augmentation params
self.photo_aug = ColorJitter(
brightness=0.4, contrast=0.4, saturation=0.4, hue=0.5 / 3.14
)
self.asymmetric_color_aug_prob = 0.2
self.eraser_aug_prob = 0.5
self.pwc_aug = pwc_aug
if self.pwc_aug:
print("[Using pwc-style spatial augmentation]")
def color_transform(self, img1, img2):
"""Photometric augmentation"""
# asymmetric
if np.random.rand() < self.asymmetric_color_aug_prob:
img1 = np.array(self.photo_aug(Image.fromarray(img1)), dtype=np.uint8)
img2 = np.array(self.photo_aug(Image.fromarray(img2)), dtype=np.uint8)
# symmetric
else:
image_stack = np.concatenate([img1, img2], axis=0)
image_stack = np.array(
self.photo_aug(Image.fromarray(image_stack)), dtype=np.uint8
)
img1, img2 = np.split(image_stack, 2, axis=0)
return img1, img2
def eraser_transform(self, img1, img2, bounds=[50, 100]):
"""Occlusion augmentation"""
ht, wd = img1.shape[:2]
if np.random.rand() < self.eraser_aug_prob:
mean_color = np.mean(img2.reshape(-1, 3), axis=0)
for _ in range(np.random.randint(1, 3)):
x0 = np.random.randint(0, wd)
y0 = np.random.randint(0, ht)
dx = np.random.randint(bounds[0], bounds[1])
dy = np.random.randint(bounds[0], bounds[1])
img2[y0 : y0 + dy, x0 : x0 + dx, :] = mean_color
return img1, img2
def spatial_transform(self, img1, img2, flow):
# randomly sample scale
ht, wd = img1.shape[:2]
min_scale = np.maximum(
(self.crop_size[0] + 8) / float(ht), (self.crop_size[1] + 8) / float(wd)
)
scale = 2 ** np.random.uniform(self.min_scale, self.max_scale)
scale_x = scale
scale_y = scale
if np.random.rand() < self.stretch_prob:
scale_x *= 2 ** np.random.uniform(-self.max_stretch, self.max_stretch)
scale_y *= 2 ** np.random.uniform(-self.max_stretch, self.max_stretch)
scale_x = np.clip(scale_x, min_scale, None)
scale_y = np.clip(scale_y, min_scale, None)
if np.random.rand() < self.spatial_aug_prob:
# rescale the images
img1 = cv2.resize(
img1, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
img2 = cv2.resize(
img2, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
flow = cv2.resize(
flow, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
flow = flow * [scale_x, scale_y]
if self.do_flip:
if np.random.rand() < self.h_flip_prob: # h-flip
img1 = img1[:, ::-1]
img2 = img2[:, ::-1]
flow = flow[:, ::-1] * [-1.0, 1.0]
if np.random.rand() < self.v_flip_prob: # v-flip
img1 = img1[::-1, :]
img2 = img2[::-1, :]
flow = flow[::-1, :] * [1.0, -1.0]
if img1.shape[0] == self.crop_size[0]:
y0 = 0
else:
y0 = np.random.randint(0, img1.shape[0] - self.crop_size[0])
if img1.shape[1] == self.crop_size[1]:
x0 = 0
else:
x0 = np.random.randint(0, img1.shape[1] - self.crop_size[1])
img1 = img1[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
img2 = img2[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
flow = flow[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
return img1, img2, flow
def __call__(self, img1, img2, flow):
img1, img2 = self.color_transform(img1, img2)
img1, img2 = self.eraser_transform(img1, img2)
if self.pwc_aug:
th, tw = self.crop_size
schedule = [0.5, 1.0] # initial coeff, final_coeff, half life
difficulty = np.random.uniform(0, 1)
schedule_coeff = schedule[0] + (schedule[1] - schedule[0]) * (
2 / (1 + np.exp(-1.0986 * difficulty)) - 1
)
spatial_augmentor = flow_transforms.SpatialAug(
[th, tw],
scale=[0.4, 0.03, 0.2],
rot=[0.4, 0.03],
trans=[0.4, 0.03],
squeeze=[0.3, 0.0],
schedule_coeff=schedule_coeff,
order=1,
black=False,
)
flow = np.concatenate(
[flow, np.ones((flow.shape[0], flow.shape[1], 1))], axis=-1
)
augmented, flow_valid = spatial_augmentor([img1, img2], flow)
flow = flow_valid[:, :, :2]
img1 = augmented[0]
img2 = augmented[1]
else:
img1, img2, flow = self.spatial_transform(img1, img2, flow)
img1 = np.ascontiguousarray(img1)
img2 = np.ascontiguousarray(img2)
flow = np.ascontiguousarray(flow)
return img1, img2, flow
class SparseFlowAugmentor:
def __init__(self, crop_size, min_scale=-0.2, max_scale=0.5, do_flip=False):
# spatial augmentation params
self.crop_size = crop_size
self.min_scale = min_scale
self.max_scale = max_scale
self.spatial_aug_prob = 0.8
self.stretch_prob = 0.8
self.max_stretch = 0.2
# flip augmentation params
self.do_flip = do_flip
self.h_flip_prob = 0.5
self.v_flip_prob = 0.1
# photometric augmentation params
self.photo_aug = ColorJitter(
brightness=0.3, contrast=0.3, saturation=0.3, hue=0.3 / 3.14
)
self.asymmetric_color_aug_prob = 0.2
self.eraser_aug_prob = 0.5
def color_transform(self, img1, img2):
image_stack = np.concatenate([img1, img2], axis=0)
image_stack = np.array(
self.photo_aug(Image.fromarray(image_stack)), dtype=np.uint8
)
img1, img2 = np.split(image_stack, 2, axis=0)
return img1, img2
def eraser_transform(self, img1, img2):
ht, wd = img1.shape[:2]
if np.random.rand() < self.eraser_aug_prob:
mean_color = np.mean(img2.reshape(-1, 3), axis=0)
for _ in range(np.random.randint(1, 3)):
x0 = np.random.randint(0, wd)
y0 = np.random.randint(0, ht)
dx = np.random.randint(50, 100)
dy = np.random.randint(50, 100)
img2[y0 : y0 + dy, x0 : x0 + dx, :] = mean_color
return img1, img2
def resize_sparse_flow_map(self, flow, valid, fx=1.0, fy=1.0):
ht, wd = flow.shape[:2]
coords = np.meshgrid(np.arange(wd), np.arange(ht))
coords = np.stack(coords, axis=-1)
coords = coords.reshape(-1, 2).astype(np.float32)
flow = flow.reshape(-1, 2).astype(np.float32)
valid = valid.reshape(-1).astype(np.float32)
coords0 = coords[valid >= 1]
flow0 = flow[valid >= 1]
ht1 = int(round(ht * fy))
wd1 = int(round(wd * fx))
coords1 = coords0 * [fx, fy]
flow1 = flow0 * [fx, fy]
xx = np.round(coords1[:, 0]).astype(np.int32)
yy = np.round(coords1[:, 1]).astype(np.int32)
v = (xx > 0) & (xx < wd1) & (yy > 0) & (yy < ht1)
xx = xx[v]
yy = yy[v]
flow1 = flow1[v]
flow_img = np.zeros([ht1, wd1, 2], dtype=np.float32)
valid_img = np.zeros([ht1, wd1], dtype=np.int32)
flow_img[yy, xx] = flow1
valid_img[yy, xx] = 1
return flow_img, valid_img
def spatial_transform(self, img1, img2, flow, valid):
pad_t = 0
pad_b = 0
pad_l = 0
pad_r = 0
if self.crop_size[0] > img1.shape[0]:
pad_b = self.crop_size[0] - img1.shape[0]
if self.crop_size[1] > img1.shape[1]:
pad_r = self.crop_size[1] - img1.shape[1]
if pad_b != 0 or pad_r != 0:
img1 = np.pad(
img1,
((pad_t, pad_b), (pad_l, pad_r), (0, 0)),
"constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
img2 = np.pad(
img2,
((pad_t, pad_b), (pad_l, pad_r), (0, 0)),
"constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
flow = np.pad(
flow,
((pad_t, pad_b), (pad_l, pad_r), (0, 0)),
"constant",
constant_values=((0, 0), (0, 0), (0, 0)),
)
valid = np.pad(
valid,
((pad_t, pad_b), (pad_l, pad_r)),
"constant",
constant_values=((0, 0), (0, 0)),
)
# randomly sample scale
ht, wd = img1.shape[:2]
min_scale = np.maximum(
(self.crop_size[0] + 1) / float(ht), (self.crop_size[1] + 1) / float(wd)
)
scale = 2 ** np.random.uniform(self.min_scale, self.max_scale)
scale_x = np.clip(scale, min_scale, None)
scale_y = np.clip(scale, min_scale, None)
if np.random.rand() < self.spatial_aug_prob:
# rescale the images
img1 = cv2.resize(
img1, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
img2 = cv2.resize(
img2, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
flow, valid = self.resize_sparse_flow_map(
flow, valid, fx=scale_x, fy=scale_y
)
if self.do_flip:
if np.random.rand() < 0.5: # h-flip
img1 = img1[:, ::-1]
img2 = img2[:, ::-1]
flow = flow[:, ::-1] * [-1.0, 1.0]
valid = valid[:, ::-1]
margin_y = 20
margin_x = 50
y0 = np.random.randint(0, img1.shape[0] - self.crop_size[0] + margin_y)
x0 = np.random.randint(-margin_x, img1.shape[1] - self.crop_size[1] + margin_x)
y0 = np.clip(y0, 0, img1.shape[0] - self.crop_size[0])
x0 = np.clip(x0, 0, img1.shape[1] - self.crop_size[1])
img1 = img1[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
img2 = img2[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
flow = flow[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
valid = valid[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
return img1, img2, flow, valid
def __call__(self, img1, img2, flow, valid):
img1, img2 = self.color_transform(img1, img2)
img1, img2 = self.eraser_transform(img1, img2)
img1, img2, flow, valid = self.spatial_transform(img1, img2, flow, valid)
img1 = np.ascontiguousarray(img1)
img2 = np.ascontiguousarray(img2)
flow = np.ascontiguousarray(flow)
valid = np.ascontiguousarray(valid)
return img1, img2, flow, valid
@@ -0,0 +1,577 @@
import numpy as np
import torch
import torch.utils.data as data
import torch.nn.functional as F
import os
import math
import random
from glob import glob
import os.path as osp
from utils import frame_utils
from utils.augmentor import FlowAugmentor, SparseFlowAugmentor
# from utils import flow_transforms
from torchvision.utils import save_image
from utils import flow_viz
import cv2
from utils.utils import coords_grid, bilinear_sampler
class FlowDataset(data.Dataset):
def __init__(self, aug_params=None, sparse=False):
self.augmentor = None
self.sparse = sparse
if aug_params is not None:
if sparse:
self.augmentor = SparseFlowAugmentor(**aug_params)
else:
self.augmentor = FlowAugmentor(**aug_params)
self.is_test = False
self.init_seed = False
self.flow_list = []
self.image_list = []
self.extra_info = []
def __getitem__(self, index):
# print(self.flow_list[index])
if self.is_test:
img1 = frame_utils.read_gen(self.image_list[index][0], test=self.is_test)
img2 = frame_utils.read_gen(self.image_list[index][1], test=self.is_test)
img1 = np.array(img1).astype(np.uint8)[..., :3]
img2 = np.array(img2).astype(np.uint8)[..., :3]
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
return img1, img2, self.extra_info[index]
if not self.init_seed:
worker_info = torch.utils.data.get_worker_info()
if worker_info is not None:
torch.manual_seed(worker_info.id)
np.random.seed(worker_info.id)
random.seed(worker_info.id)
self.init_seed = True
index = index % len(self.image_list)
valid = None
if self.sparse:
flow, valid = frame_utils.readFlowKITTI(self.flow_list[index])
else:
flow = frame_utils.read_gen(self.flow_list[index])
img1 = frame_utils.read_gen(self.image_list[index][0])
img2 = frame_utils.read_gen(self.image_list[index][1])
flow = np.array(flow).astype(np.float32)
img1 = np.array(img1).astype(np.uint8)
img2 = np.array(img2).astype(np.uint8)
# grayscale images
if len(img1.shape) == 2:
img1 = np.tile(img1[..., None], (1, 1, 3))
img2 = np.tile(img2[..., None], (1, 1, 3))
else:
img1 = img1[..., :3]
img2 = img2[..., :3]
if self.augmentor is not None:
if self.sparse:
img1, img2, flow, valid = self.augmentor(img1, img2, flow, valid)
else:
img1, img2, flow = self.augmentor(img1, img2, flow)
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
flow = torch.from_numpy(flow).permute(2, 0, 1).float()
if valid is not None:
valid = torch.from_numpy(valid)
else:
valid = (flow[0].abs() < 1000) & (flow[1].abs() < 1000)
return img1, img2, flow, valid.float()
def __rmul__(self, v):
self.flow_list = v * self.flow_list
self.image_list = v * self.image_list
return self
def __len__(self):
return len(self.image_list)
class MpiSintel_submission(FlowDataset):
def __init__(
self, aug_params=None, split="test", root="datasets/Sintel", dstype="clean"
):
super(MpiSintel_submission, self).__init__(aug_params)
flow_root = osp.join(root, split, "flow")
image_root = osp.join(root, split, dstype)
if split == "test":
self.is_test = True
for scene in os.listdir(image_root):
image_list = sorted(glob(osp.join(image_root, scene, "*.png")))
for i in range(len(image_list) - 1):
self.image_list += [[image_list[i], image_list[i + 1]]]
self.extra_info += [(scene, i)] # scene and frame_id
if split != "test":
self.flow_list += sorted(glob(osp.join(flow_root, scene, "*.flo")))
class MpiSintel(FlowDataset):
def __init__(
self, aug_params=None, split="training", root="datasets/Sintel", dstype="clean"
):
super(MpiSintel, self).__init__(aug_params)
root = "s3://"
self.image_list = []
with open("./flow_dataset/Sintel/Sintel_" + dstype + "_png.txt") as f:
images = f.readlines()
for img1, img2 in zip(images[0::2], images[1::2]):
self.image_list.append([root + img1.strip(), root + img2.strip()])
self.flow_list = []
with open("./flow_dataset/Sintel/Sintel_" + dstype + "_flo.txt") as f:
flows = f.readlines()
for flow in flows:
self.flow_list.append(root + flow.strip())
assert len(self.image_list) == len(self.flow_list)
self.extra_info = []
with open("./flow_dataset/Sintel/Sintel_" + dstype + "_extra_info.txt") as f:
info = f.readlines()
for scene, id in zip(info[0::2], info[1::2]):
self.extra_info.append((scene.strip(), int(id.strip())))
# flow_root = osp.join(root, split, 'flow')
# image_root = osp.join(root, split, dstype)
# if split == 'test':
# self.is_test = True
# for scene in os.listdir(image_root):
# image_list = sorted(glob(osp.join(image_root, scene, '*.png')))
# for i in range(len(image_list)-1):
# self.image_list += [ [image_list[i], image_list[i+1]] ]
# self.extra_info += [ (scene, i) ] # scene and frame_id
# if split != 'test':
# self.flow_list += sorted(glob(osp.join(flow_root, scene, '*.flo')))
class FlyingChairs(FlowDataset):
def __init__(
self, aug_params=None, split="train", root="datasets/FlyingChairs_release/data"
):
super(FlyingChairs, self).__init__(aug_params)
root = "s3://"
with open("./flow_dataset/flying_chairs/flyingchairs_ppm.txt") as f:
images = f.readlines()
images = [root + img.strip() for img in images]
with open("./flow_dataset/flying_chairs/flyingchairs_flo.txt") as f:
flows = f.readlines()
flows = [root + flo.strip() for flo in flows]
# images = sorted(glob(osp.join(root, '*.ppm')))
# flows = sorted(glob(osp.join(root, '*.flo')))
assert len(images) // 2 == len(flows)
split_list = np.loadtxt("chairs_split.txt", dtype=np.int32)
for i in range(len(flows)):
xid = split_list[i]
if (split == "training" and xid == 1) or (
split == "validation" and xid == 2
):
self.flow_list += [flows[i]]
self.image_list += [[images[2 * i], images[2 * i + 1]]]
class FlyingThings3D(FlowDataset):
def __init__(
self, aug_params=None, root="datasets/FlyingThings3D", dstype="frames_cleanpass"
):
super(FlyingThings3D, self).__init__(aug_params)
root = "s3://"
self.image_list = []
with open(
"./flow_dataset/flying_things/flyingthings_" + dstype + "_png.txt"
) as f:
images = f.readlines()
for img1, img2 in zip(images[0::2], images[1::2]):
self.image_list.append([root + img1.strip(), root + img2.strip()])
self.flow_list = []
with open(
"./flow_dataset/flying_things/flyingthings_" + dstype + "_pfm.txt"
) as f:
flows = f.readlines()
for flow in flows:
self.flow_list.append(root + flow.strip())
# for cam in ['left']:
# for direction in ['into_future', 'into_past']:
# image_dirs = sorted(glob(osp.join(root, dstype, 'TRAIN/*/*')))
# image_dirs = sorted([osp.join(f, cam) for f in image_dirs])
# flow_dirs = sorted(glob(osp.join(root, 'optical_flow/TRAIN/*/*')))
# flow_dirs = sorted([osp.join(f, direction, cam) for f in flow_dirs])
# for idir, fdir in zip(image_dirs, flow_dirs):
# images = sorted(glob(osp.join(idir, '*.png')) )
# flows = sorted(glob(osp.join(fdir, '*.pfm')) )
# for i in range(len(flows)-1):
# if direction == 'into_future':
# self.image_list += [ [images[i], images[i+1]] ]
# self.flow_list += [ flows[i] ]
# elif direction == 'into_past':
# self.image_list += [ [images[i+1], images[i]] ]
# self.flow_list += [ flows[i+1] ]
class KITTI(FlowDataset):
def __init__(self, aug_params=None, split="training", root="datasets/KITTI"):
super(KITTI, self).__init__(aug_params, sparse=True)
if split == "testing":
self.is_test = True
root = "s3://"
self.image_list = []
with open("./flow_dataset/KITTI/KITTI_{}_image.txt".format(split)) as f:
images = f.readlines()
for img1, img2 in zip(images[0::2], images[1::2]):
self.image_list.append([root + img1.strip(), root + img2.strip()])
self.extra_info = []
with open("./flow_dataset/KITTI/KITTI_{}_extra_info.txt".format(split)) as f:
info = f.readlines()
for id in info:
self.extra_info.append([id.strip()])
if split == "training":
self.flow_list = []
with open("./flow_dataset/KITTI/KITTI_{}_flow.txt".format(split)) as f:
flow = f.readlines()
for flo in flow:
self.flow_list.append(root + flo.strip())
# root = osp.join(root, split)
# images1 = sorted(glob(osp.join(root, 'image_2/*_10.png')))
# images2 = sorted(glob(osp.join(root, 'image_2/*_11.png')))
# for img1, img2 in zip(images1, images2):
# frame_id = img1.split('/')[-1]
# self.extra_info += [ [frame_id] ]
# self.image_list += [ [img1, img2] ]
# if split == 'training':
# self.flow_list = sorted(glob(osp.join(root, 'flow_occ/*_10.png')))
class AutoFlow(data.Dataset):
def __init__(self, num_steps, crop_size, log_dir, root="datasets/"):
super(AutoFlow, self).__init__()
root = "s3://"
self.image_list = []
with open("./flow_dataset/AutoFlow/AutoFlow_image.txt") as f:
images = f.readlines()
for img1, img2 in zip(images[0::2], images[1::2]):
self.image_list.append([root + img1.strip(), root + img2.strip()])
self.flow_list = []
with open("./flow_dataset/AutoFlow/AutoFlow_flow.txt") as f:
flows = f.readlines()
for flow in flows:
self.flow_list.append(root + flow.strip())
self.crop_size = crop_size
self.log_dir = log_dir
self.num_steps = num_steps
self.scale = 1
self.order = 1
self.black = False
self.noise = 0
self.is_test = False
self.init_seed = False
self.iter_counts = 0
def __rmul__(self, v):
self.flow_list = v * self.flow_list
self.image_list = v * self.image_list
return self
def __len__(self):
return len(self.image_list) * 100
def __getitem__(self, index):
# print(self.flow_list[index])
if self.is_test:
img1 = frame_utils.read_gen(self.image_list[index][0], test=self.is_test)
img2 = frame_utils.read_gen(self.image_list[index][1], test=self.is_test)
img1 = np.array(img1).astype(np.uint8)[..., :3]
img2 = np.array(img2).astype(np.uint8)[..., :3]
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
return img1, img2, self.extra_info[index]
if not self.init_seed:
worker_info = torch.utils.data.get_worker_info()
if worker_info is not None:
torch.manual_seed(worker_info.id)
np.random.seed(worker_info.id)
random.seed(worker_info.id)
self.init_seed = True
index = index % len(self.image_list)
valid = None
flow = frame_utils.read_gen(self.flow_list[index])
img1 = frame_utils.read_gen(self.image_list[index][0])
img2 = frame_utils.read_gen(self.image_list[index][1])
flow = np.array(flow).astype(np.float32)
# For PWC-style augmentation, pixel values are in [0, 1]
img1 = np.array(img1).astype(np.uint8) / 255.0
img2 = np.array(img2).astype(np.uint8) / 255.0
# grayscale images
if len(img1.shape) == 2:
img1 = np.tile(img1[..., None], (1, 1, 3))
img2 = np.tile(img2[..., None], (1, 1, 3))
else:
img1 = img1[..., :3]
img2 = img2[..., :3]
iter_counts = self.iter_counts
self.iter_counts = self.iter_counts + 1
print(self.iter_counts)
th, tw = self.crop_size
schedule = [0.5, 1.0, self.num_steps] # initial coeff, final_coeff, half life
schedule_coeff = schedule[0] + (schedule[1] - schedule[0]) * (
2 / (1 + np.exp(-1.0986 * iter_counts / schedule[2])) - 1
)
co_transform = flow_transforms.Compose(
[
flow_transforms.Scale(self.scale, order=self.order),
flow_transforms.SpatialAug(
[th, tw],
scale=[0.4, 0.03, 0.2],
rot=[0.4, 0.03],
trans=[0.4, 0.03],
squeeze=[0.3, 0.0],
schedule_coeff=schedule_coeff,
order=self.order,
black=self.black,
),
flow_transforms.PCAAug(schedule_coeff=schedule_coeff),
flow_transforms.ChromaticAug(
schedule_coeff=schedule_coeff, noise=self.noise
),
]
)
flow = np.concatenate(
[flow, np.ones((flow.shape[0], flow.shape[1], 1))], axis=-1
)
augmented, flow_valid = co_transform([img1, img2], flow)
flow = flow_valid[:, :, :2]
valid = flow_valid[:, :, 2:3]
img1 = augmented[0]
img2 = augmented[1]
if np.random.binomial(1, 0.5):
# sx = int(np.random.uniform(25,100))
# sy = int(np.random.uniform(25,100))
sx = int(np.random.uniform(50, 125))
sy = int(np.random.uniform(50, 125))
# sx = int(np.random.uniform(50,150))
# sy = int(np.random.uniform(50,150))
cx = int(np.random.uniform(sx, img2.shape[0] - sx))
cy = int(np.random.uniform(sy, img2.shape[1] - sy))
img2[cx - sx : cx + sx, cy - sy : cy + sy] = np.mean(np.mean(img2, 0), 0)[
np.newaxis, np.newaxis
]
img1 = torch.from_numpy(img1).permute(2, 0, 1).float()
img2 = torch.from_numpy(img2).permute(2, 0, 1).float()
flow = torch.from_numpy(flow).permute(2, 0, 1).float()
if valid is not None:
valid = torch.from_numpy(valid).permute(2, 0, 1).float()
valid = valid[0]
else:
valid = (flow[0].abs() < 1000) & (flow[1].abs() < 1000)
return img1 * 255, img2 * 255, flow, valid.float()
class HD1K(FlowDataset):
def __init__(self, aug_params=None, root="datasets/HD1k"):
super(HD1K, self).__init__(aug_params, sparse=True)
root = "s3://"
self.image_list = []
with open("./flow_dataset/HD1K/HD1K_image.txt") as f:
images = f.readlines()
for img1, img2 in zip(images[0::2], images[1::2]):
self.image_list.append([root + img1.strip(), root + img2.strip()])
self.flow_list = []
with open("./flow_dataset/HD1K/HD1K_flow.txt") as f:
flows = f.readlines()
for flow in flows:
self.flow_list.append(root + flow.strip())
# seq_ix = 0
# while 1:
# flows = sorted(glob(os.path.join(root, 'hd1k_flow_gt', 'flow_occ/%06d_*.png' % seq_ix)))
# images = sorted(glob(os.path.join(root, 'hd1k_input', 'image_2/%06d_*.png' % seq_ix)))
# if len(flows) == 0:
# break
# for i in range(len(flows)-1):
# self.flow_list += [flows[i]]
# self.image_list += [ [images[i], images[i+1]] ]
# seq_ix += 1
def fetch_dataloader(args, TRAIN_DS="C+T+K+S+H"):
"""Create the data loader for the corresponding trainign set"""
if args.stage == "chairs":
if hasattr(args.percostformer, "pwc_aug") and args.percostformer.pwc_aug:
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.1,
"max_scale": 1.0,
"do_flip": True,
"pwc_aug": True,
}
else:
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.1,
"max_scale": 1.0,
"do_flip": True,
}
train_dataset = FlyingChairs(aug_params, split="training")
elif args.stage == "things":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.4,
"max_scale": 0.8,
"do_flip": True,
}
clean_dataset = FlyingThings3D(aug_params, dstype="frames_cleanpass")
final_dataset = FlyingThings3D(aug_params, dstype="frames_finalpass")
train_dataset = clean_dataset + final_dataset
elif args.stage == "sintel":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.2,
"max_scale": 0.6,
"do_flip": True,
}
things = FlyingThings3D(aug_params, dstype="frames_cleanpass")
sintel_clean = MpiSintel(aug_params, split="training", dstype="clean")
sintel_final = MpiSintel(aug_params, split="training", dstype="final")
if TRAIN_DS == "C+T+K+S+H":
kitti = KITTI(
{
"crop_size": args.image_size,
"min_scale": -0.3,
"max_scale": 0.5,
"do_flip": True,
}
)
hd1k = HD1K(
{
"crop_size": args.image_size,
"min_scale": -0.5,
"max_scale": 0.2,
"do_flip": True,
}
)
train_dataset = (
100 * sintel_clean
+ 100 * sintel_final
+ 200 * kitti
+ 5 * hd1k
+ things
)
elif TRAIN_DS == "C+T+K/S":
train_dataset = 100 * sintel_clean + 100 * sintel_final + things
elif args.stage == "kitti":
aug_params = {
"crop_size": args.image_size,
"min_scale": -0.2,
"max_scale": 0.4,
"do_flip": False,
}
train_dataset = KITTI(aug_params, split="training")
elif args.stage == "autoflow-pwcaug":
aug_params = {
"num_steps": args.trainer.num_steps,
"crop_size": args.image_size,
"log_dir": args.log_dir,
}
train_dataset = AutoFlow(**aug_params)
train_loader = data.DataLoader(
train_dataset,
batch_size=args.batch_size,
pin_memory=False,
shuffle=True,
num_workers=args.batch_size,
drop_last=True,
)
print("Training with %d image pairs" % len(train_dataset))
return train_loader
if __name__ == "__main__":
aug_params = {
"crop_size": [400, 720],
"min_scale": -0.2,
"max_scale": 0,
"do_flip": True,
}
aug_params["min_scale"] = -0.2
aug_params["min_stretch"] = -0.2
sintel_clean = MpiSintel(aug_params, split="training", dstype="clean")
train_loader = data.DataLoader(
sintel_clean,
batch_size=1,
pin_memory=False,
shuffle=True,
num_workers=1,
drop_last=True,
)
for i_batch, data_blob in enumerate(train_loader):
image1, image2, flow, valid = [x for x in data_blob]
print(i_batch, image1.shape)
# if i_batch==5:
# exit()
@@ -0,0 +1,657 @@
from __future__ import division
import torch
import random
import numpy as np
import numbers
import types
import scipy.ndimage as ndimage
import pdb
import torchvision
import PIL.Image as Image
import cv2
from torch.nn import functional as F
class Compose(object):
"""Composes several co_transforms together.
For example:
>>> co_transforms.Compose([
>>> co_transforms.CenterCrop(10),
>>> co_transforms.ToTensor(),
>>> ])
"""
def __init__(self, co_transforms):
self.co_transforms = co_transforms
def __call__(self, input, target):
for t in self.co_transforms:
input, target = t(input, target)
return input, target
class Scale(object):
"""Rescales the inputs and target arrays to the given 'size'.
'size' will be the size of the smaller edge.
For example, if height > width, then image will be
rescaled to (size * height / width, size)
size: size of the smaller edge
interpolation order: Default: 2 (bilinear)
"""
def __init__(self, size, order=1):
self.ratio = size
self.order = order
if order == 0:
self.code = cv2.INTER_NEAREST
elif order == 1:
self.code = cv2.INTER_LINEAR
elif order == 2:
self.code = cv2.INTER_CUBIC
def __call__(self, inputs, target):
if self.ratio == 1:
return inputs, target
h, w, _ = inputs[0].shape
ratio = self.ratio
inputs[0] = cv2.resize(
inputs[0], None, fx=ratio, fy=ratio, interpolation=cv2.INTER_LINEAR
)
inputs[1] = cv2.resize(
inputs[1], None, fx=ratio, fy=ratio, interpolation=cv2.INTER_LINEAR
)
# keep the mask same
tmp = cv2.resize(
target[:, :, 2], None, fx=ratio, fy=ratio, interpolation=cv2.INTER_NEAREST
)
target = (
cv2.resize(target, None, fx=ratio, fy=ratio, interpolation=self.code)
* ratio
)
target[:, :, 2] = tmp
return inputs, target
class SpatialAug(object):
def __init__(
self,
crop,
scale=None,
rot=None,
trans=None,
squeeze=None,
schedule_coeff=1,
order=1,
black=False,
):
self.crop = crop
self.scale = scale
self.rot = rot
self.trans = trans
self.squeeze = squeeze
self.t = np.zeros(6)
self.schedule_coeff = schedule_coeff
self.order = order
self.black = black
def to_identity(self):
self.t[0] = 1
self.t[2] = 0
self.t[4] = 0
self.t[1] = 0
self.t[3] = 1
self.t[5] = 0
def left_multiply(self, u0, u1, u2, u3, u4, u5):
result = np.zeros(6)
result[0] = self.t[0] * u0 + self.t[1] * u2
result[1] = self.t[0] * u1 + self.t[1] * u3
result[2] = self.t[2] * u0 + self.t[3] * u2
result[3] = self.t[2] * u1 + self.t[3] * u3
result[4] = self.t[4] * u0 + self.t[5] * u2 + u4
result[5] = self.t[4] * u1 + self.t[5] * u3 + u5
self.t = result
def inverse(self):
result = np.zeros(6)
a = self.t[0]
c = self.t[2]
e = self.t[4]
b = self.t[1]
d = self.t[3]
f = self.t[5]
denom = a * d - b * c
result[0] = d / denom
result[1] = -b / denom
result[2] = -c / denom
result[3] = a / denom
result[4] = (c * f - d * e) / denom
result[5] = (b * e - a * f) / denom
return result
def grid_transform(self, meshgrid, t, normalize=True, gridsize=None):
if gridsize is None:
h, w = meshgrid[0].shape
else:
h, w = gridsize
vgrid = torch.cat(
[
(meshgrid[0] * t[0] + meshgrid[1] * t[2] + t[4])[:, :, np.newaxis],
(meshgrid[0] * t[1] + meshgrid[1] * t[3] + t[5])[:, :, np.newaxis],
],
-1,
)
if normalize:
vgrid[:, :, 0] = 2.0 * vgrid[:, :, 0] / max(w - 1, 1) - 1.0
vgrid[:, :, 1] = 2.0 * vgrid[:, :, 1] / max(h - 1, 1) - 1.0
return vgrid
def __call__(self, inputs, target):
h, w, _ = inputs[0].shape
th, tw = self.crop
meshgrid = torch.meshgrid([torch.Tensor(range(th)), torch.Tensor(range(tw))])[
::-1
]
cornergrid = torch.meshgrid(
[torch.Tensor([0, th - 1]), torch.Tensor([0, tw - 1])]
)[::-1]
for i in range(50):
# im0
self.to_identity()
# TODO add mirror
if np.random.binomial(1, 0.5):
mirror = True
else:
mirror = False
##TODO
# mirror = False
if mirror:
self.left_multiply(-1, 0, 0, 1, 0.5 * tw, -0.5 * th)
else:
self.left_multiply(1, 0, 0, 1, -0.5 * tw, -0.5 * th)
scale0 = 1
scale1 = 1
squeeze0 = 1
squeeze1 = 1
if not self.rot is None:
rot0 = np.random.uniform(-self.rot[0], +self.rot[0])
rot1 = (
np.random.uniform(
-self.rot[1] * self.schedule_coeff,
self.rot[1] * self.schedule_coeff,
)
+ rot0
)
self.left_multiply(
np.cos(rot0), np.sin(rot0), -np.sin(rot0), np.cos(rot0), 0, 0
)
if not self.trans is None:
trans0 = np.random.uniform(-self.trans[0], +self.trans[0], 2)
trans1 = (
np.random.uniform(
-self.trans[1] * self.schedule_coeff,
+self.trans[1] * self.schedule_coeff,
2,
)
+ trans0
)
self.left_multiply(1, 0, 0, 1, trans0[0] * tw, trans0[1] * th)
if not self.squeeze is None:
squeeze0 = np.exp(np.random.uniform(-self.squeeze[0], self.squeeze[0]))
squeeze1 = (
np.exp(
np.random.uniform(
-self.squeeze[1] * self.schedule_coeff,
self.squeeze[1] * self.schedule_coeff,
)
)
* squeeze0
)
if not self.scale is None:
scale0 = np.exp(
np.random.uniform(
self.scale[2] - self.scale[0], self.scale[2] + self.scale[0]
)
)
scale1 = (
np.exp(
np.random.uniform(
-self.scale[1] * self.schedule_coeff,
self.scale[1] * self.schedule_coeff,
)
)
* scale0
)
self.left_multiply(
1.0 / (scale0 * squeeze0), 0, 0, 1.0 / (scale0 / squeeze0), 0, 0
)
self.left_multiply(1, 0, 0, 1, 0.5 * w, 0.5 * h)
transmat0 = self.t.copy()
# im1
self.to_identity()
if mirror:
self.left_multiply(-1, 0, 0, 1, 0.5 * tw, -0.5 * th)
else:
self.left_multiply(1, 0, 0, 1, -0.5 * tw, -0.5 * th)
if not self.rot is None:
self.left_multiply(
np.cos(rot1), np.sin(rot1), -np.sin(rot1), np.cos(rot1), 0, 0
)
if not self.trans is None:
self.left_multiply(1, 0, 0, 1, trans1[0] * tw, trans1[1] * th)
self.left_multiply(
1.0 / (scale1 * squeeze1), 0, 0, 1.0 / (scale1 / squeeze1), 0, 0
)
self.left_multiply(1, 0, 0, 1, 0.5 * w, 0.5 * h)
transmat1 = self.t.copy()
transmat1_inv = self.inverse()
if self.black:
# black augmentation, allowing 0 values in the input images
# https://github.com/lmb-freiburg/flownet2/blob/master/src/caffe/layers/black_augmentation_layer.cu
break
else:
if (
(
self.grid_transform(
cornergrid, transmat0, gridsize=[float(h), float(w)]
).abs()
> 1
).sum()
+ (
self.grid_transform(
cornergrid, transmat1, gridsize=[float(h), float(w)]
).abs()
> 1
).sum()
) == 0:
break
if i == 49:
print("max_iter in augmentation")
self.to_identity()
self.left_multiply(1, 0, 0, 1, -0.5 * tw, -0.5 * th)
self.left_multiply(1, 0, 0, 1, 0.5 * w, 0.5 * h)
transmat0 = self.t.copy()
transmat1 = self.t.copy()
# do the real work
vgrid = self.grid_transform(meshgrid, transmat0, gridsize=[float(h), float(w)])
inputs_0 = F.grid_sample(
torch.Tensor(inputs[0]).permute(2, 0, 1)[np.newaxis], vgrid[np.newaxis]
)[0].permute(1, 2, 0)
if self.order == 0:
target_0 = F.grid_sample(
torch.Tensor(target).permute(2, 0, 1)[np.newaxis],
vgrid[np.newaxis],
mode="nearest",
)[0].permute(1, 2, 0)
else:
target_0 = F.grid_sample(
torch.Tensor(target).permute(2, 0, 1)[np.newaxis], vgrid[np.newaxis]
)[0].permute(1, 2, 0)
mask_0 = target[:, :, 2:3].copy()
mask_0[mask_0 == 0] = np.nan
if self.order == 0:
mask_0 = F.grid_sample(
torch.Tensor(mask_0).permute(2, 0, 1)[np.newaxis],
vgrid[np.newaxis],
mode="nearest",
)[0].permute(1, 2, 0)
else:
mask_0 = F.grid_sample(
torch.Tensor(mask_0).permute(2, 0, 1)[np.newaxis], vgrid[np.newaxis]
)[0].permute(1, 2, 0)
mask_0[torch.isnan(mask_0)] = 0
vgrid = self.grid_transform(meshgrid, transmat1, gridsize=[float(h), float(w)])
inputs_1 = F.grid_sample(
torch.Tensor(inputs[1]).permute(2, 0, 1)[np.newaxis], vgrid[np.newaxis]
)[0].permute(1, 2, 0)
# flow
pos = target_0[:, :, :2] + self.grid_transform(
meshgrid, transmat0, normalize=False
)
pos = self.grid_transform(pos.permute(2, 0, 1), transmat1_inv, normalize=False)
if target_0.shape[2] >= 4:
# scale
exp = target_0[:, :, 3:] * scale1 / scale0
target = torch.cat(
[
(pos[:, :, 0] - meshgrid[0]).unsqueeze(-1),
(pos[:, :, 1] - meshgrid[1]).unsqueeze(-1),
mask_0,
exp,
],
-1,
)
else:
target = torch.cat(
[
(pos[:, :, 0] - meshgrid[0]).unsqueeze(-1),
(pos[:, :, 1] - meshgrid[1]).unsqueeze(-1),
mask_0,
],
-1,
)
# target_0[:,:,2].unsqueeze(-1) ], -1)
inputs = [np.asarray(inputs_0), np.asarray(inputs_1)]
target = np.asarray(target)
return inputs, target
class pseudoPCAAug(object):
"""
Chromatic Eigen Augmentation: https://github.com/lmb-freiburg/flownet2/blob/master/src/caffe/layers/data_augmentation_layer.cu
This version is faster.
"""
def __init__(self, schedule_coeff=1):
self.augcolor = torchvision.transforms.ColorJitter(
brightness=0.4, contrast=0.4, saturation=0.5, hue=0.5 / 3.14
)
def __call__(self, inputs, target):
inputs[0] = (
np.asarray(self.augcolor(Image.fromarray(np.uint8(inputs[0] * 255))))
/ 255.0
)
inputs[1] = (
np.asarray(self.augcolor(Image.fromarray(np.uint8(inputs[1] * 255))))
/ 255.0
)
return inputs, target
class PCAAug(object):
"""
Chromatic Eigen Augmentation: https://github.com/lmb-freiburg/flownet2/blob/master/src/caffe/layers/data_augmentation_layer.cu
"""
def __init__(
self,
lmult_pow=[0.4, 0, -0.2],
lmult_mult=[
0.4,
0,
0,
],
lmult_add=[
0.03,
0,
0,
],
sat_pow=[
0.4,
0,
0,
],
sat_mult=[0.5, 0, -0.3],
sat_add=[
0.03,
0,
0,
],
col_pow=[
0.4,
0,
0,
],
col_mult=[
0.2,
0,
0,
],
col_add=[
0.02,
0,
0,
],
ladd_pow=[
0.4,
0,
0,
],
ladd_mult=[
0.4,
0,
0,
],
ladd_add=[
0.04,
0,
0,
],
col_rotate=[
1.0,
0,
0,
],
schedule_coeff=1,
):
# no mean
self.pow_nomean = [1, 1, 1]
self.add_nomean = [0, 0, 0]
self.mult_nomean = [1, 1, 1]
self.pow_withmean = [1, 1, 1]
self.add_withmean = [0, 0, 0]
self.mult_withmean = [1, 1, 1]
self.lmult_pow = 1
self.lmult_mult = 1
self.lmult_add = 0
self.col_angle = 0
if not ladd_pow is None:
self.pow_nomean[0] = np.exp(np.random.normal(ladd_pow[2], ladd_pow[0]))
if not col_pow is None:
self.pow_nomean[1] = np.exp(np.random.normal(col_pow[2], col_pow[0]))
self.pow_nomean[2] = np.exp(np.random.normal(col_pow[2], col_pow[0]))
if not ladd_add is None:
self.add_nomean[0] = np.random.normal(ladd_add[2], ladd_add[0])
if not col_add is None:
self.add_nomean[1] = np.random.normal(col_add[2], col_add[0])
self.add_nomean[2] = np.random.normal(col_add[2], col_add[0])
if not ladd_mult is None:
self.mult_nomean[0] = np.exp(np.random.normal(ladd_mult[2], ladd_mult[0]))
if not col_mult is None:
self.mult_nomean[1] = np.exp(np.random.normal(col_mult[2], col_mult[0]))
self.mult_nomean[2] = np.exp(np.random.normal(col_mult[2], col_mult[0]))
# with mean
if not sat_pow is None:
self.pow_withmean[1] = np.exp(
np.random.uniform(sat_pow[2] - sat_pow[0], sat_pow[2] + sat_pow[0])
)
self.pow_withmean[2] = self.pow_withmean[1]
if not sat_add is None:
self.add_withmean[1] = np.random.uniform(
sat_add[2] - sat_add[0], sat_add[2] + sat_add[0]
)
self.add_withmean[2] = self.add_withmean[1]
if not sat_mult is None:
self.mult_withmean[1] = np.exp(
np.random.uniform(sat_mult[2] - sat_mult[0], sat_mult[2] + sat_mult[0])
)
self.mult_withmean[2] = self.mult_withmean[1]
if not lmult_pow is None:
self.lmult_pow = np.exp(
np.random.uniform(
lmult_pow[2] - lmult_pow[0], lmult_pow[2] + lmult_pow[0]
)
)
if not lmult_mult is None:
self.lmult_mult = np.exp(
np.random.uniform(
lmult_mult[2] - lmult_mult[0], lmult_mult[2] + lmult_mult[0]
)
)
if not lmult_add is None:
self.lmult_add = np.random.uniform(
lmult_add[2] - lmult_add[0], lmult_add[2] + lmult_add[0]
)
if not col_rotate is None:
self.col_angle = np.random.uniform(
col_rotate[2] - col_rotate[0], col_rotate[2] + col_rotate[0]
)
# eigen vectors
self.eigvec = np.reshape(
[0.51, 0.56, 0.65, 0.79, 0.01, -0.62, 0.35, -0.83, 0.44], [3, 3]
).transpose()
def __call__(self, inputs, target):
inputs[0] = self.pca_image(inputs[0])
inputs[1] = self.pca_image(inputs[1])
return inputs, target
def pca_image(self, rgb):
eig = np.dot(rgb, self.eigvec)
max_rgb = np.clip(rgb, 0, np.inf).max((0, 1))
min_rgb = rgb.min((0, 1))
mean_rgb = rgb.mean((0, 1))
max_abs_eig = np.abs(eig).max((0, 1))
max_l = np.sqrt(np.sum(max_abs_eig * max_abs_eig))
mean_eig = np.dot(mean_rgb, self.eigvec)
# no-mean stuff
eig -= mean_eig[np.newaxis, np.newaxis]
for c in range(3):
if max_abs_eig[c] > 1e-2:
mean_eig[c] /= max_abs_eig[c]
eig[:, :, c] = eig[:, :, c] / max_abs_eig[c]
eig[:, :, c] = (
np.power(np.abs(eig[:, :, c]), self.pow_nomean[c])
* ((eig[:, :, c] > 0) - 0.5)
* 2
)
eig[:, :, c] = eig[:, :, c] + self.add_nomean[c]
eig[:, :, c] = eig[:, :, c] * self.mult_nomean[c]
eig += mean_eig[np.newaxis, np.newaxis]
# withmean stuff
if max_abs_eig[0] > 1e-2:
eig[:, :, 0] = (
np.power(np.abs(eig[:, :, 0]), self.pow_withmean[0])
* ((eig[:, :, 0] > 0) - 0.5)
* 2
)
eig[:, :, 0] = eig[:, :, 0] + self.add_withmean[0]
eig[:, :, 0] = eig[:, :, 0] * self.mult_withmean[0]
s = np.sqrt(eig[:, :, 1] * eig[:, :, 1] + eig[:, :, 2] * eig[:, :, 2])
smask = s > 1e-2
s1 = np.power(s, self.pow_withmean[1])
s1 = np.clip(s1 + self.add_withmean[1], 0, np.inf)
s1 = s1 * self.mult_withmean[1]
s1 = s1 * smask + s * (1 - smask)
# color angle
if self.col_angle != 0:
temp1 = (
np.cos(self.col_angle) * eig[:, :, 1]
- np.sin(self.col_angle) * eig[:, :, 2]
)
temp2 = (
np.sin(self.col_angle) * eig[:, :, 1]
+ np.cos(self.col_angle) * eig[:, :, 2]
)
eig[:, :, 1] = temp1
eig[:, :, 2] = temp2
# to origin magnitude
for c in range(3):
if max_abs_eig[c] > 1e-2:
eig[:, :, c] = eig[:, :, c] * max_abs_eig[c]
if max_l > 1e-2:
l1 = np.sqrt(
eig[:, :, 0] * eig[:, :, 0]
+ eig[:, :, 1] * eig[:, :, 1]
+ eig[:, :, 2] * eig[:, :, 2]
)
l1 = l1 / max_l
eig[:, :, 1][smask] = (eig[:, :, 1] / s * s1)[smask]
eig[:, :, 2][smask] = (eig[:, :, 2] / s * s1)[smask]
# eig[:,:,1] = (eig[:,:,1] / s * s1) * smask + eig[:,:,1] * (1-smask)
# eig[:,:,2] = (eig[:,:,2] / s * s1) * smask + eig[:,:,2] * (1-smask)
if max_l > 1e-2:
l = np.sqrt(
eig[:, :, 0] * eig[:, :, 0]
+ eig[:, :, 1] * eig[:, :, 1]
+ eig[:, :, 2] * eig[:, :, 2]
)
l1 = np.power(l1, self.lmult_pow)
l1 = np.clip(l1 + self.lmult_add, 0, np.inf)
l1 = l1 * self.lmult_mult
l1 = l1 * max_l
lmask = l > 1e-2
eig[lmask] = (eig / l[:, :, np.newaxis] * l1[:, :, np.newaxis])[lmask]
for c in range(3):
eig[:, :, c][lmask] = (np.clip(eig[:, :, c], -np.inf, max_abs_eig[c]))[
lmask
]
# for c in range(3):
# # eig[:,:,c][lmask] = (eig[:,:,c] / l * l1)[lmask] * lmask + eig[:,:,c] * (1-lmask)
# eig[:,:,c][lmask] = (eig[:,:,c] / l * l1)[lmask]
# eig[:,:,c] = (np.clip(eig[:,:,c], -np.inf, max_abs_eig[c])) * lmask + eig[:,:,c] * (1-lmask)
return np.clip(np.dot(eig, self.eigvec.transpose()), 0, 1)
class ChromaticAug(object):
"""
Chromatic augmentation: https://github.com/lmb-freiburg/flownet2/blob/master/src/caffe/layers/data_augmentation_layer.cu
"""
def __init__(
self,
noise=0.06,
gamma=0.02,
brightness=0.02,
contrast=0.02,
color=0.02,
schedule_coeff=1,
):
self.noise = np.random.uniform(0, noise)
self.gamma = np.exp(np.random.normal(0, gamma * schedule_coeff))
self.brightness = np.random.normal(0, brightness * schedule_coeff)
self.contrast = np.exp(np.random.normal(0, contrast * schedule_coeff))
self.color = np.exp(np.random.normal(0, color * schedule_coeff, 3))
def __call__(self, inputs, target):
inputs[1] = self.chrom_aug(inputs[1])
# noise
inputs[0] += np.random.normal(0, self.noise, inputs[0].shape)
inputs[1] += np.random.normal(0, self.noise, inputs[0].shape)
return inputs, target
def chrom_aug(self, rgb):
# color change
mean_in = rgb.sum(-1)
rgb = rgb * self.color[np.newaxis, np.newaxis]
brightness_coeff = mean_in / (rgb.sum(-1) + 0.01)
rgb = np.clip(rgb * brightness_coeff[:, :, np.newaxis], 0, 1)
# gamma
rgb = np.power(rgb, self.gamma)
# brightness
rgb += self.brightness
# contrast
rgb = 0.5 + (rgb - 0.5) * self.contrast
rgb = np.clip(rgb, 0, 1)
return
@@ -0,0 +1,136 @@
# Flow visualization code used from https://github.com/tomrunia/OpticalFlow_Visualization
# MIT License
#
# Copyright (c) 2018 Tom Runia
#
# 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 conditions.
#
# Author: Tom Runia
# Date Created: 2018-08-03
import numpy as np
def make_colorwheel():
"""
Generates a color wheel for optical flow visualization as presented in:
Baker et al. "A Database and Evaluation Methodology for Optical Flow" (ICCV, 2007)
URL: http://vision.middlebury.edu/flow/flowEval-iccv07.pdf
Code follows the original C++ source code of Daniel Scharstein.
Code follows the the Matlab source code of Deqing Sun.
Returns:
np.ndarray: Color wheel
"""
RY = 15
YG = 6
GC = 4
CB = 11
BM = 13
MR = 6
ncols = RY + YG + GC + CB + BM + MR
colorwheel = np.zeros((ncols, 3))
col = 0
# RY
colorwheel[0:RY, 0] = 255
colorwheel[0:RY, 1] = np.floor(255 * np.arange(0, RY) / RY)
col = col + RY
# YG
colorwheel[col : col + YG, 0] = 255 - np.floor(255 * np.arange(0, YG) / YG)
colorwheel[col : col + YG, 1] = 255
col = col + YG
# GC
colorwheel[col : col + GC, 1] = 255
colorwheel[col : col + GC, 2] = np.floor(255 * np.arange(0, GC) / GC)
col = col + GC
# CB
colorwheel[col : col + CB, 1] = 255 - np.floor(255 * np.arange(CB) / CB)
colorwheel[col : col + CB, 2] = 255
col = col + CB
# BM
colorwheel[col : col + BM, 2] = 255
colorwheel[col : col + BM, 0] = np.floor(255 * np.arange(0, BM) / BM)
col = col + BM
# MR
colorwheel[col : col + MR, 2] = 255 - np.floor(255 * np.arange(MR) / MR)
colorwheel[col : col + MR, 0] = 255
return colorwheel
def flow_uv_to_colors(u, v, convert_to_bgr=False):
"""
Applies the flow color wheel to (possibly clipped) flow components u and v.
According to the C++ source code of Daniel Scharstein
According to the Matlab source code of Deqing Sun
Args:
u (np.ndarray): Input horizontal flow of shape [H,W]
v (np.ndarray): Input vertical flow of shape [H,W]
convert_to_bgr (bool, optional): Convert output image to BGR. Defaults to False.
Returns:
np.ndarray: Flow visualization image of shape [H,W,3]
"""
flow_image = np.zeros((u.shape[0], u.shape[1], 3), np.uint8)
colorwheel = make_colorwheel() # shape [55x3]
ncols = colorwheel.shape[0]
rad = np.sqrt(np.square(u) + np.square(v))
a = np.arctan2(-v, -u) / np.pi
fk = (a + 1) / 2 * (ncols - 1)
k0 = np.floor(fk).astype(np.int32)
k1 = k0 + 1
k1[k1 == ncols] = 0
f = fk - k0
for i in range(colorwheel.shape[1]):
tmp = colorwheel[:, i]
col0 = tmp[k0] / 255.0
col1 = tmp[k1] / 255.0
col = (1 - f) * col0 + f * col1
idx = rad <= 1
col[idx] = 1 - rad[idx] * (1 - col[idx])
col[~idx] = col[~idx] * 0.75 # out of range
# Note the 2-i => BGR instead of RGB
ch_idx = 2 - i if convert_to_bgr else i
flow_image[:, :, ch_idx] = np.floor(255 * col)
return flow_image
def flow_to_image(flow_uv, clip_flow=None, convert_to_bgr=False, max_flow=None):
"""
Expects a two dimensional flow image of shape.
Args:
flow_uv (np.ndarray): Flow UV image of shape [H,W,2]
clip_flow (float, optional): Clip maximum of flow values. Defaults to None.
convert_to_bgr (bool, optional): Convert output image to BGR. Defaults to False.
Returns:
np.ndarray: Flow visualization image of shape [H,W,3]
"""
assert flow_uv.ndim == 3, "input flow must have three dimensions"
assert flow_uv.shape[2] == 2, "input flow must have shape [H,W,2]"
if clip_flow is not None:
flow_uv = np.clip(flow_uv, 0, clip_flow)
u = flow_uv[:, :, 0]
v = flow_uv[:, :, 1]
if max_flow is None:
rad = np.sqrt(np.square(u) + np.square(v))
rad_max = np.max(rad)
else:
rad_max = max_flow
epsilon = 1e-5
u = u / (rad_max + epsilon)
v = v / (rad_max + epsilon)
return flow_uv_to_colors(u, v, convert_to_bgr)
@@ -0,0 +1,142 @@
import numpy as np
from PIL import Image
from os.path import *
import re
import cv2
cv2.setNumThreads(0)
cv2.ocl.setUseOpenCL(False)
TAG_CHAR = np.array([202021.25], np.float32)
def readFlow(fn):
"""Read .flo file in Middlebury format"""
# Code adapted from:
# http://stackoverflow.com/questions/28013200/reading-middlebury-flow-files-with-python-bytes-array-numpy
# WARNING: this will work on little-endian architectures (eg Intel x86) only!
# print 'fn = %s'%(fn)
with open(fn, "rb") as f:
magic = np.fromfile(f, np.float32, count=1)
if 202021.25 != magic:
print("Magic number incorrect. Invalid .flo file")
return None
else:
w = np.fromfile(f, np.int32, count=1)
h = np.fromfile(f, np.int32, count=1)
# print 'Reading %d x %d flo file\n' % (w, h)
data = np.fromfile(f, np.float32, count=2 * int(w) * int(h))
# Reshape data into 3D array (columns, rows, bands)
# The reshape here is for visualization, the original code is (w,h,2)
return np.resize(data, (int(h), int(w), 2))
def readPFM(file):
file = open(file, "rb")
color = None
width = None
height = None
scale = None
endian = None
header = file.readline().rstrip()
if header == b"PF":
color = True
elif header == b"Pf":
color = False
else:
raise Exception("Not a PFM file.")
dim_match = re.match(rb"^(\d+)\s(\d+)\s$", file.readline())
if dim_match:
width, height = map(int, dim_match.groups())
else:
raise Exception("Malformed PFM header.")
scale = float(file.readline().rstrip())
if scale < 0: # little-endian
endian = "<"
scale = -scale
else:
endian = ">" # big-endian
data = np.fromfile(file, endian + "f")
shape = (height, width, 3) if color else (height, width)
data = np.reshape(data, shape)
data = np.flipud(data)
return data
def writeFlow(filename, uv, v=None):
"""Write optical flow to file.
If v is None, uv is assumed to contain both u and v channels,
stacked in depth.
Original code by Deqing Sun, adapted from Daniel Scharstein.
"""
nBands = 2
if v is None:
assert uv.ndim == 3
assert uv.shape[2] == 2
u = uv[:, :, 0]
v = uv[:, :, 1]
else:
u = uv
assert u.shape == v.shape
height, width = u.shape
f = open(filename, "wb")
# write the header
f.write(TAG_CHAR)
np.array(width).astype(np.int32).tofile(f)
np.array(height).astype(np.int32).tofile(f)
# arrange into matrix form
tmp = np.zeros((height, width * nBands))
tmp[:, np.arange(width) * 2] = u
tmp[:, np.arange(width) * 2 + 1] = v
tmp.astype(np.float32).tofile(f)
f.close()
def readFlowKITTI(filename):
flow = cv2.imread(filename, cv2.IMREAD_ANYDEPTH | cv2.IMREAD_COLOR)
flow = flow[:, :, ::-1].astype(np.float32)
flow, valid = flow[:, :, :2], flow[:, :, 2]
flow = (flow - 2**15) / 64.0
return flow, valid
def readDispKITTI(filename):
disp = cv2.imread(filename, cv2.IMREAD_ANYDEPTH) / 256.0
valid = disp > 0.0
flow = np.stack([-disp, np.zeros_like(disp)], -1)
return flow, valid
def writeFlowKITTI(filename, uv):
uv = 64.0 * uv + 2**15
valid = np.ones([uv.shape[0], uv.shape[1], 1])
uv = np.concatenate([uv, valid], axis=-1).astype(np.uint16)
cv2.imwrite(filename, uv[..., ::-1])
def read_gen(file_name, pil=False):
ext = splitext(file_name)[-1]
if ext == ".png" or ext == ".jpeg" or ext == ".ppm" or ext == ".jpg":
return Image.open(file_name)
elif ext == ".bin" or ext == ".raw":
return np.load(file_name)
elif ext == ".flo":
return readFlow(file_name).astype(np.float32)
elif ext == ".pfm":
flow = readPFM(file_name).astype(np.float32)
if len(flow.shape) == 2:
return flow
else:
return flow[:, :, :-1]
return []
@@ -0,0 +1,60 @@
from torch.utils.tensorboard import SummaryWriter
from loguru import logger as loguru_logger
class Logger:
def __init__(self, model, scheduler, cfg):
self.model = model
self.scheduler = scheduler
self.total_steps = 0
self.running_loss = {}
self.writer = None
self.cfg = cfg
def _print_training_status(self):
metrics_data = [
self.running_loss[k] / self.cfg.sum_freq
for k in sorted(self.running_loss.keys())
]
training_str = "[{:6d}, {}] ".format(
self.total_steps + 1, self.scheduler.get_last_lr()
)
metrics_str = ("{:10.4f}, " * len(metrics_data)).format(*metrics_data)
# print the training status
loguru_logger.info(training_str + metrics_str)
if self.writer is None:
if self.cfg.log_dir is None:
self.writer = SummaryWriter()
else:
self.writer = SummaryWriter(self.cfg.log_dir)
for k in self.running_loss:
self.writer.add_scalar(
k, self.running_loss[k] / self.cfg.sum_freq, self.total_steps
)
self.running_loss[k] = 0.0
def push(self, metrics):
self.total_steps += 1
for key in metrics:
if key not in self.running_loss:
self.running_loss[key] = 0.0
self.running_loss[key] += metrics[key]
if self.total_steps % self.cfg.sum_freq == self.cfg.sum_freq - 1:
self._print_training_status()
self.running_loss = {}
def write_dict(self, results):
if self.writer is None:
self.writer = SummaryWriter()
for key in results:
self.writer.add_scalar(key, results[key], self.total_steps)
def close(self):
self.writer.close()
@@ -0,0 +1,33 @@
import time
import os
import shutil
def process_transformer_cfg(cfg):
log_dir = ""
if "critical_params" in cfg:
critical_params = [cfg[key] for key in cfg.critical_params]
for name, param in zip(cfg["critical_params"], critical_params):
log_dir += "{:s}[{:s}]".format(name, str(param))
return log_dir
def process_cfg(cfg):
log_dir = "logs/" + cfg.name + "/" + cfg.transformer + "/"
critical_params = [cfg.trainer[key] for key in cfg.critical_params]
for name, param in zip(cfg["critical_params"], critical_params):
log_dir += "{:s}[{:s}]".format(name, str(param))
log_dir += process_transformer_cfg(cfg[cfg.transformer])
now = time.localtime()
now_time = "{:02d}_{:02d}_{:02d}_{:02d}".format(
now.tm_mon, now.tm_mday, now.tm_hour, now.tm_min
)
log_dir += cfg.suffix + "(" + now_time + ")"
cfg.log_dir = log_dir
os.makedirs(log_dir)
shutil.copytree("configs", f"{log_dir}/configs")
shutil.copytree("core/FlowFormer", f"{log_dir}/FlowFormer")
@@ -0,0 +1,113 @@
import torch
import torch.nn.functional as F
import numpy as np
from scipy import interpolate
class InputPadder:
"""Pads images such that dimensions are divisible by 8"""
def __init__(self, dims, mode="sintel"):
self.ht, self.wd = dims[-2:]
pad_ht = (((self.ht // 8) + 1) * 8 - self.ht) % 8
pad_wd = (((self.wd // 8) + 1) * 8 - self.wd) % 8
if mode == "sintel":
self._pad = [
pad_wd // 2,
pad_wd - pad_wd // 2,
pad_ht // 2,
pad_ht - pad_ht // 2,
]
elif mode == "kitti400":
self._pad = [0, 0, 0, 400 - self.ht]
else:
self._pad = [pad_wd // 2, pad_wd - pad_wd // 2, 0, pad_ht]
def pad(self, *inputs):
return [F.pad(x, self._pad, mode="replicate") for x in inputs]
def unpad(self, x):
ht, wd = x.shape[-2:]
c = [self._pad[2], ht - self._pad[3], self._pad[0], wd - self._pad[1]]
return x[..., c[0] : c[1], c[2] : c[3]]
def forward_interpolate(flow):
flow = flow.detach().cpu().numpy()
dx, dy = flow[0], flow[1]
ht, wd = dx.shape
x0, y0 = np.meshgrid(np.arange(wd), np.arange(ht))
x1 = x0 + dx
y1 = y0 + dy
x1 = x1.reshape(-1)
y1 = y1.reshape(-1)
dx = dx.reshape(-1)
dy = dy.reshape(-1)
valid = (x1 > 0) & (x1 < wd) & (y1 > 0) & (y1 < ht)
x1 = x1[valid]
y1 = y1[valid]
dx = dx[valid]
dy = dy[valid]
flow_x = interpolate.griddata(
(x1, y1), dx, (x0, y0), method="nearest", fill_value=0
)
flow_y = interpolate.griddata(
(x1, y1), dy, (x0, y0), method="nearest", fill_value=0
)
flow = np.stack([flow_x, flow_y], axis=0)
return torch.from_numpy(flow).float()
def bilinear_sampler(img, coords, mode="bilinear", mask=False):
"""Wrapper for grid_sample, uses pixel coordinates"""
H, W = img.shape[-2:]
xgrid, ygrid = coords.split([1, 1], dim=-1)
xgrid = 2 * xgrid / (W - 1) - 1
ygrid = 2 * ygrid / (H - 1) - 1
grid = torch.cat([xgrid, ygrid], dim=-1)
img = F.grid_sample(img, grid, align_corners=True)
if mask:
mask = (xgrid > -1) & (ygrid > -1) & (xgrid < 1) & (ygrid < 1)
return img, mask.float()
return img
def indexing(img, coords, mask=False):
"""Wrapper for grid_sample, uses pixel coordinates"""
"""
TODO: directly indexing features instead of sampling
"""
H, W = img.shape[-2:]
xgrid, ygrid = coords.split([1, 1], dim=-1)
xgrid = 2 * xgrid / (W - 1) - 1
ygrid = 2 * ygrid / (H - 1) - 1
grid = torch.cat([xgrid, ygrid], dim=-1)
img = F.grid_sample(img, grid, align_corners=True, mode="nearest")
if mask:
mask = (xgrid > -1) & (ygrid > -1) & (xgrid < 1) & (ygrid < 1)
return img, mask.float()
return img
def coords_grid(batch, ht, wd):
coords = torch.meshgrid(torch.arange(ht), torch.arange(wd))
coords = torch.stack(coords[::-1], dim=0).float()
return coords[None].repeat(batch, 1, 1, 1)
def upflow8(flow, mode="bilinear"):
new_size = (8 * flow.shape[2], 8 * flow.shape[3])
return 8 * F.interpolate(flow, size=new_size, mode=mode, align_corners=True)
@@ -0,0 +1,201 @@
import sys
sys.path.append("core")
from PIL import Image
import argparse
import os
import time
import numpy as np
import torch
import torch.nn.functional as F
import matplotlib.pyplot as plt
from configs.default import get_cfg
from configs.things_eval import get_cfg as get_things_cfg
from configs.small_things_eval import get_cfg as get_small_things_cfg
from core.utils.misc import process_cfg
import datasets
from utils import flow_viz
from utils import frame_utils
# from FlowFormer import FlowFormer
from core.FlowFormer import build_flowformer
from raft import RAFT
from utils.utils import InputPadder, forward_interpolate
@torch.no_grad()
def validate_chairs(model):
"""Perform evaluation on the FlyingChairs (test) split"""
model.eval()
epe_list = []
val_dataset = datasets.FlyingChairs(split="validation")
for val_id in range(len(val_dataset)):
image1, image2, flow_gt, _ = val_dataset[val_id]
image1 = image1[None].cuda()
image2 = image2[None].cuda()
flow_pre, _ = model(image1, image2)
epe = torch.sum((flow_pre[0].cpu() - flow_gt) ** 2, dim=0).sqrt()
epe_list.append(epe.view(-1).numpy())
epe = np.mean(np.concatenate(epe_list))
print("Validation Chairs EPE: %f" % epe)
return {"chairs": epe}
@torch.no_grad()
def validate_sintel(model):
"""Peform validation using the Sintel (train) split"""
model.eval()
results = {}
for dstype in ["clean", "final"]:
val_dataset = datasets.MpiSintel(split="training", dstype=dstype)
epe_list = []
for val_id in range(len(val_dataset)):
image1, image2, flow_gt, _ = val_dataset[val_id]
image1 = image1[None].cuda()
image2 = image2[None].cuda()
padder = InputPadder(image1.shape)
image1, image2 = padder.pad(image1, image2)
flow_pre = model(image1, image2)
flow_pre = padder.unpad(flow_pre[0]).cpu()[0]
epe = torch.sum((flow_pre - flow_gt) ** 2, dim=0).sqrt()
epe_list.append(epe.view(-1).numpy())
epe_all = np.concatenate(epe_list)
epe = np.mean(epe_all)
px1 = np.mean(epe_all < 1)
px3 = np.mean(epe_all < 3)
px5 = np.mean(epe_all < 5)
print(
"Validation (%s) EPE: %f, 1px: %f, 3px: %f, 5px: %f"
% (dstype, epe, px1, px3, px5)
)
results[dstype] = np.mean(epe_list)
return results
@torch.no_grad()
def create_sintel_submission(model, output_path="sintel_submission"):
"""Create submission for the Sintel leaderboard"""
model.eval()
for dstype in ["final", "clean"]:
test_dataset = datasets.MpiSintel(split="test", aug_params=None, dstype=dstype)
for test_id in range(len(test_dataset)):
if (test_id + 1) % 100 == 0:
print(f"{test_id} / {len(test_dataset)}")
image1, image2, (sequence, frame) = test_dataset[test_id]
image1, image2 = image1[None].cuda(), image2[None].cuda()
padder = InputPadder(image1.shape)
image1, image2 = padder.pad(image1, image2)
flow_pre = model(image1, image2)
flow_pre = padder.unpad(flow_pre[0]).cpu()
flow = flow_pre[0].permute(1, 2, 0).cpu().numpy()
output_dir = os.path.join(output_path, dstype, sequence)
output_file = os.path.join(output_dir, "frame%04d.flo" % (frame + 1))
if not os.path.exists(output_dir):
os.makedirs(output_dir)
frame_utils.writeFlow(output_file, flow)
@torch.no_grad()
def validate_kitti(model):
"""Peform validation using the KITTI-2015 (train) split"""
model.eval()
val_dataset = datasets.KITTI(split="training")
out_list, epe_list = [], []
for val_id in range(len(val_dataset)):
image1, image2, flow_gt, valid_gt = val_dataset[val_id]
image1 = image1[None].cuda()
image2 = image2[None].cuda()
padder = InputPadder(image1.shape)
image1, image2 = padder.pad(image1, image2)
flow_pre = model(image1, image2)
flow_pre = padder.unpad(flow_pre[0]).cpu()[0]
epe = torch.sum((flow_pre - flow_gt) ** 2, dim=0).sqrt()
mag = torch.sum(flow_gt**2, dim=0).sqrt()
epe = epe.view(-1)
mag = mag.view(-1)
val = valid_gt.view(-1) >= 0.5
out = ((epe > 3.0) & ((epe / mag) > 0.05)).float()
epe_list.append(epe[val].mean().item())
out_list.append(out[val].cpu().numpy())
epe_list = np.array(epe_list)
out_list = np.concatenate(out_list)
epe = np.mean(epe_list)
f1 = 100 * np.mean(out_list)
print("Validation KITTI: %f, %f" % (epe, f1))
return {"kitti-epe": epe, "kitti-f1": f1}
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", help="restore checkpoint")
parser.add_argument("--dataset", help="dataset for evaluation")
parser.add_argument("--small", action="store_true", help="use small model")
parser.add_argument(
"--mixed_precision", action="store_true", help="use mixed precision"
)
parser.add_argument(
"--alternate_corr",
action="store_true",
help="use efficent correlation implementation",
)
args = parser.parse_args()
# cfg = get_cfg()
if args.small:
cfg = get_small_things_cfg()
else:
cfg = get_things_cfg()
cfg.update(vars(args))
model = torch.nn.DataParallel(build_flowformer(cfg))
model.load_state_dict(torch.load(cfg.model))
print(args)
model.cuda()
model.eval()
# create_sintel_submission(model.module, warm_start=True)
# create_kitti_submission(model.module)
with torch.no_grad():
if args.dataset == "chairs":
validate_chairs(model.module)
elif args.dataset == "sintel":
validate_sintel(model.module)
elif args.dataset == "kitti":
validate_kitti(model.module)
elif args.dataset == "sintel_submission":
create_sintel_submission(model.module)
@@ -0,0 +1,410 @@
import sys
from attr import validate
sys.path.append("core")
from PIL import Image
import argparse
import os
import time
import numpy as np
import torch
import torch.nn.functional as F
import matplotlib.pyplot as plt
from configs.submission import get_cfg as get_submission_cfg
# from configs.kitti_submission import get_cfg as get_kitti_cfg
from configs.things_eval import get_cfg as get_things_cfg
from configs.small_things_eval import get_cfg as get_small_things_cfg
from core.utils.misc import process_cfg
import datasets
from utils import flow_viz
from utils import frame_utils
from core.FlowFormer import build_flowformer
from raft import RAFT
from utils.utils import InputPadder, forward_interpolate
import imageio
import itertools
TRAIN_SIZE = [432, 960]
class InputPadder:
"""Pads images such that dimensions are divisible by 8"""
def __init__(self, dims, mode="sintel"):
self.ht, self.wd = dims[-2:]
pad_ht = (((self.ht // 8) + 1) * 8 - self.ht) % 8
pad_wd = (((self.wd // 8) + 1) * 8 - self.wd) % 8
if mode == "sintel":
self._pad = [
pad_wd // 2,
pad_wd - pad_wd // 2,
pad_ht // 2,
pad_ht - pad_ht // 2,
]
elif mode == "kitti432":
self._pad = [0, 0, 0, 432 - self.ht]
elif mode == "kitti400":
self._pad = [0, 0, 0, 400 - self.ht]
elif mode == "kitti376":
self._pad = [0, 0, 0, 376 - self.ht]
else:
self._pad = [pad_wd // 2, pad_wd - pad_wd // 2, 0, pad_ht]
def pad(self, *inputs):
return [F.pad(x, self._pad, mode="constant", value=0.0) for x in inputs]
def unpad(self, x):
ht, wd = x.shape[-2:]
c = [self._pad[2], ht - self._pad[3], self._pad[0], wd - self._pad[1]]
return x[..., c[0] : c[1], c[2] : c[3]]
def compute_grid_indices(image_shape, patch_size=TRAIN_SIZE, min_overlap=20):
if min_overlap >= patch_size[0] or min_overlap >= patch_size[1]:
raise ValueError("!!")
hs = list(range(0, image_shape[0], patch_size[0] - min_overlap))
ws = list(range(0, image_shape[1], patch_size[1] - min_overlap))
# Make sure the final patch is flush with the image boundary
hs[-1] = image_shape[0] - patch_size[0]
ws[-1] = image_shape[1] - patch_size[1]
return [(h, w) for h in hs for w in ws]
import math
def compute_weight(
hws, image_shape, patch_size=TRAIN_SIZE, sigma=1.0, wtype="gaussian"
):
patch_num = len(hws)
h, w = torch.meshgrid(torch.arange(patch_size[0]), torch.arange(patch_size[1]))
h, w = h / float(patch_size[0]), w / float(patch_size[1])
c_h, c_w = 0.5, 0.5
h, w = h - c_h, w - c_w
weights_hw = (h**2 + w**2) ** 0.5 / sigma
denorm = 1 / (sigma * math.sqrt(2 * math.pi))
weights_hw = denorm * torch.exp(-0.5 * (weights_hw) ** 2)
weights = torch.zeros(1, patch_num, *image_shape)
for idx, (h, w) in enumerate(hws):
weights[:, idx, h : h + patch_size[0], w : w + patch_size[1]] = weights_hw
weights = weights.cuda()
patch_weights = []
for idx, (h, w) in enumerate(hws):
patch_weights.append(
weights[:, idx : idx + 1, h : h + patch_size[0], w : w + patch_size[1]]
)
return patch_weights
@torch.no_grad()
def create_sintel_submission(
model, output_path="sintel_submission_multi8_768", sigma=0.05
):
"""Create submission for the Sintel leaderboard"""
print("no warm start")
# print(f"output path: {output_path}")
IMAGE_SIZE = [436, 1024]
hws = compute_grid_indices(IMAGE_SIZE)
weights = compute_weight(hws, IMAGE_SIZE, TRAIN_SIZE, sigma)
model.eval()
for dstype in ["final", "clean"]:
test_dataset = datasets.MpiSintel_submission(
split="test", aug_params=None, dstype=dstype, root="./dataset/Sintel/test"
)
epe_list = []
for test_id in range(len(test_dataset)):
if (test_id + 1) % 100 == 0:
print(f"{test_id} / {len(test_dataset)}")
# break
image1, image2, (sequence, frame) = test_dataset[test_id]
image1, image2 = image1[None].cuda(), image2[None].cuda()
flows = 0
flow_count = 0
for idx, (h, w) in enumerate(hws):
image1_tile = image1[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
image2_tile = image2[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
flow_pre, flow_low = model(image1_tile, image2_tile)
padding = (
w,
IMAGE_SIZE[1] - w - TRAIN_SIZE[1],
h,
IMAGE_SIZE[0] - h - TRAIN_SIZE[0],
0,
0,
)
flows += F.pad(flow_pre * weights[idx], padding)
flow_count += F.pad(weights[idx], padding)
flow_pre = flows / flow_count
flow = flow_pre[0].permute(1, 2, 0).cpu().numpy()
output_dir = os.path.join(output_path, dstype, sequence)
output_file = os.path.join(output_dir, "frame%04d.flo" % (frame + 1))
if not os.path.exists(output_dir):
os.makedirs(output_dir)
frame_utils.writeFlow(output_file, flow)
@torch.no_grad()
def create_kitti_submission(model, output_path="kitti_submission", sigma=0.05):
"""Create submission for the Sintel leaderboard"""
IMAGE_SIZE = [432, 1242]
print(f"output path: {output_path}")
print(f"image size: {IMAGE_SIZE}")
print(f"training size: {TRAIN_SIZE}")
hws = compute_grid_indices(IMAGE_SIZE)
weights = compute_weight(hws, (432, 1242), TRAIN_SIZE, sigma)
model.eval()
test_dataset = datasets.KITTI(split="testing", aug_params=None)
if not os.path.exists(output_path):
os.makedirs(output_path)
for test_id in range(len(test_dataset)):
image1, image2, (frame_id,) = test_dataset[test_id]
new_shape = image1.shape[1:]
if (
new_shape[1] != IMAGE_SIZE[1]
): # fix the height=432, adaptive ajust the width
print(f"replace {IMAGE_SIZE} with {new_shape}")
IMAGE_SIZE[0] = 432
IMAGE_SIZE[1] = new_shape[1]
hws = compute_grid_indices(IMAGE_SIZE)
weights = compute_weight(hws, IMAGE_SIZE, TRAIN_SIZE, sigma)
padder = InputPadder(
image1.shape, mode="kitti432"
) # padding the image to height of 432
image1, image2 = padder.pad(image1[None].cuda(), image2[None].cuda())
flows = 0
flow_count = 0
for idx, (h, w) in enumerate(hws):
image1_tile = image1[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
image2_tile = image2[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
flow_pre, _ = model(image1_tile, image2_tile)
padding = (
w,
IMAGE_SIZE[1] - w - TRAIN_SIZE[1],
h,
IMAGE_SIZE[0] - h - TRAIN_SIZE[0],
0,
0,
)
flows += F.pad(flow_pre * weights[idx], padding)
flow_count += F.pad(weights[idx], padding)
flow_pre = flows / flow_count
flow = padder.unpad(flow_pre[0]).permute(1, 2, 0).cpu().numpy()
output_filename = os.path.join(output_path, frame_id)
frame_utils.writeFlowKITTI(output_filename, flow)
flow_img = flow_viz.flow_to_image(flow)
image = Image.fromarray(flow_img)
if not os.path.exists(f"vis_kitti_3patch"):
os.makedirs(f"vis_kitti_3patch/flow")
os.makedirs(f"vis_kitti_3patch/image")
image.save(f"vis_kitti_3patch/flow/{test_id}.png")
imageio.imwrite(
f"vis_kitti_3patch/image/{test_id}_0.png",
image1[0].cpu().permute(1, 2, 0).numpy(),
)
imageio.imwrite(
f"vis_kitti_3patch/image/{test_id}_1.png",
image2[0].cpu().permute(1, 2, 0).numpy(),
)
@torch.no_grad()
def validate_kitti(model, sigma=0.05):
IMAGE_SIZE = [376, 1242]
TRAIN_SIZE = [376, 720]
hws = compute_grid_indices(IMAGE_SIZE, TRAIN_SIZE)
weights = compute_weight(hws, IMAGE_SIZE, TRAIN_SIZE, sigma)
model.eval()
val_dataset = datasets.KITTI(split="training")
out_list, epe_list = [], []
for val_id in range(len(val_dataset)):
image1, image2, flow_gt, valid_gt = val_dataset[val_id]
new_shape = image1.shape[1:]
if new_shape[1] != IMAGE_SIZE[1]:
print(f"replace {IMAGE_SIZE} with {new_shape}")
IMAGE_SIZE[0] = 376
IMAGE_SIZE[1] = new_shape[1]
hws = compute_grid_indices(IMAGE_SIZE, TRAIN_SIZE)
weights = compute_weight(hws, IMAGE_SIZE, TRAIN_SIZE, sigma)
padder = InputPadder(image1.shape, mode="kitti376")
image1, image2 = padder.pad(image1[None].cuda(), image2[None].cuda())
flows = 0
flow_count = 0
for idx, (h, w) in enumerate(hws):
image1_tile = image1[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
image2_tile = image2[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
flow_pre, flow_low = model(image1_tile, image2_tile)
padding = (
w,
IMAGE_SIZE[1] - w - TRAIN_SIZE[1],
h,
IMAGE_SIZE[0] - h - TRAIN_SIZE[0],
0,
0,
)
flows += F.pad(flow_pre * weights[idx], padding)
flow_count += F.pad(weights[idx], padding)
flow_pre = flows / flow_count
flow = padder.unpad(flow_pre[0]).cpu()
epe = torch.sum((flow - flow_gt) ** 2, dim=0).sqrt()
mag = torch.sum(flow_gt**2, dim=0).sqrt()
epe = epe.view(-1)
mag = mag.view(-1)
val = valid_gt.view(-1) >= 0.5
out = ((epe > 3.0) & ((epe / mag) > 0.05)).float()
epe_list.append(epe[val].mean().item())
out_list.append(out[val].cpu().numpy())
epe_list = np.array(epe_list)
out_list = np.concatenate(out_list)
epe = np.mean(epe_list)
f1 = 100 * np.mean(out_list)
print("Validation KITTI: %f, %f" % (epe, f1))
return {"kitti-epe": epe, "kitti-f1": f1}
@torch.no_grad()
def validate_sintel(model, sigma=0.05):
"""Peform validation using the Sintel (train) split"""
IMAGE_SIZE = [436, 1024]
hws = compute_grid_indices(IMAGE_SIZE)
weights = compute_weight(hws, IMAGE_SIZE, TRAIN_SIZE, sigma)
model.eval()
results = {}
for dstype in ["final", "clean"]:
val_dataset = datasets.MpiSintel(split="training", dstype=dstype)
epe_list = []
for val_id in range(len(val_dataset)):
if val_id % 50 == 0:
print(val_id)
image1, image2, flow_gt, _ = val_dataset[val_id]
image1 = image1[None].cuda()
image2 = image2[None].cuda()
flows = 0
flow_count = 0
for idx, (h, w) in enumerate(hws):
image1_tile = image1[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
image2_tile = image2[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
flow_pre, _ = model(image1_tile, image2_tile, flow_init=None)
padding = (
w,
IMAGE_SIZE[1] - w - TRAIN_SIZE[1],
h,
IMAGE_SIZE[0] - h - TRAIN_SIZE[0],
0,
0,
)
flows += F.pad(flow_pre * weights[idx], padding)
flow_count += F.pad(weights[idx], padding)
flow_pre = flows / flow_count
flow_pre = flow_pre[0].cpu()
epe = torch.sum((flow_pre - flow_gt) ** 2, dim=0).sqrt()
epe_list.append(epe.view(-1).numpy())
epe_all = np.concatenate(epe_list)
epe = np.mean(epe_all)
px1 = np.mean(epe_all < 1)
px3 = np.mean(epe_all < 3)
px5 = np.mean(epe_all < 5)
print(
"Validation (%s) EPE: %f, 1px: %f, 3px: %f, 5px: %f"
% (dstype, epe, px1, px3, px5)
)
results[f"{dstype}_tile"] = np.mean(epe_list)
return results
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--model", help="load model")
parser.add_argument("--eval", help="eval benchmark")
parser.add_argument("--small", action="store_true", help="use small model")
args = parser.parse_args()
exp_func = None
cfg = None
if args.eval == "sintel_submission":
exp_func = create_sintel_submission
cfg = get_submission_cfg()
elif args.eval == "kitti_submission":
exp_func = create_kitti_submission
cfg = get_submission_cfg()
cfg.latentcostformer.decoder_depth = 24
elif args.eval == "sintel_validation":
exp_func = validate_sintel
if args.small:
cfg = get_small_things_cfg()
else:
cfg = get_things_cfg()
elif args.eval == "kitti_validation":
exp_func = validate_kitti
if args.small:
cfg = get_small_things_cfg()
else:
cfg = get_things_cfg()
cfg.latentcostformer.decoder_depth = 24
else:
print(f"EROOR: {args.eval} is not valid")
cfg.update(vars(args))
print(cfg)
model = torch.nn.DataParallel(build_flowformer(cfg))
model.load_state_dict(torch.load(cfg.model))
model.cuda()
model.eval()
exp_func(model.module)
@@ -0,0 +1,5 @@
mkdir -p checkpoints
python -u train_FlowFormer.py --name chairs --stage chairs --validation chairs
python -u train_FlowFormer.py --name things --stage things --validation sintel
python -u train_FlowFormer.py --name sintel --stage sintel --validation sintel
python -u train_FlowFormer.py --name kitti --stage kitti --validation kitti
@@ -0,0 +1,182 @@
from __future__ import print_function, division
import sys
# sys.path.append('core')
import argparse
import os
import cv2
import time
import numpy as np
import matplotlib.pyplot as plt
from pathlib import Path
import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.utils.data import DataLoader
from core import optimizer
import evaluate_FlowFormer as evaluate
import evaluate_FlowFormer_tile as evaluate_tile
import core.datasets as datasets
from core.loss import sequence_loss
from core.optimizer import fetch_optimizer
from core.utils.misc import process_cfg
from loguru import logger as loguru_logger
# from torch.utils.tensorboard import SummaryWriter
from core.utils.logger import Logger
# from core.FlowFormer import FlowFormer
from core.FlowFormer import build_flowformer
try:
from torch.cuda.amp import GradScaler
except:
# dummy GradScaler for PyTorch < 1.6
class GradScaler:
def __init__(self):
pass
def scale(self, loss):
return loss
def unscale_(self, optimizer):
pass
def step(self, optimizer):
optimizer.step()
def update(self):
pass
# torch.autograd.set_detect_anomaly(True)
def count_parameters(model):
return sum(p.numel() for p in model.parameters() if p.requires_grad)
def train(cfg):
model = nn.DataParallel(build_flowformer(cfg))
loguru_logger.info("Parameter Count: %d" % count_parameters(model))
if cfg.restore_ckpt is not None:
print("[Loading ckpt from {}]".format(cfg.restore_ckpt))
model.load_state_dict(torch.load(cfg.restore_ckpt), strict=True)
model.cuda()
model.train()
train_loader = datasets.fetch_dataloader(cfg)
optimizer, scheduler = fetch_optimizer(model, cfg.trainer)
total_steps = 0
scaler = GradScaler(enabled=cfg.mixed_precision)
logger = Logger(model, scheduler, cfg)
add_noise = False
should_keep_training = True
while should_keep_training:
for i_batch, data_blob in enumerate(train_loader):
optimizer.zero_grad()
image1, image2, flow, valid = [x.cuda() for x in data_blob]
if cfg.add_noise:
stdv = np.random.uniform(0.0, 5.0)
image1 = (image1 + stdv * torch.randn(*image1.shape).cuda()).clamp(
0.0, 255.0
)
image2 = (image2 + stdv * torch.randn(*image2.shape).cuda()).clamp(
0.0, 255.0
)
output = {}
flow_predictions = model(image1, image2, output)
loss, metrics = sequence_loss(flow_predictions, flow, valid, cfg)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), cfg.trainer.clip)
scaler.step(optimizer)
scheduler.step()
scaler.update()
metrics.update(output)
logger.push(metrics)
### change evaluate to functions
if total_steps % cfg.val_freq == cfg.val_freq - 1:
PATH = "%s/%d_%s.pth" % (cfg.log_dir, total_steps + 1, cfg.name)
# torch.save(model.state_dict(), PATH)
results = {}
for val_dataset in cfg.validation:
if val_dataset == "chairs":
results.update(evaluate.validate_chairs(model.module))
elif val_dataset == "sintel":
results.update(evaluate.validate_sintel(model.module))
elif val_dataset == "kitti":
results.update(evaluate.validate_kitti(model.module))
logger.write_dict(results)
model.train()
total_steps += 1
if total_steps > cfg.trainer.num_steps:
should_keep_training = False
break
logger.close()
PATH = cfg.log_dir + "/final"
torch.save(model.state_dict(), PATH)
PATH = f"checkpoints/{cfg.stage}.pth"
torch.save(model.state_dict(), PATH)
return PATH
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--name", default="flowformer", help="name your experiment")
parser.add_argument("--stage", help="determines which dataset to use for training")
parser.add_argument("--validation", type=str, nargs="+")
parser.add_argument(
"--mixed_precision", action="store_true", help="use mixed precision"
)
args = parser.parse_args()
if args.stage == "chairs":
from configs.default import get_cfg
elif args.stage == "things":
from configs.things import get_cfg
elif args.stage == "sintel":
from configs.sintel import get_cfg
elif args.stage == "kitti":
from configs.kitti import get_cfg
elif args.stage == "autoflow":
from configs.autoflow import get_cfg
cfg = get_cfg()
cfg.update(vars(args))
process_cfg(cfg)
loguru_logger.add(str(Path(cfg.log_dir) / "log.txt"), encoding="utf8")
loguru_logger.info(cfg)
torch.manual_seed(1234)
np.random.seed(1234)
if not os.path.isdir("checkpoints"):
os.mkdir("checkpoints")
train(cfg)
@@ -0,0 +1,238 @@
import sys
sys.path.append("core")
from PIL import Image
from glob import glob
import argparse
import os
import time
import numpy as np
import torch
import torch.nn.functional as F
import matplotlib.pyplot as plt
from configs.submission import get_cfg
from core.utils.misc import process_cfg
import datasets
from utils import flow_viz
from utils import frame_utils
import cv2
import math
import os.path as osp
from core.FlowFormer import build_flowformer
from utils.utils import InputPadder, forward_interpolate
import itertools
TRAIN_SIZE = [432, 960]
def compute_grid_indices(image_shape, patch_size=TRAIN_SIZE, min_overlap=20):
if min_overlap >= TRAIN_SIZE[0] or min_overlap >= TRAIN_SIZE[1]:
raise ValueError(
f"Overlap should be less than size of patch (got {min_overlap}"
f"for patch size {patch_size})."
)
if image_shape[0] == TRAIN_SIZE[0]:
hs = list(range(0, image_shape[0], TRAIN_SIZE[0]))
else:
hs = list(range(0, image_shape[0], TRAIN_SIZE[0] - min_overlap))
if image_shape[1] == TRAIN_SIZE[1]:
ws = list(range(0, image_shape[1], TRAIN_SIZE[1]))
else:
ws = list(range(0, image_shape[1], TRAIN_SIZE[1] - min_overlap))
# Make sure the final patch is flush with the image boundary
hs[-1] = image_shape[0] - patch_size[0]
ws[-1] = image_shape[1] - patch_size[1]
return [(h, w) for h in hs for w in ws]
def compute_weight(
hws, image_shape, patch_size=TRAIN_SIZE, sigma=1.0, wtype="gaussian"
):
patch_num = len(hws)
h, w = torch.meshgrid(torch.arange(patch_size[0]), torch.arange(patch_size[1]))
h, w = h / float(patch_size[0]), w / float(patch_size[1])
c_h, c_w = 0.5, 0.5
h, w = h - c_h, w - c_w
weights_hw = (h**2 + w**2) ** 0.5 / sigma
denorm = 1 / (sigma * math.sqrt(2 * math.pi))
weights_hw = denorm * torch.exp(-0.5 * (weights_hw) ** 2)
weights = torch.zeros(1, patch_num, *image_shape)
for idx, (h, w) in enumerate(hws):
weights[:, idx, h : h + patch_size[0], w : w + patch_size[1]] = weights_hw
weights = weights.cuda()
patch_weights = []
for idx, (h, w) in enumerate(hws):
patch_weights.append(
weights[:, idx : idx + 1, h : h + patch_size[0], w : w + patch_size[1]]
)
return patch_weights
def compute_flow(model, image1, image2, weights=None):
print(f"computing flow...")
image_size = image1.shape[1:]
image1, image2 = image1[None].cuda(), image2[None].cuda()
hws = compute_grid_indices(image_size)
if weights is None: # no tile
padder = InputPadder(image1.shape)
image1, image2 = padder.pad(image1, image2)
flow_pre, _ = model(image1, image2)
flow_pre = padder.unpad(flow_pre)
flow = flow_pre[0].permute(1, 2, 0).cpu().numpy()
else: # tile
flows = 0
flow_count = 0
for idx, (h, w) in enumerate(hws):
image1_tile = image1[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
image2_tile = image2[:, :, h : h + TRAIN_SIZE[0], w : w + TRAIN_SIZE[1]]
flow_pre, _ = model(image1_tile, image2_tile)
padding = (
w,
image_size[1] - w - TRAIN_SIZE[1],
h,
image_size[0] - h - TRAIN_SIZE[0],
0,
0,
)
flows += F.pad(flow_pre * weights[idx], padding)
flow_count += F.pad(weights[idx], padding)
flow_pre = flows / flow_count
flow = flow_pre[0].permute(1, 2, 0).cpu().numpy()
return flow
def compute_adaptive_image_size(image_size):
target_size = TRAIN_SIZE
scale0 = target_size[0] / image_size[0]
scale1 = target_size[1] / image_size[1]
if scale0 > scale1:
scale = scale0
else:
scale = scale1
image_size = (int(image_size[1] * scale), int(image_size[0] * scale))
return image_size
def prepare_image(root_dir, viz_root_dir, fn1, fn2, keep_size):
print(f"preparing image...")
print(f"root dir = {root_dir}, fn = {fn1}")
image1 = frame_utils.read_gen(osp.join(root_dir, fn1))
image2 = frame_utils.read_gen(osp.join(root_dir, fn2))
image1 = np.array(image1).astype(np.uint8)[..., :3]
image2 = np.array(image2).astype(np.uint8)[..., :3]
if not keep_size:
dsize = compute_adaptive_image_size(image1.shape[0:2])
image1 = cv2.resize(image1, dsize=dsize, interpolation=cv2.INTER_CUBIC)
image2 = cv2.resize(image2, dsize=dsize, interpolation=cv2.INTER_CUBIC)
image1 = torch.from_numpy(image1).permute(2, 0, 1).float()
image2 = torch.from_numpy(image2).permute(2, 0, 1).float()
dirname = osp.dirname(fn1)
filename = osp.splitext(osp.basename(fn1))[0]
viz_dir = osp.join(viz_root_dir, dirname)
if not osp.exists(viz_dir):
os.makedirs(viz_dir)
viz_fn = osp.join(viz_dir, filename + ".png")
return image1, image2, viz_fn
def build_model():
print(f"building model...")
cfg = get_cfg()
model = torch.nn.DataParallel(build_flowformer(cfg))
model.load_state_dict(torch.load(cfg.model))
model.cuda()
model.eval()
return model
def visualize_flow(root_dir, viz_root_dir, model, img_pairs, keep_size):
weights = None
for img_pair in img_pairs:
fn1, fn2 = img_pair
print(f"processing {fn1}, {fn2}...")
image1, image2, viz_fn = prepare_image(
root_dir, viz_root_dir, fn1, fn2, keep_size
)
flow = compute_flow(model, image1, image2, weights)
flow_img = flow_viz.flow_to_image(flow)
cv2.imwrite(viz_fn, flow_img[:, :, [2, 1, 0]])
def process_sintel(sintel_dir):
img_pairs = []
for scene in os.listdir(sintel_dir):
dirname = osp.join(sintel_dir, scene)
image_list = sorted(glob(osp.join(dirname, "*.png")))
for i in range(len(image_list) - 1):
img_pairs.append((image_list[i], image_list[i + 1]))
return img_pairs
def generate_pairs(dirname, start_idx, end_idx):
img_pairs = []
for idx in range(start_idx, end_idx):
img1 = osp.join(dirname, f"{idx:06}.png")
img2 = osp.join(dirname, f"{idx+1:06}.png")
# img1 = f'{idx:06}.png'
# img2 = f'{idx+1:06}.png'
img_pairs.append((img1, img2))
return img_pairs
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--eval_type", default="sintel")
parser.add_argument("--root_dir", default=".")
parser.add_argument("--sintel_dir", default="datasets/Sintel/test/clean")
parser.add_argument("--seq_dir", default="demo_data/mihoyo")
parser.add_argument(
"--start_idx", type=int, default=1
) # starting index of the image sequence
parser.add_argument(
"--end_idx", type=int, default=1200
) # ending index of the image sequence
parser.add_argument("--viz_root_dir", default="viz_results")
parser.add_argument(
"--keep_size", action="store_true"
) # keep the image size, or the image will be adaptively resized.
args = parser.parse_args()
root_dir = args.root_dir
viz_root_dir = args.viz_root_dir
model = build_model()
if args.eval_type == "sintel":
img_pairs = process_sintel(args.sintel_dir)
elif args.eval_type == "seq":
img_pairs = generate_pairs(args.seq_dir, args.start_idx, args.end_idx)
with torch.no_grad():
visualize_flow(root_dir, viz_root_dir, model, img_pairs, args.keep_size)
+253
View File
@@ -0,0 +1,253 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# motif: https://github.com/sichun233746/MoTIF
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
import torch
import torch.nn as nn
import torch.nn.functional as F
from .configs import GIMMConfig
from .modules.coord_sampler import CoordSampler3D
from .modules.fi_components import LateralBlock
from .modules.hyponet import HypoNet
from .modules.fi_utils import warp
from .modules.softsplat import softsplat
class GIMM(nn.Module):
Config = GIMMConfig
def __init__(self, config: GIMMConfig):
super().__init__()
self.config = config = config.copy()
self.hyponet_config = config.hyponet
self.coord_sampler = CoordSampler3D(config.coord_range)
self.fwarp_type = config.fwarp_type
# Motion Encoder
channel = 32
in_dim = 2
self.cnn_encoder = nn.Sequential(
nn.Conv2d(in_dim, channel // 2, 3, 1, 1, bias=True, groups=1),
nn.Conv2d(channel // 2, channel, 3, 1, 1, bias=True, groups=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
LateralBlock(channel),
LateralBlock(channel),
LateralBlock(channel),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(
channel, channel // 2, 3, 1, 1, padding_mode="reflect", bias=True
),
)
# Latent Refiner
channel = 64
in_dim = 64
self.res_conv = nn.Sequential(
nn.Conv2d(in_dim, channel // 2, 3, 1, 1, bias=True, groups=1),
nn.Conv2d(channel // 2, channel, 3, 1, 1, bias=True, groups=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
LateralBlock(channel),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(
channel, channel // 2, 3, 1, 1, padding_mode="reflect", bias=True
),
)
self.g_filter = torch.nn.Parameter(
torch.FloatTensor(
[
[1.0 / 16.0, 1.0 / 8.0, 1.0 / 16.0],
[1.0 / 8.0, 1.0 / 4.0, 1.0 / 8.0],
[1.0 / 16.0, 1.0 / 8.0, 1.0 / 16.0],
]
).reshape(1, 1, 1, 3, 3),
requires_grad=False,
)
self.alpha_v = torch.nn.Parameter(torch.FloatTensor([1]), requires_grad=True)
self.alpha_fe = torch.nn.Parameter(torch.FloatTensor([1]), requires_grad=True)
self.hyponet = HypoNet(config.hyponet, add_coord_dim=32)
def cal_splatting_weights(self, raft_flow01, raft_flow10):
batch_size = raft_flow01.shape[0]
raft_flows = torch.cat([raft_flow01, raft_flow10], dim=0)
## flow variance metric
sqaure_mean, mean_square = torch.split(
F.conv3d(
F.pad(
torch.cat([raft_flows**2, raft_flows], 1),
(1, 1, 1, 1),
mode="reflect",
).unsqueeze(1),
self.g_filter,
).squeeze(1),
2,
dim=1,
)
var = (
(sqaure_mean - mean_square**2)
.clamp(1e-9, None)
.sqrt()
.mean(1)
.unsqueeze(1)
)
var01 = var[:batch_size]
var10 = var[batch_size:]
## flow warp metirc
f01_warp = -warp(raft_flow10, raft_flow01)
f10_warp = -warp(raft_flow01, raft_flow10)
err01 = (
torch.nn.functional.l1_loss(
input=f01_warp, target=raft_flow01, reduction="none"
)
.mean(1)
.unsqueeze(1)
)
err02 = (
torch.nn.functional.l1_loss(
input=f10_warp, target=raft_flow10, reduction="none"
)
.mean(1)
.unsqueeze(1)
)
weights1 = 1 / (1 + err01 * self.alpha_fe) + 1 / (1 + var01 * self.alpha_v)
weights2 = 1 / (1 + err02 * self.alpha_fe) + 1 / (1 + var10 * self.alpha_v)
return weights1, weights2
def forward(
self, xs, coord=None, keep_xs_shape=True, ori_flow=None, timesteps=None
):
coord = self.sample_coord_input(xs) if coord is None else coord
raft_flow01 = ori_flow[:, :, 0]
raft_flow10 = ori_flow[:, :, 1]
# calculate splatting metrics
weights1, weights2 = self.cal_splatting_weights(raft_flow01, raft_flow10)
# b,c,h,w
pixel_latent_0 = self.cnn_encoder(xs[:, :, 0])
pixel_latent_1 = self.cnn_encoder(xs[:, :, 1])
pixel_latent = []
modulation_params_dict = None
strtype = self.fwarp_type
if isinstance(timesteps, list):
assert isinstance(coord, list)
assert len(timesteps) == len(coord)
for i, cur_t in enumerate(timesteps):
cur_t = cur_t.reshape(-1, 1, 1, 1)
tmp_pixel_latent_0 = softsplat(
tenIn=pixel_latent_0,
tenFlow=raft_flow01 * cur_t,
tenMetric=weights1,
strMode=strtype + "-zeroeps",
)
tmp_pixel_latent_1 = softsplat(
tenIn=pixel_latent_1,
tenFlow=raft_flow10 * (1 - cur_t),
tenMetric=weights2,
strMode=strtype + "-zeroeps",
)
tmp_pixel_latent = torch.cat(
[tmp_pixel_latent_0, tmp_pixel_latent_1], dim=1
)
tmp_pixel_latent = tmp_pixel_latent + self.res_conv(
torch.cat([pixel_latent_0, pixel_latent_1, tmp_pixel_latent], dim=1)
)
pixel_latent.append(tmp_pixel_latent.permute(0, 2, 3, 1))
all_outputs = []
for idx, c in enumerate(coord):
outputs = self.hyponet(
c,
modulation_params_dict=modulation_params_dict,
pixel_latent=pixel_latent[idx],
)
if keep_xs_shape:
permute_idx_range = [i for i in range(1, xs.ndim - 1)]
outputs = outputs.permute(0, -1, *permute_idx_range)
all_outputs.append(outputs)
return all_outputs
else:
cur_t = timesteps.reshape(-1, 1, 1, 1)
tmp_pixel_latent_0 = softsplat(
tenIn=pixel_latent_0,
tenFlow=raft_flow01 * cur_t,
tenMetric=weights1,
strMode=strtype + "-zeroeps",
)
tmp_pixel_latent_1 = softsplat(
tenIn=pixel_latent_1,
tenFlow=raft_flow10 * (1 - cur_t),
tenMetric=weights2,
strMode=strtype + "-zeroeps",
)
tmp_pixel_latent = torch.cat(
[tmp_pixel_latent_0, tmp_pixel_latent_1], dim=1
)
tmp_pixel_latent = tmp_pixel_latent + self.res_conv(
torch.cat([pixel_latent_0, pixel_latent_1, tmp_pixel_latent], dim=1)
)
pixel_latent = tmp_pixel_latent.permute(0, 2, 3, 1)
# predict all pixels of coord after applying the modulation_parms into hyponet
outputs = self.hyponet(
coord,
modulation_params_dict=modulation_params_dict,
pixel_latent=pixel_latent,
)
if keep_xs_shape:
permute_idx_range = [i for i in range(1, xs.ndim - 1)]
outputs = outputs.permute(0, -1, *permute_idx_range)
return outputs
def compute_loss(self, preds, targets, reduction="mean", single=False):
assert reduction in ["mean", "sum", "none"]
batch_size = preds.shape[0]
sample_mses = 0
assert preds.shape[2] == 1
assert targets.shape[2] == 1
for i in range(preds.shape[2]):
sample_mses += torch.reshape(
(preds[:, :, i] - targets[:, :, i]) ** 2, (batch_size, -1)
).mean(dim=-1)
sample_mses = sample_mses / preds.shape[2]
if reduction == "mean":
total_loss = sample_mses.mean()
psnr = (-10 * torch.log10(sample_mses)).mean()
elif reduction == "sum":
total_loss = sample_mses.sum()
psnr = (-10 * torch.log10(sample_mses)).sum()
else:
total_loss = sample_mses
psnr = -10 * torch.log10(sample_mses)
return {"loss_total": total_loss, "mse": total_loss, "psnr": psnr}
def sample_coord_input(
self,
batch_size,
s_shape,
t_ids,
coord_range=None,
upsample_ratio=1.0,
device=None,
):
assert device is not None
assert coord_range is None
coord_inputs = self.coord_sampler(
batch_size, s_shape, t_ids, coord_range, upsample_ratio, device
)
return coord_inputs
+468
View File
@@ -0,0 +1,468 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# amt: https://github.com/MCG-NKU/AMT
# motif: https://github.com/sichun233746/MoTIF
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
import torch
import torch.nn as nn
import torch.nn.functional as F
from .configs import GIMMVFIConfig
from .modules.coord_sampler import CoordSampler3D
from .modules.hyponet import HypoNet
from .modules.fi_components import *
#from .flowformer import initialize_Flowformer
from .modules.fi_utils import normalize_flow, unnormalize_flow, warp, resize
from .raft.corr import BidirCorrBlock
from .modules.softsplat import softsplat
class GIMMVFI_F(nn.Module):
Config = GIMMVFIConfig
def __init__(self, config: GIMMVFIConfig):
super().__init__()
self.config = config = config.copy()
self.hyponet_config = config.hyponet
self.raft_iter = config.raft_iter
######### Encoder and Decoder Settings #########
#self.flow_estimator = initialize_Flowformer()
f_dims = [256, 128]
skip_channels = f_dims[-1] // 2
self.num_flows = 3
self.amt_init_decoder = NewInitDecoder(f_dims[0], skip_channels)
self.amt_final_decoder = NewMultiFlowDecoder(f_dims[1], skip_channels)
self.amt_update4_low = self._get_updateblock(f_dims[0] // 2, 2.0)
self.amt_update4_high = self._get_updateblock(f_dims[0] // 2, None)
self.amt_comb_block = nn.Sequential(
nn.Conv2d(3 * self.num_flows, 6 * self.num_flows, 7, 1, 3),
nn.PReLU(6 * self.num_flows),
nn.Conv2d(6 * self.num_flows, 3, 7, 1, 3),
)
################ GIMM settings #################
self.coord_sampler = CoordSampler3D(config.coord_range)
self.g_filter = torch.nn.Parameter(
torch.FloatTensor(
[
[1.0 / 16.0, 1.0 / 8.0, 1.0 / 16.0],
[1.0 / 8.0, 1.0 / 4.0, 1.0 / 8.0],
[1.0 / 16.0, 1.0 / 8.0, 1.0 / 16.0],
]
).reshape(1, 1, 1, 3, 3),
requires_grad=False,
)
self.fwarp_type = config.fwarp_type
self.alpha_v = torch.nn.Parameter(torch.FloatTensor([1]), requires_grad=True)
self.alpha_fe = torch.nn.Parameter(torch.FloatTensor([1]), requires_grad=True)
channel = 32
in_dim = 2
self.cnn_encoder = nn.Sequential(
nn.Conv2d(in_dim, channel // 2, 3, 1, 1, bias=True, groups=1),
nn.Conv2d(channel // 2, channel, 3, 1, 1, bias=True, groups=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
LateralBlock(channel),
LateralBlock(channel),
LateralBlock(channel),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(
channel, channel // 2, 3, 1, 1, padding_mode="reflect", bias=True
),
)
channel = 64
in_dim = 64
self.res_conv = nn.Sequential(
nn.Conv2d(in_dim, channel // 2, 3, 1, 1, bias=True, groups=1),
nn.Conv2d(channel // 2, channel, 3, 1, 1, bias=True, groups=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
LateralBlock(channel),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(
channel, channel // 2, 3, 1, 1, padding_mode="reflect", bias=True
),
)
self.hyponet = HypoNet(config.hyponet, add_coord_dim=32)
def _get_updateblock(self, cdim, scale_factor=None):
return BasicUpdateBlock(
cdim=cdim,
hidden_dim=192,
flow_dim=64,
corr_dim=256,
corr_dim2=192,
fc_dim=188,
scale_factor=scale_factor,
corr_levels=4,
radius=4,
)
def cal_bidirection_flow(self, im0, im1):
f01, features0, fnet0 = self.flow_estimator(
im0, im1, return_feat=True, iters=None
)
f10, features1, fnet1 = self.flow_estimator(
im1, im0, return_feat=True, iters=None
)
f01 = f01[0]
f10 = f10[0]
corr_fn = BidirCorrBlock(fnet0, fnet1, radius=4)
flow01 = f01.unsqueeze(2)
flow10 = f10.unsqueeze(2)
noraml_flows = torch.cat([flow01, -flow10], dim=2)
noraml_flows, flow_scalers = normalize_flow(noraml_flows)
ori_flows = torch.cat([flow01, flow10], dim=2)
return (
noraml_flows,
ori_flows,
flow_scalers,
features0,
features1,
corr_fn,
torch.cat([f01.unsqueeze(2), f10.unsqueeze(2)], dim=2),
)
def predict_flow(self, f, coord, t, flows):
raft_flow01 = flows[:, :, 0].detach()
raft_flow10 = flows[:, :, 1].detach()
# calculate splatting metrics
weights1, weights2 = self.cal_splatting_weights(raft_flow01, raft_flow10)
strtype = self.fwarp_type + "-zeroeps"
# b,c,h,w
pixel_latent_0 = self.cnn_encoder(f[:, :, 0])
pixel_latent_1 = self.cnn_encoder(f[:, :, 1])
pixel_latent = []
for i, cur_t in enumerate(t):
cur_t = cur_t.reshape(-1, 1, 1, 1)
tmp_pixel_latent_0 = softsplat(
tenIn=pixel_latent_0,
tenFlow=raft_flow01 * cur_t,
tenMetric=weights1,
strMode=strtype,
)
tmp_pixel_latent_1 = softsplat(
tenIn=pixel_latent_1,
tenFlow=raft_flow10 * (1 - cur_t),
tenMetric=weights2,
strMode=strtype,
)
tmp_pixel_latent = torch.cat(
[tmp_pixel_latent_0, tmp_pixel_latent_1], dim=1
)
tmp_pixel_latent = tmp_pixel_latent + self.res_conv(
torch.cat([pixel_latent_0, pixel_latent_1, tmp_pixel_latent], dim=1)
)
pixel_latent.append(tmp_pixel_latent.permute(0, 2, 3, 1))
all_outputs = []
permute_idx_range = [i for i in range(1, f.ndim - 1)]
for idx, c in enumerate(coord):
assert c[0][0, 0, 0, 0, 0] == t[idx][0].squeeze()
assert isinstance(c, tuple)
if c[1] is None:
outputs = self.hyponet(
c, modulation_params_dict=None, pixel_latent=pixel_latent[idx]
).permute(0, -1, *permute_idx_range)
else:
outputs = self.hyponet(
c, modulation_params_dict=None, pixel_latent=pixel_latent[idx]
)
all_outputs.append(outputs)
return all_outputs
def warp_w_mask(self, img0, img1, ft0, ft1, mask, scale=1):
ft0 = scale * resize(ft0, scale_factor=scale)
ft1 = scale * resize(ft1, scale_factor=scale)
mask = resize(mask, scale_factor=scale).sigmoid()
img0_warp = warp(img0, ft0)
img1_warp = warp(img1, ft1)
img_warp = mask * img0_warp + (1 - mask) * img1_warp
return img_warp
def frame_synthesize(
self, img_xs, flow_t, features0, features1, corr_fn, cur_t, full_img=None
):
"""
flow_t: b,2,h,w
cur_t: b,1,1,1
"""
batch_size = img_xs.shape[0]
img0 = 2 * img_xs[:, :, 0] - 1.0
img1 = 2 * img_xs[:, :, 1] - 1.0
##################### update the predicted flow #####################
## initialize coordinates for looking up
lookup_coord = self.flow_estimator.build_coord(img_xs[:, :, 0]).to(
img_xs[:, :, 0].device
)
flow_t0_fullsize = flow_t * (-cur_t)
flow_t1_fullsize = flow_t * (1.0 - cur_t)
inv = 1 / 4
flow_t0_inr4 = inv * resize(flow_t0_fullsize, inv)
flow_t1_inr4 = inv * resize(flow_t1_fullsize, inv)
############################# scale 1/4 #############################
# i. Initialize feature t at scale 1/4
flowt0_4, flowt1_4, ft_4_ = self.amt_init_decoder(
features0[-1],
features1[-1],
flow_t0_inr4,
flow_t1_inr4,
img0=img0,
img1=img1,
)
mask_4_, ft_4_ = ft_4_[:, :1], ft_4_[:, 1:]
img_warp_4 = self.warp_w_mask(img0, img1, flowt0_4, flowt1_4, mask_4_, scale=4)
img_warp_4 = (img_warp_4 + 1.0) / 2
img_warp_4 = torch.clamp(img_warp_4, 0, 1)
corr_4, flow_4_lr = self._amt_corr_scale_lookup(
corr_fn, lookup_coord, flowt0_4, flowt1_4, cur_t, downsample=2
)
delta_ft_4_, delta_flow_4 = self.amt_update4_low(ft_4_, flow_4_lr, corr_4)
delta_flow0_4, delta_flow1_4 = torch.chunk(delta_flow_4, 2, 1)
flowt0_4 = flowt0_4 + delta_flow0_4
flowt1_4 = flowt1_4 + delta_flow1_4
ft_4_ = ft_4_ + delta_ft_4_
# iii. residue update with lookup corr
corr_4 = resize(corr_4, scale_factor=2.0)
flow_4 = torch.cat([flowt0_4, flowt1_4], dim=1)
delta_ft_4_, delta_flow_4 = self.amt_update4_high(ft_4_, flow_4, corr_4)
flowt0_4 = flowt0_4 + delta_flow_4[:, :2]
flowt1_4 = flowt1_4 + delta_flow_4[:, 2:4]
ft_4_ = ft_4_ + delta_ft_4_
############################# scale 1/1 #############################
flowt0_1, flowt1_1, mask, img_res = self.amt_final_decoder(
ft_4_,
features0[0],
features1[0],
flowt0_4,
flowt1_4,
mask=mask_4_,
img0=img0,
img1=img1,
)
if full_img is not None:
img0 = 2 * full_img[:, :, 0] - 1.0
img1 = 2 * full_img[:, :, 1] - 1.0
inv = img1.shape[2] / flowt0_1.shape[2]
flowt0_1 = inv * resize(flowt0_1, scale_factor=inv)
flowt1_1 = inv * resize(flowt1_1, scale_factor=inv)
flow_t0_fullsize = inv * resize(flow_t0_fullsize, scale_factor=inv)
flow_t1_fullsize = inv * resize(flow_t1_fullsize, scale_factor=inv)
mask = resize(mask, scale_factor=inv)
img_res = resize(img_res, scale_factor=inv)
imgt_pred = multi_flow_combine(
self.amt_comb_block, img0, img1, flowt0_1, flowt1_1, mask, img_res, None
)
imgt_pred = torch.clamp(imgt_pred, 0, 1)
######################################################################
flowt0_1 = flowt0_1.reshape(
batch_size, self.num_flows, 2, img0.shape[-2], img0.shape[-1]
)
flowt1_1 = flowt1_1.reshape(
batch_size, self.num_flows, 2, img0.shape[-2], img0.shape[-1]
)
flowt0_pred = [flowt0_1, flowt0_4]
flowt1_pred = [flowt1_1, flowt1_4]
other_pred = [img_warp_4]
return imgt_pred, flowt0_pred, flowt1_pred, other_pred
def forward(self, img_xs, coord=None, t=None, ds_factor=None):
assert isinstance(t, list)
assert isinstance(coord, list)
assert len(t) == len(coord)
full_size_img = None
if ds_factor is not None:
full_size_img = img_xs.clone()
img_xs = torch.cat(
[
resize(img_xs[:, :, 0], scale_factor=ds_factor).unsqueeze(2),
resize(img_xs[:, :, 1], scale_factor=ds_factor).unsqueeze(2),
],
dim=2,
)
(
normal_flows,
flows,
flow_scalers,
features0,
features1,
corr_fn,
preserved_raft_flows,
) = self.cal_bidirection_flow(255 * img_xs[:, :, 0], 255 * img_xs[:, :, 1])
assert coord is not None
# List of flows
normal_inr_flows = self.predict_flow(normal_flows, coord, t, flows)
############ Unnormalize the predicted/reconstructed flow ############
start_idx = 0
if coord[0][1] is not None:
# Subsmapled flows for reconstruction supervision in the GIMM module
# In such case, first two coords in the list are subsampled for supervision up-mentioned
# Normalized flow_t towards positive t-axis
assert len(coord) > 2
flow_t = [
unnormalize_flow(normal_inr_flows[i], flow_scalers).squeeze()
for i in range(2, len(coord))
]
start_idx = 2
else:
flow_t = [
unnormalize_flow(normal_inr_flows[i], flow_scalers).squeeze()
for i in range(len(coord))
]
imgt_preds, flowt0_preds, flowt1_preds, all_others = [], [], [], []
for idx in range(start_idx, len(coord)):
cur_flow_t = flow_t[idx - start_idx]
cur_t = t[idx].reshape(-1, 1, 1, 1)
if cur_flow_t.ndim != 4:
cur_flow_t = cur_flow_t.unsqueeze(0)
assert cur_flow_t.ndim == 4
imgt_pred, flowt0_pred, flowt1_pred, others = self.frame_synthesize(
img_xs,
cur_flow_t,
features0,
features1,
corr_fn,
cur_t,
full_img=full_size_img,
)
imgt_preds.append(imgt_pred)
flowt0_preds.append(flowt0_pred)
flowt1_preds.append(flowt1_pred)
all_others.append(others)
return {
"imgt_pred": imgt_preds,
"other_pred": all_others,
"flowt0_pred": flowt0_preds,
"flowt1_pred": flowt1_preds,
"raft_flow": preserved_raft_flows,
"ninrflow": normal_inr_flows,
"nflow": normal_flows,
"flowt": flow_t,
}
def warp_frame(self, frame, flow):
return warp(frame, flow)
def sample_coord_input(
self,
batch_size,
s_shape,
t_ids,
coord_range=None,
upsample_ratio=1.0,
device=None,
):
assert device is not None
assert coord_range is None
coord_inputs = self.coord_sampler(
batch_size, s_shape, t_ids, coord_range, upsample_ratio, device
)
return coord_inputs
def cal_splatting_weights(self, raft_flow01, raft_flow10):
batch_size = raft_flow01.shape[0]
raft_flows = torch.cat([raft_flow01, raft_flow10], dim=0)
## flow variance metric
sqaure_mean, mean_square = torch.split(
F.conv3d(
F.pad(
torch.cat([raft_flows**2, raft_flows], 1),
(1, 1, 1, 1),
mode="reflect",
).unsqueeze(1),
self.g_filter,
).squeeze(1),
2,
dim=1,
)
var = (
(sqaure_mean - mean_square**2)
.clamp(1e-9, None)
.sqrt()
.mean(1)
.unsqueeze(1)
)
var01 = var[:batch_size]
var10 = var[batch_size:]
## flow warp metirc
f01_warp = -warp(raft_flow10, raft_flow01)
f10_warp = -warp(raft_flow01, raft_flow10)
err01 = (
torch.nn.functional.l1_loss(
input=f01_warp, target=raft_flow01, reduction="none"
)
.mean(1)
.unsqueeze(1)
)
err02 = (
torch.nn.functional.l1_loss(
input=f10_warp, target=raft_flow10, reduction="none"
)
.mean(1)
.unsqueeze(1)
)
weights1 = 1 / (1 + err01 * self.alpha_fe) + 1 / (1 + var01 * self.alpha_v)
weights2 = 1 / (1 + err02 * self.alpha_fe) + 1 / (1 + var10 * self.alpha_v)
return weights1, weights2
def _amt_corr_scale_lookup(self, corr_fn, coord, flow0, flow1, embt, downsample=1):
# convert t -> 0 to 0 -> 1 | convert t -> 1 to 1 -> 0
# based on linear assumption
t0_scale = 1.0 / embt
t1_scale = 1.0 / (1.0 - embt)
if downsample != 1:
inv = 1 / downsample
flow0 = inv * resize(flow0, scale_factor=inv)
flow1 = inv * resize(flow1, scale_factor=inv)
corr0, corr1 = corr_fn(coord + flow1 * t1_scale, coord + flow0 * t0_scale)
corr = torch.cat([corr0, corr1], dim=1)
flow = torch.cat([flow0, flow1], dim=1)
return corr, flow
+508
View File
@@ -0,0 +1,508 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# amt: https://github.com/MCG-NKU/AMT
# motif: https://github.com/sichun233746/MoTIF
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
import torch
import torch.nn as nn
from .configs import GIMMVFIConfig
from .modules.coord_sampler import CoordSampler3D
from .modules.hyponet import HypoNet
from .modules.fi_components import *
from .raft import initialize_RAFT
from .modules.fi_utils import (
normalize_flow,
unnormalize_flow,
warp,
resize,
build_coord,
)
import torch.nn.functional as F
from .raft.corr import BidirCorrBlock
from .modules.softsplat import softsplat
class GIMMVFI_R(nn.Module):
Config = GIMMVFIConfig
def __init__(self, config: GIMMVFIConfig):
super().__init__()
self.config = config = config.copy()
self.hyponet_config = config.hyponet
self.raft_iter = 20
######### Encoder and Decoder Settings #########
#self.flow_estimator = initialize_RAFT()
cur_f_dims = [128, 96]
f_dims = [256, 128]
self.dtype = torch.float32
skip_channels = f_dims[-1] // 2
self.num_flows = 3
self.amt_last_cproj = nn.Conv2d(cur_f_dims[0], f_dims[0], 1)
self.amt_second_last_cproj = nn.Conv2d(cur_f_dims[1], f_dims[1], 1)
self.amt_fproj = nn.Conv2d(f_dims[0], f_dims[0], 1)
self.amt_init_decoder = NewInitDecoder(f_dims[0], skip_channels)
self.amt_final_decoder = NewMultiFlowDecoder(f_dims[1], skip_channels)
self.amt_update4_low = self._get_updateblock(f_dims[0] // 2, 2.0)
self.amt_update4_high = self._get_updateblock(f_dims[0] // 2, None)
self.amt_comb_block = nn.Sequential(
nn.Conv2d(3 * self.num_flows, 6 * self.num_flows, 7, 1, 3),
nn.PReLU(6 * self.num_flows),
nn.Conv2d(6 * self.num_flows, 3, 7, 1, 3),
)
################ GIMM settings #################
self.coord_sampler = CoordSampler3D(config.coord_range)
self.g_filter = torch.nn.Parameter(
torch.FloatTensor(
[
[1.0 / 16.0, 1.0 / 8.0, 1.0 / 16.0],
[1.0 / 8.0, 1.0 / 4.0, 1.0 / 8.0],
[1.0 / 16.0, 1.0 / 8.0, 1.0 / 16.0],
]
).reshape(1, 1, 1, 3, 3),
requires_grad=False,
)
self.fwarp_type = config.fwarp_type
self.alpha_v = torch.nn.Parameter(torch.FloatTensor([1]), requires_grad=True)
self.alpha_fe = torch.nn.Parameter(torch.FloatTensor([1]), requires_grad=True)
channel = 32
in_dim = 2
self.cnn_encoder = nn.Sequential(
nn.Conv2d(in_dim, channel // 2, 3, 1, 1, bias=True, groups=1),
nn.Conv2d(channel // 2, channel, 3, 1, 1, bias=True, groups=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
LateralBlock(channel),
LateralBlock(channel),
LateralBlock(channel),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(
channel, channel // 2, 3, 1, 1, padding_mode="reflect", bias=True
),
)
channel = 64
in_dim = 64
self.res_conv = nn.Sequential(
nn.Conv2d(in_dim, channel // 2, 3, 1, 1, bias=True, groups=1),
nn.Conv2d(channel // 2, channel, 3, 1, 1, bias=True, groups=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
LateralBlock(channel),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(
channel, channel // 2, 3, 1, 1, padding_mode="reflect", bias=True
),
)
self.hyponet = HypoNet(config.hyponet, add_coord_dim=32)
def _get_updateblock(self, cdim, scale_factor=None):
return BasicUpdateBlock(
cdim=cdim,
hidden_dim=192,
flow_dim=64,
corr_dim=256,
corr_dim2=192,
fc_dim=188,
scale_factor=scale_factor,
corr_levels=4,
radius=4,
)
def cal_bidirection_flow(self, im0, im1, iters=20):
f01, features0, fnet0 = self.flow_estimator(
im0, im1, return_feat=True, iters=20
)
f10, features1, fnet1 = self.flow_estimator(
im1, im0, return_feat=True, iters=20
)
corr_fn = BidirCorrBlock(self.amt_fproj(fnet0), self.amt_fproj(fnet1), radius=4)
features0 = [
self.amt_second_last_cproj(features0[0]),
self.amt_last_cproj(features0[1]),
]
features1 = [
self.amt_second_last_cproj(features1[0]),
self.amt_last_cproj(features1[1]),
]
flow01 = f01.unsqueeze(2)
flow10 = f10.unsqueeze(2)
noraml_flows = torch.cat([flow01, -flow10], dim=2)
noraml_flows, flow_scalers = normalize_flow(noraml_flows)
ori_flows = torch.cat([flow01, flow10], dim=2)
return (
noraml_flows,
ori_flows,
flow_scalers,
features0,
features1,
corr_fn,
torch.cat([f01.unsqueeze(2), f10.unsqueeze(2)], dim=2),
)
def predict_flow(self, f, coord, t, flows):
raft_flow01 = flows[:, :, 0].detach()
raft_flow10 = flows[:, :, 1].detach()
# calculate splatting metrics
weights1, weights2 = self.cal_splatting_weights(raft_flow01, raft_flow10)
strtype = self.fwarp_type + "-zeroeps"
# b,c,h,w
pixel_latent_0 = self.cnn_encoder(f[:, :, 0])
pixel_latent_1 = self.cnn_encoder(f[:, :, 1])
pixel_latent = []
for i, cur_t in enumerate(t):
cur_t = cur_t.reshape(-1, 1, 1, 1)
tmp_pixel_latent_0 = softsplat(
tenIn=pixel_latent_0,
tenFlow=raft_flow01 * cur_t,
tenMetric=weights1,
strMode=strtype,
)
tmp_pixel_latent_1 = softsplat(
tenIn=pixel_latent_1,
tenFlow=raft_flow10 * (1 - cur_t),
tenMetric=weights2,
strMode=strtype,
)
tmp_pixel_latent = torch.cat(
[tmp_pixel_latent_0, tmp_pixel_latent_1], dim=1
)
tmp_pixel_latent = tmp_pixel_latent + self.res_conv(
torch.cat([pixel_latent_0, pixel_latent_1, tmp_pixel_latent], dim=1)
)
pixel_latent.append(tmp_pixel_latent.permute(0, 2, 3, 1))
all_outputs = []
permute_idx_range = [i for i in range(1, f.ndim - 1)]
for idx, c in enumerate(coord):
assert c[0][0, 0, 0, 0, 0] == t[idx][0].squeeze()
assert isinstance(c, tuple)
if c[1] is None:
outputs = self.hyponet(
c, modulation_params_dict=None, pixel_latent=pixel_latent[idx]
).permute(0, -1, *permute_idx_range)
else:
outputs = self.hyponet(
c, modulation_params_dict=None, pixel_latent=pixel_latent[idx]
)
all_outputs.append(outputs)
return all_outputs
def warp_w_mask(self, img0, img1, ft0, ft1, mask, scale=1):
ft0 = scale * resize(ft0, scale_factor=scale)
ft1 = scale * resize(ft1, scale_factor=scale)
mask = resize(mask, scale_factor=scale).sigmoid()
img0_warp = warp(img0, ft0)
img1_warp = warp(img1, ft1)
img_warp = mask * img0_warp + (1 - mask) * img1_warp
return img_warp
def frame_synthesize(
self, img_xs, flow_t, features0, features1, corr_fn, cur_t, full_img=None
):
"""
flow_t: b,2,h,w
cur_t: b,1,1,1
"""
batch_size = img_xs.shape[0] # b,c,t,h,w
img0 = 2 * img_xs[:, :, 0] - 1.0
img1 = 2 * img_xs[:, :, 1] - 1.0
##################### update the predicted flow #####################
##initialize coordinates for looking up
lookup_coord = build_coord(img_xs[:, :, 0]).to(
img_xs[:, :, 0].device
) # H//8,W//8
flow_t0_fullsize = flow_t * (-cur_t)
flow_t1_fullsize = flow_t * (1.0 - cur_t)
inv = 1 / 4
flow_t0_inr4 = inv * resize(flow_t0_fullsize, inv)
flow_t1_inr4 = inv * resize(flow_t1_fullsize, inv)
############################# scale 1/4 #############################
# i. Initialize feature t at scale 1/4
flowt0_4, flowt1_4, ft_4_ = self.amt_init_decoder(
features0[-1],
features1[-1],
flow_t0_inr4,
flow_t1_inr4,
img0=img0,
img1=img1,
)
features0, features1 = features0[:-1], features1[:-1]
mask_4_, ft_4_ = ft_4_[:, :1], ft_4_[:, 1:]
img_warp_4 = self.warp_w_mask(img0, img1, flowt0_4, flowt1_4, mask_4_, scale=4)
img_warp_4 = (img_warp_4 + 1.0) / 2
img_warp_4 = torch.clamp(img_warp_4, 0, 1)
corr_4, flow_4_lr = self._amt_corr_scale_lookup(
corr_fn, lookup_coord, flowt0_4, flowt1_4, cur_t, downsample=2
)
delta_ft_4_, delta_flow_4 = self.amt_update4_low(ft_4_, flow_4_lr, corr_4)
delta_flow0_4, delta_flow1_4 = torch.chunk(delta_flow_4, 2, 1)
flowt0_4 = flowt0_4 + delta_flow0_4
flowt1_4 = flowt1_4 + delta_flow1_4
ft_4_ = ft_4_ + delta_ft_4_
# iii. residue update with lookup corr
corr_4 = resize(corr_4, scale_factor=2.0)
flow_4 = torch.cat([flowt0_4, flowt1_4], dim=1)
delta_ft_4_, delta_flow_4 = self.amt_update4_high(ft_4_, flow_4, corr_4)
flowt0_4 = flowt0_4 + delta_flow_4[:, :2]
flowt1_4 = flowt1_4 + delta_flow_4[:, 2:4]
ft_4_ = ft_4_ + delta_ft_4_
############################# scale 1/1 #############################
flowt0_1, flowt1_1, mask, img_res = self.amt_final_decoder(
ft_4_,
features0[0],
features1[0],
flowt0_4,
flowt1_4,
mask=mask_4_,
img0=img0,
img1=img1,
)
if full_img is not None:
img0 = 2 * full_img[:, :, 0] - 1.0
img1 = 2 * full_img[:, :, 1] - 1.0
inv = img1.shape[2] / flowt0_1.shape[2]
flowt0_1 = inv * resize(flowt0_1, scale_factor=inv)
flowt1_1 = inv * resize(flowt1_1, scale_factor=inv)
flow_t0_fullsize = inv * resize(flow_t0_fullsize, scale_factor=inv)
flow_t1_fullsize = inv * resize(flow_t1_fullsize, scale_factor=inv)
mask = resize(mask, scale_factor=inv)
img_res = resize(img_res, scale_factor=inv)
imgt_pred = multi_flow_combine(
self.amt_comb_block, img0, img1, flowt0_1, flowt1_1, mask, img_res, None
)
imgt_pred = torch.clamp(imgt_pred, 0, 1)
######################################################################
flowt0_1 = flowt0_1.reshape(
batch_size, self.num_flows, 2, img0.shape[-2], img0.shape[-1]
)
flowt1_1 = flowt1_1.reshape(
batch_size, self.num_flows, 2, img0.shape[-2], img0.shape[-1]
)
flowt0_pred = [flowt0_1, flowt0_4]
flowt1_pred = [flowt1_1, flowt1_4]
other_pred = [img_warp_4]
return imgt_pred, flowt0_pred, flowt1_pred, other_pred
def forward(self, img_xs, coord=None, t=None, iters=None, ds_factor=None):
assert isinstance(t, list)
assert isinstance(coord, list)
assert len(t) == len(coord)
full_size_img = None
if ds_factor is not None:
full_size_img = img_xs.clone()
img_xs = torch.cat(
[
resize(img_xs[:, :, 0], scale_factor=ds_factor).unsqueeze(2),
resize(img_xs[:, :, 1], scale_factor=ds_factor).unsqueeze(2),
],
dim=2,
)
iters = self.raft_iter if iters is None else iters
(
normal_flows,
flows,
flow_scalers,
features0,
features1,
corr_fn,
preserved_raft_flows,
) = self.cal_bidirection_flow(
255 * img_xs[:, :, 0], 255 * img_xs[:, :, 1], iters=iters
)
assert coord is not None
# List of flows
normal_inr_flows = self.predict_flow(normal_flows, coord, t, flows)
############ Unnormalize the predicted/reconstructed flow ############
start_idx = 0
if coord[0][1] is not None:
# Subsmapled flows for reconstruction supervision in the GIMM module
# In such case, by default, first two coords are subsampled for supervision up-mentioned
# normalized flow_t versus positive t-axis
assert len(coord) > 2
flow_t = [
unnormalize_flow(normal_inr_flows[i], flow_scalers).squeeze()
for i in range(2, len(coord))
]
start_idx = 2
else:
flow_t = [
unnormalize_flow(normal_inr_flows[i], flow_scalers).squeeze()
for i in range(len(coord))
]
imgt_preds, flowt0_preds, flowt1_preds, all_others = [], [], [], []
for idx in range(start_idx, len(coord)):
cur_flow_t = flow_t[idx - start_idx]
cur_t = t[idx].reshape(-1, 1, 1, 1)
if cur_flow_t.ndim != 4:
cur_flow_t = cur_flow_t.unsqueeze(0)
assert cur_flow_t.ndim == 4
imgt_pred, flowt0_pred, flowt1_pred, others = self.frame_synthesize(
img_xs,
cur_flow_t,
features0,
features1,
corr_fn,
cur_t,
full_img=full_size_img,
)
imgt_preds.append(imgt_pred)
flowt0_preds.append(flowt0_pred)
flowt1_preds.append(flowt1_pred)
all_others.append(others)
return {
"imgt_pred": imgt_preds,
"other_pred": all_others,
"flowt0_pred": flowt0_preds,
"flowt1_pred": flowt1_preds,
"raft_flow": preserved_raft_flows,
"ninrflow": normal_inr_flows,
"nflow": normal_flows,
"flowt": flow_t,
}
def warp_frame(self, frame, flow):
return warp(frame, flow)
def compute_psnr(self, preds, targets, reduction="mean"):
assert reduction in ["mean", "sum", "none"]
batch_size = preds.shape[0]
sample_mses = torch.reshape((preds - targets) ** 2, (batch_size, -1)).mean(
dim=-1
)
if reduction == "mean":
psnr = (-10 * torch.log10(sample_mses)).mean()
elif reduction == "sum":
psnr = (-10 * torch.log10(sample_mses)).sum()
else:
psnr = -10 * torch.log10(sample_mses)
return psnr
def sample_coord_input(
self,
batch_size,
s_shape,
t_ids,
coord_range=None,
upsample_ratio=1.0,
device=None,
):
assert device is not None
assert coord_range is None
coord_inputs = self.coord_sampler(
batch_size, s_shape, t_ids, coord_range, upsample_ratio, device
)
return coord_inputs
def cal_splatting_weights(self, raft_flow01, raft_flow10):
batch_size = raft_flow01.shape[0]
raft_flows = torch.cat([raft_flow01, raft_flow10], dim=0)
## flow variance metric
sqaure_mean, mean_square = torch.split(
F.conv3d(
F.pad(
torch.cat([raft_flows**2, raft_flows], 1),
(1, 1, 1, 1),
mode="reflect",
).unsqueeze(1),
self.g_filter,
).squeeze(1),
2,
dim=1,
)
var = (
(sqaure_mean - mean_square**2)
.clamp(1e-9, None)
.sqrt()
.mean(1)
.unsqueeze(1)
)
var01 = var[:batch_size]
var10 = var[batch_size:]
## flow warp metirc
f01_warp = -warp(raft_flow10, raft_flow01)
f10_warp = -warp(raft_flow01, raft_flow10)
err01 = (
torch.nn.functional.l1_loss(
input=f01_warp, target=raft_flow01, reduction="none"
)
.mean(1)
.unsqueeze(1)
)
err02 = (
torch.nn.functional.l1_loss(
input=f10_warp, target=raft_flow10, reduction="none"
)
.mean(1)
.unsqueeze(1)
)
weights1 = 1 / (1 + err01 * self.alpha_fe) + 1 / (1 + var01 * self.alpha_v)
weights2 = 1 / (1 + err02 * self.alpha_fe) + 1 / (1 + var10 * self.alpha_v)
return weights1, weights2
def _amt_corr_scale_lookup(self, corr_fn, coord, flow0, flow1, embt, downsample=1):
# convert t -> 0 to 0 -> 1 | convert t -> 1 to 1 -> 0
# based on linear assumption
t0_scale = 1.0 / embt
t1_scale = 1.0 / (1.0 - embt)
if downsample != 1:
inv = 1 / downsample
flow0 = inv * resize(flow0, scale_factor=inv)
flow1 = inv * resize(flow1, scale_factor=inv)
corr0, corr1 = corr_fn(coord + flow1 * t1_scale, coord + flow0 * t0_scale)
corr = torch.cat([corr0, corr1], dim=1)
flow = torch.cat([flow0, flow1], dim=1)
return corr, flow
@@ -0,0 +1,91 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
import torch
import torch.nn as nn
class CoordSampler3D(nn.Module):
def __init__(self, coord_range, t_coord_only=False):
super().__init__()
self.coord_range = coord_range
self.t_coord_only = t_coord_only
def shape2coordinate(
self,
batch_size,
spatial_shape,
t_ids,
coord_range=(-1.0, 1.0),
upsample_ratio=1,
device=None,
):
coords = []
assert isinstance(t_ids, list)
_coords = torch.tensor(t_ids, device=device) / 1.0
coords.append(_coords.to(torch.float32))
for num_s in spatial_shape:
num_s = int(num_s * upsample_ratio)
_coords = (0.5 + torch.arange(num_s, device=device)) / num_s
_coords = coord_range[0] + (coord_range[1] - coord_range[0]) * _coords
coords.append(_coords)
coords = torch.meshgrid(*coords, indexing="ij")
coords = torch.stack(coords, dim=-1)
ones_like_shape = (1,) * coords.ndim
coords = coords.unsqueeze(0).repeat(batch_size, *ones_like_shape)
return coords # (B,T,H,W,3)
def batchshape2coordinate(
self,
batch_size,
spatial_shape,
t_ids,
coord_range=(-1.0, 1.0),
upsample_ratio=1,
device=None,
):
coords = []
_coords = torch.tensor(1, device=device)
coords.append(_coords.to(torch.float32))
for num_s in spatial_shape:
num_s = int(num_s * upsample_ratio)
_coords = (0.5 + torch.arange(num_s, device=device)) / num_s
_coords = coord_range[0] + (coord_range[1] - coord_range[0]) * _coords
coords.append(_coords)
coords = torch.meshgrid(*coords, indexing="ij")
coords = torch.stack(coords, dim=-1)
ones_like_shape = (1,) * coords.ndim
# Now coords b,1,h,w,3, coords[...,0]=1.
coords = coords.unsqueeze(0).repeat(batch_size, *ones_like_shape)
# assign per-sample timestep within the batch
coords[..., :1] = coords[..., :1] * t_ids.reshape(-1, 1, 1, 1, 1)
return coords
def forward(
self,
batch_size,
s_shape,
t_ids,
coord_range=None,
upsample_ratio=1.0,
device=None,
):
coord_range = self.coord_range if coord_range is None else coord_range
if isinstance(t_ids, list):
coords = self.shape2coordinate(
batch_size, s_shape, t_ids, coord_range, upsample_ratio, device
)
elif isinstance(t_ids, torch.Tensor):
coords = self.batchshape2coordinate(
batch_size, s_shape, t_ids, coord_range, upsample_ratio, device
)
if self.t_coord_only:
coords = coords[..., :1]
return coords
@@ -0,0 +1,340 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# amt: https://github.com/MCG-NKU/AMT
# motif: https://github.com/sichun233746/MoTIF
# --------------------------------------------------------
import torch
import torch.nn as nn
from .fi_utils import warp, resize
class LateralBlock(nn.Module):
def __init__(self, dim):
super(LateralBlock, self).__init__()
self.layers = nn.Sequential(
nn.Conv2d(dim, dim, 3, 1, 1, bias=True),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(dim, dim, 3, 1, 1, bias=True),
)
def forward(self, x):
res = x
x = self.layers(x)
return x + res
def convrelu(
in_channels,
out_channels,
kernel_size=3,
stride=1,
padding=1,
dilation=1,
groups=1,
bias=True,
):
return nn.Sequential(
nn.Conv2d(
in_channels,
out_channels,
kernel_size,
stride,
padding,
dilation,
groups,
bias=bias,
),
nn.PReLU(out_channels),
)
def multi_flow_combine(
comb_block, img0, img1, flow0, flow1, mask=None, img_res=None, mean=None
):
assert mean is None
b, c, h, w = flow0.shape
num_flows = c // 2
flow0 = flow0.reshape(b, num_flows, 2, h, w).reshape(-1, 2, h, w)
flow1 = flow1.reshape(b, num_flows, 2, h, w).reshape(-1, 2, h, w)
mask = (
mask.reshape(b, num_flows, 1, h, w).reshape(-1, 1, h, w)
if mask is not None
else None
)
img_res = (
img_res.reshape(b, num_flows, 3, h, w).reshape(-1, 3, h, w)
if img_res is not None
else 0
)
img0 = torch.stack([img0] * num_flows, 1).reshape(-1, 3, h, w)
img1 = torch.stack([img1] * num_flows, 1).reshape(-1, 3, h, w)
mean = (
torch.stack([mean] * num_flows, 1).reshape(-1, 1, 1, 1)
if mean is not None
else 0
)
img0_warp = warp(img0, flow0)
img1_warp = warp(img1, flow1)
img_warps = mask * img0_warp + (1 - mask) * img1_warp + mean + img_res
img_warps = img_warps.reshape(b, num_flows, 3, h, w)
res = comb_block(img_warps.view(b, -1, h, w))
imgt_pred = img_warps.mean(1) + res
imgt_pred = (imgt_pred + 1.0) / 2
return imgt_pred
class ResBlock(nn.Module):
def __init__(self, in_channels, side_channels, bias=True):
super(ResBlock, self).__init__()
self.side_channels = side_channels
self.conv1 = nn.Sequential(
nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
),
nn.PReLU(in_channels),
)
self.conv2 = nn.Sequential(
nn.Conv2d(
side_channels,
side_channels,
kernel_size=3,
stride=1,
padding=1,
bias=bias,
),
nn.PReLU(side_channels),
)
self.conv3 = nn.Sequential(
nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
),
nn.PReLU(in_channels),
)
self.conv4 = nn.Sequential(
nn.Conv2d(
side_channels,
side_channels,
kernel_size=3,
stride=1,
padding=1,
bias=bias,
),
nn.PReLU(side_channels),
)
self.conv5 = nn.Conv2d(
in_channels, in_channels, kernel_size=3, stride=1, padding=1, bias=bias
)
self.prelu = nn.PReLU(in_channels)
def forward(self, x):
out = self.conv1(x)
res_feat = out[:, : -self.side_channels, ...]
side_feat = out[:, -self.side_channels :, :, :]
side_feat = self.conv2(side_feat)
out = self.conv3(torch.cat([res_feat, side_feat], 1))
res_feat = out[:, : -self.side_channels, ...]
side_feat = out[:, -self.side_channels :, :, :]
side_feat = self.conv4(side_feat)
out = self.conv5(torch.cat([res_feat, side_feat], 1))
out = self.prelu(x + out)
return out
class BasicUpdateBlock(nn.Module):
def __init__(
self,
cdim,
hidden_dim,
flow_dim,
corr_dim,
corr_dim2,
fc_dim,
corr_levels=4,
radius=3,
scale_factor=None,
out_num=1,
):
super(BasicUpdateBlock, self).__init__()
cor_planes = corr_levels * (2 * radius + 1) ** 2
self.scale_factor = scale_factor
self.convc1 = nn.Conv2d(2 * cor_planes, corr_dim, 1, padding=0)
self.convc2 = nn.Conv2d(corr_dim, corr_dim2, 3, padding=1)
self.convf1 = nn.Conv2d(4, flow_dim * 2, 7, padding=3)
self.convf2 = nn.Conv2d(flow_dim * 2, flow_dim, 3, padding=1)
self.conv = nn.Conv2d(flow_dim + corr_dim2, fc_dim, 3, padding=1)
self.gru = nn.Sequential(
nn.Conv2d(fc_dim + 4 + cdim, hidden_dim, 3, padding=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(hidden_dim, hidden_dim, 3, padding=1),
)
self.feat_head = nn.Sequential(
nn.Conv2d(hidden_dim, hidden_dim, 3, padding=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(hidden_dim, cdim, 3, padding=1),
)
self.flow_head = nn.Sequential(
nn.Conv2d(hidden_dim, hidden_dim, 3, padding=1),
nn.LeakyReLU(negative_slope=0.1, inplace=True),
nn.Conv2d(hidden_dim, 4 * out_num, 3, padding=1),
)
self.lrelu = nn.LeakyReLU(negative_slope=0.1, inplace=True)
def forward(self, net, flow, corr):
net = (
resize(net, 1 / self.scale_factor) if self.scale_factor is not None else net
)
cor = self.lrelu(self.convc1(corr))
cor = self.lrelu(self.convc2(cor))
flo = self.lrelu(self.convf1(flow))
flo = self.lrelu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
inp = self.lrelu(self.conv(cor_flo))
inp = torch.cat([inp, flow, net], dim=1)
out = self.gru(inp)
delta_net = self.feat_head(out)
delta_flow = self.flow_head(out)
if self.scale_factor is not None:
delta_net = resize(delta_net, scale_factor=self.scale_factor)
delta_flow = self.scale_factor * resize(
delta_flow, scale_factor=self.scale_factor
)
return delta_net, delta_flow
def get_bn():
return nn.BatchNorm2d
class NewInitDecoder(nn.Module):
def __init__(self, in_ch, skip_ch):
super().__init__()
norm_layer = get_bn()
self.upsample = nn.Sequential(
nn.PixelShuffle(2),
convrelu(in_ch // 4, in_ch // 4, 5, 1, 2),
convrelu(in_ch // 4, in_ch // 4),
convrelu(in_ch // 4, in_ch // 4),
convrelu(in_ch // 4, in_ch // 4),
convrelu(in_ch // 4, in_ch // 2),
nn.Conv2d(in_ch // 2, in_ch // 2, kernel_size=1),
norm_layer(in_ch // 2),
nn.ReLU(inplace=True),
)
in_ch = in_ch // 2
self.convblock = nn.Sequential(
convrelu(in_ch * 2 + 16, in_ch, kernel_size=1, padding=0),
ResBlock(in_ch, skip_ch),
ResBlock(in_ch, skip_ch),
ResBlock(in_ch, skip_ch),
nn.Conv2d(in_ch, in_ch + 5, 3, 1, 1, 1, 1, True),
)
def forward(self, f0, f1, flow0_in, flow1_in, img0=None, img1=None):
f0 = self.upsample(f0)
f1 = self.upsample(f1)
f0_warp_ks = warp(f0, flow0_in)
f1_warp_ks = warp(f1, flow1_in)
f_in = torch.cat([f0_warp_ks, f1_warp_ks, flow0_in, flow1_in], dim=1)
assert img0 is not None
assert img1 is not None
scale_factor = f_in.shape[2] / img0.shape[2]
img0 = resize(img0, scale_factor=scale_factor)
img1 = resize(img1, scale_factor=scale_factor)
warped_img0 = warp(img0, flow0_in)
warped_img1 = warp(img1, flow1_in)
f_in = torch.cat([f_in, img0, img1, warped_img0, warped_img1], dim=1)
out = self.convblock(f_in)
ft_ = out[:, 4:, ...]
flow0 = flow0_in + out[:, :2, ...]
flow1 = flow1_in + out[:, 2:4, ...]
return flow0, flow1, ft_
class NewMultiFlowDecoder(nn.Module):
def __init__(self, in_ch, skip_ch, num_flows=3):
super(NewMultiFlowDecoder, self).__init__()
norm_layer = get_bn()
self.upsample = nn.Sequential(
nn.PixelShuffle(2),
nn.PixelShuffle(2),
convrelu(in_ch // (4 * 4), in_ch // 4, 5, 1, 2),
convrelu(in_ch // 4, in_ch // 4),
convrelu(in_ch // 4, in_ch // 4),
convrelu(in_ch // 4, in_ch // 4),
convrelu(in_ch // 4, in_ch // 2),
nn.Conv2d(in_ch // 2, in_ch // 2, kernel_size=1),
norm_layer(in_ch // 2),
nn.ReLU(inplace=True),
)
self.num_flows = num_flows
ch_factor = 2
self.convblock = nn.Sequential(
convrelu(in_ch * ch_factor + 17, in_ch * ch_factor),
ResBlock(in_ch * ch_factor, skip_ch),
ResBlock(in_ch * ch_factor, skip_ch),
ResBlock(in_ch * ch_factor, skip_ch),
nn.Conv2d(in_ch * ch_factor, 8 * num_flows, kernel_size=3, padding=1),
)
def forward(self, ft_, f0, f1, flow0, flow1, mask=None, img0=None, img1=None):
f0 = self.upsample(f0)
# print([f1.shape,f0.shape])
f1 = self.upsample(f1)
n = self.num_flows
flow0 = 4.0 * resize(flow0, scale_factor=4.0)
flow1 = 4.0 * resize(flow1, scale_factor=4.0)
ft_ = resize(ft_, scale_factor=4.0)
mask = resize(mask, scale_factor=4.0)
f0_warp = warp(f0, flow0)
f1_warp = warp(f1, flow1)
f_in = torch.cat([ft_, f0_warp, f1_warp, flow0, flow1], 1)
assert mask is not None
f_in = torch.cat([f_in, mask], 1)
assert img0 is not None
assert img1 is not None
warped_img0 = warp(img0, flow0)
warped_img1 = warp(img1, flow1)
f_in = torch.cat([f_in, img0, img1, warped_img0, warped_img1], dim=1)
out = self.convblock(f_in)
delta_flow0, delta_flow1, delta_mask, img_res = torch.split(
out, [2 * n, 2 * n, n, 3 * n], 1
)
mask = delta_mask + mask.repeat(1, self.num_flows, 1, 1)
mask = torch.sigmoid(mask)
flow0 = delta_flow0 + flow0.repeat(1, self.num_flows, 1, 1)
flow1 = delta_flow1 + flow1.repeat(1, self.num_flows, 1, 1)
return flow0, flow1, mask, img_res
@@ -0,0 +1,82 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# raft: https://github.com/princeton-vl/RAFT
# ema-vfi: https://github.com/MCG-NJU/EMA-VFI
# --------------------------------------------------------
import torch
import torch.nn.functional as F
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
backwarp_tenGrid = {}
def warp(tenInput, tenFlow):
k = (str(tenFlow.device), str(tenFlow.size()))
if k not in backwarp_tenGrid:
tenHorizontal = (
torch.linspace(-1.0, 1.0, tenFlow.shape[3], device=device)
.view(1, 1, 1, tenFlow.shape[3])
.expand(tenFlow.shape[0], -1, tenFlow.shape[2], -1)
)
tenVertical = (
torch.linspace(-1.0, 1.0, tenFlow.shape[2], device=device)
.view(1, 1, tenFlow.shape[2], 1)
.expand(tenFlow.shape[0], -1, -1, tenFlow.shape[3])
)
backwarp_tenGrid[k] = torch.cat([tenHorizontal, tenVertical], 1).to(device)
tenFlow = torch.cat(
[
tenFlow[:, 0:1, :, :] / ((tenInput.shape[3] - 1.0) / 2.0),
tenFlow[:, 1:2, :, :] / ((tenInput.shape[2] - 1.0) / 2.0),
],
1,
)
g = (backwarp_tenGrid[k] + tenFlow).permute(0, 2, 3, 1)
return torch.nn.functional.grid_sample(
input=tenInput,
grid=g,
mode="bilinear",
padding_mode="border",
align_corners=True,
)
def normalize_flow(flows):
# FIXME: MULTI-DIMENSION
flow_scaler = torch.max(torch.abs(flows).flatten(1), dim=-1)[0].reshape(
-1, 1, 1, 1, 1
)
flows = flows / flow_scaler # [-1,1]
# # Adapt to [0,1]
flows = (flows + 1.0) / 2.0
return flows, flow_scaler
def unnormalize_flow(flows, flow_scaler):
return (flows * 2.0 - 1.0) * flow_scaler
def resize(x, scale_factor):
return F.interpolate(
x, scale_factor=scale_factor, mode="bilinear", align_corners=False
)
def coords_grid(batch, ht, wd):
coords = torch.meshgrid(torch.arange(ht), torch.arange(wd))
coords = torch.stack(coords[::-1], dim=0).float()
return coords[None].repeat(batch, 1, 1, 1)
def build_coord(img):
N, C, H, W = img.shape
coords = coords_grid(N, H // 8, W // 8)
return coords
@@ -0,0 +1,198 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
import einops
import torch
import torch.nn as nn
import torch.nn.functional as F
from omegaconf import OmegaConf
from ..configs import HypoNetConfig
from .utils import create_params_with_init, create_activation
class HypoNet(nn.Module):
r"""
The Hyponetwork with a coordinate-based MLP to be modulated.
"""
def __init__(self, config: HypoNetConfig, add_coord_dim=32):
super().__init__()
self.config = config
self.use_bias = config.use_bias
self.init_config = config.initialization
self.num_layer = config.n_layer
self.hidden_dims = config.hidden_dim
self.add_coord_dim = add_coord_dim
if len(self.hidden_dims) == 1:
self.hidden_dims = OmegaConf.to_object(self.hidden_dims) * (
self.num_layer - 1
) # exclude output layer
else:
assert len(self.hidden_dims) == self.num_layer - 1
if self.config.activation.type == "siren":
assert self.init_config.weight_init_type == "siren"
assert self.init_config.bias_init_type == "siren"
# after computes the shape of trainable parameters, initialize them
self.params_dict = None
self.params_shape_dict = self.compute_params_shape()
self.activation = create_activation(self.config.activation)
self.build_base_params_dict(self.config.initialization)
self.output_bias = config.output_bias
self.normalize_weight = config.normalize_weight
self.ignore_base_param_dict = {name: False for name in self.params_dict}
@staticmethod
def subsample_coords(coords, subcoord_idx=None):
if subcoord_idx is None:
return coords
batch_size = coords.shape[0]
sub_coords = []
coords = coords.view(batch_size, -1, coords.shape[-1])
for idx in range(batch_size):
sub_coords.append(coords[idx : idx + 1, subcoord_idx[idx]])
sub_coords = torch.cat(sub_coords, dim=0)
return sub_coords
def forward(self, coord, modulation_params_dict=None, pixel_latent=None):
sub_idx = None
if isinstance(coord, tuple):
coord, sub_idx = coord[0], coord[1]
if modulation_params_dict is not None:
self.check_valid_param_keys(modulation_params_dict)
batch_size, coord_shape, input_dim = (
coord.shape[0],
coord.shape[1:-1],
coord.shape[-1],
)
coord = coord.view(batch_size, -1, input_dim) # flatten the coordinates
assert pixel_latent is not None
pixel_latent = F.interpolate(
pixel_latent.permute(0, 3, 1, 2),
size=(coord_shape[1], coord_shape[2]),
mode="bilinear",
).permute(0, 2, 3, 1)
pixel_latent_dim = pixel_latent.shape[-1]
pixel_latent = pixel_latent.view(batch_size, -1, pixel_latent_dim)
hidden = coord
hidden = torch.cat([pixel_latent, hidden], dim=-1)
hidden = self.subsample_coords(hidden, sub_idx)
for idx in range(self.config.n_layer):
param_key = f"linear_wb{idx}"
base_param = einops.repeat(
self.params_dict[param_key], "n m -> b n m", b=batch_size
)
if (modulation_params_dict is not None) and (
param_key in modulation_params_dict.keys()
):
modulation_param = modulation_params_dict[param_key]
else:
if self.config.use_bias:
modulation_param = torch.ones_like(base_param[:, :-1])
else:
modulation_param = torch.ones_like(base_param)
if self.config.use_bias:
ones = torch.ones(*hidden.shape[:-1], 1, device=hidden.device)
hidden = torch.cat([hidden, ones], dim=-1)
base_param_w, base_param_b = (
base_param[:, :-1, :],
base_param[:, -1:, :],
)
if self.ignore_base_param_dict[param_key]:
base_param_w = 1.0
param_w = base_param_w * modulation_param
if self.normalize_weight:
param_w = F.normalize(param_w, dim=1)
modulated_param = torch.cat([param_w, base_param_b], dim=1)
else:
if self.ignore_base_param_dict[param_key]:
base_param = 1.0
if self.normalize_weight:
modulated_param = F.normalize(base_param * modulation_param, dim=1)
else:
modulated_param = base_param * modulation_param
# print([param_key,hidden.shape,modulated_param.shape])
hidden = torch.bmm(hidden, modulated_param)
if idx < (self.config.n_layer - 1):
hidden = self.activation(hidden)
outputs = hidden + self.output_bias
if sub_idx is None:
outputs = outputs.view(batch_size, *coord_shape, -1)
return outputs
def compute_params_shape(self):
"""
Computes the shape of MLP parameters.
The computed shapes are used to build the initial weights by `build_base_params_dict`.
"""
config = self.config
use_bias = self.use_bias
param_shape_dict = dict()
fan_in = config.input_dim
add_dim = self.add_coord_dim
fan_in = fan_in + add_dim
fan_in = fan_in + 1 if use_bias else fan_in
for i in range(config.n_layer - 1):
fan_out = self.hidden_dims[i]
param_shape_dict[f"linear_wb{i}"] = (fan_in, fan_out)
fan_in = fan_out + 1 if use_bias else fan_out
param_shape_dict[f"linear_wb{config.n_layer-1}"] = (fan_in, config.output_dim)
return param_shape_dict
def build_base_params_dict(self, init_config):
assert self.params_shape_dict
params_dict = nn.ParameterDict()
for idx, (name, shape) in enumerate(self.params_shape_dict.items()):
is_first = idx == 0
params = create_params_with_init(
shape,
init_type=init_config.weight_init_type,
include_bias=self.use_bias,
bias_init_type=init_config.bias_init_type,
is_first=is_first,
siren_w0=self.config.activation.siren_w0, # valid only for siren
)
params = nn.Parameter(params)
params_dict[name] = params
self.set_params_dict(params_dict)
def check_valid_param_keys(self, params_dict):
predefined_params_keys = self.params_shape_dict.keys()
for param_key in params_dict.keys():
if param_key in predefined_params_keys:
continue
else:
raise KeyError
def set_params_dict(self, params_dict):
self.check_valid_param_keys(params_dict)
self.params_dict = params_dict
@@ -0,0 +1,42 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
from torch import nn
import torch
# define siren layer & Siren model
class Sine(nn.Module):
"""Sine activation with scaling.
Args:
w0 (float): Omega_0 parameter from SIREN paper.
"""
def __init__(self, w0=1.0):
super().__init__()
self.w0 = w0
def forward(self, x):
return torch.sin(self.w0 * x)
# Damping activation from http://arxiv.org/abs/2306.15242
class Damping(nn.Module):
"""Sine activation with sublinear factor
Args:
w0 (float): Omega_0 parameter from SIREN paper.
"""
def __init__(self, w0=1.0):
super().__init__()
self.w0 = w0
def forward(self, x):
x = torch.clamp(x, min=1e-30)
return torch.sin(self.w0 * x) * torch.sqrt(x.abs())
@@ -0,0 +1,52 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
from typing import List, Optional
from dataclasses import dataclass, field
from omegaconf import MISSING
@dataclass
class HypoNetActivationConfig:
type: str = "relu"
siren_w0: Optional[float] = 30.0
@dataclass
class HypoNetInitConfig:
weight_init_type: Optional[str] = "kaiming_uniform"
bias_init_type: Optional[str] = "zero"
@dataclass
class HypoNetConfig:
type: str = "mlp"
n_layer: int = 5
hidden_dim: List[int] = MISSING
use_bias: bool = True
input_dim: int = 2
output_dim: int = 3
output_bias: float = 0.5
activation: HypoNetActivationConfig = field(default_factory=HypoNetActivationConfig)
initialization: HypoNetInitConfig = field(default_factory=HypoNetInitConfig)
normalize_weight: bool = True
linear_interpo: bool = False
@dataclass
class CoordSamplerConfig:
data_type: str = "image"
t_coord_only: bool = False
coord_range: List[float] = MISSING
time_range: List[float] = MISSING
train_strategy: Optional[str] = MISSING
val_strategy: Optional[str] = MISSING
patch_size: Optional[int] = MISSING
@@ -0,0 +1,666 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# softmax-splatting: https://github.com/sniklaus/softmax-splatting
# --------------------------------------------------------
import collections
import cupy
import os
import re
import torch
import typing
##########################################################
objCudacache = {}
def cuda_int32(intIn: int):
return cupy.int32(intIn)
# end
def cuda_float32(fltIn: float):
return cupy.float32(fltIn)
# end
def cuda_kernel(strFunction: str, strKernel: str, objVariables: typing.Dict):
if "device" not in objCudacache:
objCudacache["device"] = torch.cuda.get_device_name()
# end
strKey = strFunction
for strVariable in objVariables:
objValue = objVariables[strVariable]
strKey += strVariable
if objValue is None:
continue
elif type(objValue) == int:
strKey += str(objValue)
elif type(objValue) == float:
strKey += str(objValue)
elif type(objValue) == bool:
strKey += str(objValue)
elif type(objValue) == str:
strKey += objValue
elif type(objValue) == torch.Tensor:
strKey += str(objValue.dtype)
strKey += str(objValue.shape)
strKey += str(objValue.stride())
elif True:
print(strVariable, type(objValue))
assert False
# end
# end
strKey += objCudacache["device"]
if strKey not in objCudacache:
for strVariable in objVariables:
objValue = objVariables[strVariable]
if objValue is None:
continue
elif type(objValue) == int:
strKernel = strKernel.replace("{{" + strVariable + "}}", str(objValue))
elif type(objValue) == float:
strKernel = strKernel.replace("{{" + strVariable + "}}", str(objValue))
elif type(objValue) == bool:
strKernel = strKernel.replace("{{" + strVariable + "}}", str(objValue))
elif type(objValue) == str:
strKernel = strKernel.replace("{{" + strVariable + "}}", objValue)
elif type(objValue) == torch.Tensor and objValue.dtype == torch.uint8:
strKernel = strKernel.replace("{{type}}", "unsigned char")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.float16:
strKernel = strKernel.replace("{{type}}", "half")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.float32:
strKernel = strKernel.replace("{{type}}", "float")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.float64:
strKernel = strKernel.replace("{{type}}", "double")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.int32:
strKernel = strKernel.replace("{{type}}", "int")
elif type(objValue) == torch.Tensor and objValue.dtype == torch.int64:
strKernel = strKernel.replace("{{type}}", "long")
elif type(objValue) == torch.Tensor:
print(strVariable, objValue.dtype)
assert False
elif True:
print(strVariable, type(objValue))
assert False
# end
# end
while True:
objMatch = re.search(r"(SIZE_)([0-4])(\()([^\)]*)(\))", strKernel)
if objMatch is None:
break
# end
intArg = int(objMatch.group(2))
strTensor = objMatch.group(4)
intSizes = objVariables[strTensor].size()
strKernel = strKernel.replace(
objMatch.group(),
str(
intSizes[intArg]
if torch.is_tensor(intSizes[intArg]) == False
else intSizes[intArg].item()
),
)
# end
while True:
objMatch = re.search(r"(OFFSET_)([0-4])(\()", strKernel)
if objMatch is None:
break
# end
intStart = objMatch.span()[1]
intStop = objMatch.span()[1]
intParentheses = 1
while True:
intParentheses += 1 if strKernel[intStop] == "(" else 0
intParentheses -= 1 if strKernel[intStop] == ")" else 0
if intParentheses == 0:
break
# end
intStop += 1
# end
intArgs = int(objMatch.group(2))
strArgs = strKernel[intStart:intStop].split(",")
assert intArgs == len(strArgs) - 1
strTensor = strArgs[0]
intStrides = objVariables[strTensor].stride()
strIndex = []
for intArg in range(intArgs):
strIndex.append(
"(("
+ strArgs[intArg + 1].replace("{", "(").replace("}", ")").strip()
+ ")*"
+ str(
intStrides[intArg]
if torch.is_tensor(intStrides[intArg]) == False
else intStrides[intArg].item()
)
+ ")"
)
# end
strKernel = strKernel.replace(
"OFFSET_" + str(intArgs) + "(" + strKernel[intStart:intStop] + ")",
"(" + str.join("+", strIndex) + ")",
)
# end
while True:
objMatch = re.search(r"(VALUE_)([0-4])(\()", strKernel)
if objMatch is None:
break
# end
intStart = objMatch.span()[1]
intStop = objMatch.span()[1]
intParentheses = 1
while True:
intParentheses += 1 if strKernel[intStop] == "(" else 0
intParentheses -= 1 if strKernel[intStop] == ")" else 0
if intParentheses == 0:
break
# end
intStop += 1
# end
intArgs = int(objMatch.group(2))
strArgs = strKernel[intStart:intStop].split(",")
assert intArgs == len(strArgs) - 1
strTensor = strArgs[0]
intStrides = objVariables[strTensor].stride()
strIndex = []
for intArg in range(intArgs):
strIndex.append(
"(("
+ strArgs[intArg + 1].replace("{", "(").replace("}", ")").strip()
+ ")*"
+ str(
intStrides[intArg]
if torch.is_tensor(intStrides[intArg]) == False
else intStrides[intArg].item()
)
+ ")"
)
# end
strKernel = strKernel.replace(
"VALUE_" + str(intArgs) + "(" + strKernel[intStart:intStop] + ")",
strTensor + "[" + str.join("+", strIndex) + "]",
)
# end
objCudacache[strKey] = {"strFunction": strFunction, "strKernel": strKernel}
# end
return strKey
# end
@cupy.memoize(for_each_device=True)
def cuda_launch(strKey: str):
if "CUDA_HOME" not in os.environ:
os.environ["CUDA_HOME"] = cupy.cuda.get_cuda_path()
strKernel = objCudacache[strKey]["strKernel"]
strFunction = objCudacache[strKey]["strFunction"]
return cupy.RawModule(
code=strKernel,
options=(
"-I " + os.environ["CUDA_HOME"],
"-I " + os.environ["CUDA_HOME"] + "/include",
),
).get_function(strFunction)
##########################################################
def softsplat(tenIn, tenFlow, tenMetric, strMode, return_norm=False):
assert strMode.split("-")[0] in ["sum", "avg", "linear", "softmax"]
if strMode == "sum":
assert tenMetric is None
if strMode == "avg":
assert tenMetric is None
if strMode.split("-")[0] == "linear":
assert tenMetric is not None
if strMode.split("-")[0] == "softmax":
assert tenMetric is not None
if strMode == "avg":
tenIn = torch.cat(
[
tenIn,
tenIn.new_ones([tenIn.shape[0], 1, tenIn.shape[2], tenIn.shape[3]]),
],
1,
)
elif strMode.split("-")[0] == "linear":
tenIn = torch.cat([tenIn * tenMetric, tenMetric], 1)
elif strMode.split("-")[0] == "softmax":
tenIn = torch.cat([tenIn * tenMetric.exp(), tenMetric.exp()], 1)
# end
if torch.isnan(tenIn).any():
print("NaN values detected during training in tenIn. Exiting.")
assert False
tenOut = softsplat_func.apply(tenIn, tenFlow)
if torch.isnan(tenOut).any():
print("NaN values detected during training in tenOut_1. Exiting.")
assert False
if strMode.split("-")[0] in ["avg", "linear", "softmax"]:
tenNormalize = tenOut[:, -1:, :, :]
if len(strMode.split("-")) == 1:
tenNormalize = tenNormalize + 0.0000001
elif strMode.split("-")[1] == "addeps":
tenNormalize = tenNormalize + 0.0000001
elif strMode.split("-")[1] == "zeroeps":
tenNormalize[tenNormalize == 0.0] = 1.0
elif strMode.split("-")[1] == "clipeps":
tenNormalize = tenNormalize.clip(0.0000001, None)
# end
if return_norm:
return tenOut[:, :-1, :, :], tenNormalize
tenOut = tenOut[:, :-1, :, :] / tenNormalize
if torch.isnan(tenOut).any():
print("NaN values detected during training in tenOut_2. Exiting.")
assert False
# end
return tenOut
# end
class softsplat_func(torch.autograd.Function):
@staticmethod
@torch.amp.custom_fwd(device_type="cuda", cast_inputs=torch.float32)
def forward(self, tenIn, tenFlow):
tenOut = tenIn.new_zeros(
[tenIn.shape[0], tenIn.shape[1], tenIn.shape[2], tenIn.shape[3]]
)
if tenIn.is_cuda == True:
cuda_launch(
cuda_kernel(
"softsplat_out",
"""
extern "C" __global__ void __launch_bounds__(512) softsplat_out(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenFlow,
{{type}}* __restrict__ tenOut
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) / SIZE_1(tenOut) ) % SIZE_0(tenOut);
const int intC = ( intIndex / SIZE_3(tenOut) / SIZE_2(tenOut) ) % SIZE_1(tenOut);
const int intY = ( intIndex / SIZE_3(tenOut) ) % SIZE_2(tenOut);
const int intX = ( intIndex ) % SIZE_3(tenOut);
assert(SIZE_1(tenFlow) == 2);
{{type}} fltX = ({{type}}) (intX) + VALUE_4(tenFlow, intN, 0, intY, intX);
{{type}} fltY = ({{type}}) (intY) + VALUE_4(tenFlow, intN, 1, intY, intX);
if (isfinite(fltX) == false) { return; }
if (isfinite(fltY) == false) { return; }
{{type}} fltIn = VALUE_4(tenIn, intN, intC, intY, intX);
int intNorthwestX = (int) (floor(fltX));
int intNorthwestY = (int) (floor(fltY));
int intNortheastX = intNorthwestX + 1;
int intNortheastY = intNorthwestY;
int intSouthwestX = intNorthwestX;
int intSouthwestY = intNorthwestY + 1;
int intSoutheastX = intNorthwestX + 1;
int intSoutheastY = intNorthwestY + 1;
{{type}} fltNorthwest = (({{type}}) (intSoutheastX) - fltX) * (({{type}}) (intSoutheastY) - fltY);
{{type}} fltNortheast = (fltX - ({{type}}) (intSouthwestX)) * (({{type}}) (intSouthwestY) - fltY);
{{type}} fltSouthwest = (({{type}}) (intNortheastX) - fltX) * (fltY - ({{type}}) (intNortheastY));
{{type}} fltSoutheast = (fltX - ({{type}}) (intNorthwestX)) * (fltY - ({{type}}) (intNorthwestY));
if ((intNorthwestX >= 0) && (intNorthwestX < SIZE_3(tenOut)) && (intNorthwestY >= 0) && (intNorthwestY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intNorthwestY, intNorthwestX)], fltIn * fltNorthwest);
}
if ((intNortheastX >= 0) && (intNortheastX < SIZE_3(tenOut)) && (intNortheastY >= 0) && (intNortheastY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intNortheastY, intNortheastX)], fltIn * fltNortheast);
}
if ((intSouthwestX >= 0) && (intSouthwestX < SIZE_3(tenOut)) && (intSouthwestY >= 0) && (intSouthwestY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intSouthwestY, intSouthwestX)], fltIn * fltSouthwest);
}
if ((intSoutheastX >= 0) && (intSoutheastX < SIZE_3(tenOut)) && (intSoutheastY >= 0) && (intSoutheastY < SIZE_2(tenOut))) {
atomicAdd(&tenOut[OFFSET_4(tenOut, intN, intC, intSoutheastY, intSoutheastX)], fltIn * fltSoutheast);
}
} }
""",
{"tenIn": tenIn, "tenFlow": tenFlow, "tenOut": tenOut},
)
)(
grid=tuple([int((tenOut.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenOut.nelement()),
tenIn.data_ptr(),
tenFlow.data_ptr(),
tenOut.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
elif tenIn.is_cuda != True:
assert False
# end
self.save_for_backward(tenIn, tenFlow)
return tenOut
# end
@staticmethod
@torch.amp.custom_bwd(device_type="cuda")
def backward(self, tenOutgrad):
tenIn, tenFlow = self.saved_tensors
tenOutgrad = tenOutgrad.contiguous()
assert tenOutgrad.is_cuda == True
tenIngrad = (
tenIn.new_zeros(
[tenIn.shape[0], tenIn.shape[1], tenIn.shape[2], tenIn.shape[3]]
)
if self.needs_input_grad[0] == True
else None
)
tenFlowgrad = (
tenFlow.new_zeros(
[tenFlow.shape[0], tenFlow.shape[1], tenFlow.shape[2], tenFlow.shape[3]]
)
if self.needs_input_grad[1] == True
else None
)
if tenIngrad is not None:
cuda_launch(
cuda_kernel(
"softsplat_ingrad",
"""
extern "C" __global__ void __launch_bounds__(512) softsplat_ingrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenFlow,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenFlowgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenIngrad) / SIZE_2(tenIngrad) / SIZE_1(tenIngrad) ) % SIZE_0(tenIngrad);
const int intC = ( intIndex / SIZE_3(tenIngrad) / SIZE_2(tenIngrad) ) % SIZE_1(tenIngrad);
const int intY = ( intIndex / SIZE_3(tenIngrad) ) % SIZE_2(tenIngrad);
const int intX = ( intIndex ) % SIZE_3(tenIngrad);
assert(SIZE_1(tenFlow) == 2);
{{type}} fltIngrad = 0.0f;
{{type}} fltX = ({{type}}) (intX) + VALUE_4(tenFlow, intN, 0, intY, intX);
{{type}} fltY = ({{type}}) (intY) + VALUE_4(tenFlow, intN, 1, intY, intX);
if (isfinite(fltX) == false) { return; }
if (isfinite(fltY) == false) { return; }
int intNorthwestX = (int) (floor(fltX));
int intNorthwestY = (int) (floor(fltY));
int intNortheastX = intNorthwestX + 1;
int intNortheastY = intNorthwestY;
int intSouthwestX = intNorthwestX;
int intSouthwestY = intNorthwestY + 1;
int intSoutheastX = intNorthwestX + 1;
int intSoutheastY = intNorthwestY + 1;
{{type}} fltNorthwest = (({{type}}) (intSoutheastX) - fltX) * (({{type}}) (intSoutheastY) - fltY);
{{type}} fltNortheast = (fltX - ({{type}}) (intSouthwestX)) * (({{type}}) (intSouthwestY) - fltY);
{{type}} fltSouthwest = (({{type}}) (intNortheastX) - fltX) * (fltY - ({{type}}) (intNortheastY));
{{type}} fltSoutheast = (fltX - ({{type}}) (intNorthwestX)) * (fltY - ({{type}}) (intNorthwestY));
if ((intNorthwestX >= 0) && (intNorthwestX < SIZE_3(tenOutgrad)) && (intNorthwestY >= 0) && (intNorthwestY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intNorthwestY, intNorthwestX) * fltNorthwest;
}
if ((intNortheastX >= 0) && (intNortheastX < SIZE_3(tenOutgrad)) && (intNortheastY >= 0) && (intNortheastY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intNortheastY, intNortheastX) * fltNortheast;
}
if ((intSouthwestX >= 0) && (intSouthwestX < SIZE_3(tenOutgrad)) && (intSouthwestY >= 0) && (intSouthwestY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intSouthwestY, intSouthwestX) * fltSouthwest;
}
if ((intSoutheastX >= 0) && (intSoutheastX < SIZE_3(tenOutgrad)) && (intSoutheastY >= 0) && (intSoutheastY < SIZE_2(tenOutgrad))) {
fltIngrad += VALUE_4(tenOutgrad, intN, intC, intSoutheastY, intSoutheastX) * fltSoutheast;
}
tenIngrad[intIndex] = fltIngrad;
} }
""",
{
"tenIn": tenIn,
"tenFlow": tenFlow,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenFlowgrad": tenFlowgrad,
},
)
)(
grid=tuple([int((tenIngrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenIngrad.nelement()),
tenIn.data_ptr(),
tenFlow.data_ptr(),
tenOutgrad.data_ptr(),
tenIngrad.data_ptr(),
None,
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
# end
if tenFlowgrad is not None:
cuda_launch(
cuda_kernel(
"softsplat_flowgrad",
"""
extern "C" __global__ void __launch_bounds__(512) softsplat_flowgrad(
const int n,
const {{type}}* __restrict__ tenIn,
const {{type}}* __restrict__ tenFlow,
const {{type}}* __restrict__ tenOutgrad,
{{type}}* __restrict__ tenIngrad,
{{type}}* __restrict__ tenFlowgrad
) { for (int intIndex = (blockIdx.x * blockDim.x) + threadIdx.x; intIndex < n; intIndex += blockDim.x * gridDim.x) {
const int intN = ( intIndex / SIZE_3(tenFlowgrad) / SIZE_2(tenFlowgrad) / SIZE_1(tenFlowgrad) ) % SIZE_0(tenFlowgrad);
const int intC = ( intIndex / SIZE_3(tenFlowgrad) / SIZE_2(tenFlowgrad) ) % SIZE_1(tenFlowgrad);
const int intY = ( intIndex / SIZE_3(tenFlowgrad) ) % SIZE_2(tenFlowgrad);
const int intX = ( intIndex ) % SIZE_3(tenFlowgrad);
assert(SIZE_1(tenFlow) == 2);
{{type}} fltFlowgrad = 0.0f;
{{type}} fltX = ({{type}}) (intX) + VALUE_4(tenFlow, intN, 0, intY, intX);
{{type}} fltY = ({{type}}) (intY) + VALUE_4(tenFlow, intN, 1, intY, intX);
if (isfinite(fltX) == false) { return; }
if (isfinite(fltY) == false) { return; }
int intNorthwestX = (int) (floor(fltX));
int intNorthwestY = (int) (floor(fltY));
int intNortheastX = intNorthwestX + 1;
int intNortheastY = intNorthwestY;
int intSouthwestX = intNorthwestX;
int intSouthwestY = intNorthwestY + 1;
int intSoutheastX = intNorthwestX + 1;
int intSoutheastY = intNorthwestY + 1;
{{type}} fltNorthwest = 0.0f;
{{type}} fltNortheast = 0.0f;
{{type}} fltSouthwest = 0.0f;
{{type}} fltSoutheast = 0.0f;
if (intC == 0) {
fltNorthwest = (({{type}}) (-1.0f)) * (({{type}}) (intSoutheastY) - fltY);
fltNortheast = (({{type}}) (+1.0f)) * (({{type}}) (intSouthwestY) - fltY);
fltSouthwest = (({{type}}) (-1.0f)) * (fltY - ({{type}}) (intNortheastY));
fltSoutheast = (({{type}}) (+1.0f)) * (fltY - ({{type}}) (intNorthwestY));
} else if (intC == 1) {
fltNorthwest = (({{type}}) (intSoutheastX) - fltX) * (({{type}}) (-1.0f));
fltNortheast = (fltX - ({{type}}) (intSouthwestX)) * (({{type}}) (-1.0f));
fltSouthwest = (({{type}}) (intNortheastX) - fltX) * (({{type}}) (+1.0f));
fltSoutheast = (fltX - ({{type}}) (intNorthwestX)) * (({{type}}) (+1.0f));
}
for (int intChannel = 0; intChannel < SIZE_1(tenOutgrad); intChannel += 1) {
{{type}} fltIn = VALUE_4(tenIn, intN, intChannel, intY, intX);
if ((intNorthwestX >= 0) && (intNorthwestX < SIZE_3(tenOutgrad)) && (intNorthwestY >= 0) && (intNorthwestY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intNorthwestY, intNorthwestX) * fltIn * fltNorthwest;
}
if ((intNortheastX >= 0) && (intNortheastX < SIZE_3(tenOutgrad)) && (intNortheastY >= 0) && (intNortheastY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intNortheastY, intNortheastX) * fltIn * fltNortheast;
}
if ((intSouthwestX >= 0) && (intSouthwestX < SIZE_3(tenOutgrad)) && (intSouthwestY >= 0) && (intSouthwestY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intSouthwestY, intSouthwestX) * fltIn * fltSouthwest;
}
if ((intSoutheastX >= 0) && (intSoutheastX < SIZE_3(tenOutgrad)) && (intSoutheastY >= 0) && (intSoutheastY < SIZE_2(tenOutgrad))) {
fltFlowgrad += VALUE_4(tenOutgrad, intN, intChannel, intSoutheastY, intSoutheastX) * fltIn * fltSoutheast;
}
}
tenFlowgrad[intIndex] = fltFlowgrad;
} }
""",
{
"tenIn": tenIn,
"tenFlow": tenFlow,
"tenOutgrad": tenOutgrad,
"tenIngrad": tenIngrad,
"tenFlowgrad": tenFlowgrad,
},
)
)(
grid=tuple([int((tenFlowgrad.nelement() + 512 - 1) / 512), 1, 1]),
block=tuple([512, 1, 1]),
args=[
cuda_int32(tenFlowgrad.nelement()),
tenIn.data_ptr(),
tenFlow.data_ptr(),
tenOutgrad.data_ptr(),
None,
tenFlowgrad.data_ptr(),
],
stream=collections.namedtuple("Stream", "ptr")(
torch.cuda.current_stream().cuda_stream
),
)
# end
return tenIngrad, tenFlowgrad
# end
# end
@@ -0,0 +1,76 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# ginr-ipc: https://github.com/kakaobrain/ginr-ipc
# --------------------------------------------------------
import math
import torch
import torch.nn as nn
from .layers import Sine, Damping
def convert_int_to_list(size, len_list=2):
if isinstance(size, int):
return [size] * len_list
else:
assert len(size) == len_list
return size
def initialize_params(params, init_type, **kwargs):
fan_in, fan_out = params.shape[0], params.shape[1]
if init_type is None or init_type == "normal":
nn.init.normal_(params)
elif init_type == "kaiming_uniform":
nn.init.kaiming_uniform_(params, a=math.sqrt(5))
elif init_type == "uniform_fan_in":
bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0
nn.init.uniform_(params, -bound, bound)
elif init_type == "zero":
nn.init.zeros_(params)
elif "siren" == init_type:
assert "siren_w0" in kwargs.keys() and "is_first" in kwargs.keys()
w0 = kwargs["siren_w0"]
if kwargs["is_first"]:
w_std = 1 / fan_in
else:
w_std = math.sqrt(6.0 / fan_in) / w0
nn.init.uniform_(params, -w_std, w_std)
else:
raise NotImplementedError
def create_params_with_init(
shape, init_type="normal", include_bias=False, bias_init_type="zero", **kwargs
):
if not include_bias:
params = torch.empty([shape[0], shape[1]])
initialize_params(params, init_type, **kwargs)
return params
else:
params = torch.empty([shape[0] - 1, shape[1]])
bias = torch.empty([1, shape[1]])
initialize_params(params, init_type, **kwargs)
initialize_params(bias, bias_init_type, **kwargs)
return torch.cat([params, bias], dim=0)
def create_activation(config):
if config.type == "relu":
activation = nn.ReLU()
elif config.type == "siren":
activation = Sine(config.siren_w0)
elif config.type == "silu":
activation = nn.SiLU()
elif config.type == "damp":
activation = Damping(config.siren_w0)
else:
raise NotImplementedError
return activation
@@ -0,0 +1,24 @@
from .raft import RAFT
import argparse
import torch
from .extractor import BasicEncoder
def initialize_RAFT(model_path="pretrained_ckpt/raft-things.pth", device="cuda"):
"""Initializes the RAFT model."""
args = argparse.ArgumentParser()
args.raft_model = model_path
args.small = False
args.mixed_precision = False
args.alternate_corr = False
model = RAFT(args)
ckpt = torch.load(args.raft_model, map_location="cpu")
def convert(param):
return {k.replace("module.", ""): v for k, v in param.items() if "module" in k}
ckpt = convert(ckpt)
model.load_state_dict(ckpt, strict=True)
print("load raft from " + model_path)
return model
+175
View File
@@ -0,0 +1,175 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
# All rights reserved.
# This source code is licensed under the license found in the
# LICENSE file in the root directory of this source tree.
# --------------------------------------------------------
# References:
# amt: https://github.com/MCG-NKU/AMT
# raft: https://github.com/princeton-vl/RAFT
# --------------------------------------------------------
import torch
import torch.nn.functional as F
from .utils.utils import bilinear_sampler, coords_grid
try:
import alt_cuda_corr
except:
# alt_cuda_corr is not compiled
pass
class BidirCorrBlock:
def __init__(self, fmap1, fmap2, num_levels=4, radius=4):
self.num_levels = num_levels
self.radius = radius
self.corr_pyramid = []
self.corr_pyramid_T = []
corr = BidirCorrBlock.corr(fmap1, fmap2)
batch, h1, w1, dim, h2, w2 = corr.shape
corr_T = corr.clone().permute(0, 4, 5, 3, 1, 2)
corr = corr.reshape(batch * h1 * w1, dim, h2, w2)
corr_T = corr_T.reshape(batch * h2 * w2, dim, h1, w1)
self.corr_pyramid.append(corr)
self.corr_pyramid_T.append(corr_T)
for _ in range(self.num_levels - 1):
corr = F.avg_pool2d(corr, 2, stride=2)
corr_T = F.avg_pool2d(corr_T, 2, stride=2)
self.corr_pyramid.append(corr)
self.corr_pyramid_T.append(corr_T)
def __call__(self, coords0, coords1):
r = self.radius
coords0 = coords0.permute(0, 2, 3, 1)
coords1 = coords1.permute(0, 2, 3, 1)
assert (
coords0.shape == coords1.shape
), f"coords0 shape: [{coords0.shape}] is not equal to [{coords1.shape}]"
batch, h1, w1, _ = coords0.shape
out_pyramid = []
out_pyramid_T = []
for i in range(self.num_levels):
corr = self.corr_pyramid[i]
corr_T = self.corr_pyramid_T[i]
dx = torch.linspace(-r, r, 2 * r + 1, device=coords0.device)
dy = torch.linspace(-r, r, 2 * r + 1, device=coords0.device)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1)
delta_lvl = delta.view(1, 2 * r + 1, 2 * r + 1, 2)
centroid_lvl_0 = coords0.reshape(batch * h1 * w1, 1, 1, 2) / 2**i
centroid_lvl_1 = coords1.reshape(batch * h1 * w1, 1, 1, 2) / 2**i
coords_lvl_0 = centroid_lvl_0 + delta_lvl
coords_lvl_1 = centroid_lvl_1 + delta_lvl
corr = bilinear_sampler(corr, coords_lvl_0)
corr_T = bilinear_sampler(corr_T, coords_lvl_1)
corr = corr.view(batch, h1, w1, -1)
corr_T = corr_T.view(batch, h1, w1, -1)
out_pyramid.append(corr)
out_pyramid_T.append(corr_T)
out = torch.cat(out_pyramid, dim=-1)
out_T = torch.cat(out_pyramid_T, dim=-1)
return (
out.permute(0, 3, 1, 2).contiguous().float(),
out_T.permute(0, 3, 1, 2).contiguous().float(),
)
@staticmethod
def corr(fmap1, fmap2):
batch, dim, ht, wd = fmap1.shape
fmap1 = fmap1.view(batch, dim, ht * wd)
fmap2 = fmap2.view(batch, dim, ht * wd)
corr = torch.matmul(fmap1.transpose(1, 2), fmap2)
corr = corr.view(batch, ht, wd, 1, ht, wd)
return corr / torch.sqrt(torch.tensor(dim).float())
class AlternateCorrBlock:
def __init__(self, fmap1, fmap2, num_levels=4, radius=4):
self.num_levels = num_levels
self.radius = radius
self.pyramid = [(fmap1, fmap2)]
for i in range(self.num_levels):
fmap1 = F.avg_pool2d(fmap1, 2, stride=2)
fmap2 = F.avg_pool2d(fmap2, 2, stride=2)
self.pyramid.append((fmap1, fmap2))
def __call__(self, coords):
coords = coords.permute(0, 2, 3, 1)
B, H, W, _ = coords.shape
dim = self.pyramid[0][0].shape[1]
corr_list = []
for i in range(self.num_levels):
r = self.radius
fmap1_i = self.pyramid[0][0].permute(0, 2, 3, 1).contiguous()
fmap2_i = self.pyramid[i][1].permute(0, 2, 3, 1).contiguous()
coords_i = (coords / 2**i).reshape(B, 1, H, W, 2).contiguous()
(corr,) = alt_cuda_corr.forward(fmap1_i, fmap2_i, coords_i, r)
corr_list.append(corr.squeeze(1))
corr = torch.stack(corr_list, dim=1)
corr = corr.reshape(B, -1, H, W)
return corr / torch.sqrt(torch.tensor(dim).float())
class CorrBlock:
def __init__(self, fmap1, fmap2, num_levels=4, radius=4):
self.num_levels = num_levels
self.radius = radius
self.corr_pyramid = []
# all pairs correlation
corr = CorrBlock.corr(fmap1, fmap2)
batch, h1, w1, dim, h2, w2 = corr.shape
corr = corr.reshape(batch * h1 * w1, dim, h2, w2)
self.corr_pyramid.append(corr)
for i in range(self.num_levels - 1):
corr = F.avg_pool2d(corr, 2, stride=2)
self.corr_pyramid.append(corr)
def __call__(self, coords):
r = self.radius
coords = coords.permute(0, 2, 3, 1)
batch, h1, w1, _ = coords.shape
out_pyramid = []
for i in range(self.num_levels):
corr = self.corr_pyramid[i]
dx = torch.linspace(-r, r, 2 * r + 1, device=coords.device)
dy = torch.linspace(-r, r, 2 * r + 1, device=coords.device)
delta = torch.stack(torch.meshgrid(dy, dx), axis=-1)
centroid_lvl = coords.reshape(batch * h1 * w1, 1, 1, 2) / 2**i
delta_lvl = delta.view(1, 2 * r + 1, 2 * r + 1, 2)
coords_lvl = centroid_lvl + delta_lvl
corr = bilinear_sampler(corr, coords_lvl)
corr = corr.view(batch, h1, w1, -1)
out_pyramid.append(corr)
out = torch.cat(out_pyramid, dim=-1)
return out.permute(0, 3, 1, 2).contiguous().float()
@staticmethod
def corr(fmap1, fmap2):
batch, dim, ht, wd = fmap1.shape
fmap1 = fmap1.view(batch, dim, ht * wd)
fmap2 = fmap2.view(batch, dim, ht * wd)
corr = torch.matmul(fmap1.transpose(1, 2), fmap2)
corr = corr.view(batch, ht, wd, 1, ht, wd)
return corr / torch.sqrt(torch.tensor(dim).float())
+293
View File
@@ -0,0 +1,293 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class ResidualBlock(nn.Module):
def __init__(self, in_planes, planes, norm_fn="group", stride=1):
super(ResidualBlock, self).__init__()
self.conv1 = nn.Conv2d(
in_planes, planes, kernel_size=3, padding=1, stride=stride
)
self.conv2 = nn.Conv2d(planes, planes, kernel_size=3, padding=1)
self.relu = nn.ReLU(inplace=True)
num_groups = planes // 8
if norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
if not stride == 1:
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
elif norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(planes)
self.norm2 = nn.BatchNorm2d(planes)
if not stride == 1:
self.norm3 = nn.BatchNorm2d(planes)
elif norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(planes)
self.norm2 = nn.InstanceNorm2d(planes)
if not stride == 1:
self.norm3 = nn.InstanceNorm2d(planes)
elif norm_fn == "none":
self.norm1 = nn.Sequential()
self.norm2 = nn.Sequential()
if not stride == 1:
self.norm3 = nn.Sequential()
if stride == 1:
self.downsample = None
else:
self.downsample = nn.Sequential(
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm3
)
def forward(self, x):
y = x
y = self.relu(self.norm1(self.conv1(y)))
y = self.relu(self.norm2(self.conv2(y)))
if self.downsample is not None:
x = self.downsample(x)
return self.relu(x + y)
class BottleneckBlock(nn.Module):
def __init__(self, in_planes, planes, norm_fn="group", stride=1):
super(BottleneckBlock, self).__init__()
self.conv1 = nn.Conv2d(in_planes, planes // 4, kernel_size=1, padding=0)
self.conv2 = nn.Conv2d(
planes // 4, planes // 4, kernel_size=3, padding=1, stride=stride
)
self.conv3 = nn.Conv2d(planes // 4, planes, kernel_size=1, padding=0)
self.relu = nn.ReLU(inplace=True)
num_groups = planes // 8
if norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // 4)
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes // 4)
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
if not stride == 1:
self.norm4 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
elif norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(planes // 4)
self.norm2 = nn.BatchNorm2d(planes // 4)
self.norm3 = nn.BatchNorm2d(planes)
if not stride == 1:
self.norm4 = nn.BatchNorm2d(planes)
elif norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(planes // 4)
self.norm2 = nn.InstanceNorm2d(planes // 4)
self.norm3 = nn.InstanceNorm2d(planes)
if not stride == 1:
self.norm4 = nn.InstanceNorm2d(planes)
elif norm_fn == "none":
self.norm1 = nn.Sequential()
self.norm2 = nn.Sequential()
self.norm3 = nn.Sequential()
if not stride == 1:
self.norm4 = nn.Sequential()
if stride == 1:
self.downsample = None
else:
self.downsample = nn.Sequential(
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm4
)
def forward(self, x):
y = x
y = self.relu(self.norm1(self.conv1(y)))
y = self.relu(self.norm2(self.conv2(y)))
y = self.relu(self.norm3(self.conv3(y)))
if self.downsample is not None:
x = self.downsample(x)
return self.relu(x + y)
class BasicEncoder(nn.Module):
def __init__(self, output_dim=128, norm_fn="batch", dropout=0.0, only_feat=False):
super(BasicEncoder, self).__init__()
self.norm_fn = norm_fn
self.only_feat = only_feat
if self.norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=64)
elif self.norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(64)
elif self.norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(64)
elif self.norm_fn == "none":
self.norm1 = nn.Sequential()
self.conv1 = nn.Conv2d(3, 64, kernel_size=7, stride=2, padding=3)
self.relu1 = nn.ReLU(inplace=True)
self.in_planes = 64
self.layer1 = self._make_layer(64, stride=1)
self.layer2 = self._make_layer(96, stride=2)
self.layer3 = self._make_layer(128, stride=2)
if not self.only_feat:
# output convolution
self.conv2 = nn.Conv2d(128, output_dim, kernel_size=1)
self.dropout = None
if dropout > 0:
self.dropout = nn.Dropout2d(p=dropout)
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
if m.weight is not None:
nn.init.constant_(m.weight, 1)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def _make_layer(self, dim, stride=1):
layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride)
layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1)
layers = (layer1, layer2)
self.in_planes = dim
return nn.Sequential(*layers)
def forward(self, x, return_feature=False, mif=False):
features = []
# if input is list, combine batch dimension
is_list = isinstance(x, tuple) or isinstance(x, list)
if is_list:
batch_dim = x[0].shape[0]
x = torch.cat(x, dim=0)
x_2 = F.interpolate(x, scale_factor=1 / 2, mode="bilinear", align_corners=False)
x_4 = F.interpolate(x, scale_factor=1 / 4, mode="bilinear", align_corners=False)
def f1(feat):
feat = self.conv1(feat)
feat = self.norm1(feat)
feat = self.relu1(feat)
feat = self.layer1(feat)
return feat
x = f1(x)
features.append(x)
x = self.layer2(x)
if mif:
x_2_2 = f1(x_2)
features.append(torch.cat([x, x_2_2], dim=1))
else:
features.append(x)
x = self.layer3(x)
if mif:
x_2_4 = self.layer2(x_2_2)
x_4_4 = f1(x_4)
features.append(torch.cat([x, x_2_4, x_4_4], dim=1))
else:
features.append(x)
if not self.only_feat:
x = self.conv2(x)
if self.training and self.dropout is not None:
x = self.dropout(x)
if is_list:
x = torch.split(x, [batch_dim, batch_dim], dim=0)
features = [torch.split(f, [batch_dim, batch_dim], dim=0) for f in features]
if return_feature:
return x, features
else:
return x
class SmallEncoder(nn.Module):
def __init__(self, output_dim=128, norm_fn="batch", dropout=0.0):
super(SmallEncoder, self).__init__()
self.norm_fn = norm_fn
if self.norm_fn == "group":
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=32)
elif self.norm_fn == "batch":
self.norm1 = nn.BatchNorm2d(32)
elif self.norm_fn == "instance":
self.norm1 = nn.InstanceNorm2d(32)
elif self.norm_fn == "none":
self.norm1 = nn.Sequential()
self.conv1 = nn.Conv2d(3, 32, kernel_size=7, stride=2, padding=3)
self.relu1 = nn.ReLU(inplace=True)
self.in_planes = 32
self.layer1 = self._make_layer(32, stride=1)
self.layer2 = self._make_layer(64, stride=2)
self.layer3 = self._make_layer(96, stride=2)
self.dropout = None
if dropout > 0:
self.dropout = nn.Dropout2d(p=dropout)
self.conv2 = nn.Conv2d(96, output_dim, kernel_size=1)
for m in self.modules():
if isinstance(m, nn.Conv2d):
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
if m.weight is not None:
nn.init.constant_(m.weight, 1)
if m.bias is not None:
nn.init.constant_(m.bias, 0)
def _make_layer(self, dim, stride=1):
layer1 = BottleneckBlock(self.in_planes, dim, self.norm_fn, stride=stride)
layer2 = BottleneckBlock(dim, dim, self.norm_fn, stride=1)
layers = (layer1, layer2)
self.in_planes = dim
return nn.Sequential(*layers)
def forward(self, x):
# if input is list, combine batch dimension
is_list = isinstance(x, tuple) or isinstance(x, list)
if is_list:
batch_dim = x[0].shape[0]
x = torch.cat(x, dim=0)
x = self.conv1(x)
x = self.norm1(x)
x = self.relu1(x)
x = self.layer1(x)
x = self.layer2(x)
x = self.layer3(x)
x = self.conv2(x)
if self.training and self.dropout is not None:
x = self.dropout(x)
if is_list:
x = torch.split(x, [batch_dim, batch_dim], dim=0)
return x
@@ -0,0 +1,238 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from .update import BasicUpdateBlock, SmallUpdateBlock
from .extractor import BasicEncoder, SmallEncoder
from .corr import BidirCorrBlock, AlternateCorrBlock
from .utils.utils import bilinear_sampler, coords_grid, upflow8
try:
autocast = torch.cuda.amp.autocast
except:
# dummy autocast for PyTorch < 1.6
class autocast:
def __init__(self, enabled):
pass
def __enter__(self):
pass
def __exit__(self, *args):
pass
# BiRAFT
class RAFT(nn.Module):
def __init__(self, args):
super(RAFT, self).__init__()
self.args = args
if args.small:
self.hidden_dim = hdim = 96
self.context_dim = cdim = 64
args.corr_levels = 4
args.corr_radius = 3
self.corr_levels = 4
self.corr_radius = 3
else:
self.hidden_dim = hdim = 128
self.context_dim = cdim = 128
args.corr_levels = 4
args.corr_radius = 4
self.corr_levels = 4
self.corr_radius = 4
if "dropout" not in args._get_kwargs():
self.args.dropout = 0
if "alternate_corr" not in args._get_kwargs():
self.args.alternate_corr = False
# feature network, context network, and update block
if args.small:
self.fnet = SmallEncoder(
output_dim=128, norm_fn="instance", dropout=args.dropout
)
self.cnet = SmallEncoder(
output_dim=hdim + cdim, norm_fn="none", dropout=args.dropout
)
self.update_block = SmallUpdateBlock(self.args, hidden_dim=hdim)
else:
self.fnet = BasicEncoder(
output_dim=256, norm_fn="instance", dropout=args.dropout
)
self.cnet = BasicEncoder(
output_dim=hdim + cdim, norm_fn="batch", dropout=args.dropout
)
self.update_block = BasicUpdateBlock(self.args, hidden_dim=hdim)
def freeze_bn(self):
for m in self.modules():
if isinstance(m, nn.BatchNorm2d):
m.eval()
def build_coord(self, img):
N, C, H, W = img.shape
coords = coords_grid(N, H // 8, W // 8, device=img.device)
return coords
def initialize_flow(self, img, img2):
"""Flow is represented as difference between two coordinate grids flow = coords1 - coords0"""
assert img.shape == img2.shape
N, C, H, W = img.shape
coords01 = coords_grid(N, H // 8, W // 8, device=img.device)
coords02 = coords_grid(N, H // 8, W // 8, device=img.device)
coords1 = coords_grid(N, H // 8, W // 8, device=img.device)
coords2 = coords_grid(N, H // 8, W // 8, device=img.device)
# optical flow computed as difference: flow = coords1 - coords0
return coords01, coords02, coords1, coords2
def upsample_flow(self, flow, mask):
"""Upsample flow field [H/8, W/8, 2] -> [H, W, 2] using convex combination"""
N, _, H, W = flow.shape
mask = mask.view(N, 1, 9, 8, 8, H, W)
mask = torch.softmax(mask, dim=2)
up_flow = F.unfold(8 * flow, [3, 3], padding=1)
up_flow = up_flow.view(N, 2, 9, 1, 1, H, W)
up_flow = torch.sum(mask * up_flow, dim=2)
up_flow = up_flow.permute(0, 1, 4, 2, 5, 3)
return up_flow.reshape(N, 2, 8 * H, 8 * W)
def get_corr_fn(self, image1, image2, projector=None):
# run the feature network
with autocast(enabled=self.args.mixed_precision):
fmaps, feats = self.fnet([image1, image2], return_feature=True)
fmap1, fmap2 = fmaps
fmap1 = fmap1.float()
fmap2 = fmap2.float()
corr_fn1 = None
if self.args.alternate_corr:
corr_fn = AlternateCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
if projector is not None:
corr_fn1 = AlternateCorrBlock(
projector(feats[-1][0]),
projector(feats[-1][1]),
radius=self.args.corr_radius,
)
else:
corr_fn = BidirCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
if projector is not None:
corr_fn1 = BidirCorrBlock(
projector(feats[-1][0]),
projector(feats[-1][1]),
radius=self.args.corr_radius,
)
if corr_fn1 is None:
return corr_fn, corr_fn
else:
return corr_fn, corr_fn1
def get_corr_fn_from_feat(self, fmap1, fmap2):
fmap1 = fmap1.float()
fmap2 = fmap2.float()
if self.args.alternate_corr:
corr_fn = AlternateCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
else:
corr_fn = BidirCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
return corr_fn
def forward(
self,
image1,
image2,
iters=12,
flow_init=None,
upsample=True,
test_mode=False,
corr_fn=None,
mif=False,
):
"""Estimate optical flow between pair of frames"""
assert flow_init is None
image1 = 2 * (image1 / 255.0) - 1.0
image2 = 2 * (image2 / 255.0) - 1.0
image1 = image1.contiguous()
image2 = image2.contiguous()
hdim = self.hidden_dim
cdim = self.context_dim
if corr_fn is None:
corr_fn, _ = self.get_corr_fn(image1, image2)
# # run the feature network
# with autocast(enabled=self.args.mixed_precision):
# fmap1, fmap2 = self.fnet([image1, image2])
# fmap1 = fmap1.float()
# fmap2 = fmap2.float()
# if self.args.alternate_corr:
# corr_fn = AlternateCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
# else:
# corr_fn = BidirCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
# run the context network
with autocast(enabled=self.args.mixed_precision):
# for image1
cnet1, features1 = self.cnet(image1, return_feature=True, mif=mif)
net1, inp1 = torch.split(cnet1, [hdim, cdim], dim=1)
net1 = torch.tanh(net1)
inp1 = torch.relu(inp1)
# for image2
cnet2, features2 = self.cnet(image2, return_feature=True, mif=mif)
net2, inp2 = torch.split(cnet2, [hdim, cdim], dim=1)
net2 = torch.tanh(net2)
inp2 = torch.relu(inp2)
coords01, coords02, coords1, coords2 = self.initialize_flow(image1, image2)
# if flow_init is not None:
# coords1 = coords1 + flow_init
# flow_predictions1 = []
# flow_predictions2 = []
for itr in range(iters):
coords1 = coords1.detach()
coords2 = coords2.detach()
corr1, corr2 = corr_fn(coords1, coords2) # index correlation volume
flow1 = coords1 - coords01
flow2 = coords2 - coords02
with autocast(enabled=self.args.mixed_precision):
net1, up_mask1, delta_flow1 = self.update_block(
net1, inp1, corr1, flow1
)
net2, up_mask2, delta_flow2 = self.update_block(
net2, inp2, corr2, flow2
)
# F(t+1) = F(t) + \Delta(t)
coords1 = coords1 + delta_flow1
coords2 = coords2 + delta_flow2
flow_low1 = coords1 - coords01
flow_low2 = coords2 - coords02
# upsample predictions
if up_mask1 is None:
flow_up1 = upflow8(coords1 - coords01)
flow_up2 = upflow8(coords2 - coords02)
else:
flow_up1 = self.upsample_flow(coords1 - coords01, up_mask1)
flow_up2 = self.upsample_flow(coords2 - coords02, up_mask2)
# flow_predictions.append(flow_up)
return flow_up1, flow_up2, flow_low1, flow_low2, features1, features2
# if test_mode:
# return coords1 - coords0, flow_up
# return flow_predictions
+169
View File
@@ -0,0 +1,169 @@
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from .update import BasicUpdateBlock, SmallUpdateBlock
from .extractor import BasicEncoder, SmallEncoder
from .corr import CorrBlock, AlternateCorrBlock
from .utils.utils import bilinear_sampler, coords_grid, upflow8
try:
autocast = torch.cuda.amp.autocast
except:
# dummy autocast for PyTorch < 1.6
class autocast:
def __init__(self, enabled):
pass
def __enter__(self):
pass
def __exit__(self, *args):
pass
class RAFT(nn.Module):
def __init__(self, args):
super(RAFT, self).__init__()
self.args = args
if args.small:
self.hidden_dim = hdim = 96
self.context_dim = cdim = 64
args.corr_levels = 4
args.corr_radius = 3
self.corr_levels = 4
self.corr_radius = 3
else:
self.hidden_dim = hdim = 128
self.context_dim = cdim = 128
args.corr_levels = 4
args.corr_radius = 4
self.corr_levels = 4
self.corr_radius = 4
if "dropout" not in args._get_kwargs():
self.args.dropout = 0
if "alternate_corr" not in args._get_kwargs():
self.args.alternate_corr = False
# feature network, context network, and update block
if args.small:
self.fnet = SmallEncoder(
output_dim=128, norm_fn="instance", dropout=args.dropout
)
self.cnet = SmallEncoder(
output_dim=hdim + cdim, norm_fn="none", dropout=args.dropout
)
self.update_block = SmallUpdateBlock(self.args, hidden_dim=hdim)
else:
self.fnet = BasicEncoder(
output_dim=256, norm_fn="instance", dropout=args.dropout
)
self.cnet = BasicEncoder(
output_dim=hdim + cdim, norm_fn="batch", dropout=args.dropout
)
self.update_block = BasicUpdateBlock(self.args, hidden_dim=hdim)
def freeze_bn(self):
for m in self.modules():
if isinstance(m, nn.BatchNorm2d):
m.eval()
def initialize_flow(self, img):
"""Flow is represented as difference between two coordinate grids flow = coords1 - coords0"""
N, C, H, W = img.shape
coords0 = coords_grid(N, H // 8, W // 8, device=img.device)
coords1 = coords_grid(N, H // 8, W // 8, device=img.device)
# optical flow computed as difference: flow = coords1 - coords0
return coords0, coords1
def upsample_flow(self, flow, mask):
"""Upsample flow field [H/8, W/8, 2] -> [H, W, 2] using convex combination"""
N, _, H, W = flow.shape
mask = mask.view(N, 1, 9, 8, 8, H, W)
mask = torch.softmax(mask, dim=2)
up_flow = F.unfold(8 * flow, [3, 3], padding=1)
up_flow = up_flow.view(N, 2, 9, 1, 1, H, W)
up_flow = torch.sum(mask * up_flow, dim=2)
up_flow = up_flow.permute(0, 1, 4, 2, 5, 3)
return up_flow.reshape(N, 2, 8 * H, 8 * W)
def forward(
self,
image1,
image2,
iters=12,
flow_init=None,
upsample=True,
test_mode=False,
return_feat=True,
):
"""Estimate optical flow between pair of frames"""
image1 = 2 * (image1 / 255.0) - 1.0
image2 = 2 * (image2 / 255.0) - 1.0
image1 = image1.contiguous()
image2 = image2.contiguous()
hdim = self.hidden_dim
cdim = self.context_dim
# run the feature network
with autocast(enabled=self.args.mixed_precision):
fmap1, fmap2 = self.fnet([image1, image2])
fmap1 = fmap1.float()
fmap2 = fmap2.float()
if self.args.alternate_corr:
corr_fn = AlternateCorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
else:
corr_fn = CorrBlock(fmap1, fmap2, radius=self.args.corr_radius)
# run the context network
with autocast(enabled=self.args.mixed_precision):
cnet, feats = self.cnet(image1, return_feature=True)
net, inp = torch.split(cnet, [hdim, cdim], dim=1)
net = torch.tanh(net)
inp = torch.relu(inp)
coords0, coords1 = self.initialize_flow(image1)
if flow_init is not None:
coords1 = coords1 + flow_init
flow_predictions = []
for itr in range(iters):
coords1 = coords1.detach()
corr = corr_fn(coords1) # index correlation volume
flow = coords1 - coords0
with autocast(enabled=self.args.mixed_precision):
net, up_mask, delta_flow = self.update_block(net, inp, corr, flow)
# F(t+1) = F(t) + \Delta(t)
coords1 = coords1 + delta_flow
# upsample predictions
if up_mask is None:
flow_up = upflow8(coords1 - coords0)
else:
flow_up = self.upsample_flow(coords1 - coords0, up_mask)
flow_predictions.append(flow_up)
if test_mode:
return coords1 - coords0, flow_up
if return_feat:
return flow_up, feats[1:], fmap1
return flow_predictions
+154
View File
@@ -0,0 +1,154 @@
import torch
import torch.nn as nn
import torch.nn.functional as F
class FlowHead(nn.Module):
def __init__(self, input_dim=128, hidden_dim=256):
super(FlowHead, self).__init__()
self.conv1 = nn.Conv2d(input_dim, hidden_dim, 3, padding=1)
self.conv2 = nn.Conv2d(hidden_dim, 2, 3, padding=1)
self.relu = nn.ReLU(inplace=True)
def forward(self, x):
return self.conv2(self.relu(self.conv1(x)))
class ConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(ConvGRU, self).__init__()
self.convz = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convr = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
self.convq = nn.Conv2d(hidden_dim + input_dim, hidden_dim, 3, padding=1)
def forward(self, h, x):
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz(hx))
r = torch.sigmoid(self.convr(hx))
q = torch.tanh(self.convq(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class SepConvGRU(nn.Module):
def __init__(self, hidden_dim=128, input_dim=192 + 128):
super(SepConvGRU, self).__init__()
self.convz1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convr1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convq1 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (1, 5), padding=(0, 2)
)
self.convz2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convr2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
self.convq2 = nn.Conv2d(
hidden_dim + input_dim, hidden_dim, (5, 1), padding=(2, 0)
)
def forward(self, h, x):
# horizontal
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz1(hx))
r = torch.sigmoid(self.convr1(hx))
q = torch.tanh(self.convq1(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
# vertical
hx = torch.cat([h, x], dim=1)
z = torch.sigmoid(self.convz2(hx))
r = torch.sigmoid(self.convr2(hx))
q = torch.tanh(self.convq2(torch.cat([r * h, x], dim=1)))
h = (1 - z) * h + z * q
return h
class SmallMotionEncoder(nn.Module):
def __init__(self, args):
super(SmallMotionEncoder, self).__init__()
cor_planes = args.corr_levels * (2 * args.corr_radius + 1) ** 2
self.convc1 = nn.Conv2d(cor_planes, 96, 1, padding=0)
self.convf1 = nn.Conv2d(2, 64, 7, padding=3)
self.convf2 = nn.Conv2d(64, 32, 3, padding=1)
self.conv = nn.Conv2d(128, 80, 3, padding=1)
def forward(self, flow, corr):
cor = F.relu(self.convc1(corr))
flo = F.relu(self.convf1(flow))
flo = F.relu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
out = F.relu(self.conv(cor_flo))
return torch.cat([out, flow], dim=1)
class BasicMotionEncoder(nn.Module):
def __init__(self, args):
super(BasicMotionEncoder, self).__init__()
cor_planes = args.corr_levels * (2 * args.corr_radius + 1) ** 2
self.convc1 = nn.Conv2d(cor_planes, 256, 1, padding=0)
self.convc2 = nn.Conv2d(256, 192, 3, padding=1)
self.convf1 = nn.Conv2d(2, 128, 7, padding=3)
self.convf2 = nn.Conv2d(128, 64, 3, padding=1)
self.conv = nn.Conv2d(64 + 192, 128 - 2, 3, padding=1)
def forward(self, flow, corr):
cor = F.relu(self.convc1(corr))
cor = F.relu(self.convc2(cor))
flo = F.relu(self.convf1(flow))
flo = F.relu(self.convf2(flo))
cor_flo = torch.cat([cor, flo], dim=1)
out = F.relu(self.conv(cor_flo))
return torch.cat([out, flow], dim=1)
class SmallUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=96):
super(SmallUpdateBlock, self).__init__()
self.encoder = SmallMotionEncoder(args)
self.gru = ConvGRU(hidden_dim=hidden_dim, input_dim=82 + 64)
self.flow_head = FlowHead(hidden_dim, hidden_dim=128)
def forward(self, net, inp, corr, flow):
motion_features = self.encoder(flow, corr)
inp = torch.cat([inp, motion_features], dim=1)
net = self.gru(net, inp)
delta_flow = self.flow_head(net)
return net, None, delta_flow
class BasicUpdateBlock(nn.Module):
def __init__(self, args, hidden_dim=128, input_dim=128):
super(BasicUpdateBlock, self).__init__()
self.args = args
self.encoder = BasicMotionEncoder(args)
self.gru = SepConvGRU(hidden_dim=hidden_dim, input_dim=128 + hidden_dim)
self.flow_head = FlowHead(hidden_dim, hidden_dim=256)
self.mask = nn.Sequential(
nn.Conv2d(128, 256, 3, padding=1),
nn.ReLU(inplace=True),
nn.Conv2d(256, 64 * 9, 1, padding=0),
)
def forward(self, net, inp, corr, flow, upsample=True):
motion_features = self.encoder(flow, corr)
inp = torch.cat([inp, motion_features], dim=1)
net = self.gru(net, inp)
delta_flow = self.flow_head(net)
# scale mask to balence gradients
mask = 0.25 * self.mask(net)
return net, mask, delta_flow
@@ -0,0 +1,266 @@
import numpy as np
import random
import math
from PIL import Image
import cv2
cv2.setNumThreads(0)
cv2.ocl.setUseOpenCL(False)
import torch
from torchvision.transforms import ColorJitter
import torch.nn.functional as F
class FlowAugmentor:
def __init__(self, crop_size, min_scale=-0.2, max_scale=0.5, do_flip=True):
# spatial augmentation params
self.crop_size = crop_size
self.min_scale = min_scale
self.max_scale = max_scale
self.spatial_aug_prob = 0.8
self.stretch_prob = 0.8
self.max_stretch = 0.2
# flip augmentation params
self.do_flip = do_flip
self.h_flip_prob = 0.5
self.v_flip_prob = 0.1
# photometric augmentation params
self.photo_aug = ColorJitter(
brightness=0.4, contrast=0.4, saturation=0.4, hue=0.5 / 3.14
)
self.asymmetric_color_aug_prob = 0.2
self.eraser_aug_prob = 0.5
def color_transform(self, img1, img2):
"""Photometric augmentation"""
# asymmetric
if np.random.rand() < self.asymmetric_color_aug_prob:
img1 = np.array(self.photo_aug(Image.fromarray(img1)), dtype=np.uint8)
img2 = np.array(self.photo_aug(Image.fromarray(img2)), dtype=np.uint8)
# symmetric
else:
image_stack = np.concatenate([img1, img2], axis=0)
image_stack = np.array(
self.photo_aug(Image.fromarray(image_stack)), dtype=np.uint8
)
img1, img2 = np.split(image_stack, 2, axis=0)
return img1, img2
def eraser_transform(self, img1, img2, bounds=[50, 100]):
"""Occlusion augmentation"""
ht, wd = img1.shape[:2]
if np.random.rand() < self.eraser_aug_prob:
mean_color = np.mean(img2.reshape(-1, 3), axis=0)
for _ in range(np.random.randint(1, 3)):
x0 = np.random.randint(0, wd)
y0 = np.random.randint(0, ht)
dx = np.random.randint(bounds[0], bounds[1])
dy = np.random.randint(bounds[0], bounds[1])
img2[y0 : y0 + dy, x0 : x0 + dx, :] = mean_color
return img1, img2
def spatial_transform(self, img1, img2, flow):
# randomly sample scale
ht, wd = img1.shape[:2]
min_scale = np.maximum(
(self.crop_size[0] + 8) / float(ht), (self.crop_size[1] + 8) / float(wd)
)
scale = 2 ** np.random.uniform(self.min_scale, self.max_scale)
scale_x = scale
scale_y = scale
if np.random.rand() < self.stretch_prob:
scale_x *= 2 ** np.random.uniform(-self.max_stretch, self.max_stretch)
scale_y *= 2 ** np.random.uniform(-self.max_stretch, self.max_stretch)
scale_x = np.clip(scale_x, min_scale, None)
scale_y = np.clip(scale_y, min_scale, None)
if np.random.rand() < self.spatial_aug_prob:
# rescale the images
img1 = cv2.resize(
img1, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
img2 = cv2.resize(
img2, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
flow = cv2.resize(
flow, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
flow = flow * [scale_x, scale_y]
if self.do_flip:
if np.random.rand() < self.h_flip_prob: # h-flip
img1 = img1[:, ::-1]
img2 = img2[:, ::-1]
flow = flow[:, ::-1] * [-1.0, 1.0]
if np.random.rand() < self.v_flip_prob: # v-flip
img1 = img1[::-1, :]
img2 = img2[::-1, :]
flow = flow[::-1, :] * [1.0, -1.0]
y0 = np.random.randint(0, img1.shape[0] - self.crop_size[0])
x0 = np.random.randint(0, img1.shape[1] - self.crop_size[1])
img1 = img1[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
img2 = img2[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
flow = flow[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
return img1, img2, flow
def __call__(self, img1, img2, flow):
img1, img2 = self.color_transform(img1, img2)
img1, img2 = self.eraser_transform(img1, img2)
img1, img2, flow = self.spatial_transform(img1, img2, flow)
img1 = np.ascontiguousarray(img1)
img2 = np.ascontiguousarray(img2)
flow = np.ascontiguousarray(flow)
return img1, img2, flow
class SparseFlowAugmentor:
def __init__(self, crop_size, min_scale=-0.2, max_scale=0.5, do_flip=False):
# spatial augmentation params
self.crop_size = crop_size
self.min_scale = min_scale
self.max_scale = max_scale
self.spatial_aug_prob = 0.8
self.stretch_prob = 0.8
self.max_stretch = 0.2
# flip augmentation params
self.do_flip = do_flip
self.h_flip_prob = 0.5
self.v_flip_prob = 0.1
# photometric augmentation params
self.photo_aug = ColorJitter(
brightness=0.3, contrast=0.3, saturation=0.3, hue=0.3 / 3.14
)
self.asymmetric_color_aug_prob = 0.2
self.eraser_aug_prob = 0.5
def color_transform(self, img1, img2):
image_stack = np.concatenate([img1, img2], axis=0)
image_stack = np.array(
self.photo_aug(Image.fromarray(image_stack)), dtype=np.uint8
)
img1, img2 = np.split(image_stack, 2, axis=0)
return img1, img2
def eraser_transform(self, img1, img2):
ht, wd = img1.shape[:2]
if np.random.rand() < self.eraser_aug_prob:
mean_color = np.mean(img2.reshape(-1, 3), axis=0)
for _ in range(np.random.randint(1, 3)):
x0 = np.random.randint(0, wd)
y0 = np.random.randint(0, ht)
dx = np.random.randint(50, 100)
dy = np.random.randint(50, 100)
img2[y0 : y0 + dy, x0 : x0 + dx, :] = mean_color
return img1, img2
def resize_sparse_flow_map(self, flow, valid, fx=1.0, fy=1.0):
ht, wd = flow.shape[:2]
coords = np.meshgrid(np.arange(wd), np.arange(ht))
coords = np.stack(coords, axis=-1)
coords = coords.reshape(-1, 2).astype(np.float32)
flow = flow.reshape(-1, 2).astype(np.float32)
valid = valid.reshape(-1).astype(np.float32)
coords0 = coords[valid >= 1]
flow0 = flow[valid >= 1]
ht1 = int(round(ht * fy))
wd1 = int(round(wd * fx))
coords1 = coords0 * [fx, fy]
flow1 = flow0 * [fx, fy]
xx = np.round(coords1[:, 0]).astype(np.int32)
yy = np.round(coords1[:, 1]).astype(np.int32)
v = (xx > 0) & (xx < wd1) & (yy > 0) & (yy < ht1)
xx = xx[v]
yy = yy[v]
flow1 = flow1[v]
flow_img = np.zeros([ht1, wd1, 2], dtype=np.float32)
valid_img = np.zeros([ht1, wd1], dtype=np.int32)
flow_img[yy, xx] = flow1
valid_img[yy, xx] = 1
return flow_img, valid_img
def spatial_transform(self, img1, img2, flow, valid):
# randomly sample scale
ht, wd = img1.shape[:2]
min_scale = np.maximum(
(self.crop_size[0] + 1) / float(ht), (self.crop_size[1] + 1) / float(wd)
)
scale = 2 ** np.random.uniform(self.min_scale, self.max_scale)
scale_x = np.clip(scale, min_scale, None)
scale_y = np.clip(scale, min_scale, None)
if np.random.rand() < self.spatial_aug_prob:
# rescale the images
img1 = cv2.resize(
img1, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
img2 = cv2.resize(
img2, None, fx=scale_x, fy=scale_y, interpolation=cv2.INTER_LINEAR
)
flow, valid = self.resize_sparse_flow_map(
flow, valid, fx=scale_x, fy=scale_y
)
if self.do_flip:
if np.random.rand() < 0.5: # h-flip
img1 = img1[:, ::-1]
img2 = img2[:, ::-1]
flow = flow[:, ::-1] * [-1.0, 1.0]
valid = valid[:, ::-1]
margin_y = 20
margin_x = 50
y0 = np.random.randint(0, img1.shape[0] - self.crop_size[0] + margin_y)
x0 = np.random.randint(-margin_x, img1.shape[1] - self.crop_size[1] + margin_x)
y0 = np.clip(y0, 0, img1.shape[0] - self.crop_size[0])
x0 = np.clip(x0, 0, img1.shape[1] - self.crop_size[1])
img1 = img1[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
img2 = img2[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
flow = flow[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
valid = valid[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
return img1, img2, flow, valid
def __call__(self, img1, img2, flow, valid):
img1, img2 = self.color_transform(img1, img2)
img1, img2 = self.eraser_transform(img1, img2)
img1, img2, flow, valid = self.spatial_transform(img1, img2, flow, valid)
img1 = np.ascontiguousarray(img1)
img2 = np.ascontiguousarray(img2)
flow = np.ascontiguousarray(flow)
valid = np.ascontiguousarray(valid)
return img1, img2, flow, valid
@@ -0,0 +1,133 @@
# Flow visualization code used from https://github.com/tomrunia/OpticalFlow_Visualization
# MIT License
#
# Copyright (c) 2018 Tom Runia
#
# 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 conditions.
#
# Author: Tom Runia
# Date Created: 2018-08-03
import numpy as np
def make_colorwheel():
"""
Generates a color wheel for optical flow visualization as presented in:
Baker et al. "A Database and Evaluation Methodology for Optical Flow" (ICCV, 2007)
URL: http://vision.middlebury.edu/flow/flowEval-iccv07.pdf
Code follows the original C++ source code of Daniel Scharstein.
Code follows the the Matlab source code of Deqing Sun.
Returns:
np.ndarray: Color wheel
"""
RY = 15
YG = 6
GC = 4
CB = 11
BM = 13
MR = 6
ncols = RY + YG + GC + CB + BM + MR
colorwheel = np.zeros((ncols, 3))
col = 0
# RY
colorwheel[0:RY, 0] = 255
colorwheel[0:RY, 1] = np.floor(255 * np.arange(0, RY) / RY)
col = col + RY
# YG
colorwheel[col : col + YG, 0] = 255 - np.floor(255 * np.arange(0, YG) / YG)
colorwheel[col : col + YG, 1] = 255
col = col + YG
# GC
colorwheel[col : col + GC, 1] = 255
colorwheel[col : col + GC, 2] = np.floor(255 * np.arange(0, GC) / GC)
col = col + GC
# CB
colorwheel[col : col + CB, 1] = 255 - np.floor(255 * np.arange(CB) / CB)
colorwheel[col : col + CB, 2] = 255
col = col + CB
# BM
colorwheel[col : col + BM, 2] = 255
colorwheel[col : col + BM, 0] = np.floor(255 * np.arange(0, BM) / BM)
col = col + BM
# MR
colorwheel[col : col + MR, 2] = 255 - np.floor(255 * np.arange(MR) / MR)
colorwheel[col : col + MR, 0] = 255
return colorwheel
def flow_uv_to_colors(u, v, convert_to_bgr=False):
"""
Applies the flow color wheel to (possibly clipped) flow components u and v.
According to the C++ source code of Daniel Scharstein
According to the Matlab source code of Deqing Sun
Args:
u (np.ndarray): Input horizontal flow of shape [H,W]
v (np.ndarray): Input vertical flow of shape [H,W]
convert_to_bgr (bool, optional): Convert output image to BGR. Defaults to False.
Returns:
np.ndarray: Flow visualization image of shape [H,W,3]
"""
flow_image = np.zeros((u.shape[0], u.shape[1], 3), np.uint8)
colorwheel = make_colorwheel() # shape [55x3]
ncols = colorwheel.shape[0]
rad = np.sqrt(np.square(u) + np.square(v))
a = np.arctan2(-v, -u) / np.pi
fk = (a + 1) / 2 * (ncols - 1)
k0 = np.floor(fk).astype(np.int32)
k1 = k0 + 1
k1[k1 == ncols] = 0
f = fk - k0
for i in range(colorwheel.shape[1]):
tmp = colorwheel[:, i]
col0 = tmp[k0] / 255.0
col1 = tmp[k1] / 255.0
col = (1 - f) * col0 + f * col1
idx = rad <= 1
col[idx] = 1 - rad[idx] * (1 - col[idx])
col[~idx] = col[~idx] * 0.75 # out of range
# Note the 2-i => BGR instead of RGB
ch_idx = 2 - i if convert_to_bgr else i
flow_image[:, :, ch_idx] = np.floor(255 * col)
return flow_image
def flow_to_image(flow_uv, clip_flow=None, convert_to_bgr=False):
"""
Expects a two dimensional flow image of shape.
Args:
flow_uv (np.ndarray): Flow UV image of shape [H,W,2]
clip_flow (float, optional): Clip maximum of flow values. Defaults to None.
convert_to_bgr (bool, optional): Convert output image to BGR. Defaults to False.
Returns:
np.ndarray: Flow visualization image of shape [H,W,3]
"""
assert flow_uv.ndim == 3, "input flow must have three dimensions"
assert flow_uv.shape[2] == 2, "input flow must have shape [H,W,2]"
if clip_flow is not None:
flow_uv = np.clip(flow_uv, 0, clip_flow)
u = flow_uv[:, :, 0]
v = flow_uv[:, :, 1]
rad = np.sqrt(np.square(u) + np.square(v))
rad_max = np.max(rad)
epsilon = 1e-5
u = u / (rad_max + epsilon)
v = v / (rad_max + epsilon)
return flow_uv_to_colors(u, v, convert_to_bgr)
@@ -0,0 +1,142 @@
import numpy as np
from PIL import Image
from os.path import *
import re
import cv2
cv2.setNumThreads(0)
cv2.ocl.setUseOpenCL(False)
TAG_CHAR = np.array([202021.25], np.float32)
def readFlow(fn):
"""Read .flo file in Middlebury format"""
# Code adapted from:
# http://stackoverflow.com/questions/28013200/reading-middlebury-flow-files-with-python-bytes-array-numpy
# WARNING: this will work on little-endian architectures (eg Intel x86) only!
# print 'fn = %s'%(fn)
with open(fn, "rb") as f:
magic = np.fromfile(f, np.float32, count=1)
if 202021.25 != magic:
print("Magic number incorrect. Invalid .flo file")
return None
else:
w = np.fromfile(f, np.int32, count=1)
h = np.fromfile(f, np.int32, count=1)
# print 'Reading %d x %d flo file\n' % (w, h)
data = np.fromfile(f, np.float32, count=2 * int(w) * int(h))
# Reshape data into 3D array (columns, rows, bands)
# The reshape here is for visualization, the original code is (w,h,2)
return np.resize(data, (int(h), int(w), 2))
def readPFM(file):
file = open(file, "rb")
color = None
width = None
height = None
scale = None
endian = None
header = file.readline().rstrip()
if header == b"PF":
color = True
elif header == b"Pf":
color = False
else:
raise Exception("Not a PFM file.")
dim_match = re.match(rb"^(\d+)\s(\d+)\s$", file.readline())
if dim_match:
width, height = map(int, dim_match.groups())
else:
raise Exception("Malformed PFM header.")
scale = float(file.readline().rstrip())
if scale < 0: # little-endian
endian = "<"
scale = -scale
else:
endian = ">" # big-endian
data = np.fromfile(file, endian + "f")
shape = (height, width, 3) if color else (height, width)
data = np.reshape(data, shape)
data = np.flipud(data)
return data
def writeFlow(filename, uv, v=None):
"""Write optical flow to file.
If v is None, uv is assumed to contain both u and v channels,
stacked in depth.
Original code by Deqing Sun, adapted from Daniel Scharstein.
"""
nBands = 2
if v is None:
assert uv.ndim == 3
assert uv.shape[2] == 2
u = uv[:, :, 0]
v = uv[:, :, 1]
else:
u = uv
assert u.shape == v.shape
height, width = u.shape
f = open(filename, "wb")
# write the header
f.write(TAG_CHAR)
np.array(width).astype(np.int32).tofile(f)
np.array(height).astype(np.int32).tofile(f)
# arrange into matrix form
tmp = np.zeros((height, width * nBands))
tmp[:, np.arange(width) * 2] = u
tmp[:, np.arange(width) * 2 + 1] = v
tmp.astype(np.float32).tofile(f)
f.close()
def readFlowKITTI(filename):
flow = cv2.imread(filename, cv2.IMREAD_ANYDEPTH | cv2.IMREAD_COLOR)
flow = flow[:, :, ::-1].astype(np.float32)
flow, valid = flow[:, :, :2], flow[:, :, 2]
flow = (flow - 2**15) / 64.0
return flow, valid
def readDispKITTI(filename):
disp = cv2.imread(filename, cv2.IMREAD_ANYDEPTH) / 256.0
valid = disp > 0.0
flow = np.stack([-disp, np.zeros_like(disp)], -1)
return flow, valid
def writeFlowKITTI(filename, uv):
uv = 64.0 * uv + 2**15
valid = np.ones([uv.shape[0], uv.shape[1], 1])
uv = np.concatenate([uv, valid], axis=-1).astype(np.uint16)
cv2.imwrite(filename, uv[..., ::-1])
def read_gen(file_name, pil=False):
ext = splitext(file_name)[-1]
if ext == ".png" or ext == ".jpeg" or ext == ".ppm" or ext == ".jpg":
return Image.open(file_name)
elif ext == ".bin" or ext == ".raw":
return np.load(file_name)
elif ext == ".flo":
return readFlow(file_name).astype(np.float32)
elif ext == ".pfm":
flow = readPFM(file_name).astype(np.float32)
if len(flow.shape) == 2:
return flow
else:
return flow[:, :, :-1]
return []
@@ -0,0 +1,93 @@
import torch
import torch.nn.functional as F
import numpy as np
from scipy import interpolate
class InputPadder:
"""Pads images such that dimensions are divisible by 8"""
def __init__(self, dims, mode="sintel"):
self.ht, self.wd = dims[-2:]
pad_ht = (((self.ht // 8) + 1) * 8 - self.ht) % 8
pad_wd = (((self.wd // 8) + 1) * 8 - self.wd) % 8
if mode == "sintel":
self._pad = [
pad_wd // 2,
pad_wd - pad_wd // 2,
pad_ht // 2,
pad_ht - pad_ht // 2,
]
else:
self._pad = [pad_wd // 2, pad_wd - pad_wd // 2, 0, pad_ht]
def pad(self, *inputs):
return [F.pad(x, self._pad, mode="replicate") for x in inputs]
def unpad(self, x):
ht, wd = x.shape[-2:]
c = [self._pad[2], ht - self._pad[3], self._pad[0], wd - self._pad[1]]
return x[..., c[0] : c[1], c[2] : c[3]]
def forward_interpolate(flow):
flow = flow.detach().cpu().numpy()
dx, dy = flow[0], flow[1]
ht, wd = dx.shape
x0, y0 = np.meshgrid(np.arange(wd), np.arange(ht))
x1 = x0 + dx
y1 = y0 + dy
x1 = x1.reshape(-1)
y1 = y1.reshape(-1)
dx = dx.reshape(-1)
dy = dy.reshape(-1)
valid = (x1 > 0) & (x1 < wd) & (y1 > 0) & (y1 < ht)
x1 = x1[valid]
y1 = y1[valid]
dx = dx[valid]
dy = dy[valid]
flow_x = interpolate.griddata(
(x1, y1), dx, (x0, y0), method="nearest", fill_value=0
)
flow_y = interpolate.griddata(
(x1, y1), dy, (x0, y0), method="nearest", fill_value=0
)
flow = np.stack([flow_x, flow_y], axis=0)
return torch.from_numpy(flow).float()
def bilinear_sampler(img, coords, mode="bilinear", mask=False):
"""Wrapper for grid_sample, uses pixel coordinates"""
H, W = img.shape[-2:]
xgrid, ygrid = coords.split([1, 1], dim=-1)
xgrid = 2 * xgrid / (W - 1) - 1
ygrid = 2 * ygrid / (H - 1) - 1
grid = torch.cat([xgrid, ygrid], dim=-1)
img = F.grid_sample(img, grid, align_corners=True)
if mask:
mask = (xgrid > -1) & (ygrid > -1) & (xgrid < 1) & (ygrid < 1)
return img, mask.float()
return img
def coords_grid(batch, ht, wd, device):
coords = torch.meshgrid(
torch.arange(ht, device=device), torch.arange(wd, device=device)
)
coords = torch.stack(coords[::-1], dim=0).float()
return coords[None].repeat(batch, 1, 1, 1)
def upflow8(flow, mode="bilinear"):
new_size = (8 * flow.shape[2], 8 * flow.shape[3])
return 8 * F.interpolate(flow, size=new_size, mode=mode, align_corners=True)
+136
View File
@@ -0,0 +1,136 @@
# Flow visualization code used from https://github.com/tomrunia/OpticalFlow_Visualization
# MIT License
#
# Copyright (c) 2018 Tom Runia
#
# 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 conditions.
#
# Author: Tom Runia
# Date Created: 2018-08-03
import numpy as np
def make_colorwheel():
"""
Generates a color wheel for optical flow visualization as presented in:
Baker et al. "A Database and Evaluation Methodology for Optical Flow" (ICCV, 2007)
URL: http://vision.middlebury.edu/flow/flowEval-iccv07.pdf
Code follows the original C++ source code of Daniel Scharstein.
Code follows the the Matlab source code of Deqing Sun.
Returns:
np.ndarray: Color wheel
"""
RY = 15
YG = 6
GC = 4
CB = 11
BM = 13
MR = 6
ncols = RY + YG + GC + CB + BM + MR
colorwheel = np.zeros((ncols, 3))
col = 0
# RY
colorwheel[0:RY, 0] = 255
colorwheel[0:RY, 1] = np.floor(255 * np.arange(0, RY) / RY)
col = col + RY
# YG
colorwheel[col : col + YG, 0] = 255 - np.floor(255 * np.arange(0, YG) / YG)
colorwheel[col : col + YG, 1] = 255
col = col + YG
# GC
colorwheel[col : col + GC, 1] = 255
colorwheel[col : col + GC, 2] = np.floor(255 * np.arange(0, GC) / GC)
col = col + GC
# CB
colorwheel[col : col + CB, 1] = 255 - np.floor(255 * np.arange(CB) / CB)
colorwheel[col : col + CB, 2] = 255
col = col + CB
# BM
colorwheel[col : col + BM, 2] = 255
colorwheel[col : col + BM, 0] = np.floor(255 * np.arange(0, BM) / BM)
col = col + BM
# MR
colorwheel[col : col + MR, 2] = 255 - np.floor(255 * np.arange(MR) / MR)
colorwheel[col : col + MR, 0] = 255
return colorwheel
def flow_uv_to_colors(u, v, convert_to_bgr=False):
"""
Applies the flow color wheel to (possibly clipped) flow components u and v.
According to the C++ source code of Daniel Scharstein
According to the Matlab source code of Deqing Sun
Args:
u (np.ndarray): Input horizontal flow of shape [H,W]
v (np.ndarray): Input vertical flow of shape [H,W]
convert_to_bgr (bool, optional): Convert output image to BGR. Defaults to False.
Returns:
np.ndarray: Flow visualization image of shape [H,W,3]
"""
flow_image = np.zeros((u.shape[0], u.shape[1], 3), np.uint8)
colorwheel = make_colorwheel() # shape [55x3]
ncols = colorwheel.shape[0]
rad = np.sqrt(np.square(u) + np.square(v))
a = np.arctan2(-v, -u) / np.pi
fk = (a + 1) / 2 * (ncols - 1)
k0 = np.floor(fk).astype(np.int32)
k1 = k0 + 1
k1[k1 == ncols] = 0
f = fk - k0
for i in range(colorwheel.shape[1]):
tmp = colorwheel[:, i]
col0 = tmp[k0] / 255.0
col1 = tmp[k1] / 255.0
col = (1 - f) * col0 + f * col1
idx = rad <= 1
col[idx] = 1 - rad[idx] * (1 - col[idx])
col[~idx] = col[~idx] * 0.75 # out of range
# Note the 2-i => BGR instead of RGB
ch_idx = 2 - i if convert_to_bgr else i
flow_image[:, :, ch_idx] = np.floor(255 * col)
return flow_image
def flow_to_image(flow_uv, clip_flow=None, convert_to_bgr=False, max_flow=None):
"""
Expects a two dimensional flow image of shape.
Args:
flow_uv (np.ndarray): Flow UV image of shape [H,W,2]
clip_flow (float, optional): Clip maximum of flow values. Defaults to None.
convert_to_bgr (bool, optional): Convert output image to BGR. Defaults to False.
Returns:
np.ndarray: Flow visualization image of shape [H,W,3]
"""
assert flow_uv.ndim == 3, "input flow must have three dimensions"
assert flow_uv.shape[2] == 2, "input flow must have shape [H,W,2]"
if clip_flow is not None:
flow_uv = np.clip(flow_uv, 0, clip_flow)
u = flow_uv[:, :, 0]
v = flow_uv[:, :, 1]
if max_flow is None:
rad = np.sqrt(np.square(u) + np.square(v))
rad_max = np.max(rad)
else:
rad_max = max_flow
epsilon = 1e-5
u = u / (rad_max + epsilon)
v = v / (rad_max + epsilon)
return flow_uv_to_colors(u, v, convert_to_bgr)
+52
View File
@@ -0,0 +1,52 @@
from easydict import EasyDict as edict
import torch.nn.functional as F
class InputPadder:
"""Pads images such that dimensions are divisible by divisor"""
def __init__(self, dims, divisor=16):
self.ht, self.wd = dims[-2:]
pad_ht = (((self.ht // divisor) + 1) * divisor - self.ht) % divisor
pad_wd = (((self.wd // divisor) + 1) * divisor - self.wd) % divisor
self._pad = [
pad_wd // 2,
pad_wd - pad_wd // 2,
pad_ht // 2,
pad_ht - pad_ht // 2,
]
def pad(self, *inputs):
if len(inputs) == 1:
return F.pad(inputs[0], self._pad, mode="replicate")
else:
return [F.pad(x, self._pad, mode="replicate") for x in inputs]
def unpad(self, *inputs):
if len(inputs) == 1:
return self._unpad(inputs[0])
else:
return [self._unpad(x) for x in inputs]
def _unpad(self, x):
ht, wd = x.shape[-2:]
c = [self._pad[2], ht - self._pad[3], self._pad[0], wd - self._pad[1]]
return x[..., c[0] : c[1], c[2] : c[3]]
def easydict_to_dict(obj):
if not isinstance(obj, edict):
return obj
else:
return {k: easydict_to_dict(v) for k, v in obj.items()}
class RaftArgs:
def __init__(self, small, mixed_precision, alternate_corr):
self.small = small
self.mixed_precision = mixed_precision
self.alternate_corr = alternate_corr
def _get_kwargs(self):
return {
"small": self.small,
"mixed_precision": self.mixed_precision,
"alternate_corr": self.alternate_corr
}
+243
View File
@@ -0,0 +1,243 @@
import os
import torch
import folder_paths
import yaml
import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file
from PIL import Image
from omegaconf import OmegaConf
from tqdm import tqdm
import numpy as np
import cv2
from .gimmvfi.generalizable_INR.gimmvfi_r import GIMMVFI_R
from .gimmvfi.generalizable_INR.gimmvfi_f import GIMMVFI_F
from .gimmvfi.generalizable_INR.configs import GIMMVFIConfig
from .gimmvfi.generalizable_INR.raft import RAFT
from .gimmvfi.generalizable_INR.flowformer.core.FlowFormer.LatentCostFormer.transformer import FlowFormer
from .gimmvfi.generalizable_INR.flowformer.configs.submission import get_cfg
from .gimmvfi.utils.flow_viz import flow_to_image
from .gimmvfi.utils.utils import InputPadder, RaftArgs, easydict_to_dict
import logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
log = logging.getLogger(__name__)
script_directory = os.path.dirname(os.path.abspath(__file__))
class DownloadAndLoadGIMMVFIModel:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ([
"gimmvfi_r_arb_lpips_fp32.safetensors",
"gimmvfi_f_arb_lpips_fp32.safetensors"
],),
},
}
RETURN_TYPES = ("GIMMVIF_MODEL",)
RETURN_NAMES = ("gimmvfi_model",)
FUNCTION = "loadmodel"
CATEGORY = "GIMM-VFI"
def loadmodel(self, model):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
download_path = os.path.join(folder_paths.models_dir, 'interpolation', 'gimm-vfi')
model_path = os.path.join(download_path, model)
if not os.path.exists(model_path):
log.info(f"Downloading GMMI-VFI model to: {model_path}")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id="Kijai/GIMM-VFI_safetensors",
allow_patterns=[f"*{model}*"],
local_dir=download_path,
local_dir_use_symlinks=False,
)
if "gimmvfi_r" in model:
config_path = os.path.join(script_directory, "configs", "gimmvfi", "gimmvfi_r_arb.yaml")
flow_model = "raft-things_fp32.safetensors"
elif "gimmvfi_f" in model:
config_path = os.path.join(script_directory, "configs", "gimmvfi", "gimmvfi_f_arb.yaml")
flow_model = "flowformer_sintel_fp32.safetensors"
flow_model_path = os.path.join(folder_paths.models_dir, 'interpolation', 'gimm-vfi', flow_model)
if not os.path.exists(flow_model_path):
log.info(f"Downloading RAFT model to: {flow_model_path}")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id="Kijai/GIMM-VFI_safetensors",
allow_patterns=[f"*{flow_model}*"],
local_dir=download_path,
local_dir_use_symlinks=False,
)
with open(config_path) as f:
config = yaml.load(f, Loader=yaml.FullLoader)
config = easydict_to_dict(config)
config = OmegaConf.create(config)
arch_defaults = GIMMVFIConfig.create(config.arch)
config = OmegaConf.merge(arch_defaults, config.arch)
# load model
if "gimmvfi_r" in model:
model = GIMMVFI_R(config)
#load RAFT
raft_args = RaftArgs(
small=False,
mixed_precision=False,
alternate_corr=False
)
raft_model = RAFT(raft_args)
raft_sd = load_torch_file(flow_model_path)
raft_model.load_state_dict(raft_sd, strict=True)
raft_model.to(device)
flow_estimator = raft_model
elif "gimmvfi_f" in model:
model = GIMMVFI_F(config)
cfg = get_cfg()
flowformer = FlowFormer(cfg.latentcostformer)
flowformer_sd = load_torch_file(flow_model_path)
flowformer.load_state_dict(flowformer_sd, strict=True)
flow_estimator = flowformer
sd = load_torch_file(model_path)
model.load_state_dict(sd, strict=False)
model.flow_estimator = flow_estimator
model = model.eval().to(device)
return (model,)
def load_image(img_path):
img = Image.open(img_path)
raw_img = np.array(img.convert("RGB"))
img = torch.from_numpy(raw_img.copy()).permute(2, 0, 1) / 255.0
return img.to(torch.float).unsqueeze(0)
#region Interpolate
class GIMMVFI_interpolate:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"gimmvfi_model": ("GIMMVIF_MODEL",),
"images": ("IMAGE", {"tooltip": "The images to interpolate between"}),
"ds_factor": ("INT", {"default": 1, "min": 1, "max": 8, "step": 1}),
"interpolation_factor": ("INT", {"default": 8, "min": 1, "max": 100, "step": 1}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
},
}
RETURN_TYPES = ("IMAGE", "IMAGE",)
RETURN_NAMES = ("images", "flow_tensors",)
FUNCTION = "interpolate"
CATEGORY = "PyramidFlowWrapper"
def interpolate(self, gimmvfi_model, images, ds_factor, interpolation_factor,seed):
mm.soft_empty_cache()
images = images.permute(0, 3, 1, 2)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
gimmvfi_model.to(device)
out_images_list = []
flows = []
start = 0
end = images.shape[0] - 1
pbar = ProgressBar(images.shape[0] - 1)
for j in tqdm(range(start, end)):
I0 = images[j].unsqueeze(0)
I2 = images[j+1].unsqueeze(0)
if j == start:
out_images_list.append(I0.squeeze(0).permute(1, 2, 0))
padder = InputPadder(I0.shape, 32)
I0, I2 = padder.pad(I0, I2)
xs = torch.cat((I0.unsqueeze(2), I2.unsqueeze(2)), dim=2).to(device, non_blocking=True)
batch_size = xs.shape[0]
s_shape = xs.shape[-2:]
coord_inputs = [
(
gimmvfi_model.sample_coord_input(
batch_size,
s_shape,
[1 / interpolation_factor * i],
device=xs.device,
upsample_ratio=ds_factor,
),
None,
)
for i in range(1, interpolation_factor)
]
timesteps = [
i * 1 / interpolation_factor * torch.ones(xs.shape[0]).to(xs.device).to(torch.float)
for i in range(1, interpolation_factor)
]
all_outputs = gimmvfi_model(xs, coord_inputs, t=timesteps, ds_factor=ds_factor)
out_frames = [padder.unpad(im) for im in all_outputs["imgt_pred"]]
out_flowts = [padder.unpad(f) for f in all_outputs["flowt"]]
flowt_imgs = [
flow_to_image(
flowt.squeeze().detach().cpu().permute(1, 2, 0).numpy(),
convert_to_bgr=True,
)
for flowt in out_flowts
]
I1_pred_img = [
(I1_pred[0].detach().cpu().permute(1, 2, 0))
for I1_pred in out_frames
]
for i in range(interpolation_factor - 1):
out_images_list.append(I1_pred_img[i])
flows.append(flowt_imgs[i])
out_images_list.append(
((padder.unpad(I2)).squeeze().detach().cpu().permute(1, 2, 0))
)
pbar.update(1)
image_tensors = torch.stack(out_images_list)
image_tensors = image_tensors.cpu().float()
rgb_images = [cv2.cvtColor(flow, cv2.COLOR_BGR2RGB) for flow in flows]
flow_tensors = torch.stack([torch.from_numpy(image) for image in rgb_images])
flow_tensors = flow_tensors / 255.0
flow_tensors = flow_tensors.cpu().float()
return (image_tensors, flow_tensors)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadGIMMVFIModel": DownloadAndLoadGIMMVFIModel,
"GIMMVFI_interpolate": GIMMVFI_interpolate,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadGIMMVFIModel": "(Down)Load GIMMVFI Model",
"GIMMVFI_interpolate": "GIMM-VFI Interpolate",
}
+4
View File
@@ -0,0 +1,4 @@
opencv-python
numpy
pillow
cupy-cuda12x>=13.3.0