Добавлены пропсы конвейера и стереодвижки, задействованные в прогоне
assets/conveyors (274 МБ) - ленты и угловая секция NVIDIA, на которые ссылается сцена относительным путём. Раньше исключались как перекачиваемые, но без них сцена не композится из коробки. cv/ - код стереодвижков, которые вызывает control_test, без весов: * defom-stereo - рабочий бейзлайн (DEFOM vitl, вход 480, iters 24) * crestereo - второй движок, точнее по габаритам (MAE 23.5 против 32.8 мм) * fast-foundationstereo - проверялся, в бейзлайн не вошёл * circular_section.py - показатель кругового сечения, перенесён в measure_plane.py: выравнивает облако по СОБСТВЕННЫМ главным осям и режет на пяти высотах вдоль каждой. Три самодельные версии (мировые оси, одно сечение) давали хуже; результаты проверки на эталонной геометрии - в circular_section_results.json Веса по-прежнему не в репозитории - источники в MODELS.md. Наборы кадров прежних прогонов (cv/flow_*, 1.26 ГБ) исключены: это выход, а не исходники. Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,2 @@
|
||||
# Auto detect text files and perform LF normalization
|
||||
* text=auto
|
||||
@@ -0,0 +1,156 @@
|
||||
# Byte-compiled / optimized / DLL files
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*$py.class
|
||||
|
||||
# C extensions
|
||||
*.so
|
||||
|
||||
# Distribution / packaging
|
||||
.Python
|
||||
build/
|
||||
develop-eggs/
|
||||
dist/
|
||||
downloads/
|
||||
eggs/
|
||||
.eggs/
|
||||
lib/
|
||||
lib64/
|
||||
parts/
|
||||
sdist/
|
||||
var/
|
||||
wheels/
|
||||
share/python-wheels/
|
||||
*.egg-info/
|
||||
.installed.cfg
|
||||
*.egg
|
||||
MANIFEST
|
||||
|
||||
# PyInstaller
|
||||
# Usually these files are written by a python script from a template
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
pip-delete-this-directory.txt
|
||||
|
||||
# Unit test / coverage reports
|
||||
htmlcov/
|
||||
.tox/
|
||||
.nox/
|
||||
.coverage
|
||||
.coverage.*
|
||||
.cache
|
||||
nosetests.xml
|
||||
coverage.xml
|
||||
*.cover
|
||||
*.py,cover
|
||||
.hypothesis/
|
||||
.pytest_cache/
|
||||
cover/
|
||||
|
||||
# Translations
|
||||
*.mo
|
||||
*.pot
|
||||
|
||||
# Django stuff:
|
||||
*.log
|
||||
local_settings.py
|
||||
db.sqlite3
|
||||
db.sqlite3-journal
|
||||
|
||||
# Flask stuff:
|
||||
instance/
|
||||
.webassets-cache
|
||||
|
||||
# Scrapy stuff:
|
||||
.scrapy
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
|
||||
# PyBuilder
|
||||
.pybuilder/
|
||||
target/
|
||||
|
||||
# Jupyter Notebook
|
||||
.ipynb_checkpoints
|
||||
|
||||
# IPython
|
||||
profile_default/
|
||||
ipython_config.py
|
||||
|
||||
# pyenv
|
||||
# For a library or package, you might want to ignore these files since the code is
|
||||
# intended to run in multiple environments; otherwise, check them in:
|
||||
# .python-version
|
||||
|
||||
# pipenv
|
||||
# According to pypa/pipenv#598, it is recommended to include Pipfile.lock in version control.
|
||||
# However, in case of collaboration, if having platform-specific dependencies or dependencies
|
||||
# having no cross-platform support, pipenv may install dependencies that don't work, or not
|
||||
# install all needed dependencies.
|
||||
#Pipfile.lock
|
||||
|
||||
# poetry
|
||||
# Similar to Pipfile.lock, it is generally recommended to include poetry.lock in version control.
|
||||
# This is especially recommended for binary packages to ensure reproducibility, and is more
|
||||
# commonly ignored for libraries.
|
||||
# https://python-poetry.org/docs/basic-usage/#commit-your-poetrylock-file-to-version-control
|
||||
#poetry.lock
|
||||
|
||||
# PEP 582; used by e.g. github.com/David-OConnor/pyflow
|
||||
__pypackages__/
|
||||
|
||||
# Celery stuff
|
||||
celerybeat-schedule
|
||||
celerybeat.pid
|
||||
|
||||
# SageMath parsed files
|
||||
*.sage.py
|
||||
|
||||
# Environments
|
||||
.env
|
||||
.venv
|
||||
env/
|
||||
venv/
|
||||
ENV/
|
||||
env.bak/
|
||||
venv.bak/
|
||||
|
||||
# Spyder project settings
|
||||
.spyderproject
|
||||
.spyproject
|
||||
|
||||
# Rope project settings
|
||||
.ropeproject
|
||||
|
||||
# mkdocs documentation
|
||||
/site
|
||||
|
||||
# mypy
|
||||
.mypy_cache/
|
||||
.dmypy.json
|
||||
dmypy.json
|
||||
|
||||
# Pyre type checker
|
||||
.pyre/
|
||||
|
||||
# pytype static type analyzer
|
||||
.pytype/
|
||||
|
||||
# Cython debug symbols
|
||||
cython_debug/
|
||||
|
||||
# PyCharm
|
||||
# JetBrains specific template is maintainted in a separate JetBrains.gitignore that can
|
||||
# be found at https://github.com/github/gitignore/blob/main/Global/JetBrains.gitignore
|
||||
# and can be added to the global gitignore or merged into this file. For a more nuclear
|
||||
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
|
||||
#.idea/
|
||||
|
||||
vis_results/
|
||||
models/*
|
||||
test_data/*
|
||||
@@ -0,0 +1,46 @@
|
||||
# CREStereo-Pytorch
|
||||
Non-official Pytorch implementation of the CREStereo (CVPR 2022 Oral) model converted from the original MegEngine implementation.
|
||||
|
||||

|
||||
|
||||
**update 2023/01/03**:
|
||||
- enable DistributedDataParallel (DDP) training, training time is much faster than before.
|
||||
|
||||
```shell
|
||||
# train DDP
|
||||
# change 'dist' to True in /cfgs/train.yaml file
|
||||
python -m torch.distributed.launch --nproc_per_node=8 train.py
|
||||
# train DP
|
||||
# change 'dist' to False in /cfgs/train.yaml file
|
||||
python train.py
|
||||
```
|
||||
|
||||
# Important
|
||||
- This is just an effort to try to implement the CREStereo model into Pytorch from MegEngine due to the issues of the framework to convert to other formats (https://github.com/megvii-research/CREStereo/issues/3).
|
||||
- I am not the author of the paper, and I am don't fully understand what the model is doing. Therefore, there might be small differences with the original model that might impact the performance.
|
||||
- I have not added any license, since the repository uses code from different repositories. Check the License section below for more detail.
|
||||
|
||||
# Pretrained model
|
||||
- Download the model from [here](https://drive.google.com/file/d/1D2s1v4VhJlNz98FQpFxf_kBAKQVN_7xo/view?usp=sharing) and save it into the **[models](https://github.com/ibaiGorordo/CREStereo-Pytorch/tree/main/models)** folder.
|
||||
- The model was converted from the original **[MegEngine weights](https://drive.google.com/file/d/1Wx_-zDQh7BUFBmN9im_26DFpnf3AkXj4/view)** using the `convert_weights.py` script. Place the MegEngine weights (crestereo_eth3d.mge) file into the **[models](https://github.com/ibaiGorordo/CREStereo-Pytorch/tree/main/models)** folder before the conversion.
|
||||
|
||||
# ONNX Conversion
|
||||
- After either downloading the pretrained weights or training your own model, you will have a `models/crestereo_eth3d.pth` file. If you want to run your model with ONNX, you need to run the convert_to_onnx.py script. The script has two parts:
|
||||
1. Convert the model to an ONNX model that takes in left, right images as well as an initial flow estimate (takes a few seconds)
|
||||
2. Convert the model to an ONNX model that takes in left, right images and NO initial flow estimate (takes several minutes and requires pytorch >= 1.12)
|
||||
(afaik) You will need both models to get the same results as you do from test_model.py.
|
||||
- Run the test_onnx_model.py script to verify your models work as expected!
|
||||
- NOTE: although the test_model.py script works with any size images as input, once you have converted your
|
||||
Pytorch model into ONNX models, you must provide them with the image sizes used at conversion time or it will not work.
|
||||
|
||||
# Licences:
|
||||
- CREStereo (Apache License 2.0): https://github.com/megvii-research/CREStereo/blob/master/LICENSE
|
||||
- RAFT (BSD 3-Clause):https://github.com/princeton-vl/RAFT/blob/master/LICENSE
|
||||
- LoFTR (Apache License 2.0):https://github.com/zju3dv/LoFTR/blob/master/LICENSE
|
||||
|
||||
# References:
|
||||
- CREStereo: https://github.com/megvii-research/CREStereo
|
||||
- RAFT: https://github.com/princeton-vl/RAFT
|
||||
- LoFTR: https://github.com/zju3dv/LoFTR
|
||||
- Grid sample replacement: https://zenn.dev/pinto0309/scraps/7d4032067d0160
|
||||
- torch2mge: https://github.com/MegEngine/torch2mge
|
||||
@@ -0,0 +1,20 @@
|
||||
seed: 0
|
||||
mixed_precision: false
|
||||
base_lr: 4.0e-4
|
||||
|
||||
nr_gpus: 8
|
||||
batch_size: 4
|
||||
n_total_epoch: 600
|
||||
minibatch_per_epoch: 500
|
||||
|
||||
loadmodel: ~
|
||||
log_dir: "./train_log"
|
||||
model_save_freq_epoch: 1
|
||||
|
||||
max_disp: 256
|
||||
image_width: 512
|
||||
image_height: 384
|
||||
training_data_path: "./stereo_trainset/crestereo"
|
||||
|
||||
log_level: "logging.INFO"
|
||||
dist: True # True for DDP, False for DP
|
||||
@@ -0,0 +1,49 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
import cv2
|
||||
from imread_from_url import imread_from_url
|
||||
|
||||
from nets import Model
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
model_path = "models/crestereo_eth3d.pth"
|
||||
|
||||
model = Model(max_disp=256, mixed_precision=False, test_mode=True)
|
||||
model.load_state_dict(torch.load(model_path), strict=True)
|
||||
model.eval()
|
||||
|
||||
in_h, in_w = (480, 640)
|
||||
t1_half = torch.rand(1, 3, in_h//2, in_w//2)
|
||||
t2_half = torch.rand(1, 3, in_h//2, in_w//2)
|
||||
|
||||
t1 = torch.rand(1, 3, in_h, in_w)
|
||||
t2 = torch.rand(1, 3, in_h, in_w)
|
||||
flow_init = torch.rand(1, 2, in_h//2, in_w//2)
|
||||
|
||||
# Export the model
|
||||
torch.onnx.export(model,
|
||||
(t1, t2, flow_init),
|
||||
"crestereo.onnx", # where to save the model (can be a file or file-like object)
|
||||
export_params=True, # store the trained parameter weights inside the model file
|
||||
opset_version=12, # the ONNX version to export the model to
|
||||
do_constant_folding=True, # whether to execute constant folding for optimization
|
||||
input_names = ['left', 'right','flow_init'], # the model's input names
|
||||
output_names = ['output'])
|
||||
|
||||
# Export the model without init_flow (it takes a lot of time)
|
||||
# !! Does not work prior to pytorch 1.12 (confirmed working on pytorch 2.0.0)
|
||||
# Ref: https://github.com/pytorch/pytorch/pull/73760
|
||||
torch.onnx.export(model,
|
||||
(t1_half, t2_half),
|
||||
"crestereo_without_flow.onnx", # where to save the model (can be a file or file-like object)
|
||||
export_params=True, # store the trained parameter weights inside the model file
|
||||
opset_version=12, # the ONNX version to export the model to
|
||||
do_constant_folding=True, # whether to execute constant folding for optimization
|
||||
input_names = ['left', 'right'], # the model's input names
|
||||
output_names = ['output'])
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,26 @@
|
||||
import copy
|
||||
import torch
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
from nets import Model
|
||||
|
||||
# Read Megengine parameters
|
||||
pretrained_dict = mge.load("models/crestereo_eth3d.mge")
|
||||
|
||||
model = Model(max_disp=256, mixed_precision=False, test_mode=True)
|
||||
model.eval()
|
||||
|
||||
state_dict = model.state_dict()
|
||||
for key, value in pretrained_dict['state_dict'].items():
|
||||
|
||||
print(f"Converting {key}")
|
||||
# Fix shape mismatch
|
||||
if value.shape[0] == 1:
|
||||
value = np.squeeze(value)
|
||||
|
||||
state_dict[key] = torch.tensor(value)
|
||||
|
||||
output_path = "models/crestereo_eth3d.pth"
|
||||
torch.save(state_dict, output_path)
|
||||
print(f"\nModel saved to: {output_path}")
|
||||
@@ -0,0 +1,215 @@
|
||||
import os
|
||||
import cv2
|
||||
import glob
|
||||
import numpy as np
|
||||
from PIL import Image, ImageEnhance
|
||||
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class Augmentor:
|
||||
def __init__(
|
||||
self,
|
||||
image_height=384,
|
||||
image_width=512,
|
||||
max_disp=256,
|
||||
scale_min=0.6,
|
||||
scale_max=1.0,
|
||||
seed=0,
|
||||
):
|
||||
super().__init__()
|
||||
self.image_height = image_height
|
||||
self.image_width = image_width
|
||||
self.max_disp = max_disp
|
||||
self.scale_min = scale_min
|
||||
self.scale_max = scale_max
|
||||
self.rng = np.random.RandomState(seed)
|
||||
|
||||
def chromatic_augmentation(self, img):
|
||||
random_brightness = np.random.uniform(0.8, 1.2)
|
||||
random_contrast = np.random.uniform(0.8, 1.2)
|
||||
random_gamma = np.random.uniform(0.8, 1.2)
|
||||
|
||||
img = Image.fromarray(img)
|
||||
|
||||
enhancer = ImageEnhance.Brightness(img)
|
||||
img = enhancer.enhance(random_brightness)
|
||||
enhancer = ImageEnhance.Contrast(img)
|
||||
img = enhancer.enhance(random_contrast)
|
||||
|
||||
gamma_map = [
|
||||
255 * 1.0 * pow(ele / 255.0, random_gamma) for ele in range(256)
|
||||
] * 3
|
||||
img = img.point(gamma_map) # use PIL's point-function to accelerate this part
|
||||
|
||||
img_ = np.array(img)
|
||||
|
||||
return img_
|
||||
|
||||
def __call__(self, left_img, right_img, left_disp):
|
||||
# 1. chromatic augmentation
|
||||
left_img = self.chromatic_augmentation(left_img)
|
||||
right_img = self.chromatic_augmentation(right_img)
|
||||
|
||||
# 2. spatial augmentation
|
||||
# 2.1) rotate & vertical shift for right image
|
||||
if self.rng.binomial(1, 0.5):
|
||||
angle, pixel = 0.1, 2
|
||||
px = self.rng.uniform(-pixel, pixel)
|
||||
ag = self.rng.uniform(-angle, angle)
|
||||
image_center = (
|
||||
self.rng.uniform(0, right_img.shape[0]),
|
||||
self.rng.uniform(0, right_img.shape[1]),
|
||||
)
|
||||
rot_mat = cv2.getRotationMatrix2D(image_center, ag, 1.0)
|
||||
right_img = cv2.warpAffine(
|
||||
right_img, rot_mat, right_img.shape[1::-1], flags=cv2.INTER_LINEAR
|
||||
)
|
||||
trans_mat = np.float32([[1, 0, 0], [0, 1, px]])
|
||||
right_img = cv2.warpAffine(
|
||||
right_img, trans_mat, right_img.shape[1::-1], flags=cv2.INTER_LINEAR
|
||||
)
|
||||
|
||||
# 2.2) random resize
|
||||
resize_scale = self.rng.uniform(self.scale_min, self.scale_max)
|
||||
|
||||
left_img = cv2.resize(
|
||||
left_img,
|
||||
None,
|
||||
fx=resize_scale,
|
||||
fy=resize_scale,
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
)
|
||||
right_img = cv2.resize(
|
||||
right_img,
|
||||
None,
|
||||
fx=resize_scale,
|
||||
fy=resize_scale,
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
)
|
||||
|
||||
disp_mask = (left_disp < float(self.max_disp / resize_scale)) & (left_disp > 0)
|
||||
disp_mask = disp_mask.astype("float32")
|
||||
disp_mask = cv2.resize(
|
||||
disp_mask,
|
||||
None,
|
||||
fx=resize_scale,
|
||||
fy=resize_scale,
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
)
|
||||
|
||||
left_disp = (
|
||||
cv2.resize(
|
||||
left_disp,
|
||||
None,
|
||||
fx=resize_scale,
|
||||
fy=resize_scale,
|
||||
interpolation=cv2.INTER_LINEAR,
|
||||
)
|
||||
* resize_scale
|
||||
)
|
||||
|
||||
# 2.3) random crop
|
||||
h, w, c = left_img.shape
|
||||
dx = w - self.image_width
|
||||
dy = h - self.image_height
|
||||
dy = self.rng.randint(min(0, dy), max(0, dy) + 1)
|
||||
dx = self.rng.randint(min(0, dx), max(0, dx) + 1)
|
||||
|
||||
M = np.float32([[1.0, 0.0, -dx], [0.0, 1.0, -dy]])
|
||||
left_img = cv2.warpAffine(
|
||||
left_img,
|
||||
M,
|
||||
(self.image_width, self.image_height),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderValue=0,
|
||||
)
|
||||
right_img = cv2.warpAffine(
|
||||
right_img,
|
||||
M,
|
||||
(self.image_width, self.image_height),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderValue=0,
|
||||
)
|
||||
left_disp = cv2.warpAffine(
|
||||
left_disp,
|
||||
M,
|
||||
(self.image_width, self.image_height),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderValue=0,
|
||||
)
|
||||
disp_mask = cv2.warpAffine(
|
||||
disp_mask,
|
||||
M,
|
||||
(self.image_width, self.image_height),
|
||||
flags=cv2.INTER_LINEAR,
|
||||
borderValue=0,
|
||||
)
|
||||
|
||||
# 3. add random occlusion to right image
|
||||
if self.rng.binomial(1, 0.5):
|
||||
sx = int(self.rng.uniform(50, 100))
|
||||
sy = int(self.rng.uniform(50, 100))
|
||||
cx = int(self.rng.uniform(sx, right_img.shape[0] - sx))
|
||||
cy = int(self.rng.uniform(sy, right_img.shape[1] - sy))
|
||||
right_img[cx - sx : cx + sx, cy - sy : cy + sy] = np.mean(
|
||||
np.mean(right_img, 0), 0
|
||||
)[np.newaxis, np.newaxis]
|
||||
|
||||
return left_img, right_img, left_disp, disp_mask
|
||||
|
||||
|
||||
class CREStereoDataset(Dataset):
|
||||
def __init__(self, root):
|
||||
super().__init__()
|
||||
self.imgs = glob.glob(os.path.join(root, "**/*_left.jpg"), recursive=True)
|
||||
self.augmentor = Augmentor(
|
||||
image_height=384,
|
||||
image_width=512,
|
||||
max_disp=256,
|
||||
scale_min=0.6,
|
||||
scale_max=1.0,
|
||||
seed=0,
|
||||
)
|
||||
self.rng = np.random.RandomState(0)
|
||||
|
||||
def get_disp(self, path):
|
||||
disp = cv2.imread(path, cv2.IMREAD_UNCHANGED)
|
||||
return disp.astype(np.float32) / 32
|
||||
|
||||
def __getitem__(self, index):
|
||||
# find path
|
||||
left_path = self.imgs[index]
|
||||
prefix = left_path[: left_path.rfind("_")]
|
||||
right_path = prefix + "_right.jpg"
|
||||
left_disp_path = prefix + "_left.disp.png"
|
||||
right_disp_path = prefix + "_right.disp.png"
|
||||
|
||||
# read img, disp
|
||||
left_img = cv2.imread(left_path, cv2.IMREAD_COLOR)
|
||||
right_img = cv2.imread(right_path, cv2.IMREAD_COLOR)
|
||||
left_disp = self.get_disp(left_disp_path)
|
||||
right_disp = self.get_disp(right_disp_path)
|
||||
|
||||
if self.rng.binomial(1, 0.5):
|
||||
left_img, right_img = np.fliplr(right_img), np.fliplr(left_img)
|
||||
left_disp, right_disp = np.fliplr(right_disp), np.fliplr(left_disp)
|
||||
left_disp[left_disp == np.inf] = 0
|
||||
|
||||
# augmentaion
|
||||
left_img, right_img, left_disp, disp_mask = self.augmentor(
|
||||
left_img, right_img, left_disp
|
||||
)
|
||||
|
||||
left_img = left_img.transpose(2, 0, 1).astype("uint8")
|
||||
right_img = right_img.transpose(2, 0, 1).astype("uint8")
|
||||
|
||||
return {
|
||||
"left": left_img,
|
||||
"right": right_img,
|
||||
"disparity": left_disp,
|
||||
"mask": disp_mask,
|
||||
}
|
||||
|
||||
def __len__(self):
|
||||
return len(self.imgs)
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 388 KiB |
@@ -0,0 +1,44 @@
|
||||
import pickle
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
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 test_bilinear_sampler():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/bilinear_sampler_test.pickle', 'rb') as f:
|
||||
right_feature_prev, coords, right_feature = pickle.load(f)
|
||||
|
||||
right_feature_prev = torch.tensor(right_feature_prev.numpy())
|
||||
coords = torch.tensor(coords.numpy())
|
||||
right_feature = right_feature.numpy()
|
||||
|
||||
# Test Pytorch
|
||||
right_feature_pytorch = bilinear_sampler(right_feature_prev, coords).numpy()
|
||||
|
||||
error = np.mean(right_feature_pytorch-right_feature)
|
||||
print(f"test_coords_grid - Avg. Error: {error}, \n \
|
||||
Original shape: {coords.numpy().shape},\n \
|
||||
Obtained shape: {right_feature_pytorch.shape}, Expected shape: {right_feature.shape}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
test_bilinear_sampler()
|
||||
@@ -0,0 +1,29 @@
|
||||
import pickle
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def coords_grid(batch, ht, wd, device):
|
||||
coords = torch.meshgrid(torch.arange(ht, device=device), torch.arange(wd, device=device), indexing='ij')
|
||||
coords = torch.stack(coords[::-1], dim=0).float()
|
||||
return coords[None].repeat(batch, 1, 1, 1)
|
||||
|
||||
def test_coords_grid():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/coords_grid_test.pickle', 'rb') as f:
|
||||
batch, ht, wd, coords = pickle.load(f)
|
||||
|
||||
coords = coords.numpy()
|
||||
|
||||
# Test Pytorch
|
||||
coords_pytorch = coords_grid(batch, ht, wd, 'cpu').numpy()
|
||||
|
||||
error = np.mean(coords_pytorch-coords)
|
||||
print(f"test_coords_grid - Avg. Error: {error}, \n \
|
||||
Obtained shape: {coords_pytorch.shape}, Expected shape: {coords.shape}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
test_coords_grid()
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
Binary file not shown.
@@ -0,0 +1,51 @@
|
||||
import pickle
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def manual_pad(x, pady, padx):
|
||||
|
||||
pad = (padx, padx, pady, pady)
|
||||
return F.pad(torch.tensor(x), pad, "replicate")
|
||||
|
||||
|
||||
def test_pad_1_1():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/manual_pad_test1_1.pickle', 'rb') as f:
|
||||
right_feature, pady, padx, right_pad = pickle.load(f)
|
||||
|
||||
right_feature = right_feature.numpy()
|
||||
right_pad = right_pad.numpy()
|
||||
|
||||
# Test Pytorch
|
||||
right_pad_pytorch = manual_pad(right_feature, pady, padx).numpy()
|
||||
|
||||
error = np.mean(right_pad_pytorch-right_pad)
|
||||
print(f"test_pad_1_1 - Avg. Error: {error}, \n \
|
||||
Orig. shape: {right_feature.shape}, \n \
|
||||
Padded shape: {right_pad_pytorch.shape}, Expected shape: {right_pad.shape}")
|
||||
|
||||
def test_pad_0_4():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/manual_pad_test0_4.pickle', 'rb') as f:
|
||||
right_feature, pady, padx, right_pad = pickle.load(f)
|
||||
|
||||
right_feature = right_feature.numpy()
|
||||
right_pad = right_pad.numpy()
|
||||
|
||||
# Test Pytorch
|
||||
right_pad_pytorch = manual_pad(right_feature, pady, padx).numpy()
|
||||
|
||||
error = np.mean(right_pad_pytorch-right_pad)
|
||||
print(f"test_pad_0_4 - Avg. Error: {error}, \n \
|
||||
Orig. shape: {right_feature.shape}, \n \
|
||||
Padded shape: {right_pad_pytorch.shape}, Expected shape: {right_pad.shape}")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
test_pad_1_1()
|
||||
|
||||
test_pad_0_4()
|
||||
@@ -0,0 +1,30 @@
|
||||
import pickle
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def test_meshgrid():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/meshgrid_np_test.pkl', 'rb') as f:
|
||||
rx, dilatex, ry, dilatey, x_grid, y_grid = pickle.load(f)
|
||||
|
||||
x_grid = x_grid.numpy()
|
||||
y_grid = y_grid.numpy()
|
||||
|
||||
# Test Pytorch
|
||||
x_grid_pytorch, y_grid_pytorch = torch.meshgrid(torch.arange(-rx, rx + 1, dilatex, device='cpu'),
|
||||
torch.arange(-ry, ry + 1, dilatey, device='cpu'), indexing='xy')
|
||||
|
||||
|
||||
error_x = np.mean(x_grid_pytorch.numpy()-x_grid)
|
||||
error_y = np.mean(y_grid_pytorch.numpy()-y_grid)
|
||||
print(f"test_meshgrid (X) - Avg. Error: {error_x}, \n \
|
||||
Obtained shape: {x_grid_pytorch.numpy().shape}, Expected shape: {x_grid.shape}")
|
||||
print(f"test_meshgrid (Y) - Avg. Error: {error_y}, \n \
|
||||
Obtained shape: {y_grid_pytorch.numpy().shape}, Expected shape: {y_grid.shape}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
test_meshgrid()
|
||||
@@ -0,0 +1,31 @@
|
||||
import pickle
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def test_offset():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/offset_test.pkl', 'rb') as f:
|
||||
x_grid, y_grid, reshape_shape, transpose_order, expand_size, repeat_size, repeat_axis, offsets = pickle.load(f)
|
||||
|
||||
x_grid = torch.tensor(x_grid.numpy())
|
||||
y_grid = torch.tensor(y_grid.numpy())
|
||||
offsets_mge = offsets.numpy()
|
||||
N = repeat_size
|
||||
|
||||
# Test Pytorch
|
||||
offsets = torch.stack((x_grid, y_grid))
|
||||
offsets = offsets.reshape(2, -1).permute(1, 0)
|
||||
for d in sorted((0, 2, 3)):
|
||||
offsets = offsets.unsqueeze(d)
|
||||
offsets = offsets.repeat_interleave(N, dim=0)
|
||||
|
||||
error = np.mean(offsets.numpy()-offsets_mge)
|
||||
print(f"test_offset - Avg. Error: {error}, \n \
|
||||
Obtained shape: {offsets.numpy().shape}, Expected shape: {offsets_mge.shape}")
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
test_offset()
|
||||
@@ -0,0 +1,47 @@
|
||||
import pickle
|
||||
import numpy as np
|
||||
import megengine as mge
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
def test_split():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/split_test.pkl', 'rb') as f:
|
||||
left_feature, size, axis, lefts = pickle.load(f)
|
||||
|
||||
left_feature = torch.tensor(left_feature.numpy())
|
||||
|
||||
# Test Pytorch
|
||||
lefts_pytorch = torch.split(left_feature, left_feature.shape[axis]//size, dim=axis)
|
||||
|
||||
for i, (left_pytorch, left) in enumerate(zip(lefts_pytorch, lefts)):
|
||||
|
||||
error = np.mean(left_pytorch.numpy()-left.numpy())
|
||||
print(f"test_split {i} - Avg. Error: {error}, \n \
|
||||
Obtained shape: {left_pytorch.numpy().shape}, Expected shape: {left.numpy().shape}\n")
|
||||
|
||||
def test_split_list():
|
||||
# Getting back the megengine objects:
|
||||
with open('test_data/split_test_list.pkl', 'rb') as f:
|
||||
fmap1, size, axis, net, inp = pickle.load(f)
|
||||
|
||||
fmap1 = torch.tensor(fmap1.numpy())
|
||||
net = net.numpy()
|
||||
inp = inp.numpy()
|
||||
|
||||
# Test Pytorch
|
||||
net_pytorch, inp_pytorch = torch.split(fmap1, [size[0],size[0]], dim=axis)
|
||||
|
||||
error_net = np.mean(net_pytorch.numpy()-net)
|
||||
error_inp = np.mean(inp_pytorch.numpy()-inp)
|
||||
print(f"test_split_list (net) - Avg. Error: {error_net}, \n \
|
||||
Obtained shape: {net_pytorch.numpy().shape}, Expected shape: {net.shape}\n")
|
||||
print(f"test_split_list (inp) - Avg. Error: {error_inp}, \n \
|
||||
Obtained shape: {inp_pytorch.numpy().shape}, Expected shape: {inp.shape}\n")
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
test_split()
|
||||
test_split_list()
|
||||
@@ -0,0 +1 @@
|
||||
from .crestereo import CREStereo as Model
|
||||
@@ -0,0 +1,2 @@
|
||||
from .transformer import LocalFeatureTransformer
|
||||
from .position_encoding import PositionEncodingSine
|
||||
@@ -0,0 +1,81 @@
|
||||
"""
|
||||
Linear Transformer proposed in "Transformers are RNNs: Fast Autoregressive Transformers with Linear Attention"
|
||||
Modified from: https://github.com/idiap/fast-transformers/blob/master/fast_transformers/attention/linear_attention.py
|
||||
"""
|
||||
|
||||
import torch
|
||||
from torch.nn import Module, Dropout
|
||||
|
||||
|
||||
def elu_feature_map(x):
|
||||
return torch.nn.functional.elu(x) + 1
|
||||
|
||||
|
||||
class LinearAttention(Module):
|
||||
def __init__(self, eps=1e-6):
|
||||
super().__init__()
|
||||
self.feature_map = elu_feature_map
|
||||
self.eps = eps
|
||||
|
||||
def forward(self, queries, keys, values, q_mask=None, kv_mask=None):
|
||||
""" Multi-Head linear attention proposed in "Transformers are RNNs"
|
||||
Args:
|
||||
queries: [N, L, H, D]
|
||||
keys: [N, S, H, D]
|
||||
values: [N, S, H, D]
|
||||
q_mask: [N, L]
|
||||
kv_mask: [N, S]
|
||||
Returns:
|
||||
queried_values: (N, L, H, D)
|
||||
"""
|
||||
Q = self.feature_map(queries)
|
||||
K = self.feature_map(keys)
|
||||
|
||||
# set padded position to zero
|
||||
if q_mask is not None:
|
||||
Q = Q * q_mask[:, :, None, None]
|
||||
if kv_mask is not None:
|
||||
K = K * kv_mask[:, :, None, None]
|
||||
values = values * kv_mask[:, :, None, None]
|
||||
|
||||
v_length = values.size(1)
|
||||
values = values / v_length # prevent fp16 overflow
|
||||
KV = torch.einsum("nshd,nshv->nhdv", K, values) # (S,D)' @ S,V
|
||||
Z = 1 / (torch.einsum("nlhd,nhd->nlh", Q, K.sum(dim=1)) + self.eps)
|
||||
queried_values = torch.einsum("nlhd,nhdv,nlh->nlhv", Q, KV, Z) * v_length
|
||||
|
||||
return queried_values.contiguous()
|
||||
|
||||
|
||||
class FullAttention(Module):
|
||||
def __init__(self, use_dropout=False, attention_dropout=0.1):
|
||||
super().__init__()
|
||||
self.use_dropout = use_dropout
|
||||
self.dropout = Dropout(attention_dropout)
|
||||
|
||||
def forward(self, queries, keys, values, q_mask=None, kv_mask=None):
|
||||
""" Multi-head scaled dot-product attention, a.k.a full attention.
|
||||
Args:
|
||||
queries: [N, L, H, D]
|
||||
keys: [N, S, H, D]
|
||||
values: [N, S, H, D]
|
||||
q_mask: [N, L]
|
||||
kv_mask: [N, S]
|
||||
Returns:
|
||||
queried_values: (N, L, H, D)
|
||||
"""
|
||||
|
||||
# Compute the unnormalized attention and apply the masks
|
||||
QK = torch.einsum("nlhd,nshd->nlsh", queries, keys)
|
||||
if kv_mask is not None:
|
||||
QK.masked_fill_(~(q_mask[:, :, None, None] * kv_mask[:, None, :, None]), float('-inf'))
|
||||
|
||||
# Compute the attention and the weighted average
|
||||
softmax_temp = 1. / queries.size(3)**.5 # sqrt(D)
|
||||
A = torch.softmax(softmax_temp * QK, dim=2)
|
||||
if self.use_dropout:
|
||||
A = self.dropout(A)
|
||||
|
||||
queried_values = torch.einsum("nlsh,nshd->nlhd", A, values)
|
||||
|
||||
return queried_values.contiguous()
|
||||
@@ -0,0 +1,41 @@
|
||||
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), temp_bug_fix=False):
|
||||
"""
|
||||
Args:
|
||||
max_shape (tuple): for 1/8 featmap, the max length of 256 corresponds to 2048 pixels
|
||||
temp_bug_fix (bool): As noted in this [issue](https://github.com/zju3dv/LoFTR/issues/41),
|
||||
the original implementation of LoFTR includes a bug in the pos-enc impl, which has little impact
|
||||
on the final performance. For now, we keep both impls for backward compatability.
|
||||
We will remove the buggy impl after re-training all variants of our released models.
|
||||
"""
|
||||
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)
|
||||
if temp_bug_fix:
|
||||
div_term = torch.exp(torch.arange(0, d_model//2, 2).float() * (-math.log(10000.0) / (d_model//2)))
|
||||
else: # a buggy implementation (for backward compatability only)
|
||||
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), persistent=False) # [1, C, H, W]
|
||||
|
||||
def forward(self, x):
|
||||
"""
|
||||
Args:
|
||||
x: [N, C, H, W]
|
||||
"""
|
||||
return x + self.pe[:, :, :x.size(2), :x.size(3)].to(x.device)
|
||||
@@ -0,0 +1,100 @@
|
||||
import copy
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from .linear_attention import LinearAttention, FullAttention
|
||||
|
||||
#Ref: https://github.com/zju3dv/LoFTR/blob/master/src/loftr/loftr_module/transformer.py
|
||||
class LoFTREncoderLayer(nn.Module):
|
||||
def __init__(self,
|
||||
d_model,
|
||||
nhead,
|
||||
attention='linear'):
|
||||
super(LoFTREncoderLayer, self).__init__()
|
||||
|
||||
self.dim = d_model // nhead
|
||||
self.nhead = nhead
|
||||
|
||||
# multi-head attention
|
||||
self.q_proj = nn.Linear(d_model, d_model, bias=False)
|
||||
self.k_proj = nn.Linear(d_model, d_model, bias=False)
|
||||
self.v_proj = nn.Linear(d_model, d_model, bias=False)
|
||||
self.attention = LinearAttention() if attention == 'linear' else FullAttention()
|
||||
self.merge = nn.Linear(d_model, d_model, bias=False)
|
||||
|
||||
# feed-forward network
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(d_model*2, d_model*2, bias=False),
|
||||
nn.ReLU(),
|
||||
nn.Linear(d_model*2, d_model, bias=False),
|
||||
)
|
||||
|
||||
# norm and dropout
|
||||
self.norm1 = nn.LayerNorm(d_model)
|
||||
self.norm2 = nn.LayerNorm(d_model)
|
||||
|
||||
def forward(self, x, source, x_mask=None, source_mask=None):
|
||||
"""
|
||||
Args:
|
||||
x (torch.Tensor): [N, L, C]
|
||||
source (torch.Tensor): [N, S, C]
|
||||
x_mask (torch.Tensor): [N, L] (optional)
|
||||
source_mask (torch.Tensor): [N, S] (optional)
|
||||
"""
|
||||
bs = x.size(0)
|
||||
query, key, value = x, source, source
|
||||
|
||||
# multi-head attention
|
||||
query = self.q_proj(query).view(bs, -1, self.nhead, self.dim) # [N, L, (H, D)]
|
||||
key = self.k_proj(key).view(bs, -1, self.nhead, self.dim) # [N, S, (H, D)]
|
||||
value = self.v_proj(value).view(bs, -1, self.nhead, self.dim)
|
||||
message = self.attention(query, key, value, q_mask=x_mask, kv_mask=source_mask) # [N, L, (H, D)]
|
||||
message = self.merge(message.view(bs, -1, self.nhead*self.dim)) # [N, L, C]
|
||||
message = self.norm1(message)
|
||||
|
||||
# feed-forward network
|
||||
message = self.mlp(torch.cat([x, message], dim=2))
|
||||
message = self.norm2(message)
|
||||
|
||||
return x + message
|
||||
|
||||
|
||||
class LocalFeatureTransformer(nn.Module):
|
||||
"""A Local Feature Transformer (LoFTR) module."""
|
||||
|
||||
def __init__(self, d_model, nhead, layer_names, attention):
|
||||
super(LocalFeatureTransformer, self).__init__()
|
||||
|
||||
self.d_model = d_model
|
||||
self.nhead = nhead
|
||||
self.layer_names = layer_names
|
||||
encoder_layer = LoFTREncoderLayer(d_model, nhead, attention)
|
||||
self.layers = nn.ModuleList([copy.deepcopy(encoder_layer) for _ in range(len(self.layer_names))])
|
||||
self._reset_parameters()
|
||||
|
||||
def _reset_parameters(self):
|
||||
for p in self.parameters():
|
||||
if p.dim() > 1:
|
||||
nn.init.xavier_uniform_(p)
|
||||
|
||||
def forward(self, feat0, feat1, mask0=None, mask1=None):
|
||||
"""
|
||||
Args:
|
||||
feat0 (torch.Tensor): [N, L, C]
|
||||
feat1 (torch.Tensor): [N, S, C]
|
||||
mask0 (torch.Tensor): [N, L] (optional)
|
||||
mask1 (torch.Tensor): [N, S] (optional)
|
||||
"""
|
||||
assert self.d_model == feat0.size(2), "the feature number of src and transformer must be equal"
|
||||
|
||||
for layer, name in zip(self.layers, self.layer_names):
|
||||
|
||||
if name == 'self':
|
||||
feat0 = layer(feat0, feat0, mask0, mask0)
|
||||
feat1 = layer(feat1, feat1, mask1, mask1)
|
||||
elif name == 'cross':
|
||||
feat0 = layer(feat0, feat1, mask0, mask1)
|
||||
feat1 = layer(feat1, feat0, mask1, mask0)
|
||||
else:
|
||||
raise KeyError
|
||||
|
||||
return feat0, feat1
|
||||
@@ -0,0 +1,148 @@
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .utils import bilinear_sampler, coords_grid, manual_pad
|
||||
|
||||
class AGCL:
|
||||
"""
|
||||
Implementation of Adaptive Group Correlation Layer (AGCL).
|
||||
"""
|
||||
|
||||
def __init__(self, fmap1, fmap2, att=None):
|
||||
self.fmap1 = fmap1
|
||||
self.fmap2 = fmap2
|
||||
|
||||
self.att = att
|
||||
|
||||
self.coords = coords_grid(fmap1.shape[0], fmap1.shape[2], fmap1.shape[3], fmap1.device)
|
||||
|
||||
def __call__(self, flow, extra_offset, small_patch=False, iter_mode=False):
|
||||
if iter_mode:
|
||||
corr = self.corr_iter(self.fmap1, self.fmap2, flow, small_patch)
|
||||
else:
|
||||
corr = self.corr_att_offset(
|
||||
self.fmap1, self.fmap2, flow, extra_offset, small_patch
|
||||
)
|
||||
return corr
|
||||
|
||||
def get_correlation(self, left_feature, right_feature, psize=(3, 3), dilate=(1, 1)):
|
||||
|
||||
N, C, H, W = left_feature.shape
|
||||
|
||||
di_y, di_x = dilate[0], dilate[1]
|
||||
pady, padx = psize[0] // 2 * di_y, psize[1] // 2 * di_x
|
||||
|
||||
right_pad = manual_pad(right_feature, pady, padx)
|
||||
|
||||
corr_list = []
|
||||
for h in range(0, pady * 2 + 1, di_y):
|
||||
for w in range(0, padx * 2 + 1, di_x):
|
||||
right_crop = right_pad[:, :, h : h + H, w : w + W]
|
||||
assert right_crop.shape == left_feature.shape
|
||||
corr = torch.mean(left_feature * right_crop, dim=1, keepdims=True)
|
||||
corr_list.append(corr)
|
||||
|
||||
corr_final = torch.cat(corr_list, dim=1)
|
||||
|
||||
return corr_final
|
||||
|
||||
def corr_iter(self, left_feature, right_feature, flow, small_patch):
|
||||
|
||||
coords = self.coords + flow
|
||||
coords = coords.permute(0, 2, 3, 1)
|
||||
right_feature = bilinear_sampler(right_feature, coords)
|
||||
|
||||
if small_patch:
|
||||
psize_list = [(3, 3), (3, 3), (3, 3), (3, 3)]
|
||||
dilate_list = [(1, 1), (1, 1), (1, 1), (1, 1)]
|
||||
else:
|
||||
psize_list = [(1, 9), (1, 9), (1, 9), (1, 9)]
|
||||
dilate_list = [(1, 1), (1, 1), (1, 1), (1, 1)]
|
||||
|
||||
N, C, H, W = left_feature.shape
|
||||
lefts = torch.split(left_feature, left_feature.shape[1]//4, dim=1)
|
||||
rights = torch.split(right_feature, right_feature.shape[1]//4, dim=1)
|
||||
|
||||
corrs = []
|
||||
for i in range(len(psize_list)):
|
||||
corr = self.get_correlation(
|
||||
lefts[i], rights[i], psize_list[i], dilate_list[i]
|
||||
)
|
||||
corrs.append(corr)
|
||||
|
||||
final_corr = torch.cat(corrs, dim=1)
|
||||
|
||||
return final_corr
|
||||
|
||||
def corr_att_offset(
|
||||
self, left_feature, right_feature, flow, extra_offset, small_patch
|
||||
):
|
||||
|
||||
N, C, H, W = left_feature.shape
|
||||
|
||||
if self.att is not None:
|
||||
left_feature = left_feature.permute(0, 2, 3, 1).reshape(N, H * W, C) # 'n c h w -> n (h w) c'
|
||||
right_feature = right_feature.permute(0, 2, 3, 1).reshape(N, H * W, C) # 'n c h w -> n (h w) c'
|
||||
# 'n (h w) c -> n c h w'
|
||||
left_feature, right_feature = self.att(left_feature, right_feature)
|
||||
# 'n (h w) c -> n c h w'
|
||||
left_feature, right_feature = [
|
||||
x.reshape(N, H, W, C).permute(0, 3, 1, 2)
|
||||
for x in [left_feature, right_feature]
|
||||
]
|
||||
|
||||
lefts = torch.split(left_feature, left_feature.shape[1]//4, dim=1)
|
||||
rights = torch.split(right_feature, right_feature.shape[1]//4, dim=1)
|
||||
|
||||
C = C // 4
|
||||
|
||||
if small_patch:
|
||||
psize_list = [(3, 3), (3, 3), (3, 3), (3, 3)]
|
||||
dilate_list = [(1, 1), (1, 1), (1, 1), (1, 1)]
|
||||
else:
|
||||
psize_list = [(1, 9), (1, 9), (1, 9), (1, 9)]
|
||||
dilate_list = [(1, 1), (1, 1), (1, 1), (1, 1)]
|
||||
|
||||
search_num = 9
|
||||
extra_offset = extra_offset.reshape(N, search_num, 2, H, W).permute(0, 1, 3, 4, 2) # [N, search_num, 1, 1, 2]
|
||||
|
||||
corrs = []
|
||||
for i in range(len(psize_list)):
|
||||
left_feature, right_feature = lefts[i], rights[i]
|
||||
psize, dilate = psize_list[i], dilate_list[i]
|
||||
|
||||
psizey, psizex = psize[0], psize[1]
|
||||
dilatey, dilatex = dilate[0], dilate[1]
|
||||
|
||||
ry = psizey // 2 * dilatey
|
||||
rx = psizex // 2 * dilatex
|
||||
x_grid, y_grid = torch.meshgrid(torch.arange(-rx, rx + 1, dilatex, device=self.fmap1.device),
|
||||
torch.arange(-ry, ry + 1, dilatey, device=self.fmap1.device), indexing='xy')
|
||||
|
||||
offsets = torch.stack((x_grid, y_grid))
|
||||
offsets = offsets.reshape(2, -1).permute(1, 0)
|
||||
for d in sorted((0, 2, 3)):
|
||||
offsets = offsets.unsqueeze(d)
|
||||
offsets = offsets.repeat_interleave(N, dim=0)
|
||||
offsets = offsets + extra_offset
|
||||
|
||||
coords = self.coords + flow # [N, 2, H, W]
|
||||
coords = coords.permute(0, 2, 3, 1) # [N, H, W, 2]
|
||||
coords = torch.unsqueeze(coords, 1) + offsets
|
||||
coords = coords.reshape(N, -1, W, 2) # [N, search_num*H, W, 2]
|
||||
|
||||
right_feature = bilinear_sampler(
|
||||
right_feature, coords
|
||||
) # [N, C, search_num*H, W]
|
||||
right_feature = right_feature.reshape(N, C, -1, H, W) # [N, C, search_num, H, W]
|
||||
left_feature = left_feature.unsqueeze(2).repeat_interleave(right_feature.shape[2], dim=2)
|
||||
|
||||
corr = torch.mean(left_feature * right_feature, dim=1)
|
||||
|
||||
corrs.append(corr)
|
||||
|
||||
final_corr = torch.cat(corrs, dim=1)
|
||||
|
||||
return final_corr
|
||||
@@ -0,0 +1,258 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .update import BasicUpdateBlock
|
||||
from .extractor import BasicEncoder
|
||||
from .corr import AGCL
|
||||
|
||||
from .attention import PositionEncodingSine, LocalFeatureTransformer
|
||||
|
||||
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
|
||||
|
||||
#Ref: https://github.com/princeton-vl/RAFT/blob/master/core/raft.py
|
||||
class CREStereo(nn.Module):
|
||||
def __init__(self, max_disp=192, mixed_precision=False, test_mode=False):
|
||||
super(CREStereo, self).__init__()
|
||||
|
||||
self.max_flow = max_disp
|
||||
self.mixed_precision = mixed_precision
|
||||
self.test_mode = test_mode
|
||||
|
||||
self.hidden_dim = 128
|
||||
self.context_dim = 128
|
||||
self.dropout = 0
|
||||
|
||||
self.fnet = BasicEncoder(output_dim=256, norm_fn='instance', dropout=self.dropout)
|
||||
self.update_block = BasicUpdateBlock(hidden_dim=self.hidden_dim, cor_planes=4 * 9, mask_size=4)
|
||||
|
||||
# loftr
|
||||
self.self_att_fn = LocalFeatureTransformer(
|
||||
d_model=256, nhead=8, layer_names=["self"] * 1, attention="linear"
|
||||
)
|
||||
self.cross_att_fn = LocalFeatureTransformer(
|
||||
d_model=256, nhead=8, layer_names=["cross"] * 1, attention="linear"
|
||||
)
|
||||
|
||||
# adaptive search
|
||||
self.search_num = 9
|
||||
self.conv_offset_16 = nn.Conv2d(
|
||||
256, self.search_num * 2, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
self.conv_offset_8 = nn.Conv2d(
|
||||
256, self.search_num * 2, kernel_size=3, stride=1, padding=1
|
||||
)
|
||||
self.range_16 = 1
|
||||
self.range_8 = 1
|
||||
|
||||
def freeze_bn(self):
|
||||
for m in self.modules():
|
||||
if isinstance(m, nn.BatchNorm2d):
|
||||
m.eval()
|
||||
|
||||
def convex_upsample(self, flow, mask, rate=4):
|
||||
""" Upsample flow field [H/8, W/8, 2] -> [H, W, 2] using convex combination """
|
||||
N, _, H, W = flow.shape
|
||||
# print(flow.shape, mask.shape, rate)
|
||||
mask = mask.view(N, 1, 9, rate, rate, H, W)
|
||||
mask = torch.softmax(mask, dim=2)
|
||||
|
||||
up_flow = F.unfold(rate * 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, rate*H, rate*W)
|
||||
|
||||
def zero_init(self, fmap):
|
||||
N, C, H, W = fmap.shape
|
||||
_x = torch.zeros([N, 1, H, W], dtype=torch.float32)
|
||||
_y = torch.zeros([N, 1, H, W], dtype=torch.float32)
|
||||
zero_flow = torch.cat((_x, _y), dim=1).to(fmap.device)
|
||||
return zero_flow
|
||||
|
||||
def forward(self, image1, image2, flow_init=None, iters=10, 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.mixed_precision):
|
||||
fmap1, fmap2 = self.fnet([image1, image2])
|
||||
|
||||
fmap1 = fmap1.float()
|
||||
fmap2 = fmap2.float()
|
||||
|
||||
with autocast(enabled=self.mixed_precision):
|
||||
|
||||
# 1/4 -> 1/8
|
||||
# feature
|
||||
fmap1_dw8 = F.avg_pool2d(fmap1, 2, stride=2)
|
||||
fmap2_dw8 = F.avg_pool2d(fmap2, 2, stride=2)
|
||||
|
||||
# offset
|
||||
offset_dw8 = self.conv_offset_8(fmap1_dw8)
|
||||
offset_dw8 = self.range_8 * (torch.sigmoid(offset_dw8) - 0.5) * 2.0
|
||||
|
||||
# context
|
||||
net, inp = torch.split(fmap1, [hdim,hdim], dim=1)
|
||||
net = torch.tanh(net)
|
||||
inp = F.relu(inp)
|
||||
net_dw8 = F.avg_pool2d(net, 2, stride=2)
|
||||
inp_dw8 = F.avg_pool2d(inp, 2, stride=2)
|
||||
|
||||
# 1/4 -> 1/16
|
||||
# feature
|
||||
fmap1_dw16 = F.avg_pool2d(fmap1, 4, stride=4)
|
||||
fmap2_dw16 = F.avg_pool2d(fmap2, 4, stride=4)
|
||||
offset_dw16 = self.conv_offset_16(fmap1_dw16)
|
||||
offset_dw16 = self.range_16 * (torch.sigmoid(offset_dw16) - 0.5) * 2.0
|
||||
|
||||
# context
|
||||
net_dw16 = F.avg_pool2d(net, 4, stride=4)
|
||||
inp_dw16 = F.avg_pool2d(inp, 4, stride=4)
|
||||
|
||||
# positional encoding and self-attention
|
||||
pos_encoding_fn_small = PositionEncodingSine(
|
||||
d_model=256, max_shape=(image1.shape[2] // 16, image1.shape[3] // 16)
|
||||
)
|
||||
# 'n c h w -> n (h w) c'
|
||||
x_tmp = pos_encoding_fn_small(fmap1_dw16)
|
||||
fmap1_dw16 = x_tmp.permute(0, 2, 3, 1).reshape(x_tmp.shape[0], x_tmp.shape[2] * x_tmp.shape[3], x_tmp.shape[1])
|
||||
# 'n c h w -> n (h w) c'
|
||||
x_tmp = pos_encoding_fn_small(fmap2_dw16)
|
||||
fmap2_dw16 = x_tmp.permute(0, 2, 3, 1).reshape(x_tmp.shape[0], x_tmp.shape[2] * x_tmp.shape[3], x_tmp.shape[1])
|
||||
|
||||
fmap1_dw16, fmap2_dw16 = self.self_att_fn(fmap1_dw16, fmap2_dw16)
|
||||
fmap1_dw16, fmap2_dw16 = [
|
||||
x.reshape(x.shape[0], image1.shape[2] // 16, -1, x.shape[2]).permute(0, 3, 1, 2)
|
||||
for x in [fmap1_dw16, fmap2_dw16]
|
||||
]
|
||||
|
||||
corr_fn = AGCL(fmap1, fmap2)
|
||||
corr_fn_dw8 = AGCL(fmap1_dw8, fmap2_dw8)
|
||||
corr_fn_att_dw16 = AGCL(fmap1_dw16, fmap2_dw16, att=self.cross_att_fn)
|
||||
|
||||
# Cascaded refinement (1/16 + 1/8 + 1/4)
|
||||
predictions = []
|
||||
flow = None
|
||||
flow_up = None
|
||||
if flow_init is not None:
|
||||
scale = fmap1.shape[2] / flow_init.shape[2]
|
||||
flow = -scale * F.interpolate(
|
||||
flow_init,
|
||||
size=(fmap1.shape[2], fmap1.shape[3]),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
else:
|
||||
# zero initialization
|
||||
flow_dw16 = self.zero_init(fmap1_dw16)
|
||||
|
||||
# Recurrent Update Module
|
||||
# RUM: 1/16
|
||||
for itr in range(iters // 2):
|
||||
if itr % 2 == 0:
|
||||
small_patch = False
|
||||
else:
|
||||
small_patch = True
|
||||
|
||||
flow_dw16 = flow_dw16.detach()
|
||||
out_corrs = corr_fn_att_dw16(
|
||||
flow_dw16, offset_dw16, small_patch=small_patch
|
||||
)
|
||||
|
||||
with autocast(enabled=self.mixed_precision):
|
||||
net_dw16, up_mask, delta_flow = self.update_block(
|
||||
net_dw16, inp_dw16, out_corrs, flow_dw16
|
||||
)
|
||||
|
||||
flow_dw16 = flow_dw16 + delta_flow
|
||||
flow = self.convex_upsample(flow_dw16, up_mask, rate=4)
|
||||
flow_up = -4 * F.interpolate(
|
||||
flow,
|
||||
size=(4 * flow.shape[2], 4 * flow.shape[3]),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
predictions.append(flow_up)
|
||||
|
||||
scale = fmap1_dw8.shape[2] / flow.shape[2]
|
||||
flow_dw8 = -scale * F.interpolate(
|
||||
flow,
|
||||
size=(fmap1_dw8.shape[2], fmap1_dw8.shape[3]),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
|
||||
# RUM: 1/8
|
||||
for itr in range(iters // 2):
|
||||
if itr % 2 == 0:
|
||||
small_patch = False
|
||||
else:
|
||||
small_patch = True
|
||||
|
||||
flow_dw8 = flow_dw8.detach()
|
||||
out_corrs = corr_fn_dw8(flow_dw8, offset_dw8, small_patch=small_patch)
|
||||
|
||||
with autocast(enabled=self.mixed_precision):
|
||||
net_dw8, up_mask, delta_flow = self.update_block(
|
||||
net_dw8, inp_dw8, out_corrs, flow_dw8
|
||||
)
|
||||
|
||||
flow_dw8 = flow_dw8 + delta_flow
|
||||
flow = self.convex_upsample(flow_dw8, up_mask, rate=4)
|
||||
flow_up = -2 * F.interpolate(
|
||||
flow,
|
||||
size=(2 * flow.shape[2], 2 * flow.shape[3]),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
predictions.append(flow_up)
|
||||
|
||||
scale = fmap1.shape[2] / flow.shape[2]
|
||||
flow = -scale * F.interpolate(
|
||||
flow,
|
||||
size=(fmap1.shape[2], fmap1.shape[3]),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
|
||||
# RUM: 1/4
|
||||
for itr in range(iters):
|
||||
if itr % 2 == 0:
|
||||
small_patch = False
|
||||
else:
|
||||
small_patch = True
|
||||
|
||||
flow = flow.detach()
|
||||
out_corrs = corr_fn(flow, None, small_patch=small_patch, iter_mode=True)
|
||||
|
||||
with autocast(enabled=self.mixed_precision):
|
||||
net, up_mask, delta_flow = self.update_block(net, inp, out_corrs, flow)
|
||||
|
||||
flow = flow + delta_flow
|
||||
flow_up = -self.convex_upsample(flow, up_mask, rate=4)
|
||||
predictions.append(flow_up)
|
||||
|
||||
if self.test_mode:
|
||||
return flow_up
|
||||
|
||||
return predictions
|
||||
@@ -0,0 +1,123 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
# Ref: https://github.com/princeton-vl/RAFT/blob/master/core/extractor.py
|
||||
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)
|
||||
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)
|
||||
self.norm3 = nn.BatchNorm2d(planes)
|
||||
|
||||
elif norm_fn == 'instance':
|
||||
self.norm1 = nn.InstanceNorm2d(planes, affine=False)
|
||||
self.norm2 = nn.InstanceNorm2d(planes, affine=False)
|
||||
self.norm3 = nn.InstanceNorm2d(planes, affine=False)
|
||||
|
||||
elif norm_fn == 'none':
|
||||
self.norm1 = nn.Sequential()
|
||||
self.norm2 = nn.Sequential()
|
||||
self.norm3 = nn.Sequential()
|
||||
|
||||
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)))
|
||||
|
||||
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, affine=False)
|
||||
|
||||
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=1)
|
||||
|
||||
# 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.dropout is not None:
|
||||
x = self.dropout(x)
|
||||
|
||||
if is_list:
|
||||
x = torch.split(x, x.shape[0]//2, dim=0)
|
||||
|
||||
return x
|
||||
@@ -0,0 +1,91 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
#Ref: https://github.com/princeton-vl/RAFT/blob/master/core/update.py
|
||||
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 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, cor_planes):
|
||||
super(BasicMotionEncoder, self).__init__()
|
||||
|
||||
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, hidden_dim, cor_planes, mask_size=8):
|
||||
super(BasicUpdateBlock, self).__init__()
|
||||
|
||||
self.encoder = BasicMotionEncoder(cor_planes)
|
||||
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, mask_size**2 *9, 1, padding=0))
|
||||
|
||||
def forward(self, net, inp, corr, flow, upsample=True):
|
||||
# print(inp.shape, corr.shape, flow.shape)
|
||||
motion_features = self.encoder(flow, corr)
|
||||
# print(motion_features.shape, inp.shape)
|
||||
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 = .25 * self.mask(net)
|
||||
return net, mask, delta_flow
|
||||
@@ -0,0 +1 @@
|
||||
from .utils import bilinear_sampler, coords_grid, manual_pad
|
||||
@@ -0,0 +1,108 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
|
||||
#Ref: https://github.com/princeton-vl/RAFT/blob/master/core/utils/utils.py
|
||||
|
||||
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)
|
||||
img = bilinear_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), indexing='ij')
|
||||
coords = torch.stack(coords[::-1], dim=0).float()
|
||||
return coords[None].repeat(batch, 1, 1, 1)
|
||||
|
||||
def manual_pad(x, pady, padx):
|
||||
|
||||
pad = (padx, padx, pady, pady)
|
||||
return F.pad(x.clone().detach(), pad, "replicate")
|
||||
|
||||
# Ref: https://zenn.dev/pinto0309/scraps/7d4032067d0160
|
||||
def bilinear_grid_sample(im, grid, align_corners=False):
|
||||
"""Given an input and a flow-field grid, computes the output using input
|
||||
values and pixel locations from grid. Supported only bilinear interpolation
|
||||
method to sample the input pixels.
|
||||
|
||||
Args:
|
||||
im (torch.Tensor): Input feature map, shape (N, C, H, W)
|
||||
grid (torch.Tensor): Point coordinates, shape (N, Hg, Wg, 2)
|
||||
align_corners {bool}: If set to True, the extrema (-1 and 1) are
|
||||
considered as referring to the center points of the input’s
|
||||
corner pixels. If set to False, they are instead considered as
|
||||
referring to the corner points of the input’s corner pixels,
|
||||
making the sampling more resolution agnostic.
|
||||
|
||||
Returns:
|
||||
torch.Tensor: A tensor with sampled points, shape (N, C, Hg, Wg)
|
||||
"""
|
||||
n, c, h, w = im.shape
|
||||
gn, gh, gw, _ = grid.shape
|
||||
assert n == gn
|
||||
|
||||
x = grid[:, :, :, 0]
|
||||
y = grid[:, :, :, 1]
|
||||
|
||||
if align_corners:
|
||||
x = ((x + 1) / 2) * (w - 1)
|
||||
y = ((y + 1) / 2) * (h - 1)
|
||||
else:
|
||||
x = ((x + 1) * w - 1) / 2
|
||||
y = ((y + 1) * h - 1) / 2
|
||||
|
||||
x = x.view(n, -1)
|
||||
y = y.view(n, -1)
|
||||
|
||||
x0 = torch.floor(x).long()
|
||||
y0 = torch.floor(y).long()
|
||||
x1 = x0 + 1
|
||||
y1 = y0 + 1
|
||||
|
||||
wa = ((x1 - x) * (y1 - y)).unsqueeze(1)
|
||||
wb = ((x1 - x) * (y - y0)).unsqueeze(1)
|
||||
wc = ((x - x0) * (y1 - y)).unsqueeze(1)
|
||||
wd = ((x - x0) * (y - y0)).unsqueeze(1)
|
||||
|
||||
# Apply default for grid_sample function zero padding
|
||||
im_padded = torch.nn.functional.pad(im, pad=[1, 1, 1, 1], mode='constant', value=0)
|
||||
padded_h = h + 2
|
||||
padded_w = w + 2
|
||||
# save points positions after padding
|
||||
x0, x1, y0, y1 = x0 + 1, x1 + 1, y0 + 1, y1 + 1
|
||||
|
||||
# Clip coordinates to padded image size
|
||||
x0 = torch.where(x0 < 0, torch.tensor(0, device=im.device), x0)
|
||||
x0 = torch.where(x0 > padded_w - 1, torch.tensor(padded_w - 1, device=im.device), x0)
|
||||
x1 = torch.where(x1 < 0, torch.tensor(0, device=im.device), x1)
|
||||
x1 = torch.where(x1 > padded_w - 1, torch.tensor(padded_w - 1, device=im.device), x1)
|
||||
y0 = torch.where(y0 < 0, torch.tensor(0, device=im.device), y0)
|
||||
y0 = torch.where(y0 > padded_h - 1, torch.tensor(padded_h - 1, device=im.device), y0)
|
||||
y1 = torch.where(y1 < 0, torch.tensor(0, device=im.device), y1)
|
||||
y1 = torch.where(y1 > padded_h - 1, torch.tensor(padded_h - 1, device=im.device), y1)
|
||||
|
||||
im_padded = im_padded.view(n, c, -1)
|
||||
|
||||
x0_y0 = (x0 + y0 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x0_y1 = (x0 + y1 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x1_y0 = (x1 + y0 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
x1_y1 = (x1 + y1 * padded_w).unsqueeze(1).expand(-1, c, -1)
|
||||
|
||||
Ia = torch.gather(im_padded, 2, x0_y0)
|
||||
Ib = torch.gather(im_padded, 2, x0_y1)
|
||||
Ic = torch.gather(im_padded, 2, x1_y0)
|
||||
Id = torch.gather(im_padded, 2, x1_y1)
|
||||
|
||||
return (Ia * wa + Ib * wb + Ic * wc + Id * wd).reshape(n, c, gh, gw)
|
||||
@@ -0,0 +1,82 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
import cv2
|
||||
from imread_from_url import imread_from_url
|
||||
|
||||
from nets import Model
|
||||
|
||||
device = 'cuda'
|
||||
|
||||
#Ref: https://github.com/megvii-research/CREStereo/blob/master/test.py
|
||||
def inference(left, right, model, n_iter=20):
|
||||
|
||||
print("Model Forwarding...")
|
||||
imgL = left.transpose(2, 0, 1)
|
||||
imgR = right.transpose(2, 0, 1)
|
||||
imgL = np.ascontiguousarray(imgL[None, :, :, :])
|
||||
imgR = np.ascontiguousarray(imgR[None, :, :, :])
|
||||
|
||||
imgL = torch.tensor(imgL.astype("float32")).to(device)
|
||||
imgR = torch.tensor(imgR.astype("float32")).to(device)
|
||||
|
||||
imgL_dw2 = F.interpolate(
|
||||
imgL,
|
||||
size=(imgL.shape[2] // 2, imgL.shape[3] // 2),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
imgR_dw2 = F.interpolate(
|
||||
imgR,
|
||||
size=(imgL.shape[2] // 2, imgL.shape[3] // 2),
|
||||
mode="bilinear",
|
||||
align_corners=True,
|
||||
)
|
||||
# print(imgR_dw2.shape)
|
||||
with torch.inference_mode():
|
||||
pred_flow_dw2 = model(imgL_dw2, imgR_dw2, iters=n_iter, flow_init=None)
|
||||
|
||||
pred_flow = model(imgL, imgR, iters=n_iter, flow_init=pred_flow_dw2)
|
||||
pred_disp = torch.squeeze(pred_flow[:, 0, :, :]).cpu().detach().numpy()
|
||||
|
||||
return pred_disp
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
left_img = imread_from_url("https://raw.githubusercontent.com/megvii-research/CREStereo/master/img/test/left.png")
|
||||
right_img = imread_from_url("https://raw.githubusercontent.com/megvii-research/CREStereo/master/img/test/right.png")
|
||||
|
||||
in_h, in_w = left_img.shape[:2]
|
||||
|
||||
# Resize image in case the GPU memory overflows
|
||||
eval_h, eval_w = (in_h,in_w)
|
||||
assert eval_h%8 == 0, "input height should be divisible by 8"
|
||||
assert eval_w%8 == 0, "input width should be divisible by 8"
|
||||
|
||||
imgL = cv2.resize(left_img, (eval_w, eval_h), interpolation=cv2.INTER_LINEAR)
|
||||
imgR = cv2.resize(right_img, (eval_w, eval_h), interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
model_path = "models/crestereo_eth3d.pth"
|
||||
|
||||
model = Model(max_disp=256, mixed_precision=False, test_mode=True)
|
||||
model.load_state_dict(torch.load(model_path), strict=True)
|
||||
model.to(device)
|
||||
model.eval()
|
||||
|
||||
pred = inference(imgL, imgR, model, n_iter=20)
|
||||
|
||||
t = float(in_w) / float(eval_w)
|
||||
disp = cv2.resize(pred, (in_w, in_h), interpolation=cv2.INTER_LINEAR) * t
|
||||
|
||||
disp_vis = (disp - disp.min()) / (disp.max() - disp.min()) * 255.0
|
||||
disp_vis = disp_vis.astype("uint8")
|
||||
disp_vis = cv2.applyColorMap(disp_vis, cv2.COLORMAP_INFERNO)
|
||||
|
||||
combined_img = np.hstack((left_img, disp_vis))
|
||||
cv2.namedWindow("output", cv2.WINDOW_NORMAL)
|
||||
cv2.imshow("output", combined_img)
|
||||
cv2.imwrite("output.jpg", disp_vis)
|
||||
cv2.waitKey(0)
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
import numpy as np
|
||||
import cv2
|
||||
|
||||
import onnxruntime
|
||||
|
||||
# Ref: https://github.com/megvii-research/CREStereo/blob/master/test.py
|
||||
def inference(left, right, model, no_flow_model):
|
||||
# Get onnx model layer names (see convert_to_onnx.py for what these are)
|
||||
input1_name = model.get_inputs()[0].name
|
||||
input2_name = model.get_inputs()[1].name
|
||||
input3_name = model.get_inputs()[2].name
|
||||
output_name = model.get_outputs()[0].name
|
||||
|
||||
# Decimate the image to half the original size for flow estimation network
|
||||
imgL_dw2 = cv2.resize(
|
||||
left, (left.shape[1] // 2, left.shape[0] // 2), interpolation=cv2.INTER_LINEAR)
|
||||
imgR_dw2 = cv2.resize(
|
||||
right, (right.shape[1] // 2, right.shape[0] // 2), interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
# Reshape inputs to match what is expected
|
||||
imgL = left.transpose(2, 0, 1)
|
||||
imgR = right.transpose(2, 0, 1)
|
||||
imgL = np.ascontiguousarray(imgL[None, :, :, :]).astype("float32")
|
||||
imgR = np.ascontiguousarray(imgR[None, :, :, :]).astype("float32")
|
||||
|
||||
imgL_dw2 = imgL_dw2.transpose(2, 0, 1)
|
||||
imgR_dw2 = imgR_dw2.transpose(2, 0, 1)
|
||||
imgL_dw2 = np.ascontiguousarray(imgL_dw2[None, :, :, :]).astype("float32")
|
||||
imgR_dw2 = np.ascontiguousarray(imgR_dw2[None, :, :, :]).astype("float32")
|
||||
|
||||
print("Model Forwarding...")
|
||||
# First pass it just to get the flow
|
||||
pred_flow_dw2 = no_flow_model.run(
|
||||
[output_name], {input1_name: imgL_dw2, input2_name: imgR_dw2})[0]
|
||||
# Second pass gets us the disparity
|
||||
pred_disp = model.run([output_name], {
|
||||
input1_name: imgL, input2_name: imgR, input3_name: pred_flow_dw2})[0]
|
||||
|
||||
return np.squeeze(pred_disp[:, 0, :, :])
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
left_img = cv2.imread("left.png")
|
||||
right_img = cv2.imread("right.png")
|
||||
|
||||
in_h, in_w = left_img.shape[:2]
|
||||
|
||||
# Resize images
|
||||
eval_h, eval_w = (in_h, in_w)
|
||||
assert eval_h % 8 == 0, "input height should be divisible by 8"
|
||||
assert eval_w % 8 == 0, "input width should be divisible by 8"
|
||||
|
||||
imgL = cv2.resize(left_img, (eval_w, eval_h),
|
||||
interpolation=cv2.INTER_LINEAR)
|
||||
imgR = cv2.resize(right_img, (eval_w, eval_h),
|
||||
interpolation=cv2.INTER_LINEAR)
|
||||
|
||||
no_flow_model_path = "models/crestereo_without_flow.onnx"
|
||||
model_path = "models/crestereo.onnx"
|
||||
|
||||
model = onnxruntime.InferenceSession(model_path)
|
||||
no_flow_model = onnxruntime.InferenceSession(no_flow_model_path)
|
||||
|
||||
pred = inference(imgL, imgR, model, no_flow_model)
|
||||
|
||||
t = float(in_w) / float(eval_w)
|
||||
disp = cv2.resize(pred, (eval_w, eval_h),
|
||||
interpolation=cv2.INTER_LINEAR) * t
|
||||
disp_vis = (disp - disp.min()) / (disp.max() - disp.min()) * 255.0
|
||||
disp_vis = disp_vis.astype("uint8")
|
||||
disp_vis = cv2.applyColorMap(disp_vis, cv2.COLORMAP_INFERNO)
|
||||
|
||||
combined_img = np.hstack((left_img, disp_vis))
|
||||
cv2.namedWindow("output", cv2.WINDOW_NORMAL)
|
||||
cv2.imshow("output", combined_img)
|
||||
cv2.imwrite("output.jpg", disp_vis)
|
||||
cv2.waitKey(0)
|
||||
@@ -0,0 +1,492 @@
|
||||
import argparse
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
import time
|
||||
import logging
|
||||
from collections import namedtuple
|
||||
from itertools import repeat
|
||||
|
||||
import yaml
|
||||
from tensorboardX import SummaryWriter
|
||||
|
||||
from nets import Model
|
||||
from dataset import CREStereoDataset
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.optim as optim
|
||||
import torch.backends.cudnn as cudnn
|
||||
import torch.distributed as dist
|
||||
from torch.nn.parallel import DistributedDataParallel as DDP
|
||||
from torch.utils.data import DataLoader, RandomSampler
|
||||
|
||||
|
||||
def parse_yaml(file_path: str) -> namedtuple:
|
||||
"""Parse yaml configuration file and return the object in `namedtuple`."""
|
||||
with open(file_path, "rb") as f:
|
||||
cfg: dict = yaml.safe_load(f)
|
||||
args = namedtuple("train_args", cfg.keys())(*cfg.values())
|
||||
# save cfg into train_log
|
||||
ensure_dir(args.log_dir)
|
||||
dst_file = os.path.join(args.log_dir, file_path.split('/')[-1])
|
||||
shutil.copy2(file_path, dst_file)
|
||||
return args
|
||||
|
||||
|
||||
def format_time(elapse):
|
||||
elapse = int(elapse)
|
||||
hour = elapse // 3600
|
||||
minute = elapse % 3600 // 60
|
||||
seconds = elapse % 60
|
||||
return "{:02d}:{:02d}:{:02d}".format(hour, minute, seconds)
|
||||
|
||||
|
||||
def ensure_dir(path):
|
||||
if not os.path.exists(path):
|
||||
os.makedirs(path, exist_ok=True)
|
||||
|
||||
|
||||
def adjust_learning_rate(optimizer, epoch):
|
||||
|
||||
warm_up = 0.02
|
||||
const_range = 0.6
|
||||
min_lr_rate = 0.05
|
||||
|
||||
if epoch <= args.n_total_epoch * warm_up:
|
||||
lr = (1 - min_lr_rate) * args.base_lr / (
|
||||
args.n_total_epoch * warm_up
|
||||
) * epoch + min_lr_rate * args.base_lr
|
||||
elif args.n_total_epoch * warm_up < epoch <= args.n_total_epoch * const_range:
|
||||
lr = args.base_lr
|
||||
else:
|
||||
lr = (min_lr_rate - 1) * args.base_lr / (
|
||||
(1 - const_range) * args.n_total_epoch
|
||||
) * epoch + (1 - min_lr_rate * const_range) / (1 - const_range) * args.base_lr
|
||||
|
||||
for param_group in optimizer.param_groups:
|
||||
param_group['lr'] = lr
|
||||
|
||||
def sequence_loss(flow_preds, flow_gt, valid, gamma=0.8):
|
||||
'''
|
||||
valid: (2, 384, 512) (B, H, W) -> (B, 1, H, W)
|
||||
flow_preds[0]: (B, 2, H, W)
|
||||
flow_gt: (B, 2, H, W)
|
||||
'''
|
||||
n_predictions = len(flow_preds)
|
||||
flow_loss = 0.0
|
||||
for i in range(n_predictions):
|
||||
i_weight = gamma ** (n_predictions - i - 1)
|
||||
i_loss = torch.abs(flow_preds[i] - flow_gt)
|
||||
flow_loss += i_weight * (valid.unsqueeze(1) * i_loss).mean()
|
||||
|
||||
return flow_loss
|
||||
|
||||
def repeater(data_loader):
|
||||
for loader in repeat(data_loader):
|
||||
for data in loader:
|
||||
yield data
|
||||
|
||||
def train_dist(args, world_size):
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--local_rank",type=int)
|
||||
FLAGS = parser.parse_args()
|
||||
local_rank = FLAGS.local_rank
|
||||
# directory check
|
||||
log_model_dir = os.path.join(args.log_dir, "models")
|
||||
ensure_dir(log_model_dir)
|
||||
|
||||
# distributed init and model / optimizer
|
||||
torch.cuda.set_device(local_rank)
|
||||
dist.init_process_group(backend='nccl') # nccl is highly recommanded
|
||||
model = Model(
|
||||
max_disp=args.max_disp, mixed_precision=args.mixed_precision, test_mode=False
|
||||
)
|
||||
# sync batch norm
|
||||
model = torch.nn.SyncBatchNorm.convert_sync_batchnorm(model).to(local_rank)
|
||||
model = DDP(model, device_ids=[local_rank], output_device=local_rank)
|
||||
optimizer = optim.Adam(model.parameters(), lr=0.1, betas=(0.9, 0.999))
|
||||
|
||||
if dist.get_rank() == 0:
|
||||
# tensorboard
|
||||
tb_log = SummaryWriter(os.path.join(args.log_dir, "train.events"))
|
||||
|
||||
# worklog
|
||||
logging.basicConfig(level=eval(args.log_level))
|
||||
worklog = logging.getLogger("train_logger")
|
||||
worklog.propagate = False
|
||||
fileHandler = logging.FileHandler(
|
||||
os.path.join(args.log_dir, "worklog.txt"), mode="a", encoding="utf8"
|
||||
)
|
||||
formatter = logging.Formatter(
|
||||
fmt="%(asctime)s %(message)s", datefmt="%Y/%m/%d %H:%M:%S"
|
||||
)
|
||||
fileHandler.setFormatter(formatter)
|
||||
consoleHandler = logging.StreamHandler(sys.stdout)
|
||||
formatter = logging.Formatter(
|
||||
fmt="\x1b[32m%(asctime)s\x1b[0m %(message)s", datefmt="%Y/%m/%d %H:%M:%S"
|
||||
)
|
||||
consoleHandler.setFormatter(formatter)
|
||||
worklog.handlers = [fileHandler, consoleHandler]
|
||||
|
||||
# params stat
|
||||
worklog.info(f"Use {world_size} GPU(s)")
|
||||
worklog.info("Params: %s" % sum([p.numel() for p in model.parameters()]))
|
||||
|
||||
# load pretrained model if exist
|
||||
chk_path = os.path.join(log_model_dir, "latest.pth")
|
||||
if args.loadmodel is not None:
|
||||
chk_path = args.loadmodel
|
||||
elif not os.path.exists(chk_path):
|
||||
chk_path = None
|
||||
|
||||
if chk_path is not None:
|
||||
if dist.get_rank() == 0:
|
||||
worklog.info(f"loading model: {chk_path}")
|
||||
# map_location=torch.device('cpu') make more balance memory usage
|
||||
state_dict = torch.load(chk_path, map_location=torch.device('cpu'))
|
||||
model.module.load_state_dict(state_dict['state_dict'])
|
||||
optimizer.load_state_dict(state_dict['optim_state_dict'])
|
||||
resume_epoch_idx = state_dict["epoch"]
|
||||
resume_iters = state_dict["iters"]
|
||||
start_epoch_idx = resume_epoch_idx + 1
|
||||
start_iters = resume_iters
|
||||
else:
|
||||
start_epoch_idx = 1
|
||||
start_iters = 0
|
||||
|
||||
# datasets
|
||||
dataset = CREStereoDataset(args.training_data_path)
|
||||
# dataset = MixDataset("train",
|
||||
# data_path=args.data["train"]["data_path"],
|
||||
# fields=args.data["train"]["fields"],
|
||||
# filelists=args.data["train"]["filelists"],
|
||||
# input_size=(args.data["train"]["input_size"][0], args.data["train"]["input_size"][1]))
|
||||
if dist.get_rank() == 0:
|
||||
worklog.info(f"Dataset size: {len(dataset)}")
|
||||
train_sampler = torch.utils.data.distributed.DistributedSampler(dataset)
|
||||
dataloader = torch.utils.data.DataLoader(dataset,
|
||||
batch_size=args.batch_size, num_workers=4, sampler=train_sampler)
|
||||
|
||||
# counter
|
||||
cur_iters = start_iters
|
||||
total_iters = args.minibatch_per_epoch * args.n_total_epoch
|
||||
t0 = time.perf_counter()
|
||||
for epoch_idx in range(start_epoch_idx, args.n_total_epoch + 1):
|
||||
dataloader.sampler.set_epoch(epoch_idx)
|
||||
# adjust learning rate
|
||||
epoch_total_train_loss = 0
|
||||
adjust_learning_rate(optimizer, epoch_idx)
|
||||
model.train()
|
||||
|
||||
t1 = time.perf_counter()
|
||||
|
||||
for batch_idx, mini_batch_data in enumerate(dataloader):
|
||||
|
||||
if batch_idx % args.minibatch_per_epoch == 0 and batch_idx != 0:
|
||||
break
|
||||
cur_iters += 1
|
||||
|
||||
# parse data
|
||||
left, right, gt_disp, valid_mask = (
|
||||
mini_batch_data["left"].to(local_rank),
|
||||
mini_batch_data["right"].to(local_rank),
|
||||
mini_batch_data["disparity"].to(local_rank),
|
||||
mini_batch_data["mask"].to(local_rank),
|
||||
)
|
||||
|
||||
t2 = time.perf_counter()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# pre-process
|
||||
gt_disp = torch.unsqueeze(gt_disp, dim=1) # [2, 384, 512] -> [2, 1, 384, 512]
|
||||
gt_flow = torch.cat([gt_disp, gt_disp * 0], dim=1) # [2, 2, 384, 512]
|
||||
|
||||
# forward
|
||||
flow_predictions = model(left, right)
|
||||
|
||||
# loss & backword
|
||||
loss = sequence_loss(
|
||||
flow_predictions, gt_flow, valid_mask, gamma=0.8
|
||||
).to(local_rank)
|
||||
|
||||
# loss stats
|
||||
loss_item = loss.data.item()
|
||||
epoch_total_train_loss += loss_item
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
t3 = time.perf_counter()
|
||||
if dist.get_rank() == 0:
|
||||
if cur_iters % 10 == 0:
|
||||
tdata = t2 - t1
|
||||
time_train_passed = t3 - t0
|
||||
time_iter_passed = t3 - t1
|
||||
step_passed = cur_iters - start_iters
|
||||
eta = (
|
||||
(total_iters - cur_iters)
|
||||
/ max(step_passed, 1e-7)
|
||||
* time_train_passed
|
||||
)
|
||||
|
||||
meta_info = list()
|
||||
meta_info.append("{:.2g} b/s".format(1.0 / time_iter_passed))
|
||||
meta_info.append("passed:{}".format(format_time(time_train_passed)))
|
||||
meta_info.append("eta:{}".format(format_time(eta)))
|
||||
meta_info.append(
|
||||
"data_time:{:.2g}".format(tdata / time_iter_passed)
|
||||
)
|
||||
|
||||
meta_info.append(
|
||||
"lr:{:.5g}".format(optimizer.param_groups[0]["lr"])
|
||||
)
|
||||
meta_info.append(
|
||||
"[{}/{}:{}/{}]".format(
|
||||
epoch_idx,
|
||||
args.n_total_epoch,
|
||||
batch_idx,
|
||||
args.minibatch_per_epoch,
|
||||
)
|
||||
)
|
||||
loss_info = list()
|
||||
loss_info.append("{}:{:.4g}".format("total_loss", loss_item))
|
||||
# exp_name = ['\n' + os.path.basename(os.getcwd())]
|
||||
|
||||
info = [",".join(meta_info+loss_info)]
|
||||
worklog.info("".join(info))
|
||||
|
||||
# minibatch loss
|
||||
tb_log.add_scalar("train/loss_batch", loss_item, cur_iters)
|
||||
tb_log.add_scalar(
|
||||
"train/lr", optimizer.param_groups[0]["lr"], cur_iters
|
||||
)
|
||||
tb_log.flush()
|
||||
|
||||
t1 = time.perf_counter()
|
||||
|
||||
if dist.get_rank() == 0:
|
||||
# epoch loss
|
||||
tb_log.add_scalar(
|
||||
"train/loss",
|
||||
epoch_total_train_loss / args.minibatch_per_epoch,
|
||||
epoch_idx,
|
||||
)
|
||||
tb_log.flush()
|
||||
|
||||
# save model params
|
||||
ckp_data = {
|
||||
"epoch": epoch_idx,
|
||||
"iters": cur_iters,
|
||||
"batch_size": args.batch_size * world_size,
|
||||
"epoch_size": args.minibatch_per_epoch,
|
||||
"train_loss": epoch_total_train_loss / args.minibatch_per_epoch,
|
||||
"state_dict": model.module.state_dict(),
|
||||
"optim_state_dict": optimizer.state_dict(),
|
||||
}
|
||||
torch.save(ckp_data, os.path.join(log_model_dir, "latest.pth"))
|
||||
if epoch_idx % args.model_save_freq_epoch == 0:
|
||||
save_path = os.path.join(log_model_dir, "epoch-%d.pth" % epoch_idx)
|
||||
worklog.info(f"Model params saved: {save_path}")
|
||||
torch.save(ckp_data, save_path)
|
||||
if dist.get_rank() == 0:
|
||||
worklog.info("Training is done, exit.")
|
||||
|
||||
def train(args, world_size):
|
||||
# directory check
|
||||
log_model_dir = os.path.join(args.log_dir, "models")
|
||||
ensure_dir(log_model_dir)
|
||||
|
||||
# model / optimizer
|
||||
model = Model(
|
||||
max_disp=args.max_disp, mixed_precision=args.mixed_precision, test_mode=False
|
||||
)
|
||||
model = nn.DataParallel(model,device_ids=[i for i in range(world_size)])
|
||||
model.cuda()
|
||||
optimizer = optim.Adam(model.parameters(), lr=0.1, betas=(0.9, 0.999))
|
||||
|
||||
tb_log = SummaryWriter(os.path.join(args.log_dir, "train.events"))
|
||||
|
||||
# worklog
|
||||
logging.basicConfig(level=eval(args.log_level))
|
||||
worklog = logging.getLogger("train_logger")
|
||||
worklog.propagate = False
|
||||
fileHandler = logging.FileHandler(
|
||||
os.path.join(args.log_dir, "worklog.txt"), mode="a", encoding="utf8"
|
||||
)
|
||||
formatter = logging.Formatter(
|
||||
fmt="%(asctime)s %(message)s", datefmt="%Y/%m/%d %H:%M:%S"
|
||||
)
|
||||
fileHandler.setFormatter(formatter)
|
||||
consoleHandler = logging.StreamHandler(sys.stdout)
|
||||
formatter = logging.Formatter(
|
||||
fmt="\x1b[32m%(asctime)s\x1b[0m %(message)s", datefmt="%Y/%m/%d %H:%M:%S"
|
||||
)
|
||||
consoleHandler.setFormatter(formatter)
|
||||
worklog.handlers = [fileHandler, consoleHandler]
|
||||
|
||||
# params stat
|
||||
worklog.info(f"Use {world_size} GPU(s)")
|
||||
worklog.info("Params: %s" % sum([p.numel() for p in model.parameters()]))
|
||||
|
||||
# load pretrained model if exist
|
||||
chk_path = os.path.join(log_model_dir, "latest.pth")
|
||||
if args.loadmodel is not None:
|
||||
chk_path = args.loadmodel
|
||||
elif not os.path.exists(chk_path):
|
||||
chk_path = None
|
||||
|
||||
if chk_path is not None:
|
||||
worklog.info(f"loading model: {chk_path}")
|
||||
state_dict = torch.load(chk_path)
|
||||
model.module.load_state_dict(state_dict['state_dict'])
|
||||
optimizer.load_state_dict(state_dict['optim_state_dict'])
|
||||
resume_epoch_idx = state_dict["epoch"]
|
||||
resume_iters = state_dict["iters"]
|
||||
start_epoch_idx = resume_epoch_idx + 1
|
||||
start_iters = resume_iters
|
||||
else:
|
||||
start_epoch_idx = 1
|
||||
start_iters = 0
|
||||
|
||||
# datasets
|
||||
dataset = CREStereoDataset(args.training_data_path)
|
||||
sampler = RandomSampler(dataset, replacement=False)
|
||||
worklog.info(f"Dataset size: {len(dataset)}")
|
||||
dataloader = DataLoader(dataset, sampler=sampler, batch_size=args.batch_size*world_size,
|
||||
num_workers=0, drop_last=True, persistent_workers=False, pin_memory=True)
|
||||
dataloader = repeater(dataloader)
|
||||
|
||||
# counter
|
||||
cur_iters = start_iters
|
||||
total_iters = args.minibatch_per_epoch * args.n_total_epoch
|
||||
t0 = time.perf_counter()
|
||||
for epoch_idx in range(start_epoch_idx, args.n_total_epoch + 1):
|
||||
|
||||
# adjust learning rate
|
||||
epoch_total_train_loss = 0
|
||||
adjust_learning_rate(optimizer, epoch_idx)
|
||||
model.train()
|
||||
|
||||
t1 = time.perf_counter()
|
||||
|
||||
# for mini_batch_data in dataloader:
|
||||
for batch_idx, mini_batch_data in enumerate(dataloader):
|
||||
|
||||
if batch_idx % args.minibatch_per_epoch == 0 and batch_idx != 0:
|
||||
break
|
||||
cur_iters += 1
|
||||
|
||||
# parse data
|
||||
left, right, gt_disp, valid_mask = (
|
||||
mini_batch_data["left"].cuda(),
|
||||
mini_batch_data["right"].cuda(),
|
||||
mini_batch_data["disparity"].cuda(),
|
||||
mini_batch_data["mask"].cuda(),
|
||||
)
|
||||
|
||||
t2 = time.perf_counter()
|
||||
optimizer.zero_grad()
|
||||
|
||||
# pre-process
|
||||
gt_disp = torch.unsqueeze(gt_disp, dim=1) # [2, 384, 512] -> [2, 1, 384, 512]
|
||||
gt_flow = torch.cat([gt_disp, gt_disp * 0], dim=1) # [2, 2, 384, 512]
|
||||
|
||||
# forward
|
||||
flow_predictions = model(left, right)
|
||||
|
||||
# loss & backword
|
||||
loss = sequence_loss(
|
||||
flow_predictions, gt_flow, valid_mask, gamma=0.8
|
||||
)
|
||||
|
||||
# loss stats
|
||||
loss_item = loss.data.item()
|
||||
epoch_total_train_loss += loss_item
|
||||
loss.backward()
|
||||
optimizer.step()
|
||||
t3 = time.perf_counter()
|
||||
|
||||
if cur_iters % 10 == 0:
|
||||
tdata = t2 - t1
|
||||
time_train_passed = t3 - t0
|
||||
time_iter_passed = t3 - t1
|
||||
step_passed = cur_iters - start_iters
|
||||
eta = (
|
||||
(total_iters - cur_iters)
|
||||
/ max(step_passed, 1e-7)
|
||||
* time_train_passed
|
||||
)
|
||||
|
||||
meta_info = list()
|
||||
meta_info.append("{:.2g} b/s".format(1.0 / time_iter_passed))
|
||||
meta_info.append("passed:{}".format(format_time(time_train_passed)))
|
||||
meta_info.append("eta:{}".format(format_time(eta)))
|
||||
meta_info.append(
|
||||
"data_time:{:.2g}".format(tdata / time_iter_passed)
|
||||
)
|
||||
|
||||
meta_info.append(
|
||||
"lr:{:.5g}".format(optimizer.param_groups[0]["lr"])
|
||||
)
|
||||
meta_info.append(
|
||||
"[{}/{}:{}/{}]".format(
|
||||
epoch_idx,
|
||||
args.n_total_epoch,
|
||||
batch_idx,
|
||||
args.minibatch_per_epoch,
|
||||
)
|
||||
)
|
||||
loss_info = [" ==> {}:{:.4g}".format("loss", loss_item)]
|
||||
# exp_name = ['\n' + os.path.basename(os.getcwd())]
|
||||
|
||||
info = [",".join(meta_info)] + loss_info
|
||||
worklog.info("".join(info))
|
||||
|
||||
# minibatch loss
|
||||
tb_log.add_scalar("train/loss_batch", loss_item, cur_iters)
|
||||
tb_log.add_scalar(
|
||||
"train/lr", optimizer.param_groups[0]["lr"], cur_iters
|
||||
)
|
||||
tb_log.flush()
|
||||
|
||||
t1 = time.perf_counter()
|
||||
|
||||
tb_log.add_scalar(
|
||||
"train/loss",
|
||||
epoch_total_train_loss / args.minibatch_per_epoch,
|
||||
epoch_idx,
|
||||
)
|
||||
tb_log.flush()
|
||||
|
||||
# save model params
|
||||
ckp_data = {
|
||||
"epoch": epoch_idx,
|
||||
"iters": cur_iters,
|
||||
"batch_size": args.batch_size*world_size,
|
||||
"epoch_size": args.minibatch_per_epoch,
|
||||
"train_loss": epoch_total_train_loss / args.minibatch_per_epoch,
|
||||
"state_dict": model.module.state_dict(),
|
||||
"optim_state_dict": optimizer.state_dict(),
|
||||
}
|
||||
torch.save(ckp_data, os.path.join(log_model_dir, "latest.pth"))
|
||||
if epoch_idx % args.model_save_freq_epoch == 0:
|
||||
save_path = os.path.join(log_model_dir, "epoch-%d.pth" % epoch_idx)
|
||||
worklog.info(f"Model params saved: {save_path}")
|
||||
torch.save(ckp_data, save_path)
|
||||
|
||||
worklog.info("Training is done, exit.")
|
||||
|
||||
def main(args):
|
||||
# initial info
|
||||
torch.manual_seed(args.seed)
|
||||
torch.cuda.manual_seed(args.seed)
|
||||
world_size = torch.cuda.device_count() # number of GPU(s)
|
||||
cudnn.benchmark = True
|
||||
if args.dist and world_size > 1:
|
||||
train_dist(args, world_size)
|
||||
else:
|
||||
train(args, world_size)
|
||||
|
||||
if __name__ == "__main__":
|
||||
# train configuration
|
||||
args = parse_yaml("cfgs/train.yaml")
|
||||
main(args)
|
||||
Reference in New Issue
Block a user