Upload 96 files
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +17 -0
- LFE_TAP/datasets/EC_dataset.py +111 -0
- LFE_TAP/datasets/EDS_dataset.py +115 -0
- LFE_TAP/datasets/TAPFormer_dataset.py +151 -0
- LFE_TAP/datasets/__pycache__/Aedat4_dataset.cpython-39.pyc +0 -0
- LFE_TAP/datasets/__pycache__/EC_dataset.cpython-39.pyc +0 -0
- LFE_TAP/datasets/__pycache__/EDS_dataset.cpython-39.pyc +0 -0
- LFE_TAP/datasets/__pycache__/MF_dataset.cpython-39.pyc +0 -0
- LFE_TAP/datasets/__pycache__/kubric_movif_dataset.cpython-39.pyc +0 -0
- LFE_TAP/datasets/__pycache__/prophesee_dataset.cpython-39.pyc +0 -0
- LFE_TAP/datasets/kubric_movif_dataset.py +778 -0
- LFE_TAP/evaluator/__pycache__/evaluation_pred.cpython-38.pyc +0 -0
- LFE_TAP/evaluator/__pycache__/evaluation_pred.cpython-39.pyc +0 -0
- LFE_TAP/evaluator/__pycache__/evaluator.cpython-38.pyc +0 -0
- LFE_TAP/evaluator/__pycache__/evaluator.cpython-39.pyc +0 -0
- LFE_TAP/evaluator/__pycache__/prediction_long.cpython-38.pyc +0 -0
- LFE_TAP/evaluator/__pycache__/prediction_long.cpython-39.pyc +0 -0
- LFE_TAP/evaluator/evaluation_pred.py +184 -0
- LFE_TAP/evaluator/evaluator.py +351 -0
- LFE_TAP/evaluator/prediction.py +311 -0
- LFE_TAP/models/__pycache__/blocks.cpython-38.pyc +0 -0
- LFE_TAP/models/__pycache__/blocks.cpython-39.pyc +0 -0
- LFE_TAP/models/__pycache__/embeddings.cpython-38.pyc +0 -0
- LFE_TAP/models/__pycache__/embeddings.cpython-39.pyc +0 -0
- LFE_TAP/models/__pycache__/etap.cpython-39.pyc +0 -0
- LFE_TAP/models/__pycache__/fusionFormer.cpython-38.pyc +0 -0
- LFE_TAP/models/__pycache__/fusionFormer.cpython-39.pyc +0 -0
- LFE_TAP/models/__pycache__/hivit.cpython-38.pyc +0 -0
- LFE_TAP/models/__pycache__/hivit.cpython-39.pyc +0 -0
- LFE_TAP/models/__pycache__/losses.cpython-39.pyc +0 -0
- LFE_TAP/models/__pycache__/tapfe.cpython-38.pyc +0 -0
- LFE_TAP/models/__pycache__/tapfe.cpython-39.pyc +0 -0
- LFE_TAP/models/blocks.py +994 -0
- LFE_TAP/models/embeddings.py +110 -0
- LFE_TAP/models/fusionFormer.py +253 -0
- LFE_TAP/models/tapformer.py +308 -0
- LFE_TAP/utils/__pycache__/dataset_utils.cpython-38.pyc +0 -0
- LFE_TAP/utils/__pycache__/dataset_utils.cpython-39.pyc +0 -0
- LFE_TAP/utils/__pycache__/feature_map_vis.cpython-38.pyc +0 -0
- LFE_TAP/utils/__pycache__/feature_map_vis.cpython-39.pyc +0 -0
- LFE_TAP/utils/__pycache__/model_utils.cpython-38.pyc +0 -0
- LFE_TAP/utils/__pycache__/model_utils.cpython-39.pyc +0 -0
- LFE_TAP/utils/__pycache__/predictor.cpython-39.pyc +0 -0
- LFE_TAP/utils/__pycache__/train_utils.cpython-39.pyc +0 -0
- LFE_TAP/utils/__pycache__/visualizer.cpython-38.pyc +0 -0
- LFE_TAP/utils/__pycache__/visualizer.cpython-39.pyc +0 -0
- LFE_TAP/utils/dataset_utils.py +115 -0
- LFE_TAP/utils/event/__pycache__/representations.cpython-39.pyc +0 -0
- LFE_TAP/utils/event/__pycache__/utils.cpython-39.pyc +0 -0
- LFE_TAP/utils/event/representations.py +312 -0
.gitattributes
CHANGED
|
@@ -33,3 +33,20 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 33 |
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
| 36 |
+
assets/11_180_380.gif filter=lfs diff=lfs merge=lfs -text
|
| 37 |
+
assets/3_143_243.gif filter=lfs diff=lfs merge=lfs -text
|
| 38 |
+
assets/indoor_fruit_410_510.gif filter=lfs diff=lfs merge=lfs -text
|
| 39 |
+
assets/indoor_fruit_guobao2_5_155.gif filter=lfs diff=lfs merge=lfs -text
|
| 40 |
+
assets/indoor_hand_move_dynamic2_300_400.gif filter=lfs diff=lfs merge=lfs -text
|
| 41 |
+
assets/outdoor_day2-1_317_417.gif filter=lfs diff=lfs merge=lfs -text
|
| 42 |
+
assets/peanuts_light_160_386.gif filter=lfs diff=lfs merge=lfs -text
|
| 43 |
+
assets/peanuts_running_2360_2460.gif filter=lfs diff=lfs merge=lfs -text
|
| 44 |
+
assets/teaser.png filter=lfs diff=lfs merge=lfs -text
|
| 45 |
+
assets/toulan_360_460.gif filter=lfs diff=lfs merge=lfs -text
|
| 46 |
+
gt_tracks/boxes_rotation_198_278_gt_track.mp4 filter=lfs diff=lfs merge=lfs -text
|
| 47 |
+
gt_tracks/boxes_translation_330_410_gt_track.mp4 filter=lfs diff=lfs merge=lfs -text
|
| 48 |
+
gt_tracks/peanuts_light_160_386_gt_track_old.mp4 filter=lfs diff=lfs merge=lfs -text
|
| 49 |
+
gt_tracks/peanuts_light_160_386.gt_track.mp4 filter=lfs diff=lfs merge=lfs -text
|
| 50 |
+
gt_tracks/peanuts_running_2360_2460_gt_track.mp4 filter=lfs diff=lfs merge=lfs -text
|
| 51 |
+
gt_tracks/rocket_earth_light_338_438_gt_track.mp4 filter=lfs diff=lfs merge=lfs -text
|
| 52 |
+
gt_tracks/ziggy_in_the_arena_1350_1650_gt_track.mp4 filter=lfs diff=lfs merge=lfs -text
|
LFE_TAP/datasets/EC_dataset.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import torch
|
| 3 |
+
import imageio
|
| 4 |
+
import re
|
| 5 |
+
import numpy as np
|
| 6 |
+
from LFE_TAP.utils.event.utils import read_input
|
| 7 |
+
from LFE_TAP.utils.dataset_utils import FrameEventData_test
|
| 8 |
+
|
| 9 |
+
class EC_dataset(torch.utils.data.Dataset):
|
| 10 |
+
def __init__(self, data_root, dt=0.0200, representation="time_surfaces_v2_5", event_template_type = "sobel"):
|
| 11 |
+
self.data_root = data_root
|
| 12 |
+
self.dt = dt
|
| 13 |
+
self.representation = representation
|
| 14 |
+
self.event_template = event_template_type
|
| 15 |
+
self.seq_names = [
|
| 16 |
+
fname
|
| 17 |
+
for fname in os.listdir(data_root)
|
| 18 |
+
if os.path.isdir(os.path.join(data_root, fname))
|
| 19 |
+
]
|
| 20 |
+
print("found %d unique seqences in %s" % (len(self.seq_names), self.data_root))
|
| 21 |
+
|
| 22 |
+
def __getitem__(self, index):
|
| 23 |
+
gotit =False
|
| 24 |
+
event_dir_path = os.path.join(str(self.data_root), self.seq_names[index], "events", f"{self.dt:.4f}", self.representation)
|
| 25 |
+
rgb_path = os.path.join(str(self.data_root), self.seq_names[index], "images_corrected")
|
| 26 |
+
track_point_path = os.path.join(str(self.data_root), self.seq_names[index], "track.gt.txt")
|
| 27 |
+
|
| 28 |
+
img_paths = sorted(os.listdir(rgb_path))
|
| 29 |
+
# img_paths = img_paths[::16]
|
| 30 |
+
event_paths = sorted(os.listdir(event_dir_path))
|
| 31 |
+
rgb_imgs = []
|
| 32 |
+
rgb_ifnew = []
|
| 33 |
+
rgb_times = []
|
| 34 |
+
rgb_imgs_plus = []
|
| 35 |
+
rgb_ind = 0
|
| 36 |
+
event_imgs = []
|
| 37 |
+
event_time = []
|
| 38 |
+
for i, img_path in enumerate(img_paths):
|
| 39 |
+
try:
|
| 40 |
+
rgb_imgs.append(imageio.v2.imread(os.path.join(rgb_path, img_path)))
|
| 41 |
+
rgb_times.append(int(re.match(r"\d+", img_path).group()))
|
| 42 |
+
except Exception as e:
|
| 43 |
+
print(f"error reading image at path:{img_path}")
|
| 44 |
+
print(f"error mrssage:{str(e)}")
|
| 45 |
+
gotit = False
|
| 46 |
+
|
| 47 |
+
for i, event_path in enumerate(event_paths):
|
| 48 |
+
try:
|
| 49 |
+
event_imgs.append(read_input(os.path.join(event_dir_path, event_path), self.representation))
|
| 50 |
+
rgb_time = rgb_times[rgb_ind] if rgb_ind < len(rgb_times) else float('inf')
|
| 51 |
+
event_time.append(int(re.match(r"\d+", event_path).group()))
|
| 52 |
+
if int(re.match(r"\d+", event_path).group()) >= rgb_time:
|
| 53 |
+
# event_time.append(rgb_times[rgb_ind])
|
| 54 |
+
rgb_imgs_plus.append(rgb_imgs[rgb_ind])
|
| 55 |
+
rgb_ind += 1
|
| 56 |
+
rgb_ifnew.append(1)
|
| 57 |
+
else:
|
| 58 |
+
# event_time.append(int(re.match(r"\d+", event_path).group())-5000)
|
| 59 |
+
rgb_imgs_plus.append(rgb_imgs[rgb_ind-1])
|
| 60 |
+
rgb_ifnew.append(0)
|
| 61 |
+
except Exception as e:
|
| 62 |
+
print(f"error reading event at path:{event_path}")
|
| 63 |
+
print(f"error mrssage:{str(e)}")
|
| 64 |
+
gotit = False
|
| 65 |
+
|
| 66 |
+
# rgb_imgs = np.stack(rgb_imgs_plus)
|
| 67 |
+
# try:
|
| 68 |
+
# event_imgs = np.stack(event_imgs)
|
| 69 |
+
# except Exception as e:
|
| 70 |
+
# print(f"error reading at path:{str(self.data_root)}")
|
| 71 |
+
# gotit = False
|
| 72 |
+
# return [], gotit
|
| 73 |
+
rgb_imgs = np.stack(rgb_imgs_plus).transpose(0, 3, 1, 2)
|
| 74 |
+
event_imgs = np.stack(event_imgs).transpose(0, 3, 1, 2)
|
| 75 |
+
traj_data = np.genfromtxt(track_point_path, delimiter=" ") # id, t, x, y
|
| 76 |
+
track_num, track_ind = np.unique(traj_data[:, 0], return_index=True)
|
| 77 |
+
query_points = traj_data[track_ind, 2:]
|
| 78 |
+
|
| 79 |
+
T, H, W, C = event_imgs.shape
|
| 80 |
+
|
| 81 |
+
segs = np.array(event_time)
|
| 82 |
+
rgb_timestamp = np.array(rgb_times)
|
| 83 |
+
img_ifnew = np.array(rgb_ifnew)
|
| 84 |
+
query_points = torch.from_numpy(query_points).float()
|
| 85 |
+
query_points = torch.cat([torch.zeros_like(query_points[:, :1]), query_points], dim=1)
|
| 86 |
+
gotit = True
|
| 87 |
+
|
| 88 |
+
sample = FrameEventData_test(
|
| 89 |
+
rgb_imgs,
|
| 90 |
+
event_imgs,
|
| 91 |
+
segs,
|
| 92 |
+
traj_data,
|
| 93 |
+
seq_name=self.seq_names[index],
|
| 94 |
+
query_points=query_points,
|
| 95 |
+
img_ifnew=img_ifnew,
|
| 96 |
+
rgb_timestamp=rgb_timestamp,
|
| 97 |
+
)
|
| 98 |
+
|
| 99 |
+
return sample, gotit
|
| 100 |
+
|
| 101 |
+
def get_a_seq(self, seq_name):
|
| 102 |
+
for i in range(len(self.seq_names)):
|
| 103 |
+
if self.seq_names[i] == seq_name:
|
| 104 |
+
sample, gotit = self.__getitem__(i)
|
| 105 |
+
return sample, gotit
|
| 106 |
+
|
| 107 |
+
print("WARNNING: did not find the sequence", seq_name)
|
| 108 |
+
return [], False
|
| 109 |
+
|
| 110 |
+
def __len__(self):
|
| 111 |
+
return len(self.seq_names)
|
LFE_TAP/datasets/EDS_dataset.py
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import torch
|
| 3 |
+
import imageio
|
| 4 |
+
import re
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
from LFE_TAP.utils.event.utils import read_input
|
| 8 |
+
from LFE_TAP.utils.dataset_utils import FrameEventData, FrameEventData_test
|
| 9 |
+
|
| 10 |
+
class EDS_dataset(torch.utils.data.Dataset):
|
| 11 |
+
def __init__(self, data_root, dt=0.0050, representation="time_surfaces_v2_5"):
|
| 12 |
+
self.data_root = data_root
|
| 13 |
+
self.dt = dt
|
| 14 |
+
self.representation = representation
|
| 15 |
+
self.seq_names = [
|
| 16 |
+
fname
|
| 17 |
+
for fname in os.listdir(data_root)
|
| 18 |
+
if os.path.isdir(os.path.join(data_root, fname))
|
| 19 |
+
]
|
| 20 |
+
print("found %d unique seqences in %s" % (len(self.seq_names), self.data_root))
|
| 21 |
+
|
| 22 |
+
def __getitem__(self, index):
|
| 23 |
+
gotit =False
|
| 24 |
+
event_dir_path = os.path.join(str(self.data_root), self.seq_names[index], "events", f"{self.dt:.4f}", self.representation)
|
| 25 |
+
rgb_path = os.path.join(str(self.data_root), self.seq_names[index], "images_corrected")
|
| 26 |
+
track_point_path = os.path.join(f"gt_tracks/{self.seq_names[index]}.gt.txt")
|
| 27 |
+
|
| 28 |
+
img_paths = sorted(os.listdir(rgb_path))
|
| 29 |
+
# img_paths = img_paths[::8]
|
| 30 |
+
event_paths = sorted([f for f in os.listdir(event_dir_path) if f.endswith('.h5')])
|
| 31 |
+
rgb_imgs = []
|
| 32 |
+
rgb_ifnew = []
|
| 33 |
+
rgb_times = []
|
| 34 |
+
rgb_imgs_plus = []
|
| 35 |
+
rgb_ind = 0
|
| 36 |
+
event_imgs = []
|
| 37 |
+
event_time = []
|
| 38 |
+
time_emmbed = []
|
| 39 |
+
for i, img_path in enumerate(img_paths):
|
| 40 |
+
try:
|
| 41 |
+
# gray
|
| 42 |
+
img = imageio.v2.imread(os.path.join(rgb_path, img_path)).reshape(480, 640, 1)
|
| 43 |
+
img = img.repeat(3, axis=2)
|
| 44 |
+
# rgb
|
| 45 |
+
# img = imageio.v2.imread(os.path.join(rgb_path, img_path))
|
| 46 |
+
|
| 47 |
+
rgb_imgs.append(img)
|
| 48 |
+
rgb_times.append(int(re.match(r"\d+", img_path).group()))
|
| 49 |
+
except Exception as e:
|
| 50 |
+
print(f"error reading image at path:{img_path}")
|
| 51 |
+
print(f"error mrssage:{str(e)}")
|
| 52 |
+
gotit = False
|
| 53 |
+
for i, event_path in enumerate(event_paths):
|
| 54 |
+
try:
|
| 55 |
+
event_imgs.append(read_input(os.path.join(event_dir_path, event_path), self.representation))
|
| 56 |
+
rgb_time = rgb_times[rgb_ind] if rgb_ind < len(rgb_times) else float('inf')
|
| 57 |
+
if int(event_path.split('.')[0]) >= rgb_time:
|
| 58 |
+
# event_imgs.append(read_input(os.path.join(event_dir_path, event_path), self.representation))
|
| 59 |
+
event_time.append(rgb_times[rgb_ind])
|
| 60 |
+
rgb_imgs_plus.append(rgb_imgs[rgb_ind])
|
| 61 |
+
time_emmbed.append(round((int(event_path.split('.')[0]) - rgb_times[rgb_ind])*1e-4, 1))
|
| 62 |
+
rgb_ind += 1
|
| 63 |
+
rgb_ifnew.append(1)
|
| 64 |
+
else:
|
| 65 |
+
# continue
|
| 66 |
+
event_time.append(int(event_path.split('.')[0]))
|
| 67 |
+
rgb_imgs_plus.append(rgb_imgs[max(rgb_ind-1,0)])
|
| 68 |
+
time_emmbed.append(round((int(event_path.split('.')[0]) - rgb_times[max(rgb_ind-1,0)])*1e-4, 1))
|
| 69 |
+
rgb_ifnew.append(0)
|
| 70 |
+
except Exception as e:
|
| 71 |
+
print(f"error reading event at path:{event_path}")
|
| 72 |
+
print(f"error mrssage:{str(e)}")
|
| 73 |
+
gotit = False
|
| 74 |
+
|
| 75 |
+
rgb_imgs = np.stack(rgb_imgs_plus).transpose(0, 3, 1, 2)
|
| 76 |
+
event_imgs = np.stack(event_imgs)
|
| 77 |
+
if self.representation != "event_stack":
|
| 78 |
+
event_imgs = event_imgs.transpose(0, 3, 1, 2)
|
| 79 |
+
traj_data = np.genfromtxt(track_point_path, delimiter=" ") # id, t, x, y
|
| 80 |
+
track_num, track_ind = np.unique(traj_data[:, 0], return_index=True)
|
| 81 |
+
query_points = traj_data[track_ind, 2:]
|
| 82 |
+
|
| 83 |
+
T, C, _, _ = event_imgs.shape
|
| 84 |
+
|
| 85 |
+
segs = np.array(event_time)
|
| 86 |
+
rgb_timestamp = np.array(rgb_times)
|
| 87 |
+
img_ifnew = np.array(rgb_ifnew)
|
| 88 |
+
query_points = torch.from_numpy(query_points).float()
|
| 89 |
+
query_points = torch.cat([torch.zeros_like(query_points[:, :1]), query_points], dim=1)
|
| 90 |
+
gotit = True
|
| 91 |
+
|
| 92 |
+
sample = FrameEventData_test(
|
| 93 |
+
rgb_imgs,
|
| 94 |
+
event_imgs,
|
| 95 |
+
segs,
|
| 96 |
+
traj_data,
|
| 97 |
+
seq_name=self.seq_names[index],
|
| 98 |
+
query_points=query_points,
|
| 99 |
+
img_ifnew=img_ifnew,
|
| 100 |
+
rgb_timestamp=rgb_timestamp,
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
return sample, gotit
|
| 104 |
+
|
| 105 |
+
def get_a_seq(self, seq_name):
|
| 106 |
+
for i in range(len(self.seq_names)):
|
| 107 |
+
if self.seq_names[i] == seq_name:
|
| 108 |
+
sample, gotit = self.__getitem__(i)
|
| 109 |
+
return sample, gotit
|
| 110 |
+
|
| 111 |
+
print("WARNNING: did not find the sequence", seq_name)
|
| 112 |
+
return [], False
|
| 113 |
+
|
| 114 |
+
def __len__(self):
|
| 115 |
+
return len(self.seq_names)
|
LFE_TAP/datasets/TAPFormer_dataset.py
ADDED
|
@@ -0,0 +1,151 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import torch
|
| 3 |
+
import imageio
|
| 4 |
+
import re
|
| 5 |
+
import numpy as np
|
| 6 |
+
from LFE_TAP.utils.event.utils import read_input
|
| 7 |
+
from LFE_TAP.utils.dataset_utils import FrameEventData, FrameEventData_test
|
| 8 |
+
|
| 9 |
+
class TAPFormer_dataset(torch.utils.data.Dataset):
|
| 10 |
+
def __init__(self, data_root, dt=0.020, representation="time_surfaces_v2_5", fix_num=None, with_gt=True):
|
| 11 |
+
self.data_root = data_root
|
| 12 |
+
self.dt = dt
|
| 13 |
+
self.representation = representation
|
| 14 |
+
self.fix_num = fix_num
|
| 15 |
+
self.with_gt = with_gt
|
| 16 |
+
self.seq_names = [
|
| 17 |
+
fname
|
| 18 |
+
for fname in os.listdir(data_root)
|
| 19 |
+
if os.path.isdir(os.path.join(data_root, fname))
|
| 20 |
+
]
|
| 21 |
+
print("found %d unique seqences in %s" % (len(self.seq_names), self.data_root))
|
| 22 |
+
|
| 23 |
+
def __getitem__(self, index):
|
| 24 |
+
gotit =False
|
| 25 |
+
if self.fix_num is not None:
|
| 26 |
+
event_dir_path = os.path.join(str(self.data_root), self.seq_names[index], "events", self.representation, f"fix_num_{self.fix_num}")
|
| 27 |
+
else:
|
| 28 |
+
event_dir_path = os.path.join(str(self.data_root), self.seq_names[index], "events", self.representation, f"{self.dt:.4f}")
|
| 29 |
+
rgb_path = os.path.join(str(self.data_root), self.seq_names[index], "images_corrected")
|
| 30 |
+
track_point_path = os.path.join(str(self.data_root), self.seq_names[index], "annotations.npy")
|
| 31 |
+
img_time_full = np.stack(np.loadtxt(os.path.join(str(self.data_root), self.seq_names[index], "image_timestamps.txt")))
|
| 32 |
+
|
| 33 |
+
img_paths = sorted(os.listdir(rgb_path), key=lambda x: int(re.search(r'\d+', x).group()))
|
| 34 |
+
event_paths = sorted(os.listdir(event_dir_path), key=lambda x: int(re.search(r'\d+', x).group()))
|
| 35 |
+
rgb_imgs = []
|
| 36 |
+
rgb_ifnew = []
|
| 37 |
+
rgb_ifnew_full = []
|
| 38 |
+
rgb_times = []
|
| 39 |
+
rgb_imgs_plus = []
|
| 40 |
+
rgb_ind = 0
|
| 41 |
+
rgb_ind_full = 0
|
| 42 |
+
event_imgs = []
|
| 43 |
+
event_time = []
|
| 44 |
+
time_emmbed = []
|
| 45 |
+
for i, img_path in enumerate(img_paths):
|
| 46 |
+
try:
|
| 47 |
+
img = imageio.v2.imread(os.path.join(rgb_path, img_path))
|
| 48 |
+
if len(img.shape) == 2:
|
| 49 |
+
img = img[:,:,np.newaxis]
|
| 50 |
+
rgb_imgs.append(img.repeat(3, axis=2))
|
| 51 |
+
else:
|
| 52 |
+
rgb_imgs.append(img)
|
| 53 |
+
rgb_times.append(int(re.match(r"\d+", img_path).group()))
|
| 54 |
+
except Exception as e:
|
| 55 |
+
print(f"error reading image at path:{img_path}")
|
| 56 |
+
print(f"error mrssage:{str(e)}")
|
| 57 |
+
gotit = False
|
| 58 |
+
# rgb_imgs_plus.append(rgb_imgs[0])
|
| 59 |
+
# event_time.append(rgb_times[0])
|
| 60 |
+
# event_imgs.append(read_input(os.path.join(self.data_root, self.seq_names[index], "events", "template", self.event_template, str(event_time[0])+".h5"), self.event_template))
|
| 61 |
+
# rgb_ifnew.append(1)
|
| 62 |
+
# rgb_ind += 1
|
| 63 |
+
for i, event_path in enumerate(event_paths):
|
| 64 |
+
try:
|
| 65 |
+
if int(event_path.split('.')[0]) < rgb_times[0]:
|
| 66 |
+
continue
|
| 67 |
+
event_imgs.append(read_input(os.path.join(event_dir_path, event_path), self.representation))
|
| 68 |
+
# event_time.append(int(event_path.split('.')[0]))
|
| 69 |
+
rgb_time = rgb_times[min(rgb_ind, len(rgb_times)-1)]
|
| 70 |
+
if int(event_path.split('.')[0]) >= rgb_time:
|
| 71 |
+
rgb_imgs_plus.append(rgb_imgs[min(rgb_ind, len(rgb_times)-1)])
|
| 72 |
+
event_time.append(rgb_times[min(rgb_ind, len(rgb_times)-1)])
|
| 73 |
+
# time_emmbed.append(round((int(event_path.split('.')[0]) - rgb_times[min(rgb_ind, len(rgb_times)-1)])*1e-4, 1))
|
| 74 |
+
rgb_ind += 1
|
| 75 |
+
rgb_ifnew.append(1)
|
| 76 |
+
else:
|
| 77 |
+
rgb_imgs_plus.append(rgb_imgs[max(rgb_ind-1,0)])
|
| 78 |
+
event_time.append(int(event_path.split('.')[0]))
|
| 79 |
+
# time_emmbed.append(round((int(event_path.split('.')[0]) - rgb_times[max(rgb_ind-1,0)])*1e-4, 1))
|
| 80 |
+
rgb_ifnew.append(0)
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
rgb_time_full = img_time_full[rgb_ind_full]
|
| 84 |
+
if int(event_path.split('.')[0]) >= rgb_time_full:
|
| 85 |
+
rgb_ind_full += 1
|
| 86 |
+
rgb_ifnew_full.append(1)
|
| 87 |
+
else:
|
| 88 |
+
rgb_ifnew_full.append(0)
|
| 89 |
+
|
| 90 |
+
except Exception as e:
|
| 91 |
+
print(f"error reading event at path:{event_path}")
|
| 92 |
+
print(f"error mrssage:{str(e)}")
|
| 93 |
+
gotit = False
|
| 94 |
+
|
| 95 |
+
rgb_imgs = np.stack(rgb_imgs_plus).transpose(0, 3, 1, 2)
|
| 96 |
+
event_imgs = np.asarray(event_imgs).transpose(0, 3, 1, 2)
|
| 97 |
+
# event_imgs = np.stack(event_imgs).transpose(0, 3, 1, 2)
|
| 98 |
+
if self.with_gt:
|
| 99 |
+
annot_dict = np.load(track_point_path, allow_pickle=True).item()
|
| 100 |
+
traj_2d = annot_dict["coords"] # N, T, 2
|
| 101 |
+
visibility = annot_dict["visibility"] # N, T, 1
|
| 102 |
+
|
| 103 |
+
traj_data = np.transpose(traj_2d, (1, 0, 2)) # N, T, 2 -> T, N, 2
|
| 104 |
+
visibility = np.transpose(np.logical_not(np.squeeze(visibility)), (1, 0)) # N, T -> T, N
|
| 105 |
+
query_points = traj_data[0]
|
| 106 |
+
|
| 107 |
+
T, H, W, C = event_imgs.shape
|
| 108 |
+
|
| 109 |
+
segs = np.array(event_time)
|
| 110 |
+
rgb_timestamp = np.array(rgb_times)
|
| 111 |
+
img_ifnew = np.array(rgb_ifnew)
|
| 112 |
+
img_ifnew_full = np.array(rgb_ifnew_full)
|
| 113 |
+
rgb_times = np.array(rgb_times) * 1e-6
|
| 114 |
+
rgb_time_full = np.array(img_time_full) * 1e-6
|
| 115 |
+
if self.with_gt:
|
| 116 |
+
traj_data = np.concatenate([rgb_time_full[:, np.newaxis, np.newaxis].repeat(len(query_points), axis=1), traj_data], axis=2)
|
| 117 |
+
|
| 118 |
+
query_points = torch.from_numpy(query_points).float()
|
| 119 |
+
query_points = torch.cat([torch.zeros_like(query_points[:, :1]), query_points], dim=1)
|
| 120 |
+
else:
|
| 121 |
+
query_points = None
|
| 122 |
+
traj_data = None
|
| 123 |
+
visibility = None
|
| 124 |
+
gotit = True
|
| 125 |
+
|
| 126 |
+
sample = FrameEventData_test(
|
| 127 |
+
rgb_imgs,
|
| 128 |
+
event_imgs,
|
| 129 |
+
segs,
|
| 130 |
+
traj_data,
|
| 131 |
+
visibility=visibility,
|
| 132 |
+
seq_name=self.seq_names[index],
|
| 133 |
+
query_points=query_points,
|
| 134 |
+
img_ifnew=img_ifnew,
|
| 135 |
+
img_ifnew_full=img_ifnew_full,
|
| 136 |
+
rgb_timestamp=rgb_timestamp,
|
| 137 |
+
)
|
| 138 |
+
|
| 139 |
+
return sample, gotit
|
| 140 |
+
|
| 141 |
+
def get_a_seq(self, seq_name):
|
| 142 |
+
for i in range(len(self.seq_names)):
|
| 143 |
+
if self.seq_names[i] == seq_name:
|
| 144 |
+
sample, gotit = self.__getitem__(i)
|
| 145 |
+
return sample, gotit
|
| 146 |
+
|
| 147 |
+
print("WARNNING: did not find the sequence", seq_name)
|
| 148 |
+
return [], False
|
| 149 |
+
|
| 150 |
+
def __len__(self):
|
| 151 |
+
return len(self.seq_names)
|
LFE_TAP/datasets/__pycache__/Aedat4_dataset.cpython-39.pyc
ADDED
|
Binary file (4.6 kB). View file
|
|
|
LFE_TAP/datasets/__pycache__/EC_dataset.cpython-39.pyc
ADDED
|
Binary file (3.4 kB). View file
|
|
|
LFE_TAP/datasets/__pycache__/EDS_dataset.cpython-39.pyc
ADDED
|
Binary file (3.97 kB). View file
|
|
|
LFE_TAP/datasets/__pycache__/MF_dataset.cpython-39.pyc
ADDED
|
Binary file (3.72 kB). View file
|
|
|
LFE_TAP/datasets/__pycache__/kubric_movif_dataset.cpython-39.pyc
ADDED
|
Binary file (19.9 kB). View file
|
|
|
LFE_TAP/datasets/__pycache__/prophesee_dataset.cpython-39.pyc
ADDED
|
Binary file (5.06 kB). View file
|
|
|
LFE_TAP/datasets/kubric_movif_dataset.py
ADDED
|
@@ -0,0 +1,778 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import torch
|
| 3 |
+
import imageio
|
| 4 |
+
import cv2
|
| 5 |
+
from PIL import Image
|
| 6 |
+
import numpy as np
|
| 7 |
+
from torchvision.transforms import ColorJitter, GaussianBlur
|
| 8 |
+
from LFE_TAP.utils.dataset_utils import FrameEventData
|
| 9 |
+
from LFE_TAP.utils.event.utils import *
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class FETAPDataset(torch.utils.data.Dataset):
|
| 13 |
+
def __init__(self, root_dir, crop_size=(512, 512), seq_len=24, traj_per_sample=512, use_augs=False, **kwargs,):
|
| 14 |
+
super(FETAPDataset, self).__init__()
|
| 15 |
+
np.random.seed(0)
|
| 16 |
+
torch.manual_seed(0)
|
| 17 |
+
self.root_dir = Path(root_dir)
|
| 18 |
+
self.crop_size = crop_size
|
| 19 |
+
self.seq_len = seq_len
|
| 20 |
+
self.traj_per_sample = traj_per_sample
|
| 21 |
+
self.use_augs = use_augs
|
| 22 |
+
|
| 23 |
+
# photometric augmentation for rgb images
|
| 24 |
+
self.photo_aug = ColorJitter(
|
| 25 |
+
brightness=0.2, contrast=0.2, saturation=0.2, hue=0.25 / 3.14
|
| 26 |
+
)
|
| 27 |
+
self.blur_aug = GaussianBlur(11, sigma=(0.1, 2.0))
|
| 28 |
+
|
| 29 |
+
self.blur_aug_prob = 0.25
|
| 30 |
+
self.color_aug_prob = 0.25
|
| 31 |
+
|
| 32 |
+
# photometric augmentation for event images
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
# occlusion augmentation
|
| 36 |
+
self.eraser_aug_prob = 0.5
|
| 37 |
+
self.eraser_bounds = [2, 100]
|
| 38 |
+
self.eraser_max = 10
|
| 39 |
+
|
| 40 |
+
# occlusion augmentation
|
| 41 |
+
self.replace_aug_prob = 0.5
|
| 42 |
+
self.replace_bounds = [2, 100]
|
| 43 |
+
self.replace_max = 10
|
| 44 |
+
|
| 45 |
+
# spatial augmentations
|
| 46 |
+
self.pad_bounds = [0, 100]
|
| 47 |
+
self.crop_size = crop_size
|
| 48 |
+
self.resize_lim = [0.25, 2.0] # sample resizes from here
|
| 49 |
+
self.resize_delta = 0.2
|
| 50 |
+
self.max_crop_offset = 50
|
| 51 |
+
|
| 52 |
+
self.do_flip = True
|
| 53 |
+
self.h_flip_prob = 0.5
|
| 54 |
+
self.v_flip_prob = 0.5
|
| 55 |
+
|
| 56 |
+
def getitem_helper(self, index):
|
| 57 |
+
return NotImplementedError
|
| 58 |
+
|
| 59 |
+
def __getitem__(self, index):
|
| 60 |
+
gotit = False
|
| 61 |
+
|
| 62 |
+
sample, gotit = self.getitem_helper(index)
|
| 63 |
+
if not gotit:
|
| 64 |
+
print("warning: sampling failed")
|
| 65 |
+
# fake sample, so we can still collate
|
| 66 |
+
sample = FrameEventData(
|
| 67 |
+
video=torch.zeros(
|
| 68 |
+
(self.seq_len, 3, self.crop_size[0], self.crop_size[1])
|
| 69 |
+
),
|
| 70 |
+
events=torch.zeros((self.seq_len, 10, self.crop_size[0], self.crop_size[1])),
|
| 71 |
+
segmentation=torch.zeros(
|
| 72 |
+
(self.seq_len, 1, self.crop_size[0], self.crop_size[1])
|
| 73 |
+
),
|
| 74 |
+
trajectory=torch.zeros((self.seq_len, self.traj_per_sample, 2)),
|
| 75 |
+
visibility=torch.zeros((self.seq_len, self.traj_per_sample)),
|
| 76 |
+
valid=torch.zeros((self.seq_len, self.traj_per_sample)),
|
| 77 |
+
)
|
| 78 |
+
|
| 79 |
+
return sample, gotit
|
| 80 |
+
|
| 81 |
+
def add_photometric_augs(self, rgbs, events, trajs, visibles, eraser=True, replace=True):
|
| 82 |
+
T, N, _ = trajs.shape
|
| 83 |
+
|
| 84 |
+
S = len(rgbs)
|
| 85 |
+
H, W = rgbs[0].shape[:2]
|
| 86 |
+
assert S == T
|
| 87 |
+
|
| 88 |
+
# 事件图像增加椒盐噪声
|
| 89 |
+
|
| 90 |
+
if eraser:
|
| 91 |
+
############ eraser transform (per image after the first) ############
|
| 92 |
+
rgbs = [rgb.astype(np.float32) for rgb in rgbs]
|
| 93 |
+
events = [event.astype(np.float32) for event in events]
|
| 94 |
+
for i in range(1, S):
|
| 95 |
+
if np.random.rand() < self.eraser_aug_prob:
|
| 96 |
+
for _ in range(
|
| 97 |
+
np.random.randint(1, self.eraser_max + 1)
|
| 98 |
+
): # number of times to occlude
|
| 99 |
+
|
| 100 |
+
xc = np.random.randint(0, W)
|
| 101 |
+
yc = np.random.randint(0, H)
|
| 102 |
+
dx = np.random.randint(
|
| 103 |
+
self.eraser_bounds[0], self.eraser_bounds[1]
|
| 104 |
+
)
|
| 105 |
+
dy = np.random.randint(
|
| 106 |
+
self.eraser_bounds[0], self.eraser_bounds[1]
|
| 107 |
+
)
|
| 108 |
+
x0 = np.clip(xc - dx / 2, 0, W - 1).round().astype(np.int32)
|
| 109 |
+
x1 = np.clip(xc + dx / 2, 0, W - 1).round().astype(np.int32)
|
| 110 |
+
y0 = np.clip(yc - dy / 2, 0, H - 1).round().astype(np.int32)
|
| 111 |
+
y1 = np.clip(yc + dy / 2, 0, H - 1).round().astype(np.int32)
|
| 112 |
+
|
| 113 |
+
mean_color_rgb = np.mean(
|
| 114 |
+
rgbs[i][y0:y1, x0:x1, :].reshape(-1, 3), axis=0
|
| 115 |
+
)
|
| 116 |
+
rgbs[i][y0:y1, x0:x1, :] = mean_color_rgb
|
| 117 |
+
|
| 118 |
+
mean_value_event = np.mean(
|
| 119 |
+
events[i][y0:y1, x0:x1, :].reshape(-1, 10), axis=0
|
| 120 |
+
)
|
| 121 |
+
events[i][y0:y1, x0:x1, :] = mean_value_event
|
| 122 |
+
|
| 123 |
+
occ_inds = np.logical_and(
|
| 124 |
+
np.logical_and(trajs[i, :, 0] >= x0, trajs[i, :, 0] < x1),
|
| 125 |
+
np.logical_and(trajs[i, :, 1] >= y0, trajs[i, :, 1] < y1),
|
| 126 |
+
)
|
| 127 |
+
visibles[i, occ_inds] = 0
|
| 128 |
+
rgbs = [rgb.astype(np.uint8) for rgb in rgbs]
|
| 129 |
+
events = [event.astype(np.uint8) for event in events]
|
| 130 |
+
|
| 131 |
+
if replace:
|
| 132 |
+
|
| 133 |
+
rgbs_alt = [
|
| 134 |
+
np.array(self.blur_aug(Image.fromarray(rgb)), dtype=np.uint8)
|
| 135 |
+
for rgb in rgbs
|
| 136 |
+
]
|
| 137 |
+
events_alt = [
|
| 138 |
+
np.array(self.blur_aug(torch.from_numpy(event)), dtype=np.uint8)
|
| 139 |
+
for event in events
|
| 140 |
+
]
|
| 141 |
+
|
| 142 |
+
############ replace transform (per image after the first) ############
|
| 143 |
+
rgbs = [rgb.astype(np.float32) for rgb in rgbs]
|
| 144 |
+
rgbs_alt = [rgb.astype(np.float32) for rgb in rgbs_alt]
|
| 145 |
+
events = [event.astype(np.float32) for event in events]
|
| 146 |
+
events_alt = [event.astype(np.float32) for event in events_alt]
|
| 147 |
+
for i in range(1, S):
|
| 148 |
+
if np.random.rand() < self.replace_aug_prob:
|
| 149 |
+
for _ in range(
|
| 150 |
+
np.random.randint(1, self.replace_max + 1)
|
| 151 |
+
): # number of times to occlude
|
| 152 |
+
xc = np.random.randint(0, W)
|
| 153 |
+
yc = np.random.randint(0, H)
|
| 154 |
+
dx = np.random.randint(
|
| 155 |
+
self.replace_bounds[0], self.replace_bounds[1]
|
| 156 |
+
)
|
| 157 |
+
dy = np.random.randint(
|
| 158 |
+
self.replace_bounds[0], self.replace_bounds[1]
|
| 159 |
+
)
|
| 160 |
+
x0 = np.clip(xc - dx / 2, 0, W - 1).round().astype(np.int32)
|
| 161 |
+
x1 = np.clip(xc + dx / 2, 0, W - 1).round().astype(np.int32)
|
| 162 |
+
y0 = np.clip(yc - dy / 2, 0, H - 1).round().astype(np.int32)
|
| 163 |
+
y1 = np.clip(yc + dy / 2, 0, H - 1).round().astype(np.int32)
|
| 164 |
+
|
| 165 |
+
wid = x1 - x0
|
| 166 |
+
hei = y1 - y0
|
| 167 |
+
y00 = np.random.randint(0, H - hei)
|
| 168 |
+
x00 = np.random.randint(0, W - wid)
|
| 169 |
+
fr = np.random.randint(0, S)
|
| 170 |
+
rep_rgb = rgbs_alt[fr][y00 : y00 + hei, x00 : x00 + wid, :]
|
| 171 |
+
rgbs[i][y0:y1, x0:x1, :] = rep_rgb
|
| 172 |
+
rep_event = events_alt[fr][y00 : y00 + hei, x00 : x00 + wid, :]
|
| 173 |
+
events[i][y0:y1, x0:x1, :] = rep_event
|
| 174 |
+
|
| 175 |
+
occ_inds = np.logical_and(
|
| 176 |
+
np.logical_and(trajs[i, :, 0] >= x0, trajs[i, :, 0] < x1),
|
| 177 |
+
np.logical_and(trajs[i, :, 1] >= y0, trajs[i, :, 1] < y1),
|
| 178 |
+
)
|
| 179 |
+
visibles[i, occ_inds] = 0
|
| 180 |
+
rgbs = [rgb.astype(np.uint8) for rgb in rgbs]
|
| 181 |
+
events = [event.astype(np.uint8) for event in events]
|
| 182 |
+
|
| 183 |
+
############ photometric augmentation ############
|
| 184 |
+
if np.random.rand() < self.color_aug_prob:
|
| 185 |
+
# random per-frame amount of aug
|
| 186 |
+
rgbs = [
|
| 187 |
+
np.array(self.photo_aug(Image.fromarray(rgb)), dtype=np.uint8)
|
| 188 |
+
for rgb in rgbs
|
| 189 |
+
]
|
| 190 |
+
|
| 191 |
+
if np.random.rand() < self.blur_aug_prob:
|
| 192 |
+
# random per-frame amount of blur
|
| 193 |
+
rgbs = [
|
| 194 |
+
np.array(self.blur_aug(Image.fromarray(rgb)), dtype=np.uint8)
|
| 195 |
+
for rgb in rgbs
|
| 196 |
+
]
|
| 197 |
+
events = [
|
| 198 |
+
np.array(self.blur_aug(torch.from_numpy(event)), dtype=np.uint8)
|
| 199 |
+
for event in events
|
| 200 |
+
]
|
| 201 |
+
|
| 202 |
+
return rgbs, events, trajs, visibles
|
| 203 |
+
|
| 204 |
+
def add_spatial_augs(self, rgbs, events, trajs, visibles):
|
| 205 |
+
T, N, __ = trajs.shape
|
| 206 |
+
|
| 207 |
+
S = len(rgbs)
|
| 208 |
+
H, W = events[0].shape[:2]
|
| 209 |
+
assert S == T
|
| 210 |
+
|
| 211 |
+
rgbs = [rgb.astype(np.float32) for rgb in rgbs]
|
| 212 |
+
events = [event.astype(np.float32) for event in events]
|
| 213 |
+
|
| 214 |
+
############ spatial transform ############
|
| 215 |
+
|
| 216 |
+
# padding
|
| 217 |
+
pad_x0 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1])
|
| 218 |
+
pad_x1 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1])
|
| 219 |
+
pad_y0 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1])
|
| 220 |
+
pad_y1 = np.random.randint(self.pad_bounds[0], self.pad_bounds[1])
|
| 221 |
+
|
| 222 |
+
rgbs = [
|
| 223 |
+
np.pad(rgb, ((pad_y0, pad_y1), (pad_x0, pad_x1), (0, 0))) for rgb in rgbs
|
| 224 |
+
]
|
| 225 |
+
events = [
|
| 226 |
+
np.pad(event, ((pad_y0, pad_y1), (pad_x0, pad_x1), (0, 0))) for event in events
|
| 227 |
+
]
|
| 228 |
+
trajs[:, :, 0] += pad_x0
|
| 229 |
+
trajs[:, :, 1] += pad_y0
|
| 230 |
+
H, W = rgbs[0].shape[:2]
|
| 231 |
+
|
| 232 |
+
# scaling + stretching
|
| 233 |
+
scale = np.random.uniform(self.resize_lim[0], self.resize_lim[1])
|
| 234 |
+
scale_x = scale
|
| 235 |
+
scale_y = scale
|
| 236 |
+
H_new = H
|
| 237 |
+
W_new = W
|
| 238 |
+
|
| 239 |
+
scale_delta_x = 0.0
|
| 240 |
+
scale_delta_y = 0.0
|
| 241 |
+
|
| 242 |
+
rgbs_scaled = []
|
| 243 |
+
events_scaled = []
|
| 244 |
+
for s in range(S):
|
| 245 |
+
if s == 1:
|
| 246 |
+
scale_delta_x = np.random.uniform(-self.resize_delta, self.resize_delta)
|
| 247 |
+
scale_delta_y = np.random.uniform(-self.resize_delta, self.resize_delta)
|
| 248 |
+
elif s > 1:
|
| 249 |
+
scale_delta_x = (
|
| 250 |
+
scale_delta_x * 0.8
|
| 251 |
+
+ np.random.uniform(-self.resize_delta, self.resize_delta) * 0.2
|
| 252 |
+
)
|
| 253 |
+
scale_delta_y = (
|
| 254 |
+
scale_delta_y * 0.8
|
| 255 |
+
+ np.random.uniform(-self.resize_delta, self.resize_delta) * 0.2
|
| 256 |
+
)
|
| 257 |
+
scale_x = scale_x + scale_delta_x
|
| 258 |
+
scale_y = scale_y + scale_delta_y
|
| 259 |
+
|
| 260 |
+
# bring h/w closer
|
| 261 |
+
scale_xy = (scale_x + scale_y) * 0.5
|
| 262 |
+
scale_x = scale_x * 0.5 + scale_xy * 0.5
|
| 263 |
+
scale_y = scale_y * 0.5 + scale_xy * 0.5
|
| 264 |
+
|
| 265 |
+
# don't get too crazy
|
| 266 |
+
scale_x = np.clip(scale_x, 0.2, 2.0)
|
| 267 |
+
scale_y = np.clip(scale_y, 0.2, 2.0)
|
| 268 |
+
|
| 269 |
+
H_new = int(H * scale_y)
|
| 270 |
+
W_new = int(W * scale_x)
|
| 271 |
+
|
| 272 |
+
# make it at least slightly bigger than the crop area,
|
| 273 |
+
# so that the random cropping can add diversity
|
| 274 |
+
H_new = np.clip(H_new, self.crop_size[0] + 10, None)
|
| 275 |
+
W_new = np.clip(W_new, self.crop_size[1] + 10, None)
|
| 276 |
+
# recompute scale in case we clipped
|
| 277 |
+
scale_x = (W_new - 1) / float(W - 1)
|
| 278 |
+
scale_y = (H_new - 1) / float(H - 1)
|
| 279 |
+
|
| 280 |
+
rgbs_scaled.append(
|
| 281 |
+
cv2.resize(rgbs[s], (W_new, H_new), interpolation=cv2.INTER_LINEAR)
|
| 282 |
+
)
|
| 283 |
+
events_scaled.append(
|
| 284 |
+
cv2.resize(events[s], (W_new, H_new), interpolation=cv2.INTER_LINEAR)
|
| 285 |
+
)
|
| 286 |
+
trajs[s, :, 0] *= scale_x
|
| 287 |
+
trajs[s, :, 1] *= scale_y
|
| 288 |
+
rgbs = rgbs_scaled
|
| 289 |
+
events = events_scaled
|
| 290 |
+
ok_inds = visibles[0, :] > 0
|
| 291 |
+
vis_trajs = trajs[:, ok_inds] # S,?,2
|
| 292 |
+
|
| 293 |
+
if vis_trajs.shape[1] > 0:
|
| 294 |
+
mid_x = np.mean(vis_trajs[0, :, 0])
|
| 295 |
+
mid_y = np.mean(vis_trajs[0, :, 1])
|
| 296 |
+
else:
|
| 297 |
+
mid_x = self.crop_size[0]
|
| 298 |
+
mid_y = self.crop_size[1]
|
| 299 |
+
|
| 300 |
+
x0 = int(mid_x - self.crop_size[1] // 2)
|
| 301 |
+
y0 = int(mid_y - self.crop_size[0] // 2)
|
| 302 |
+
|
| 303 |
+
offset_x = 0
|
| 304 |
+
offset_y = 0
|
| 305 |
+
|
| 306 |
+
for s in range(S):
|
| 307 |
+
# on each frame, shift a bit more
|
| 308 |
+
if s == 1:
|
| 309 |
+
offset_x = np.random.randint(
|
| 310 |
+
-self.max_crop_offset, self.max_crop_offset
|
| 311 |
+
)
|
| 312 |
+
offset_y = np.random.randint(
|
| 313 |
+
-self.max_crop_offset, self.max_crop_offset
|
| 314 |
+
)
|
| 315 |
+
elif s > 1:
|
| 316 |
+
offset_x = int(
|
| 317 |
+
offset_x * 0.8
|
| 318 |
+
+ np.random.randint(-self.max_crop_offset, self.max_crop_offset + 1)
|
| 319 |
+
* 0.2
|
| 320 |
+
)
|
| 321 |
+
offset_y = int(
|
| 322 |
+
offset_y * 0.8
|
| 323 |
+
+ np.random.randint(-self.max_crop_offset, self.max_crop_offset + 1)
|
| 324 |
+
* 0.2
|
| 325 |
+
)
|
| 326 |
+
x0 = x0 + offset_x
|
| 327 |
+
y0 = y0 + offset_y
|
| 328 |
+
|
| 329 |
+
H_new, W_new = rgbs[s].shape[:2]
|
| 330 |
+
if H_new == self.crop_size[0]:
|
| 331 |
+
y0 = 0
|
| 332 |
+
else:
|
| 333 |
+
y0 = min(max(0, y0), H_new - self.crop_size[0] - 1)
|
| 334 |
+
|
| 335 |
+
if W_new == self.crop_size[1]:
|
| 336 |
+
x0 = 0
|
| 337 |
+
else:
|
| 338 |
+
x0 = min(max(0, x0), W_new - self.crop_size[1] - 1)
|
| 339 |
+
|
| 340 |
+
rgbs[s] = rgbs[s][y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
|
| 341 |
+
events[s] = events[s][y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
|
| 342 |
+
trajs[s, :, 0] -= x0
|
| 343 |
+
trajs[s, :, 1] -= y0
|
| 344 |
+
|
| 345 |
+
H_new = self.crop_size[0]
|
| 346 |
+
W_new = self.crop_size[1]
|
| 347 |
+
|
| 348 |
+
# flip
|
| 349 |
+
h_flipped = False
|
| 350 |
+
v_flipped = False
|
| 351 |
+
if self.do_flip:
|
| 352 |
+
# h flip
|
| 353 |
+
if np.random.rand() < self.h_flip_prob:
|
| 354 |
+
h_flipped = True
|
| 355 |
+
rgbs = [rgb[:, ::-1] for rgb in rgbs]
|
| 356 |
+
events = [event[:, ::-1] for event in events]
|
| 357 |
+
# v flip
|
| 358 |
+
if np.random.rand() < self.v_flip_prob:
|
| 359 |
+
v_flipped = True
|
| 360 |
+
rgbs = [rgb[::-1] for rgb in rgbs]
|
| 361 |
+
events = [event[::-1] for event in events]
|
| 362 |
+
if h_flipped:
|
| 363 |
+
trajs[:, :, 0] = W_new - trajs[:, :, 0]
|
| 364 |
+
if v_flipped:
|
| 365 |
+
trajs[:, :, 1] = H_new - trajs[:, :, 1]
|
| 366 |
+
|
| 367 |
+
return rgbs, events, trajs
|
| 368 |
+
|
| 369 |
+
def crop(self, rgbs, events, trajs, clear_rgbs=None):
|
| 370 |
+
T, N, _ = trajs.shape
|
| 371 |
+
|
| 372 |
+
S = len(events)
|
| 373 |
+
H, W = events[0].shape[:2]
|
| 374 |
+
assert S == T
|
| 375 |
+
|
| 376 |
+
############ spatial transform ############
|
| 377 |
+
|
| 378 |
+
H_new = H
|
| 379 |
+
W_new = W
|
| 380 |
+
|
| 381 |
+
# simple random crop
|
| 382 |
+
y0 = 0 if self.crop_size[0] >= H_new else (H_new - self.crop_size[0]) // 2
|
| 383 |
+
# np.random.randint(0,
|
| 384 |
+
x0 = 0 if self.crop_size[1] >= W_new else np.random.randint(0, W_new - self.crop_size[1])
|
| 385 |
+
rgbs = [
|
| 386 |
+
rgb[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
|
| 387 |
+
for rgb in rgbs
|
| 388 |
+
]
|
| 389 |
+
events = [
|
| 390 |
+
event[y0: y0 + self.crop_size[0], x0: x0 + self.crop_size[1]]
|
| 391 |
+
for event in events
|
| 392 |
+
]
|
| 393 |
+
|
| 394 |
+
trajs[:, :, 0] -= x0
|
| 395 |
+
trajs[:, :, 1] -= y0
|
| 396 |
+
if clear_rgbs is not None:
|
| 397 |
+
clear_rgbs = [
|
| 398 |
+
clear_rgb[y0 : y0 + self.crop_size[0], x0 : x0 + self.crop_size[1]]
|
| 399 |
+
for clear_rgb in clear_rgbs
|
| 400 |
+
]
|
| 401 |
+
return rgbs, events, trajs, clear_rgbs
|
| 402 |
+
|
| 403 |
+
return rgbs, events, trajs
|
| 404 |
+
|
| 405 |
+
|
| 406 |
+
class KubricMovifDataset(FETAPDataset):
|
| 407 |
+
def __init__(
|
| 408 |
+
self,
|
| 409 |
+
root_dir,
|
| 410 |
+
representation="time_surfaces_v2_5",
|
| 411 |
+
event_template="sobel",
|
| 412 |
+
crop_size=(384, 512),
|
| 413 |
+
seq_len=24,
|
| 414 |
+
traj_per_sample=512,
|
| 415 |
+
sample_vis_1st_frame=False,
|
| 416 |
+
choose_long_point=False,
|
| 417 |
+
use_augs=False,
|
| 418 |
+
):
|
| 419 |
+
super(KubricMovifDataset, self).__init__(
|
| 420 |
+
root_dir=root_dir,
|
| 421 |
+
representation=representation,
|
| 422 |
+
event_template=event_template,
|
| 423 |
+
crop_size=crop_size,
|
| 424 |
+
seq_len=seq_len,
|
| 425 |
+
traj_per_sample=traj_per_sample,
|
| 426 |
+
sample_vis_1st_frame=sample_vis_1st_frame,
|
| 427 |
+
choose_long_point=choose_long_point,
|
| 428 |
+
use_augs=use_augs,
|
| 429 |
+
)
|
| 430 |
+
self.representation = representation
|
| 431 |
+
self.event_template = event_template
|
| 432 |
+
self.sample_vis_1st_frame = sample_vis_1st_frame
|
| 433 |
+
self.choose_long_point = choose_long_point
|
| 434 |
+
self.pad_bounds = [0, 25]
|
| 435 |
+
self.resize_lim = [0.75, 1.25]
|
| 436 |
+
self.resize_delta = 0.05
|
| 437 |
+
self.max_crop_offset = 15
|
| 438 |
+
self.seq_names = [
|
| 439 |
+
fname
|
| 440 |
+
for fname in os.listdir(self.root_dir) if os.path.isdir(os.path.join(self.root_dir, fname))
|
| 441 |
+
]
|
| 442 |
+
print("found %d unique seqences in %s" % (len(self.seq_names), self.root_dir))
|
| 443 |
+
|
| 444 |
+
def getitem_helper(self, index):
|
| 445 |
+
gotit = True
|
| 446 |
+
seq_name = self.seq_names[index]
|
| 447 |
+
|
| 448 |
+
npy_path = os.path.join(self.root_dir, seq_name, seq_name + ".npy")
|
| 449 |
+
rgb_dir_path = os.path.join(self.root_dir, seq_name, "frames")
|
| 450 |
+
event_dir_path = os.path.join(self.root_dir, seq_name, "events", self.representation)
|
| 451 |
+
|
| 452 |
+
rgb_files = sorted(os.listdir(rgb_dir_path))
|
| 453 |
+
event_files = sorted(os.listdir(event_dir_path))
|
| 454 |
+
rgb_imgs = []
|
| 455 |
+
img_ifnew = []
|
| 456 |
+
event_imgs = []
|
| 457 |
+
event_imgs.append(read_input(os.path.join(self.root_dir, seq_name, "events", "template", self.event_template, "000.h5"), self.event_template))
|
| 458 |
+
|
| 459 |
+
for i, img_path in enumerate(rgb_files):
|
| 460 |
+
try:
|
| 461 |
+
if i % 3 == 0:
|
| 462 |
+
rgb_imgs.append(imageio.v2.imread(os.path.join(rgb_dir_path, img_path)))
|
| 463 |
+
img_ifnew.append(1)
|
| 464 |
+
else:
|
| 465 |
+
rgb_imgs.append(rgb_imgs[-1])
|
| 466 |
+
img_ifnew.append(0)
|
| 467 |
+
except Exception as e:
|
| 468 |
+
print(f"error reading image at path:{img_path}_{rgb_dir_path}")
|
| 469 |
+
print(f"error mrssage:{str(e)}")
|
| 470 |
+
gotit = False
|
| 471 |
+
return [], gotit
|
| 472 |
+
|
| 473 |
+
for i, event_path in enumerate(event_files):
|
| 474 |
+
try:
|
| 475 |
+
event_imgs.append(read_input(os.path.join(event_dir_path, event_path), self.representation))
|
| 476 |
+
except Exception as e:
|
| 477 |
+
print(f"error reading event at path:{event_path}_{event_dir_path}")
|
| 478 |
+
print(f"error mrssage:{str(e)}")
|
| 479 |
+
gotit = False
|
| 480 |
+
return [], gotit
|
| 481 |
+
|
| 482 |
+
rgbs = np.stack(rgb_imgs)
|
| 483 |
+
events = np.stack(event_imgs)
|
| 484 |
+
annot_dict = np.load(npy_path, allow_pickle=True).item()
|
| 485 |
+
traj_2d = annot_dict["coords"]
|
| 486 |
+
visibility = annot_dict["visibility"]
|
| 487 |
+
|
| 488 |
+
assert self.seq_len == len(rgbs)
|
| 489 |
+
|
| 490 |
+
traj_2d = np.transpose(traj_2d, (1, 0, 2)) # N, T, 2 -> T, N, 2
|
| 491 |
+
visibility = np.transpose(np.logical_not(visibility), (1, 0)) # N, T -> T, N
|
| 492 |
+
if self.use_augs:
|
| 493 |
+
# rgbs, events, traj_2d, visibility = self.add_photometric_augs(rgbs, events, traj_2d, visibility)
|
| 494 |
+
rgbs, events, traj_2d = self.add_spatial_augs(rgbs, events, traj_2d, visibility)
|
| 495 |
+
else:
|
| 496 |
+
rgbs, events, traj_2d = self.crop(rgbs, events, traj_2d)
|
| 497 |
+
|
| 498 |
+
visibility[traj_2d[:, :, 0] > self.crop_size[1] - 1] = False
|
| 499 |
+
visibility[traj_2d[:, :, 0] < 0] = False
|
| 500 |
+
visibility[traj_2d[:, :, 1] > self.crop_size[0] - 1] = False
|
| 501 |
+
visibility[traj_2d[:, :, 1] < 0] = False
|
| 502 |
+
|
| 503 |
+
visibility = torch.from_numpy(visibility)
|
| 504 |
+
traj_2d = torch.from_numpy(traj_2d)
|
| 505 |
+
|
| 506 |
+
crop_tensor = torch.tensor(self.crop_size).flip(0)[None, None] / 2.0
|
| 507 |
+
close_pts_inds = torch.all(
|
| 508 |
+
torch.linalg.vector_norm(traj_2d[..., :2] - crop_tensor, dim=-1) < 1000.0,
|
| 509 |
+
dim=0,
|
| 510 |
+
)
|
| 511 |
+
traj_2d = traj_2d[:, close_pts_inds]
|
| 512 |
+
visibility = visibility[:, close_pts_inds]
|
| 513 |
+
|
| 514 |
+
visibile_pts_first_frame_inds = (visibility[0]).nonzero(as_tuple=False)
|
| 515 |
+
|
| 516 |
+
if self.sample_vis_1st_frame:
|
| 517 |
+
visiblile_pts_inds = visibile_pts_first_frame_inds
|
| 518 |
+
else:
|
| 519 |
+
visiblile_pts_mid_frame_inds = (visibility[self.seq_len // 2]).nonzero(as_tuple=False)
|
| 520 |
+
visiblile_pts_inds = torch.cat((visibile_pts_first_frame_inds, visiblile_pts_mid_frame_inds), dim=0)
|
| 521 |
+
|
| 522 |
+
point_inds = torch.randperm(len(visiblile_pts_inds))[:self.traj_per_sample]
|
| 523 |
+
if len(point_inds) < self.traj_per_sample:
|
| 524 |
+
gotit = False
|
| 525 |
+
|
| 526 |
+
if self.choose_long_point:
|
| 527 |
+
distance = np.linalg.norm(traj_2d[-1, visiblile_pts_inds, :] - traj_2d[0, visiblile_pts_inds, :], axis=-1)[:,0]
|
| 528 |
+
weight = distance / np.sum(distance)
|
| 529 |
+
point_inds = torch.tensor(np.random.choice(len(distance), size=self.traj_per_sample, p=weight))
|
| 530 |
+
visible_inds_sampled = visiblile_pts_inds[point_inds]
|
| 531 |
+
# visible_inds_sampled = visiblile_pts_inds[np.argsort(distance, axis=0)[-self.traj_per_sample:,0]]
|
| 532 |
+
else:
|
| 533 |
+
visible_inds_sampled = visiblile_pts_inds[point_inds]
|
| 534 |
+
|
| 535 |
+
if len(visible_inds_sampled.shape) == 2:
|
| 536 |
+
visible_inds_sampled = visible_inds_sampled.squeeze(1)
|
| 537 |
+
trajs = traj_2d[:, visible_inds_sampled].float()
|
| 538 |
+
visibles = visibility[:, visible_inds_sampled]
|
| 539 |
+
valid = torch.ones((self.seq_len, self.traj_per_sample))
|
| 540 |
+
|
| 541 |
+
rgbs = torch.from_numpy(np.stack(rgbs)).permute(0, 3, 1, 2).float()
|
| 542 |
+
events = torch.from_numpy(np.stack(events)).permute(0, 3, 1, 2).float()
|
| 543 |
+
seqs = torch.ones((self.seq_len, 1, self.crop_size[0], self.crop_size[1]))
|
| 544 |
+
img_ifnew = np.array(img_ifnew)
|
| 545 |
+
sample = FrameEventData(
|
| 546 |
+
video = rgbs,
|
| 547 |
+
events = events,
|
| 548 |
+
segmentation=seqs,
|
| 549 |
+
trajectory=trajs,
|
| 550 |
+
visibility=visibles,
|
| 551 |
+
valid=valid,
|
| 552 |
+
seq_name=seq_name,
|
| 553 |
+
img_ifnew=img_ifnew,
|
| 554 |
+
)
|
| 555 |
+
return sample, gotit
|
| 556 |
+
|
| 557 |
+
def __len__(self):
|
| 558 |
+
return len(self.seq_names)
|
| 559 |
+
|
| 560 |
+
|
| 561 |
+
class KubricMovifDataset_new(FETAPDataset):
|
| 562 |
+
def __init__(
|
| 563 |
+
self,
|
| 564 |
+
root_dir,
|
| 565 |
+
root_dir_fast_dataset=None,
|
| 566 |
+
representation="time_surfaces_v2_5",
|
| 567 |
+
crop_size=(384, 512),
|
| 568 |
+
seq_len=24,
|
| 569 |
+
traj_per_sample=256,
|
| 570 |
+
sample_vis_1st_frame=False,
|
| 571 |
+
choose_long_point=False,
|
| 572 |
+
use_augs=False,
|
| 573 |
+
if_test=False,
|
| 574 |
+
):
|
| 575 |
+
super(KubricMovifDataset_new, self).__init__(
|
| 576 |
+
root_dir=root_dir,
|
| 577 |
+
root_dir_fast_dataset=root_dir_fast_dataset,
|
| 578 |
+
representation=representation,
|
| 579 |
+
crop_size=crop_size,
|
| 580 |
+
seq_len=seq_len,
|
| 581 |
+
traj_per_sample=traj_per_sample,
|
| 582 |
+
sample_vis_1st_frame=sample_vis_1st_frame,
|
| 583 |
+
choose_long_point=choose_long_point,
|
| 584 |
+
use_augs=use_augs,
|
| 585 |
+
if_test=if_test,
|
| 586 |
+
)
|
| 587 |
+
self.root_dir1 = os.path.join(self.root_dir, "kubric_ori_dataset1")
|
| 588 |
+
self.root_dir2 = os.path.join(self.root_dir, "kubric_ori_dataset2")
|
| 589 |
+
self.representation = representation
|
| 590 |
+
self.sample_vis_1st_frame = sample_vis_1st_frame
|
| 591 |
+
self.choose_long_point = choose_long_point
|
| 592 |
+
self.pad_bounds = [0, 25]
|
| 593 |
+
self.resize_lim = [0.75, 1.25]
|
| 594 |
+
self.resize_delta = 0.05
|
| 595 |
+
self.max_crop_offset = 15
|
| 596 |
+
self.if_test = if_test
|
| 597 |
+
if root_dir_fast_dataset is not None:
|
| 598 |
+
root_dir_fast_dataset = Path(root_dir_fast_dataset)
|
| 599 |
+
seq_names1 = [
|
| 600 |
+
os.path.join(self.root_dir1, fname)
|
| 601 |
+
for fname in os.listdir(self.root_dir1) if os.path.isdir(os.path.join(self.root_dir1, fname))
|
| 602 |
+
]
|
| 603 |
+
seq_names2 = [
|
| 604 |
+
os.path.join(self.root_dir2, fname)
|
| 605 |
+
for fname in os.listdir(self.root_dir2) if os.path.isdir(os.path.join(self.root_dir2, fname))
|
| 606 |
+
]
|
| 607 |
+
|
| 608 |
+
seq_names_fast = [
|
| 609 |
+
os.path.join(root_dir_fast_dataset, fname)
|
| 610 |
+
for fname in os.listdir(root_dir_fast_dataset) if os.path.isdir(os.path.join(root_dir_fast_dataset, fname))
|
| 611 |
+
]
|
| 612 |
+
self.seq_names = seq_names1 + seq_names2 + seq_names_fast
|
| 613 |
+
else:
|
| 614 |
+
self.seq_names = [
|
| 615 |
+
os.path.join(self.root_dir, fname)
|
| 616 |
+
for fname in os.listdir(self.root_dir) if os.path.isdir(os.path.join(self.root_dir, fname))
|
| 617 |
+
]
|
| 618 |
+
print("found %d unique seqences in %s" % (len(self.seq_names), self.root_dir))
|
| 619 |
+
|
| 620 |
+
def getitem_helper(self, index):
|
| 621 |
+
gotit = True
|
| 622 |
+
seq_name = self.seq_names[index]
|
| 623 |
+
|
| 624 |
+
npy_path = os.path.join(seq_name, os.path.basename(seq_name) + ".npy")
|
| 625 |
+
if not os.path.exists(npy_path):
|
| 626 |
+
print(npy_path)
|
| 627 |
+
rgb_dir_path = os.path.join(seq_name, "blur_frames")
|
| 628 |
+
clear_rgb_dir_path = os.path.join(seq_name, "frames")
|
| 629 |
+
event_dir_path = os.path.join(seq_name, "events", self.representation)
|
| 630 |
+
|
| 631 |
+
rgb_files = sorted(os.listdir(rgb_dir_path))
|
| 632 |
+
clear_rgb_files = sorted(os.listdir(clear_rgb_dir_path))
|
| 633 |
+
event_files = sorted(os.listdir(event_dir_path))
|
| 634 |
+
rgb_imgs = []
|
| 635 |
+
img_ifnew = []
|
| 636 |
+
clear_rgb_imgs = []
|
| 637 |
+
event_imgs = []
|
| 638 |
+
if not self.if_test:
|
| 639 |
+
random_id = np.random.randint(0, 95-24)
|
| 640 |
+
rgb_files = rgb_files[random_id: random_id+24]
|
| 641 |
+
event_files = event_files[random_id: random_id+24]
|
| 642 |
+
clear_rgb_files = clear_rgb_files[random_id: random_id+24]
|
| 643 |
+
else:
|
| 644 |
+
rgb_files = rgb_files[:-1]
|
| 645 |
+
clear_rgb_files = clear_rgb_files[:-1]
|
| 646 |
+
|
| 647 |
+
next_insert = 0 # 记录下一个需要插入新图像的索引
|
| 648 |
+
for i, img_path in enumerate(rgb_files):
|
| 649 |
+
try:
|
| 650 |
+
if i == next_insert:
|
| 651 |
+
# 读取新图像并添加到列表
|
| 652 |
+
img = imageio.v2.imread(os.path.join(rgb_dir_path, img_path))
|
| 653 |
+
rgb_imgs.append(img)
|
| 654 |
+
img_ifnew.append(1)
|
| 655 |
+
# 随机选择下一个间隔(3或4)
|
| 656 |
+
if not self.if_test:
|
| 657 |
+
next_insert += np.random.choice([3, 4])
|
| 658 |
+
else:
|
| 659 |
+
next_insert += 4
|
| 660 |
+
else:
|
| 661 |
+
rgb_imgs.append(rgb_imgs[-1])
|
| 662 |
+
img_ifnew.append(0)
|
| 663 |
+
except Exception as e:
|
| 664 |
+
print(f"error reading image at path:{img_path}_{rgb_dir_path}")
|
| 665 |
+
print(f"error mrssage:{str(e)}")
|
| 666 |
+
gotit = False
|
| 667 |
+
return [], gotit
|
| 668 |
+
|
| 669 |
+
for i, img_path in enumerate(clear_rgb_files):
|
| 670 |
+
try:
|
| 671 |
+
clear_rgb_imgs.append(imageio.v2.imread(os.path.join(clear_rgb_dir_path, img_path)))
|
| 672 |
+
except Exception as e:
|
| 673 |
+
print(f"error reading image at path:{img_path}_{clear_rgb_dir_path}")
|
| 674 |
+
print(f"error mrssage:{str(e)}")
|
| 675 |
+
gotit = False
|
| 676 |
+
return [], gotit
|
| 677 |
+
|
| 678 |
+
for i, event_path in enumerate(event_files):
|
| 679 |
+
try:
|
| 680 |
+
event_imgs.append(read_input(os.path.join(event_dir_path, event_path), self.representation))
|
| 681 |
+
except Exception as e:
|
| 682 |
+
print(f"error reading event at path:{event_path}_{event_dir_path}")
|
| 683 |
+
print(f"error mrssage:{str(e)}")
|
| 684 |
+
gotit = False
|
| 685 |
+
return [], gotit
|
| 686 |
+
|
| 687 |
+
rgbs = np.stack(rgb_imgs)
|
| 688 |
+
clear_rgbs = np.stack(clear_rgb_imgs)
|
| 689 |
+
events = np.stack(event_imgs)
|
| 690 |
+
annot_dict = np.load(npy_path, allow_pickle=True).item()
|
| 691 |
+
if not self.if_test:
|
| 692 |
+
traj_2d = annot_dict["coords"][:,random_id: random_id+24]
|
| 693 |
+
visibility = annot_dict["visibility"][:,random_id: random_id+24]
|
| 694 |
+
else:
|
| 695 |
+
traj_2d = annot_dict["coords"][:,:-1]
|
| 696 |
+
visibility = annot_dict["visibility"][:,:-1]
|
| 697 |
+
|
| 698 |
+
assert self.seq_len == len(rgbs) == len(clear_rgbs)
|
| 699 |
+
|
| 700 |
+
traj_2d = np.transpose(traj_2d, (1, 0, 2)) # N, T, 2 -> T, N, 2
|
| 701 |
+
visibility = np.transpose(np.logical_not(visibility), (1, 0)) # N, T -> T, N
|
| 702 |
+
if self.use_augs:
|
| 703 |
+
print("new kubric dataset can't use augs!!!")
|
| 704 |
+
# rgbs, events, traj_2d, visibility = self.add_photometric_augs(rgbs, events, traj_2d, visibility)
|
| 705 |
+
rgbs, events, traj_2d = self.add_spatial_augs(rgbs, events, traj_2d, visibility)
|
| 706 |
+
else:
|
| 707 |
+
rgbs, events, traj_2d, clear_rgbs = self.crop(rgbs, events, traj_2d, clear_rgbs=clear_rgbs)
|
| 708 |
+
|
| 709 |
+
visibility[traj_2d[:, :, 0] > self.crop_size[1] - 1] = False
|
| 710 |
+
visibility[traj_2d[:, :, 0] < 0] = False
|
| 711 |
+
visibility[traj_2d[:, :, 1] > self.crop_size[0] - 1] = False
|
| 712 |
+
visibility[traj_2d[:, :, 1] < 0] = False
|
| 713 |
+
|
| 714 |
+
visibility = torch.from_numpy(visibility)
|
| 715 |
+
traj_2d = torch.from_numpy(traj_2d)
|
| 716 |
+
|
| 717 |
+
crop_tensor = torch.tensor(self.crop_size).flip(0)[None, None] / 2.0
|
| 718 |
+
close_pts_inds = torch.all(
|
| 719 |
+
torch.linalg.vector_norm(traj_2d[..., :2] - crop_tensor, dim=-1) < 1000.0,
|
| 720 |
+
dim=0,
|
| 721 |
+
)
|
| 722 |
+
traj_2d = traj_2d[:, close_pts_inds]
|
| 723 |
+
visibility = visibility[:, close_pts_inds]
|
| 724 |
+
|
| 725 |
+
visibile_pts_first_frame_inds = (visibility[0]).nonzero(as_tuple=False)
|
| 726 |
+
|
| 727 |
+
if self.sample_vis_1st_frame:
|
| 728 |
+
visiblile_pts_inds = visibile_pts_first_frame_inds
|
| 729 |
+
else:
|
| 730 |
+
visiblile_pts_mid_frame_inds = (visibility[self.seq_len // 2]).nonzero(as_tuple=False)
|
| 731 |
+
visibile_pts_last_frame_inds = (visibility[self.seq_len - 1]).nonzero(as_tuple=False)
|
| 732 |
+
visiblile_pts_inds = torch.cat((visibile_pts_first_frame_inds, visiblile_pts_mid_frame_inds, visibile_pts_last_frame_inds), dim=0)
|
| 733 |
+
|
| 734 |
+
point_inds = torch.randperm(len(visiblile_pts_inds))[:self.traj_per_sample]
|
| 735 |
+
if len(point_inds) < self.traj_per_sample and not self.if_test:
|
| 736 |
+
print(seq_name, "get point num ", len(point_inds), "less than ", self.traj_per_sample, " random_id is", random_id)
|
| 737 |
+
gotit = False
|
| 738 |
+
# shutil.rmtree(seq_name)
|
| 739 |
+
|
| 740 |
+
if self.choose_long_point:
|
| 741 |
+
distance = np.linalg.norm(traj_2d[-1, visiblile_pts_inds, :] - traj_2d[0, visiblile_pts_inds, :], axis=-1)[:,0]
|
| 742 |
+
weight = distance / np.sum(distance)
|
| 743 |
+
point_inds = torch.tensor(np.random.choice(len(distance), size=self.traj_per_sample, p=weight))
|
| 744 |
+
visible_inds_sampled = visiblile_pts_inds[point_inds]
|
| 745 |
+
# visible_inds_sampled = visiblile_pts_inds[np.argsort(distance, axis=0)[-self.traj_per_sample:,0]]
|
| 746 |
+
else:
|
| 747 |
+
visible_inds_sampled = visiblile_pts_inds[point_inds]
|
| 748 |
+
|
| 749 |
+
if len(visible_inds_sampled.shape) == 2:
|
| 750 |
+
visible_inds_sampled = visible_inds_sampled.squeeze(1)
|
| 751 |
+
trajs = traj_2d[:, visible_inds_sampled].float()
|
| 752 |
+
visibles = visibility[:, visible_inds_sampled]
|
| 753 |
+
valid = torch.ones((self.seq_len, self.traj_per_sample))
|
| 754 |
+
|
| 755 |
+
rgbs = torch.from_numpy(np.stack(rgbs)).permute(0, 3, 1, 2).float()
|
| 756 |
+
clear_rgbs = torch.from_numpy(np.stack(clear_rgbs)).permute(0, 3, 1, 2).float()
|
| 757 |
+
events = torch.from_numpy(np.stack(events)).float()
|
| 758 |
+
if "event_stack" not in self.representation:
|
| 759 |
+
events = events.permute(0, 3, 1, 2)
|
| 760 |
+
# events = torch.from_numpy(np.stack(events)).permute(0, 3, 1, 2).float()
|
| 761 |
+
seqs = torch.ones((self.seq_len, 1, self.crop_size[0], self.crop_size[1]))
|
| 762 |
+
img_ifnew = torch.Tensor(img_ifnew)
|
| 763 |
+
sample = FrameEventData(
|
| 764 |
+
video = rgbs,
|
| 765 |
+
events = events,
|
| 766 |
+
segmentation=seqs,
|
| 767 |
+
trajectory=trajs,
|
| 768 |
+
visibility=visibles,
|
| 769 |
+
valid=valid,
|
| 770 |
+
seq_name=seq_name,
|
| 771 |
+
img_ifnew=img_ifnew,
|
| 772 |
+
clear_video = clear_rgbs,
|
| 773 |
+
)
|
| 774 |
+
assert sample.img_ifnew is not None, f"发现空值样本"
|
| 775 |
+
return sample, gotit
|
| 776 |
+
|
| 777 |
+
def __len__(self):
|
| 778 |
+
return len(self.seq_names)
|
LFE_TAP/evaluator/__pycache__/evaluation_pred.cpython-38.pyc
ADDED
|
Binary file (5.1 kB). View file
|
|
|
LFE_TAP/evaluator/__pycache__/evaluation_pred.cpython-39.pyc
ADDED
|
Binary file (4.63 kB). View file
|
|
|
LFE_TAP/evaluator/__pycache__/evaluator.cpython-38.pyc
ADDED
|
Binary file (9.38 kB). View file
|
|
|
LFE_TAP/evaluator/__pycache__/evaluator.cpython-39.pyc
ADDED
|
Binary file (9.36 kB). View file
|
|
|
LFE_TAP/evaluator/__pycache__/prediction_long.cpython-38.pyc
ADDED
|
Binary file (8.3 kB). View file
|
|
|
LFE_TAP/evaluator/__pycache__/prediction_long.cpython-39.pyc
ADDED
|
Binary file (8.63 kB). View file
|
|
|
LFE_TAP/evaluator/evaluation_pred.py
ADDED
|
@@ -0,0 +1,184 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
from tqdm import tqdm
|
| 5 |
+
from typing import Tuple
|
| 6 |
+
|
| 7 |
+
from LFE_TAP.models.tapfe import TAPFormer
|
| 8 |
+
from LFE_TAP.utils.model_utils import get_points_on_a_grid, get_sift_sampled_pts, get_uniformly_sampled_pts, normalize_voxels
|
| 9 |
+
|
| 10 |
+
class EvaluationPredictor(torch.nn.Module):
|
| 11 |
+
def __init__(
|
| 12 |
+
self,
|
| 13 |
+
model: TAPFormer,
|
| 14 |
+
interp_shape: Tuple[int, int] = [384, 512],
|
| 15 |
+
grid_size: int = 5,
|
| 16 |
+
local_grid_size: int = 8,
|
| 17 |
+
single_point: bool = True,
|
| 18 |
+
sift_size: int = 0,
|
| 19 |
+
num_uniformly_sampled_pts: int = 0,
|
| 20 |
+
n_iters: int = 6,
|
| 21 |
+
local_extent: int = 50,
|
| 22 |
+
if_test: bool = False,
|
| 23 |
+
) -> None:
|
| 24 |
+
super(EvaluationPredictor, self).__init__()
|
| 25 |
+
self.model = model
|
| 26 |
+
self.interp_shape = interp_shape
|
| 27 |
+
self.grid_size = grid_size
|
| 28 |
+
self.local_grid_size = local_grid_size
|
| 29 |
+
self.single_point = single_point
|
| 30 |
+
self.sift_size = sift_size
|
| 31 |
+
self.num_uniformly_sampled_pts = num_uniformly_sampled_pts
|
| 32 |
+
self.n_iters = n_iters
|
| 33 |
+
self.local_extent = local_extent
|
| 34 |
+
self.if_test = if_test
|
| 35 |
+
self.model.eval()
|
| 36 |
+
|
| 37 |
+
def forward(self, video, events, queries, img_ifnew=None):
|
| 38 |
+
if queries is None and self.grid_size > 0:
|
| 39 |
+
grid_pts = get_points_on_a_grid(
|
| 40 |
+
self.grid_size, self.interp_shape, device="cuda"
|
| 41 |
+
)
|
| 42 |
+
queries = torch.cat(
|
| 43 |
+
[torch.ones_like(grid_pts[:, :, :1]) * 0, grid_pts],
|
| 44 |
+
dim=2,
|
| 45 |
+
)
|
| 46 |
+
self.grid_size = 0
|
| 47 |
+
|
| 48 |
+
queries = queries.clone()
|
| 49 |
+
B, T, C_r, H, W = video.shape
|
| 50 |
+
C_e = events[0].shape[1]
|
| 51 |
+
B, N, D = queries.shape
|
| 52 |
+
device = queries.device
|
| 53 |
+
|
| 54 |
+
assert D == 3
|
| 55 |
+
assert B == 1
|
| 56 |
+
|
| 57 |
+
interp_shape = self.interp_shape
|
| 58 |
+
|
| 59 |
+
if isinstance(video, torch.Tensor) and isinstance(events, torch.Tensor):
|
| 60 |
+
video = video.reshape(B * T, C_r, H, W)
|
| 61 |
+
events = events.reshape(B * T, C_e, H, W)
|
| 62 |
+
video = F.interpolate(video, tuple(interp_shape), mode="bilinear", align_corners=True)
|
| 63 |
+
events = F.interpolate(events, tuple(interp_shape), mode="bilinear", align_corners=True)
|
| 64 |
+
video = video.reshape(B, T, C_r, interp_shape[0], interp_shape[1])
|
| 65 |
+
events = events.reshape(B, T, C_e, interp_shape[0], interp_shape[1])
|
| 66 |
+
|
| 67 |
+
queries[:, :, 1] *= (interp_shape[1] - 1) / (W - 1)
|
| 68 |
+
queries[:, :, 2] *= (interp_shape[0] - 1) / (H - 1)
|
| 69 |
+
|
| 70 |
+
if self.single_point:
|
| 71 |
+
traj_e = torch.zeros((B, T, N, 2), device=device)
|
| 72 |
+
vis_e = torch.zeros((B, T, N), device=device)
|
| 73 |
+
conf_e = torch.zeros((B, T, N), device=device)
|
| 74 |
+
|
| 75 |
+
for pind in range((N)):
|
| 76 |
+
querie = queries[:, pind : pind + 1]
|
| 77 |
+
traj_e_pind, vis_e_pind, conf_e_pind = self._process_one_point(video[:,:], events[:,:], querie, img_ifnew=img_ifnew)
|
| 78 |
+
traj_e[:, :, pind:pind+1] = traj_e_pind[:, :, :1]
|
| 79 |
+
vis_e[:, :, pind:pind+1] = vis_e_pind[:, :, :1]
|
| 80 |
+
conf_e[:, :, pind:pind+1] = conf_e_pind[:, :, :1]
|
| 81 |
+
else:
|
| 82 |
+
if self.grid_size > 0:
|
| 83 |
+
xy = get_points_on_a_grid(self.grid_size, video.shape[3:])
|
| 84 |
+
xy = torch.cat([torch.zeros_like(xy[:, :, :1]), xy], dim=2).to(device)
|
| 85 |
+
queries = torch.cat([queries, xy], dim=1)
|
| 86 |
+
|
| 87 |
+
if self.num_uniformly_sampled_pts > 0:
|
| 88 |
+
xy = get_uniformly_sampled_pts(
|
| 89 |
+
self.num_uniformly_sampled_pts,
|
| 90 |
+
video.shape[1],
|
| 91 |
+
video.shape[3:],
|
| 92 |
+
device=device,
|
| 93 |
+
)
|
| 94 |
+
queries = torch.cat([queries, xy], dim=1)
|
| 95 |
+
|
| 96 |
+
sift_size = self.sift_size
|
| 97 |
+
if sift_size > 0:
|
| 98 |
+
xy = get_sift_sampled_pts(video, sift_size, T, [H, W], device=device)
|
| 99 |
+
if xy.shape[1] == sift_size:
|
| 100 |
+
queries = torch.cat([queries, xy], dim=1) #
|
| 101 |
+
else:
|
| 102 |
+
sift_size = 0
|
| 103 |
+
|
| 104 |
+
preds = self.model(rgbs=video, events=events, queries=queries, iters=self.n_iters, img_ifnew=img_ifnew)
|
| 105 |
+
traj_e, vis_e = preds[0], preds[1]
|
| 106 |
+
|
| 107 |
+
conf_e = preds[2]
|
| 108 |
+
if (
|
| 109 |
+
sift_size > 0
|
| 110 |
+
or self.grid_size > 0
|
| 111 |
+
or self.num_uniformly_sampled_pts > 0
|
| 112 |
+
):
|
| 113 |
+
traj_e = traj_e[:, :, : -self.grid_size**2 - sift_size - self.num_uniformly_sampled_pts]
|
| 114 |
+
vis_e = vis_e[:, :, : -self.grid_size**2 - sift_size - self.num_uniformly_sampled_pts]
|
| 115 |
+
if conf_e is not None:
|
| 116 |
+
conf_e = conf_e[:, :, : -self.grid_size**2 - sift_size - self.num_uniformly_sampled_pts]
|
| 117 |
+
|
| 118 |
+
# if conf_e is not None:
|
| 119 |
+
# vis_e = vis_e * conf_e
|
| 120 |
+
|
| 121 |
+
if self.if_test:
|
| 122 |
+
thr = 0.9
|
| 123 |
+
vis_e = vis_e > thr
|
| 124 |
+
for i in range(len(queries)):
|
| 125 |
+
queries_t = queries[i, : traj_e.size(2), 0].to(torch.int64)
|
| 126 |
+
arange = torch.arange(0, len(queries_t))
|
| 127 |
+
|
| 128 |
+
# overwrite the predictions with the query points
|
| 129 |
+
traj_e[i, queries_t, arange] = queries[i, : traj_e.size(2), 1:]
|
| 130 |
+
|
| 131 |
+
# correct visibilities, the query points should be visible
|
| 132 |
+
vis_e[i, queries_t, arange] = True
|
| 133 |
+
|
| 134 |
+
traj_e *= traj_e.new_tensor(
|
| 135 |
+
[(W - 1) / (self.interp_shape[1] - 1), (H - 1) / (self.interp_shape[0] - 1)]
|
| 136 |
+
)
|
| 137 |
+
|
| 138 |
+
return traj_e, vis_e, conf_e
|
| 139 |
+
|
| 140 |
+
def _process_one_point(self, video, events, query, img_ifnew):
|
| 141 |
+
B, T, C, H, W = video.shape
|
| 142 |
+
device = query.device
|
| 143 |
+
if self.local_grid_size > 0:
|
| 144 |
+
xy_target = get_points_on_a_grid(self.local_grid_size, (self.local_extent, self.local_extent), [query[0,0,2].item(), query[0,0,1].item()])
|
| 145 |
+
xy_target = torch.cat([torch.zeros_like(xy_target[:, :, :1]), xy_target], dim=2).to(device)
|
| 146 |
+
query = torch.cat([query, xy_target], dim=1)
|
| 147 |
+
|
| 148 |
+
if self.grid_size > 0:
|
| 149 |
+
xy = get_points_on_a_grid(self.grid_size, video.shape[3:])
|
| 150 |
+
xy = torch.cat([torch.zeros_like(xy[:, :, :1]), xy], dim=2).to(device)
|
| 151 |
+
query = torch.cat([query, xy], dim=1)
|
| 152 |
+
|
| 153 |
+
sift_size = self.sift_size
|
| 154 |
+
if sift_size > 0:
|
| 155 |
+
xy = get_sift_sampled_pts(video, sift_size, T, [H, W], device=device)
|
| 156 |
+
sift_size = xy.shape[1]
|
| 157 |
+
if sift_size > 0:
|
| 158 |
+
query = torch.cat([query, xy], dim=1) #
|
| 159 |
+
|
| 160 |
+
num_uniformly_sampled_pts = self.sift_size - sift_size
|
| 161 |
+
if num_uniformly_sampled_pts > 0:
|
| 162 |
+
xy2 = get_uniformly_sampled_pts(
|
| 163 |
+
num_uniformly_sampled_pts,
|
| 164 |
+
video.shape[1],
|
| 165 |
+
video.shape[3:],
|
| 166 |
+
device=device,
|
| 167 |
+
)
|
| 168 |
+
query = torch.cat([query, xy2], dim=1) #
|
| 169 |
+
|
| 170 |
+
if self.num_uniformly_sampled_pts > 0:
|
| 171 |
+
xy = get_uniformly_sampled_pts(
|
| 172 |
+
self.num_uniformly_sampled_pts,
|
| 173 |
+
video.shape[1],
|
| 174 |
+
video.shape[3:],
|
| 175 |
+
device=device,
|
| 176 |
+
)
|
| 177 |
+
query = torch.cat([query, xy], dim=1)
|
| 178 |
+
|
| 179 |
+
traj_e_pind, vis_e_pind, conf_e_pind = self.model(
|
| 180 |
+
rgbs=video, events=events, queries=query, iters=self.n_iters, img_ifnew=img_ifnew
|
| 181 |
+
)
|
| 182 |
+
|
| 183 |
+
return traj_e_pind[..., :2], vis_e_pind, conf_e_pind
|
| 184 |
+
|
LFE_TAP/evaluator/evaluator.py
ADDED
|
@@ -0,0 +1,351 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import numpy as np
|
| 3 |
+
import torch
|
| 4 |
+
import logging
|
| 5 |
+
from tqdm import tqdm
|
| 6 |
+
from typing import Optional, Mapping
|
| 7 |
+
from torch.utils.tensorboard import SummaryWriter
|
| 8 |
+
from LFE_TAP.models.tapfe import TAPFormer
|
| 9 |
+
from LFE_TAP.utils.visualizer import Visualizer
|
| 10 |
+
from LFE_TAP.utils.dataset_utils import dataclass_to_cuda_
|
| 11 |
+
|
| 12 |
+
def get_error(est_data, gt_data):
|
| 13 |
+
# discard gt which happen after last est_data
|
| 14 |
+
# gt_data = gt_data[gt_data[:, 0] <= est_data[-1, 0]]
|
| 15 |
+
|
| 16 |
+
est_t, est_x, est_y = est_data.T
|
| 17 |
+
gt_t, gt_x, gt_y = gt_data.T
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
if len(gt_t) == 0 or len(est_t) == 0:
|
| 21 |
+
return [], [], []
|
| 22 |
+
|
| 23 |
+
if len(est_t) < 2:
|
| 24 |
+
return gt_t, np.array([0]), np.array([0])
|
| 25 |
+
|
| 26 |
+
# find samples which have dt > threshold
|
| 27 |
+
error_x = np.interp(gt_t, est_t, est_x) - gt_x
|
| 28 |
+
error_y = np.interp(gt_t, est_t, est_y) - gt_y
|
| 29 |
+
|
| 30 |
+
return gt_t, error_x, error_y
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def compareTracks(tracks_pred, tracks_gt, max_distance=10):
|
| 34 |
+
pred_ids = np.unique(tracks_pred[:, 0])
|
| 35 |
+
gt_ids = np.unique(tracks_gt[:, 0])
|
| 36 |
+
pred_datas = {i: tracks_pred[tracks_pred[:, 0] == i, 1:] for i in pred_ids}
|
| 37 |
+
gt_datas = {i: tracks_gt[tracks_gt[:, 0] == i, 1:] for i in gt_ids}
|
| 38 |
+
|
| 39 |
+
error_datas = np.zeros(shape=(0, 4))
|
| 40 |
+
errors = np.zeros(shape=(0, 2))
|
| 41 |
+
|
| 42 |
+
for track_id, pred_data in tqdm(pred_datas.items(), disable=True):
|
| 43 |
+
gt_data = gt_datas[track_id]
|
| 44 |
+
init_time = gt_data[0, 0]
|
| 45 |
+
gt_data[:, 0] -= init_time
|
| 46 |
+
pred_data[:, 0] -= init_time
|
| 47 |
+
pred_data = np.concatenate((gt_data[0, :].reshape(1, -1), pred_data), axis=0)
|
| 48 |
+
|
| 49 |
+
gt_t, error_x, error_y = get_error(pred_data, gt_data)
|
| 50 |
+
|
| 51 |
+
if len(gt_t) != 0:
|
| 52 |
+
ids = (track_id * np.ones_like(error_x)).astype(int)
|
| 53 |
+
added_data = np.stack([ids, gt_t, error_x, error_y]).T
|
| 54 |
+
error_euclidean = np.sqrt(added_data[:, 2]**2 + added_data[:, 3]**2)
|
| 55 |
+
error_euclidean[0] = 0
|
| 56 |
+
idxs = np.where(error_euclidean > max_distance)[0]
|
| 57 |
+
if len(idxs) > 0:
|
| 58 |
+
feature_age = (gt_t[int(idxs[0])] - gt_t[0]) / (gt_t[-1] - gt_t[0])
|
| 59 |
+
if int(idxs[0]) <= 1:
|
| 60 |
+
continue
|
| 61 |
+
error_euclidean = error_euclidean[:idxs[0]]
|
| 62 |
+
error = np.mean(error_euclidean)
|
| 63 |
+
else:
|
| 64 |
+
feature_age = 1
|
| 65 |
+
error = np.mean(error_euclidean)
|
| 66 |
+
|
| 67 |
+
errors = np.concatenate([errors, [[error, feature_age]]])
|
| 68 |
+
error_datas = np.concatenate([error_datas, added_data])
|
| 69 |
+
|
| 70 |
+
if errors.shape[0] == 0:
|
| 71 |
+
return [], [], [0, (gt_t[1]-gt_t[0])/(gt_t[-1] - gt_t[0]), 0]
|
| 72 |
+
|
| 73 |
+
mean_error = np.mean(errors, axis=0)
|
| 74 |
+
ecpect_FA = (errors.shape[0]/len(gt_ids)) * mean_error[1]
|
| 75 |
+
mean_error = np.concatenate((mean_error, [ecpect_FA]))
|
| 76 |
+
# print("abandon ", len(gt_ids) - errors.shape[0], " points, remain ", errors.shape[0], " points")
|
| 77 |
+
# print("Mean Error: {:.4f}, Mean Feature Age: {:.4f}, Ecppect Feature Age: {:.4f}".format(mean_error[0], mean_error[1], mean_error[2]))
|
| 78 |
+
|
| 79 |
+
return error_datas, errors, mean_error
|
| 80 |
+
|
| 81 |
+
|
| 82 |
+
def compute_tapvid_metrics(
|
| 83 |
+
query_points: np.ndarray,
|
| 84 |
+
gt_occluded: np.ndarray,
|
| 85 |
+
gt_tracks: np.ndarray,
|
| 86 |
+
pred_occluded: np.ndarray,
|
| 87 |
+
pred_tracks: np.ndarray,
|
| 88 |
+
query_mode: str,
|
| 89 |
+
) -> Mapping[str, np.ndarray]:
|
| 90 |
+
"""Computes TAP-Vid metrics (Jaccard, Pts. Within Thresh, Occ. Acc.)
|
| 91 |
+
See the TAP-Vid paper for details on the metric computation. All inputs are
|
| 92 |
+
given in raster coordinates. The first three arguments should be the direct
|
| 93 |
+
outputs of the reader: the 'query_points', 'occluded', and 'target_points'.
|
| 94 |
+
The paper metrics assume these are scaled relative to 256x256 images.
|
| 95 |
+
pred_occluded and pred_tracks are your algorithm's predictions.
|
| 96 |
+
This function takes a batch of inputs, and computes metrics separately for
|
| 97 |
+
each video. The metrics for the full benchmark are a simple mean of the
|
| 98 |
+
metrics across the full set of videos. These numbers are between 0 and 1,
|
| 99 |
+
but the paper multiplies them by 100 to ease reading.
|
| 100 |
+
Args:
|
| 101 |
+
query_points: The query points, an in the format [t, y, x]. Its size is
|
| 102 |
+
[b, n, 3], where b is the batch size and n is the number of queries
|
| 103 |
+
gt_occluded: A boolean array of shape [b, n, t], where t is the number
|
| 104 |
+
of frames. True indicates that the point is occluded.
|
| 105 |
+
gt_tracks: The target points, of shape [b, n, t, 2]. Each point is
|
| 106 |
+
in the format [x, y]
|
| 107 |
+
pred_occluded: A boolean array of predicted occlusions, in the same
|
| 108 |
+
format as gt_occluded.
|
| 109 |
+
pred_tracks: An array of track predictions from your algorithm, in the
|
| 110 |
+
same format as gt_tracks.
|
| 111 |
+
query_mode: Either 'first' or 'strided', depending on how queries are
|
| 112 |
+
sampled. If 'first', we assume the prior knowledge that all points
|
| 113 |
+
before the query point are occluded, and these are removed from the
|
| 114 |
+
evaluation.
|
| 115 |
+
Returns:
|
| 116 |
+
A dict with the following keys:
|
| 117 |
+
occlusion_accuracy: Accuracy at predicting occlusion.
|
| 118 |
+
pts_within_{x} for x in [1, 2, 4, 8, 16]: Fraction of points
|
| 119 |
+
predicted to be within the given pixel threshold, ignoring occlusion
|
| 120 |
+
prediction.
|
| 121 |
+
jaccard_{x} for x in [1, 2, 4, 8, 16]: Jaccard metric for the given
|
| 122 |
+
threshold
|
| 123 |
+
average_pts_within_thresh: average across pts_within_{x}
|
| 124 |
+
average_jaccard: average across jaccard_{x}
|
| 125 |
+
"""
|
| 126 |
+
|
| 127 |
+
metrics = {}
|
| 128 |
+
eye = np.eye(gt_tracks.shape[2], dtype=np.int32)
|
| 129 |
+
|
| 130 |
+
if query_mode == "first":
|
| 131 |
+
# evaluate frames after the query frame
|
| 132 |
+
query_frame_to_eval_frames = np.cumsum(eye, axis=1) - eye
|
| 133 |
+
elif query_mode == "strided":
|
| 134 |
+
# evaluate all frames except the query frame
|
| 135 |
+
query_frame_to_eval_frames = 1 - eye
|
| 136 |
+
else:
|
| 137 |
+
raise ValueError("Unknown query mode " + query_mode)
|
| 138 |
+
|
| 139 |
+
query_frame = query_points[..., 0]
|
| 140 |
+
query_frame = np.round(query_frame).astype(np.int32)
|
| 141 |
+
evaluation_points = query_frame_to_eval_frames[query_frame] > 0
|
| 142 |
+
|
| 143 |
+
# Occlusion accuracy is simply how often the predicted occlusion equals the
|
| 144 |
+
# ground truth.
|
| 145 |
+
occ_acc = np.sum(
|
| 146 |
+
np.equal(pred_occluded, gt_occluded) & evaluation_points,
|
| 147 |
+
axis=(1, 2),
|
| 148 |
+
) / np.sum(evaluation_points)
|
| 149 |
+
metrics["occlusion_accuracy"] = occ_acc
|
| 150 |
+
|
| 151 |
+
# Next, convert the predictions and ground truth positions into pixel
|
| 152 |
+
# coordinates.
|
| 153 |
+
visible = np.logical_not(gt_occluded)
|
| 154 |
+
pred_visible = np.logical_not(pred_occluded)
|
| 155 |
+
all_frac_within = []
|
| 156 |
+
all_jaccard = []
|
| 157 |
+
for thresh in [1, 2, 4, 8, 16]:
|
| 158 |
+
# True positives are points that are within the threshold and where both
|
| 159 |
+
# the prediction and the ground truth are listed as visible.
|
| 160 |
+
within_dist = np.sum(
|
| 161 |
+
np.square(pred_tracks - gt_tracks),
|
| 162 |
+
axis=-1,
|
| 163 |
+
) < np.square(thresh)
|
| 164 |
+
is_correct = np.logical_and(within_dist, visible)
|
| 165 |
+
|
| 166 |
+
# Compute the frac_within_threshold, which is the fraction of points
|
| 167 |
+
# within the threshold among points that are visible in the ground truth,
|
| 168 |
+
# ignoring whether they're predicted to be visible.
|
| 169 |
+
count_correct = np.sum(
|
| 170 |
+
is_correct & evaluation_points,
|
| 171 |
+
axis=(1, 2),
|
| 172 |
+
)
|
| 173 |
+
count_visible_points = np.sum(visible & evaluation_points, axis=(1, 2))
|
| 174 |
+
frac_correct = count_correct / count_visible_points
|
| 175 |
+
metrics["pts_within_" + str(thresh)] = frac_correct
|
| 176 |
+
all_frac_within.append(frac_correct)
|
| 177 |
+
|
| 178 |
+
true_positives = np.sum(
|
| 179 |
+
is_correct & pred_visible & evaluation_points, axis=(1, 2)
|
| 180 |
+
)
|
| 181 |
+
|
| 182 |
+
# The denominator of the jaccard metric is the true positives plus
|
| 183 |
+
# false positives plus false negatives. However, note that true positives
|
| 184 |
+
# plus false negatives is simply the number of points in the ground truth
|
| 185 |
+
# which is easier to compute than trying to compute all three quantities.
|
| 186 |
+
# Thus we just add the number of points in the ground truth to the number
|
| 187 |
+
# of false positives.
|
| 188 |
+
#
|
| 189 |
+
# False positives are simply points that are predicted to be visible,
|
| 190 |
+
# but the ground truth is not visible or too far from the prediction.
|
| 191 |
+
gt_positives = np.sum(visible & evaluation_points, axis=(1, 2))
|
| 192 |
+
false_positives = (~visible) & pred_visible
|
| 193 |
+
false_positives = false_positives | ((~within_dist) & pred_visible)
|
| 194 |
+
false_positives = np.sum(false_positives & evaluation_points, axis=(1, 2))
|
| 195 |
+
jaccard = true_positives / (gt_positives + false_positives)
|
| 196 |
+
metrics["jaccard_" + str(thresh)] = jaccard
|
| 197 |
+
all_jaccard.append(jaccard)
|
| 198 |
+
metrics["average_jaccard"] = np.mean(
|
| 199 |
+
np.stack(all_jaccard, axis=1),
|
| 200 |
+
axis=1,
|
| 201 |
+
)
|
| 202 |
+
metrics["average_pts_within_thresh"] = np.mean(
|
| 203 |
+
np.stack(all_frac_within, axis=1),
|
| 204 |
+
axis=1,
|
| 205 |
+
)
|
| 206 |
+
return metrics
|
| 207 |
+
|
| 208 |
+
class Evaluator:
|
| 209 |
+
def __init__(self, output_dir) -> None:
|
| 210 |
+
self.output_dir = output_dir
|
| 211 |
+
os.makedirs(output_dir, exist_ok=True)
|
| 212 |
+
|
| 213 |
+
def compute_metrics(self, metrics, sample, pred_trajectory, dataset_name):
|
| 214 |
+
if isinstance(pred_trajectory, tuple):
|
| 215 |
+
pred_trajectory, pred_visibility = pred_trajectory
|
| 216 |
+
else:
|
| 217 |
+
pred_visibility = None
|
| 218 |
+
|
| 219 |
+
if "kubric" in dataset_name:
|
| 220 |
+
B, T, N, D = sample.trajectory.shape
|
| 221 |
+
traj = sample.trajectory.clone()
|
| 222 |
+
thr = 0.6
|
| 223 |
+
|
| 224 |
+
if pred_visibility is None:
|
| 225 |
+
logging.warning("visibility is NONE")
|
| 226 |
+
pred_visibility = torch.zeros_like(sample.visibility)
|
| 227 |
+
|
| 228 |
+
if not pred_visibility.dtype == torch.bool:
|
| 229 |
+
pred_visibility = pred_visibility > thr
|
| 230 |
+
|
| 231 |
+
query_points = torch.cat(
|
| 232 |
+
[
|
| 233 |
+
torch.zeros_like(sample.trajectory[:, 0, :, :1]),
|
| 234 |
+
sample.trajectory[:, 0],
|
| 235 |
+
],
|
| 236 |
+
dim=2,
|
| 237 |
+
).cpu().numpy()
|
| 238 |
+
|
| 239 |
+
pred_visibility = pred_visibility[:, :, :N]
|
| 240 |
+
pred_trajectory = pred_trajectory[:, :, :N]
|
| 241 |
+
|
| 242 |
+
gt_tracks = traj.permute(0, 2, 1, 3).cpu().numpy()
|
| 243 |
+
gt_occluded = (
|
| 244 |
+
torch.logical_not(sample.visibility.clone().permute(0, 2, 1))
|
| 245 |
+
.cpu()
|
| 246 |
+
.numpy()
|
| 247 |
+
)
|
| 248 |
+
|
| 249 |
+
pred_occluded = (
|
| 250 |
+
torch.logical_not(pred_visibility.clone().permute(0, 2, 1))
|
| 251 |
+
.cpu()
|
| 252 |
+
.numpy()
|
| 253 |
+
)
|
| 254 |
+
pred_tracks = pred_trajectory.permute(0, 2, 1, 3).cpu().numpy()
|
| 255 |
+
|
| 256 |
+
out_metrics = compute_tapvid_metrics(
|
| 257 |
+
query_points,
|
| 258 |
+
gt_occluded,
|
| 259 |
+
gt_tracks,
|
| 260 |
+
pred_occluded,
|
| 261 |
+
pred_tracks,
|
| 262 |
+
query_mode="strided" if "strided" in dataset_name else "first",
|
| 263 |
+
)
|
| 264 |
+
|
| 265 |
+
metrics[sample.seq_name[0]] = out_metrics
|
| 266 |
+
for metric_name in out_metrics.keys():
|
| 267 |
+
if "avg" not in metrics:
|
| 268 |
+
metrics["avg"] = {}
|
| 269 |
+
metrics["avg"][metric_name] = np.mean(
|
| 270 |
+
[v[metric_name] for k, v in metrics.items() if k != "avg"]
|
| 271 |
+
)
|
| 272 |
+
|
| 273 |
+
logging.info(f"Metrics: {out_metrics}")
|
| 274 |
+
logging.info(f"avg: {metrics['avg']}")
|
| 275 |
+
print("metrics", out_metrics)
|
| 276 |
+
print("avg", metrics["avg"])
|
| 277 |
+
elif "EC" or "EDS" in dataset_name:
|
| 278 |
+
traj = sample.trajectory.copy()
|
| 279 |
+
|
| 280 |
+
mean_err_avg = []
|
| 281 |
+
for i in range(1, 31):
|
| 282 |
+
error_datas, errors, mean_err = compareTracks(pred_trajectory.cpu().numpy(), traj[0], sample.segmentation, i)
|
| 283 |
+
mean_err_avg.append(mean_err)
|
| 284 |
+
mean_err_avg = np.stack(mean_err_avg)
|
| 285 |
+
mean_err_avg = np.mean(mean_err_avg, axis=0)
|
| 286 |
+
print(sample.seq_name, "deep_ev mean error:", mean_err_avg[0], " mean age:", mean_err_avg[1], "expect age:", mean_err_avg[2])
|
| 287 |
+
|
| 288 |
+
|
| 289 |
+
@torch.no_grad()
|
| 290 |
+
def evaluate_sequence(
|
| 291 |
+
self,
|
| 292 |
+
model,
|
| 293 |
+
test_dataloader: torch.utils.data.DataLoader,
|
| 294 |
+
dataset_name: str,
|
| 295 |
+
visualize_every: int = 1,
|
| 296 |
+
writer: Optional[SummaryWriter] = None,
|
| 297 |
+
step: Optional[int] = 0,
|
| 298 |
+
):
|
| 299 |
+
metrics = {}
|
| 300 |
+
|
| 301 |
+
vis = Visualizer(self.output_dir, fps=10 if "kubric" in dataset_name else 50)
|
| 302 |
+
|
| 303 |
+
for ind, sample in enumerate(tqdm(test_dataloader)):
|
| 304 |
+
if isinstance(sample, tuple):
|
| 305 |
+
sample, gotit = sample
|
| 306 |
+
if not all(gotit):
|
| 307 |
+
print(f"Skipping sample {ind} because gotit is {gotit}")
|
| 308 |
+
continue
|
| 309 |
+
|
| 310 |
+
if torch.cuda.is_available():
|
| 311 |
+
if "kubric" in dataset_name:
|
| 312 |
+
dataclass_to_cuda_(sample)
|
| 313 |
+
device = torch.device("cuda")
|
| 314 |
+
else:
|
| 315 |
+
device = torch.device("cpu")
|
| 316 |
+
|
| 317 |
+
if "kubric" in dataset_name:
|
| 318 |
+
queries = torch.cat(
|
| 319 |
+
[
|
| 320 |
+
torch.zeros_like(sample.trajectory[:, 0, :, :1]),
|
| 321 |
+
sample.trajectory[:, 0],
|
| 322 |
+
],
|
| 323 |
+
dim=2,
|
| 324 |
+
).to(device)
|
| 325 |
+
elif "EC" or "EDS" in dataset_name:
|
| 326 |
+
queries = sample.query_points
|
| 327 |
+
queries = queries.to(device)
|
| 328 |
+
|
| 329 |
+
pred_tracks = model(sample.video, sample.events, queries)
|
| 330 |
+
|
| 331 |
+
if dataset_name == "EC" or dataset_name == "EDS":
|
| 332 |
+
seq_name = sample.seq_name[0]
|
| 333 |
+
else:
|
| 334 |
+
seq_name = str(ind)
|
| 335 |
+
if ind % visualize_every == 0:
|
| 336 |
+
vis.visualize(
|
| 337 |
+
sample.video if isinstance(sample.video, torch.Tensor) else torch.from_numpy(sample.video).float(),
|
| 338 |
+
pred_tracks[0],
|
| 339 |
+
pred_tracks[1] > 0.8,
|
| 340 |
+
filename=dataset_name + "_" + seq_name,
|
| 341 |
+
writer=writer,
|
| 342 |
+
step=step,
|
| 343 |
+
)
|
| 344 |
+
self.compute_metrics(metrics, sample, pred_tracks, dataset_name)
|
| 345 |
+
return metrics
|
| 346 |
+
|
| 347 |
+
|
| 348 |
+
|
| 349 |
+
|
| 350 |
+
|
| 351 |
+
|
LFE_TAP/evaluator/prediction.py
ADDED
|
@@ -0,0 +1,311 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import time
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
import numpy as np
|
| 5 |
+
|
| 6 |
+
from LFE_TAP.models.tapfe import TAPFormer, posenc
|
| 7 |
+
from LFE_TAP.utils.model_utils import get_track_feat, normalize_voxels
|
| 8 |
+
from LFE_TAP.models.embeddings import get_1d_sincos_pos_embed_from_grid
|
| 9 |
+
|
| 10 |
+
torch.manual_seed(0)
|
| 11 |
+
starter, ender = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
|
| 12 |
+
starter1, ender1 = torch.cuda.Event(enable_timing=True), torch.cuda.Event(enable_timing=True)
|
| 13 |
+
|
| 14 |
+
class TAPFormer_online(TAPFormer):
|
| 15 |
+
def __init__(self, trained_model=None, window_size=16, stride=4, corr_radius=3, corr_levels=3,
|
| 16 |
+
backbone="basic", num_heads=8, hidden_size=384, space_depth=3, time_depth=3):
|
| 17 |
+
"""
|
| 18 |
+
Initialize TAPFormer_online model for memory-efficient inference.
|
| 19 |
+
|
| 20 |
+
Args:
|
| 21 |
+
trained_model: (optional) A trained TAPFormer model instance. If provided, parameters will be copied from it.
|
| 22 |
+
window_size, stride, corr_radius, corr_levels, backbone, num_heads, hidden_size,
|
| 23 |
+
space_depth, time_depth: Model configuration parameters. Used when trained_model is None.
|
| 24 |
+
"""
|
| 25 |
+
# If trained_model is provided, use it to initialize (backward compatibility)
|
| 26 |
+
if trained_model is not None:
|
| 27 |
+
super(TAPFormer_online, self).__init__(
|
| 28 |
+
window_size=trained_model.window_size,
|
| 29 |
+
stride=trained_model.stride,
|
| 30 |
+
corr_radius=trained_model.corr_radius,
|
| 31 |
+
corr_levels=trained_model.corr_levels,
|
| 32 |
+
backbone=trained_model.backbone,
|
| 33 |
+
hidden_size=trained_model.hidden_size,
|
| 34 |
+
space_depth=trained_model.space_depth,
|
| 35 |
+
time_depth=trained_model.time_depth
|
| 36 |
+
)
|
| 37 |
+
self.fusion_block = trained_model.fusion_block
|
| 38 |
+
self.updateformer2 = trained_model.updateformer2
|
| 39 |
+
self.corr_mlp = trained_model.corr_mlp
|
| 40 |
+
else:
|
| 41 |
+
# Direct initialization from parameters
|
| 42 |
+
super(TAPFormer_online, self).__init__(
|
| 43 |
+
window_size=window_size,
|
| 44 |
+
stride=stride,
|
| 45 |
+
corr_radius=corr_radius,
|
| 46 |
+
corr_levels=corr_levels,
|
| 47 |
+
backbone=backbone,
|
| 48 |
+
num_heads=num_heads,
|
| 49 |
+
hidden_size=hidden_size,
|
| 50 |
+
space_depth=space_depth,
|
| 51 |
+
time_depth=time_depth
|
| 52 |
+
)
|
| 53 |
+
|
| 54 |
+
time_grid = torch.linspace(0, self.window_size - 1, self.window_size).reshape(1, self.window_size, 1)
|
| 55 |
+
self.register_buffer(
|
| 56 |
+
"time_emb", get_1d_sincos_pos_embed_from_grid(self.input_dim, time_grid[0])
|
| 57 |
+
)
|
| 58 |
+
self.corr_pyramid = []
|
| 59 |
+
|
| 60 |
+
@torch.no_grad()
|
| 61 |
+
def forward(self, rgbs, events, queries, iters=6, img_ifnew=None, feat_init=None, interp_shape=(384, 512), is_train=False):
|
| 62 |
+
# starter.record()
|
| 63 |
+
if self.backbone == "image":
|
| 64 |
+
self.updateformer2 = self.updateformer
|
| 65 |
+
B, T, C, H, W = events.shape
|
| 66 |
+
_, N, _ = queries.shape
|
| 67 |
+
_, T, C_img, _, _ = rgbs.shape
|
| 68 |
+
S = self.window_size
|
| 69 |
+
step = S // 2
|
| 70 |
+
device = queries.device
|
| 71 |
+
|
| 72 |
+
queried_frames = queries[:, :, 0].long()
|
| 73 |
+
queried_coords = queries[..., 1:3]
|
| 74 |
+
queried_coords = queried_coords / self.stride
|
| 75 |
+
|
| 76 |
+
coords_predicted = torch.zeros((B, T, N, 2), device=device)
|
| 77 |
+
vis_predicted= torch.zeros((B, T, N), device=device)
|
| 78 |
+
conf_predicted = torch.zeros((B, T, N), device=device)
|
| 79 |
+
|
| 80 |
+
H_stride, W_stride = interp_shape[0] // self.stride, interp_shape[1] // self.stride
|
| 81 |
+
|
| 82 |
+
coords_init = queries[:, :, 1:].reshape(B, 1, N, 2).repeat(1, self.window_size, 1, 1) / float(self.stride)
|
| 83 |
+
|
| 84 |
+
vis_init = torch.zeros((B, S, N, 1), device=device).float()
|
| 85 |
+
conf_init = torch.zeros((B, S, N, 1), device=device).float()
|
| 86 |
+
coords_init = queried_coords.reshape(B, 1, N, 2).expand(B, S, N, 2).float()
|
| 87 |
+
|
| 88 |
+
num_windows = (T - S + step - 1) // step + 1
|
| 89 |
+
indices = range(0, step * num_windows, step)
|
| 90 |
+
|
| 91 |
+
fmaps_fusion = None
|
| 92 |
+
first_window = True
|
| 93 |
+
track_feat_pyramid = []
|
| 94 |
+
track_feat_support_pyramid = []
|
| 95 |
+
|
| 96 |
+
for ind in indices:
|
| 97 |
+
if ind > 0:
|
| 98 |
+
overlap = S - step
|
| 99 |
+
copy_over = (queried_frames < ind + overlap)[:, None, :, None] # B 1 N 1
|
| 100 |
+
coords_prev = coords_predicted[:, ind : ind + overlap] / self.stride
|
| 101 |
+
padding_tensor = coords_prev[:, -1:, :, :].expand(-1, step, -1, -1)
|
| 102 |
+
coords_prev = torch.cat([coords_prev, padding_tensor], dim=1)
|
| 103 |
+
|
| 104 |
+
vis_prev = vis_predicted[:, ind : ind + overlap, :, None].clone()
|
| 105 |
+
padding_tensor = vis_prev[:, -1:, :, :].expand(-1, step, -1, -1)
|
| 106 |
+
vis_prev = torch.cat([vis_prev, padding_tensor], dim=1)
|
| 107 |
+
|
| 108 |
+
conf_prev = conf_predicted[:, ind : ind + overlap, :, None].clone()
|
| 109 |
+
padding_tensor = conf_prev[:, -1:, :, :].expand(-1, step, -1, -1)
|
| 110 |
+
conf_prev = torch.cat([conf_prev, padding_tensor], dim=1)
|
| 111 |
+
|
| 112 |
+
coords_init = torch.where(copy_over.expand_as(coords_init), coords_prev, coords_init)
|
| 113 |
+
vis_init = torch.where(copy_over.expand_as(vis_init), vis_prev, vis_init)
|
| 114 |
+
conf_init = torch.where(copy_over.expand_as(conf_init), conf_prev, conf_init)
|
| 115 |
+
|
| 116 |
+
events_seq = events[:, ind : ind + S]
|
| 117 |
+
rgbs_seq = rgbs[:, ind : ind + S]
|
| 118 |
+
if img_ifnew is not None:
|
| 119 |
+
img_ifnew_seq = img_ifnew[ind : ind + S]
|
| 120 |
+
if not isinstance(events_seq, torch.Tensor):
|
| 121 |
+
events_seq = torch.from_numpy(events_seq).to(device).float()
|
| 122 |
+
rgbs_seq = torch.from_numpy(rgbs_seq).to(device).float()
|
| 123 |
+
events_seq = events_seq.contiguous()
|
| 124 |
+
rgbs_seq = rgbs_seq.contiguous()
|
| 125 |
+
|
| 126 |
+
if ind + S > T:
|
| 127 |
+
pad = (S - rgbs_seq.shape[1]) % S
|
| 128 |
+
rgbs_seq = rgbs_seq.reshape(B, 1, (S - pad), C_img * H * W)
|
| 129 |
+
events_seq = events_seq.reshape(B, 1, (S - pad), C * H * W)
|
| 130 |
+
padding_tensor = rgbs_seq[:, :, -1:, :].expand(B, 1, pad, C_img * H * W)
|
| 131 |
+
rgbs_seq = torch.cat([rgbs_seq, padding_tensor], dim=2)
|
| 132 |
+
padding_tensor = events_seq[:, :, -1:, :].expand(B, 1, pad, C * H * W)
|
| 133 |
+
events_seq = torch.cat([events_seq, padding_tensor], dim=2)
|
| 134 |
+
if img_ifnew is not None:
|
| 135 |
+
padding_numpy = np.ones(pad)
|
| 136 |
+
img_ifnew_seq = np.concatenate((img_ifnew_seq, padding_numpy), axis=0)
|
| 137 |
+
# rgbs_seq = rgbs_seq.reshape(B, -1, C_img, H, W)
|
| 138 |
+
# events_seq = events_seq.reshape(B, -1, C, H, W)
|
| 139 |
+
|
| 140 |
+
events_seq = events_seq.reshape(B * S, C, H, W)
|
| 141 |
+
rgbs_seq = rgbs_seq.reshape(B * S, C_img, H, W)
|
| 142 |
+
events_seq = F.interpolate(events_seq, tuple(interp_shape), mode='bilinear', align_corners=True)
|
| 143 |
+
rgbs_seq = F.interpolate(rgbs_seq, tuple(interp_shape), mode='bilinear', align_corners=True)
|
| 144 |
+
|
| 145 |
+
events_seq = 2 * events_seq - 1.0
|
| 146 |
+
rgbs_seq = 2 * (rgbs_seq / 255.0) - 1.0
|
| 147 |
+
|
| 148 |
+
dtype = rgbs_seq.dtype
|
| 149 |
+
|
| 150 |
+
if fmaps_fusion is None:
|
| 151 |
+
fmaps_pyramid = self.fusion_block(rgbs_seq, events_seq, img_ifnew_seq if img_ifnew is not None else None)
|
| 152 |
+
|
| 153 |
+
if isinstance(fmaps_pyramid, list):
|
| 154 |
+
for i, fmaps_fusion in enumerate(fmaps_pyramid):
|
| 155 |
+
fmaps_fusion = fmaps_fusion.permute(0, 2, 3, 1)
|
| 156 |
+
fmaps_fusion = fmaps_fusion / torch.sqrt(
|
| 157 |
+
torch.maximum(
|
| 158 |
+
torch.sum(torch.square(fmaps_fusion), axis=-1, keepdims=True),
|
| 159 |
+
torch.tensor(1e-12, device=fmaps_fusion.device),
|
| 160 |
+
)
|
| 161 |
+
)
|
| 162 |
+
fmaps_fusion = fmaps_fusion.permute(0, 3, 1, 2).reshape(
|
| 163 |
+
B, -1, self.latent_dim, int(H_stride / 2**i), int(W_stride / 2**i)
|
| 164 |
+
)
|
| 165 |
+
fmaps_fusion = fmaps_fusion.to(dtype)
|
| 166 |
+
fmaps_pyramid[i] = fmaps_fusion
|
| 167 |
+
else:
|
| 168 |
+
fmaps_fusion = fmaps_pyramid.permute(0, 2, 3, 1)
|
| 169 |
+
fmaps_fusion = fmaps_fusion / torch.sqrt(
|
| 170 |
+
torch.maximum(
|
| 171 |
+
torch.sum(torch.square(fmaps_fusion), axis=-1, keepdims=True),
|
| 172 |
+
torch.tensor(1e-12, device=fmaps_fusion.device),
|
| 173 |
+
)
|
| 174 |
+
)
|
| 175 |
+
fmaps_fusion = fmaps_fusion.permute(0, 3, 1, 2).reshape(
|
| 176 |
+
B, -1, self.latent_dim, H_stride, W_stride
|
| 177 |
+
)
|
| 178 |
+
else:
|
| 179 |
+
rgbs_, events_ = rgbs_seq[S//2:], events_seq[S//2:]
|
| 180 |
+
fmaps_pyramid_last = self.fusion_block(rgbs_, events_, img_ifnew_seq[S//2:] if img_ifnew is not None else None)
|
| 181 |
+
|
| 182 |
+
if isinstance(fmaps_pyramid_last, list):
|
| 183 |
+
for i, fmaps_fusion_last in enumerate(fmaps_pyramid_last):
|
| 184 |
+
fmaps_fusion_last = fmaps_fusion_last.permute(0, 2, 3, 1)
|
| 185 |
+
fmaps_fusion_last = fmaps_fusion_last / torch.sqrt(
|
| 186 |
+
torch.maximum(
|
| 187 |
+
torch.sum(torch.square(fmaps_fusion_last), axis=-1, keepdims=True),
|
| 188 |
+
torch.tensor(1e-12, device=fmaps_fusion.device),
|
| 189 |
+
)
|
| 190 |
+
)
|
| 191 |
+
fmaps_fusion_last = fmaps_fusion_last.permute(0, 3, 1, 2).reshape(
|
| 192 |
+
B, -1, self.latent_dim, int(H_stride / 2**i), int(W_stride / 2**i)
|
| 193 |
+
)
|
| 194 |
+
fmaps_fusion = torch.cat([fmaps_pyramid[i][:, S//2:], fmaps_fusion_last], dim=1)
|
| 195 |
+
fmaps_fusion = fmaps_fusion.to(dtype)
|
| 196 |
+
fmaps_pyramid[i] = fmaps_fusion
|
| 197 |
+
else:
|
| 198 |
+
fmaps_fusion_last = fmaps_pyramid_last.permute(0, 2, 3, 1)
|
| 199 |
+
fmaps_fusion_last = fmaps_fusion_last / torch.sqrt(
|
| 200 |
+
torch.maximum(
|
| 201 |
+
torch.sum(torch.square(fmaps_fusion_last), axis=-1, keepdims=True),
|
| 202 |
+
torch.tensor(1e-12, device=fmaps_fusion_last.device),
|
| 203 |
+
)
|
| 204 |
+
)
|
| 205 |
+
fmaps_fusion_last = fmaps_fusion_last.permute(0, 3, 1, 2).reshape(
|
| 206 |
+
B, -1, self.latent_dim, H_stride, W_stride
|
| 207 |
+
)
|
| 208 |
+
fmaps_fusion = torch.cat([fmaps_pyramid[0][:, S//2:], fmaps_fusion_last], dim=1)
|
| 209 |
+
fmaps_fusion = fmaps_fusion.to(dtype)
|
| 210 |
+
fmaps_pyramid = None
|
| 211 |
+
|
| 212 |
+
if not isinstance(fmaps_pyramid, list):
|
| 213 |
+
fmaps_pyramid = []
|
| 214 |
+
fmaps_pyramid.append(fmaps_fusion)
|
| 215 |
+
for i in range(self.corr_levels - 1):
|
| 216 |
+
fmaps_ = fmaps_fusion.reshape(B * S, self.latent_dim, fmaps_fusion.shape[-2], fmaps_fusion.shape[-1])
|
| 217 |
+
fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2)
|
| 218 |
+
fmaps_fusion = fmaps_.reshape(B, S, self.latent_dim, fmaps_.shape[-2], fmaps_.shape[-1])
|
| 219 |
+
fmaps_pyramid.append(fmaps_fusion)
|
| 220 |
+
if first_window:
|
| 221 |
+
for i in range(self.corr_levels):
|
| 222 |
+
track_feat, track_feat_support = get_track_feat(fmaps_pyramid[i], queried_frames, queried_coords/2**i, support_radius=self.corr_radius)
|
| 223 |
+
track_feat_pyramid.append(track_feat.repeat(1, S, 1, 1))
|
| 224 |
+
track_feat_support_pyramid.append(track_feat_support.unsqueeze(1))
|
| 225 |
+
first_window = False
|
| 226 |
+
|
| 227 |
+
|
| 228 |
+
attenstion_mask = (queried_frames < ind + S).reshape(B, 1, N) # B, 1, N
|
| 229 |
+
coords, viss, confs = self.forward_window(
|
| 230 |
+
fmaps_pyramid=[fmap for fmap in fmaps_pyramid],
|
| 231 |
+
coords=coords_init,
|
| 232 |
+
track_feat_support_pyramid=[attenstion_mask[:, None, :, :, None]*tfeat for tfeat in track_feat_support_pyramid],
|
| 233 |
+
corr_map_pyramid=self.corr_pyramid,
|
| 234 |
+
vis=vis_init,
|
| 235 |
+
conf=conf_init,
|
| 236 |
+
attenstion_mask=attenstion_mask.repeat(1, S, 1),
|
| 237 |
+
iters=iters,
|
| 238 |
+
)
|
| 239 |
+
S_trimmed = min(T - ind, S) # accounts for last window duration
|
| 240 |
+
coords_predicted[:, ind : ind + S] = coords[-1][:, :S_trimmed]
|
| 241 |
+
vis_predicted[:, ind : ind + S] = viss[-1][:, :S_trimmed]
|
| 242 |
+
conf_predicted[:, ind : ind + S] = confs[-1][:, :S_trimmed]
|
| 243 |
+
|
| 244 |
+
vis_predicted = torch.sigmoid(vis_predicted)
|
| 245 |
+
conf_predicted = torch.sigmoid(conf_predicted)
|
| 246 |
+
return coords_predicted, vis_predicted, conf_predicted
|
| 247 |
+
|
| 248 |
+
def forward_window(self, fmaps_pyramid, coords, track_feat_support_pyramid, corr_map_pyramid, vis, conf, attenstion_mask, iters=6):
|
| 249 |
+
B, S, *_ = fmaps_pyramid[0].shape
|
| 250 |
+
N = coords.shape[2]
|
| 251 |
+
r = 2 * self.corr_radius + 1
|
| 252 |
+
|
| 253 |
+
coord_preds, vis_preds, conf_preds = [], [], []
|
| 254 |
+
for it in range(iters):
|
| 255 |
+
coords = coords.detach()
|
| 256 |
+
coord_init = coords.view(B * S, N, 2)
|
| 257 |
+
corr_embs = []
|
| 258 |
+
for i in range(self.corr_levels):
|
| 259 |
+
corr_feat = self.get_correlation_feat(fmaps_pyramid[i], coord_init / 2 ** i)
|
| 260 |
+
track_feat_support = (
|
| 261 |
+
track_feat_support_pyramid[i]
|
| 262 |
+
.view(B, 1, r, r, N, self.latent_dim)
|
| 263 |
+
.squeeze(1)
|
| 264 |
+
.permute(0, 3, 1, 2, 4)
|
| 265 |
+
)
|
| 266 |
+
corr_volume = torch.einsum("btnhwc,bnijc->btnhwij", corr_feat, track_feat_support).reshape(B, S, N, r * r, r * r)
|
| 267 |
+
corr_emb = self.corr_mlp(corr_volume.reshape(B * S * N, r * r * r *r))
|
| 268 |
+
corr_embs.append(corr_emb)
|
| 269 |
+
|
| 270 |
+
corr_embs = torch.cat(corr_embs, dim=1)
|
| 271 |
+
corr_embs = corr_embs.view(B, S, N, corr_embs.shape[-1])
|
| 272 |
+
|
| 273 |
+
transformer_input = [vis, conf, corr_embs]
|
| 274 |
+
|
| 275 |
+
rel_coords_forward = coords[:, :-1] - coords[:, 1:]
|
| 276 |
+
rel_coords_backward = coords[:, 1:] - coords[:, :-1]
|
| 277 |
+
|
| 278 |
+
rel_coords_forward = torch.nn.functional.pad(rel_coords_forward, (0, 0, 0, 0, 0, 1))
|
| 279 |
+
rel_coords_backward = torch.nn.functional.pad(rel_coords_backward, (0, 0, 0, 0, 1, 0))
|
| 280 |
+
|
| 281 |
+
scale = (torch.tensor([self.model_resolution[1], self.model_resolution[0]], device=coords.device,) / self.stride)
|
| 282 |
+
rel_coords_forward = rel_coords_forward / scale
|
| 283 |
+
rel_coords_backward = rel_coords_backward / scale
|
| 284 |
+
|
| 285 |
+
rel_pos_emb_input = posenc(torch.cat([rel_coords_forward, rel_coords_backward], dim=-1), min_deg=0, max_deg=10,)
|
| 286 |
+
transformer_input.append(rel_pos_emb_input)
|
| 287 |
+
|
| 288 |
+
x = (torch.cat(transformer_input, dim=-1).permute(0, 2, 1, 3).reshape(B*N, S, -1))
|
| 289 |
+
|
| 290 |
+
x = x + self.interpolate_time_embed(x, S)
|
| 291 |
+
x = x.view(B, N, S, -1)
|
| 292 |
+
delta = self.updateformer2(x)
|
| 293 |
+
|
| 294 |
+
delta_coords = delta[..., :2].permute(0, 2, 1, 3)
|
| 295 |
+
delta_vis = delta[..., 2:3].permute(0, 2, 1, 3)
|
| 296 |
+
delta_conf = delta[..., 3:].permute(0, 2, 1, 3)
|
| 297 |
+
|
| 298 |
+
vis = vis + delta_vis
|
| 299 |
+
conf = conf + delta_conf
|
| 300 |
+
|
| 301 |
+
coords = coords + delta_coords
|
| 302 |
+
coord_preds.append(coords[..., :2] * float(self.stride))
|
| 303 |
+
|
| 304 |
+
vis_preds.append(vis[..., 0])
|
| 305 |
+
conf_preds.append(conf[..., 0])
|
| 306 |
+
return coord_preds, vis_preds, conf_preds
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
def load_parameters(self, model):
|
| 310 |
+
# 从训练模型加载参数
|
| 311 |
+
self.load_state_dict(model.state_dict())
|
LFE_TAP/models/__pycache__/blocks.cpython-38.pyc
ADDED
|
Binary file (25.6 kB). View file
|
|
|
LFE_TAP/models/__pycache__/blocks.cpython-39.pyc
ADDED
|
Binary file (25.4 kB). View file
|
|
|
LFE_TAP/models/__pycache__/embeddings.cpython-38.pyc
ADDED
|
Binary file (3.58 kB). View file
|
|
|
LFE_TAP/models/__pycache__/embeddings.cpython-39.pyc
ADDED
|
Binary file (3.55 kB). View file
|
|
|
LFE_TAP/models/__pycache__/etap.cpython-39.pyc
ADDED
|
Binary file (12.2 kB). View file
|
|
|
LFE_TAP/models/__pycache__/fusionFormer.cpython-38.pyc
ADDED
|
Binary file (15.3 kB). View file
|
|
|
LFE_TAP/models/__pycache__/fusionFormer.cpython-39.pyc
ADDED
|
Binary file (8.19 kB). View file
|
|
|
LFE_TAP/models/__pycache__/hivit.cpython-38.pyc
ADDED
|
Binary file (9.36 kB). View file
|
|
|
LFE_TAP/models/__pycache__/hivit.cpython-39.pyc
ADDED
|
Binary file (9.26 kB). View file
|
|
|
LFE_TAP/models/__pycache__/losses.cpython-39.pyc
ADDED
|
Binary file (2 kB). View file
|
|
|
LFE_TAP/models/__pycache__/tapfe.cpython-38.pyc
ADDED
|
Binary file (10.2 kB). View file
|
|
|
LFE_TAP/models/__pycache__/tapfe.cpython-39.pyc
ADDED
|
Binary file (9.27 kB). View file
|
|
|
LFE_TAP/models/blocks.py
ADDED
|
@@ -0,0 +1,994 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
import collections
|
| 5 |
+
|
| 6 |
+
from typing import Callable
|
| 7 |
+
from itertools import repeat
|
| 8 |
+
from functools import partial
|
| 9 |
+
from einops import rearrange
|
| 10 |
+
from timm.models.vision_transformer import Attention, Mlp
|
| 11 |
+
from LFE_TAP.utils.model_utils import combine_tokens, recover_tokens, bilinear_sampler
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
# From PyTorch internals
|
| 15 |
+
def _ntuple(n):
|
| 16 |
+
def parse(x):
|
| 17 |
+
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
|
| 18 |
+
return tuple(x)
|
| 19 |
+
return tuple(repeat(x, n))
|
| 20 |
+
|
| 21 |
+
return parse
|
| 22 |
+
|
| 23 |
+
def exists(val):
|
| 24 |
+
return val is not None
|
| 25 |
+
|
| 26 |
+
def default(val, d):
|
| 27 |
+
return val if exists(val) else d
|
| 28 |
+
|
| 29 |
+
to_2tuple = _ntuple(2)
|
| 30 |
+
class Mlp(nn.Module):
|
| 31 |
+
"""MLP as used in Vision Transformer, MLP-Mixer and related networks"""
|
| 32 |
+
|
| 33 |
+
def __init__(
|
| 34 |
+
self,
|
| 35 |
+
in_features,
|
| 36 |
+
hidden_features=None,
|
| 37 |
+
out_features=None,
|
| 38 |
+
act_layer=nn.GELU,
|
| 39 |
+
norm_layer=None,
|
| 40 |
+
bias=True,
|
| 41 |
+
drop=0.0,
|
| 42 |
+
use_conv=False,
|
| 43 |
+
):
|
| 44 |
+
super().__init__()
|
| 45 |
+
out_features = out_features or in_features
|
| 46 |
+
hidden_features = hidden_features or in_features
|
| 47 |
+
bias = to_2tuple(bias)
|
| 48 |
+
drop_probs = to_2tuple(drop)
|
| 49 |
+
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
|
| 50 |
+
|
| 51 |
+
self.fc1 = linear_layer(in_features, hidden_features, bias=bias[0])
|
| 52 |
+
self.act = act_layer()
|
| 53 |
+
self.drop1 = nn.Dropout(drop_probs[0])
|
| 54 |
+
self.norm = (
|
| 55 |
+
norm_layer(hidden_features) if norm_layer is not None else nn.Identity()
|
| 56 |
+
)
|
| 57 |
+
self.fc2 = linear_layer(hidden_features, out_features, bias=bias[1])
|
| 58 |
+
self.drop2 = nn.Dropout(drop_probs[1])
|
| 59 |
+
|
| 60 |
+
def forward(self, x):
|
| 61 |
+
x = self.fc1(x)
|
| 62 |
+
x = self.act(x)
|
| 63 |
+
x = self.drop1(x)
|
| 64 |
+
x = self.fc2(x)
|
| 65 |
+
x = self.drop2(x)
|
| 66 |
+
return x
|
| 67 |
+
|
| 68 |
+
|
| 69 |
+
class ResidualBlock(nn.Module):
|
| 70 |
+
def __init__(self, in_planes, planes, norm_fn="group", stride=1, dilation=1):
|
| 71 |
+
super(ResidualBlock, self).__init__()
|
| 72 |
+
|
| 73 |
+
self.conv1 = nn.Conv2d(
|
| 74 |
+
in_planes,
|
| 75 |
+
planes,
|
| 76 |
+
kernel_size=3,
|
| 77 |
+
padding=dilation,
|
| 78 |
+
dilation=dilation,
|
| 79 |
+
stride=stride,
|
| 80 |
+
padding_mode="zeros",
|
| 81 |
+
)
|
| 82 |
+
self.conv2 = nn.Conv2d(
|
| 83 |
+
planes, planes, kernel_size=3, padding=dilation, dilation=dilation, padding_mode="zeros"
|
| 84 |
+
)
|
| 85 |
+
self.relu = nn.ReLU(inplace=True)
|
| 86 |
+
|
| 87 |
+
num_groups = planes // 8
|
| 88 |
+
|
| 89 |
+
if norm_fn == "group":
|
| 90 |
+
self.norm1 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
| 91 |
+
self.norm2 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
| 92 |
+
if not (in_planes == planes and stride == 1):
|
| 93 |
+
self.norm3 = nn.GroupNorm(num_groups=num_groups, num_channels=planes)
|
| 94 |
+
|
| 95 |
+
elif norm_fn == "batch":
|
| 96 |
+
self.norm1 = nn.BatchNorm2d(planes)
|
| 97 |
+
self.norm2 = nn.BatchNorm2d(planes)
|
| 98 |
+
if not (in_planes == planes and stride == 1):
|
| 99 |
+
self.norm3 = nn.BatchNorm2d(planes)
|
| 100 |
+
|
| 101 |
+
elif norm_fn == "instance":
|
| 102 |
+
self.norm1 = nn.InstanceNorm2d(planes)
|
| 103 |
+
self.norm2 = nn.InstanceNorm2d(planes)
|
| 104 |
+
if not (in_planes == planes and stride == 1):
|
| 105 |
+
self.norm3 = nn.InstanceNorm2d(planes)
|
| 106 |
+
|
| 107 |
+
elif norm_fn == "none":
|
| 108 |
+
self.norm1 = nn.Sequential()
|
| 109 |
+
self.norm2 = nn.Sequential()
|
| 110 |
+
if not (in_planes == planes and stride == 1):
|
| 111 |
+
self.norm3 = nn.Sequential()
|
| 112 |
+
|
| 113 |
+
if in_planes == planes and stride == 1:
|
| 114 |
+
self.downsample = None
|
| 115 |
+
|
| 116 |
+
else:
|
| 117 |
+
self.downsample = nn.Sequential(
|
| 118 |
+
nn.Conv2d(in_planes, planes, kernel_size=1, stride=stride), self.norm3
|
| 119 |
+
)
|
| 120 |
+
|
| 121 |
+
def forward(self, x):
|
| 122 |
+
y = x
|
| 123 |
+
y = self.relu(self.norm1(self.conv1(y)))
|
| 124 |
+
y = self.relu(self.norm2(self.conv2(y)))
|
| 125 |
+
|
| 126 |
+
if self.downsample is not None:
|
| 127 |
+
x = self.downsample(x)
|
| 128 |
+
|
| 129 |
+
return self.relu(x + y)
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
class BasicEncoder(nn.Module):
|
| 133 |
+
def __init__(
|
| 134 |
+
self, input_dim=3, output_dim=128, stride=8, norm_fn="instance", dropout=0.0, shallow=False, in_planes=64, dilation=1,
|
| 135 |
+
):
|
| 136 |
+
super(BasicEncoder, self).__init__()
|
| 137 |
+
self.stride = stride
|
| 138 |
+
self.norm_fn = norm_fn
|
| 139 |
+
self.in_planes = in_planes
|
| 140 |
+
|
| 141 |
+
if self.norm_fn == "group":
|
| 142 |
+
self.norm1 = nn.GroupNorm(num_groups=8, num_channels=self.in_planes)
|
| 143 |
+
self.norm2 = nn.GroupNorm(num_groups=8, num_channels=output_dim * 2)
|
| 144 |
+
|
| 145 |
+
elif self.norm_fn == "batch":
|
| 146 |
+
self.norm1 = nn.BatchNorm2d(self.in_planes)
|
| 147 |
+
self.norm2 = nn.BatchNorm2d(output_dim * 2)
|
| 148 |
+
|
| 149 |
+
elif self.norm_fn == "instance":
|
| 150 |
+
self.norm1 = nn.InstanceNorm2d(self.in_planes)
|
| 151 |
+
self.norm2 = nn.InstanceNorm2d(output_dim * 2)
|
| 152 |
+
|
| 153 |
+
elif self.norm_fn == "none":
|
| 154 |
+
self.norm1 = nn.Sequential()
|
| 155 |
+
|
| 156 |
+
self.conv1 = nn.Conv2d(
|
| 157 |
+
input_dim,
|
| 158 |
+
self.in_planes,
|
| 159 |
+
kernel_size=7,
|
| 160 |
+
stride=2,
|
| 161 |
+
padding=3,
|
| 162 |
+
padding_mode="zeros",
|
| 163 |
+
)
|
| 164 |
+
self.relu1 = nn.ReLU(inplace=True)
|
| 165 |
+
|
| 166 |
+
self.shallow = shallow
|
| 167 |
+
if self.shallow:
|
| 168 |
+
# self.layer1 = ResidualBlock(self.in_planes, 64, norm_fn=self.norm_fn, stride=1)
|
| 169 |
+
# self.layer2 = ResidualBlock(64, 96, self.norm_fn, stride=2)
|
| 170 |
+
# self.layer3 = ResidualBlock(96, 128, self.norm_fn, stride=2)
|
| 171 |
+
self.layer1 = self._make_layer(64, stride=1, dilation=dilation)
|
| 172 |
+
self.layer2 = self._make_layer(96, stride=2, dilation=dilation)
|
| 173 |
+
self.layer3 = self._make_layer(128, stride=2, dilation=dilation)
|
| 174 |
+
self.conv2 = nn.Conv2d(128 + 96 + 64, output_dim, kernel_size=1)
|
| 175 |
+
else:
|
| 176 |
+
self.layer1 = self._make_layer(64, stride=1)
|
| 177 |
+
self.layer2 = self._make_layer(96, stride=2)
|
| 178 |
+
self.layer3 = self._make_layer(128, stride=2)
|
| 179 |
+
self.layer4 = self._make_layer(128, stride=2)
|
| 180 |
+
|
| 181 |
+
self.conv2 = nn.Conv2d(
|
| 182 |
+
128 + 128 + 96 + 64,
|
| 183 |
+
output_dim * 2,
|
| 184 |
+
kernel_size=3,
|
| 185 |
+
padding=1,
|
| 186 |
+
padding_mode="zeros",
|
| 187 |
+
)
|
| 188 |
+
self.relu2 = nn.ReLU(inplace=True)
|
| 189 |
+
self.conv3 = nn.Conv2d(output_dim * 2, output_dim, kernel_size=1)
|
| 190 |
+
|
| 191 |
+
self.dropout = None
|
| 192 |
+
if dropout > 0:
|
| 193 |
+
self.dropout = nn.Dropout2d(p=dropout)
|
| 194 |
+
|
| 195 |
+
for m in self.modules():
|
| 196 |
+
if isinstance(m, nn.Conv2d):
|
| 197 |
+
nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
|
| 198 |
+
elif isinstance(m, (nn.BatchNorm2d, nn.InstanceNorm2d, nn.GroupNorm)):
|
| 199 |
+
if m.weight is not None:
|
| 200 |
+
nn.init.constant_(m.weight, 1)
|
| 201 |
+
if m.bias is not None:
|
| 202 |
+
nn.init.constant_(m.bias, 0)
|
| 203 |
+
|
| 204 |
+
def _make_layer(self, dim, stride=1, dilation=1):
|
| 205 |
+
layer1 = ResidualBlock(self.in_planes, dim, self.norm_fn, stride=stride, dilation=dilation)
|
| 206 |
+
layer2 = ResidualBlock(dim, dim, self.norm_fn, stride=1, dilation=dilation)
|
| 207 |
+
layers = (layer1, layer2)
|
| 208 |
+
|
| 209 |
+
self.in_planes = dim
|
| 210 |
+
return nn.Sequential(*layers)
|
| 211 |
+
|
| 212 |
+
def forward(self, x):
|
| 213 |
+
_, _, H, W = x.shape
|
| 214 |
+
|
| 215 |
+
x = self.conv1(x)
|
| 216 |
+
x = self.norm1(x)
|
| 217 |
+
x = self.relu1(x)
|
| 218 |
+
|
| 219 |
+
if self.shallow:
|
| 220 |
+
a = self.layer1(x)
|
| 221 |
+
b = self.layer2(a)
|
| 222 |
+
c = self.layer3(b)
|
| 223 |
+
a = F.interpolate(
|
| 224 |
+
a,
|
| 225 |
+
(H // self.stride, W // self.stride),
|
| 226 |
+
mode="bilinear",
|
| 227 |
+
align_corners=True,
|
| 228 |
+
)
|
| 229 |
+
b = F.interpolate(
|
| 230 |
+
b,
|
| 231 |
+
(H // self.stride, W // self.stride),
|
| 232 |
+
mode="bilinear",
|
| 233 |
+
align_corners=True,
|
| 234 |
+
)
|
| 235 |
+
c = F.interpolate(
|
| 236 |
+
c,
|
| 237 |
+
(H // self.stride, W // self.stride),
|
| 238 |
+
mode="bilinear",
|
| 239 |
+
align_corners=True,
|
| 240 |
+
)
|
| 241 |
+
x = self.conv2(torch.cat([a, b, c], dim=1))
|
| 242 |
+
else:
|
| 243 |
+
a = self.layer1(x)
|
| 244 |
+
b = self.layer2(a)
|
| 245 |
+
c = self.layer3(b)
|
| 246 |
+
d = self.layer4(c)
|
| 247 |
+
a = F.interpolate(
|
| 248 |
+
a,
|
| 249 |
+
(H // self.stride, W // self.stride),
|
| 250 |
+
mode="bilinear",
|
| 251 |
+
align_corners=True,
|
| 252 |
+
)
|
| 253 |
+
b = F.interpolate(
|
| 254 |
+
b,
|
| 255 |
+
(H // self.stride, W // self.stride),
|
| 256 |
+
mode="bilinear",
|
| 257 |
+
align_corners=True,
|
| 258 |
+
)
|
| 259 |
+
c = F.interpolate(
|
| 260 |
+
c,
|
| 261 |
+
(H // self.stride, W // self.stride),
|
| 262 |
+
mode="bilinear",
|
| 263 |
+
align_corners=True,
|
| 264 |
+
)
|
| 265 |
+
d = F.interpolate(
|
| 266 |
+
d,
|
| 267 |
+
(H // self.stride, W // self.stride),
|
| 268 |
+
mode="bilinear",
|
| 269 |
+
align_corners=True,
|
| 270 |
+
)
|
| 271 |
+
x = self.conv2(torch.cat([a, b, c, d], dim=1))
|
| 272 |
+
x = self.norm2(x)
|
| 273 |
+
x = self.relu2(x)
|
| 274 |
+
x = self.conv3(x)
|
| 275 |
+
|
| 276 |
+
if self.training and self.dropout is not None:
|
| 277 |
+
x = self.dropout(x)
|
| 278 |
+
return x
|
| 279 |
+
|
| 280 |
+
|
| 281 |
+
class FusionBlock_basic(nn.Module):
|
| 282 |
+
def __init__(self, img_in_dim=3, event_in_dim=10, output_dim=128, stride=8, norm_fn="instance", dropout=0.0):
|
| 283 |
+
super().__init__()
|
| 284 |
+
self.imgnet = BasicEncoder(input_dim=img_in_dim, output_dim=output_dim, stride=stride, norm_fn=norm_fn, dropout=dropout)
|
| 285 |
+
self.eventnet = BasicEncoder(input_dim=event_in_dim, output_dim=output_dim, stride=stride, norm_fn=norm_fn, dropout=dropout)
|
| 286 |
+
self.conv1 = nn.Conv2d(output_dim, 192, 1, padding=0)
|
| 287 |
+
self.conv2 = nn.Conv2d(output_dim, 192, 1, padding=0)
|
| 288 |
+
self.convo = nn.Conv2d(192*2, output_dim, 3, padding=1)
|
| 289 |
+
|
| 290 |
+
# for m in self.modules():
|
| 291 |
+
# if isinstance(m, nn.Conv2d):
|
| 292 |
+
# nn.init.kaiming_normal_(m.weight, mode="fan_out", nonlinearity="relu")
|
| 293 |
+
|
| 294 |
+
def forward(self, x_i, x_e, _):
|
| 295 |
+
x_i = self.imgnet(x_i)
|
| 296 |
+
x_e = self.eventnet(x_e)
|
| 297 |
+
c1 = F.relu(self.conv1(x_i))
|
| 298 |
+
c2 = F.relu(self.conv2(x_e))
|
| 299 |
+
out = torch.cat([c1, c2], dim=1)
|
| 300 |
+
out = F.relu(self.convo(out))
|
| 301 |
+
return x_i + out
|
| 302 |
+
|
| 303 |
+
class FusionBlock(nn.Module):
|
| 304 |
+
def __init__(self, img_in_dim=3, event_in_dim=10, output_dim=128, stride=8, norm_fn="instance", dropout=0.0):
|
| 305 |
+
super(FusionBlock, self).__init__()
|
| 306 |
+
self.stride = stride
|
| 307 |
+
self.norm_fn = norm_fn
|
| 308 |
+
self.in_planes = 32
|
| 309 |
+
|
| 310 |
+
if self.norm_fn == "group":
|
| 311 |
+
self.norm1_i = nn.GroupNorm(num_groups=8, num_channels=self.in_planes)
|
| 312 |
+
self.norm1_e = nn.GroupNorm(num_groups=8, num_channels=self.in_planes)
|
| 313 |
+
self.norm2 = nn.GroupNorm(num_groups=8, num_channels=output_dim * 2)
|
| 314 |
+
|
| 315 |
+
elif self.norm_fn == "batch":
|
| 316 |
+
self.norm1_i = nn.BatchNorm2d(self.in_planes)
|
| 317 |
+
self.norm1_e = nn.BatchNorm2d(self.in_planes)
|
| 318 |
+
self.norm2 = nn.BatchNorm2d(output_dim * 2)
|
| 319 |
+
|
| 320 |
+
elif self.norm_fn == "instance":
|
| 321 |
+
self.norm1_i = nn.InstanceNorm2d(self.in_planes)
|
| 322 |
+
self.norm1_e = nn.InstanceNorm2d(self.in_planes)
|
| 323 |
+
self.norm2 = nn.InstanceNorm2d(output_dim * 2)
|
| 324 |
+
|
| 325 |
+
elif self.norm_fn == "none":
|
| 326 |
+
self.norm1_i = nn.Sequential()
|
| 327 |
+
self.norm1_e = nn.Sequential()
|
| 328 |
+
|
| 329 |
+
self.conv1_i = nn.Conv2d(img_in_dim, self.in_planes, kernel_size=7, stride=2, padding=3, padding_mode="zeros")
|
| 330 |
+
self.conv1_e = nn.Conv2d(event_in_dim, self.in_planes, kernel_size=7, stride=2, padding=3, padding_mode="zeros")
|
| 331 |
+
self.relu1 = nn.ReLU(inplace=True)
|
| 332 |
+
|
| 333 |
+
self.shallow = False
|
| 334 |
+
if self.shallow:
|
| 335 |
+
self.layer1_e = self._make_layer(self.in_planes, 64, stride=1)
|
| 336 |
+
self.layer2_e = self._make_layer(64, 96, stride=2)
|
| 337 |
+
self.layer3_e = self._make_layer(96, 128, stride=2)
|
| 338 |
+
# self.conv_half1 = self._half_conv(self.in_planes*2)
|
| 339 |
+
# self.conv_half2 = self._half_conv(64*2)
|
| 340 |
+
# self.conv_half3 = self._half_conv(96*2)
|
| 341 |
+
self.layer1_i = self._make_layer(self.in_planes*2, 64, stride=1)
|
| 342 |
+
self.layer2_i = self._make_layer(64*2, 96, stride=2)
|
| 343 |
+
self.layer3_i = self._make_layer(96*2, 128, stride=2)
|
| 344 |
+
self.conv2 = nn.Conv2d(128 + 96 + 64, output_dim, kernel_size=1)
|
| 345 |
+
else:
|
| 346 |
+
self.layer1_e = self._make_layer(self.in_planes, 64, stride=1)
|
| 347 |
+
self.layer2_e = self._make_layer(64, 96, stride=2)
|
| 348 |
+
self.layer3_e = self._make_layer(96, 128, stride=2)
|
| 349 |
+
self.layer4_e = self._make_layer(128, 128, stride=2)
|
| 350 |
+
# self.conv_half1 = self._half_conv(self.in_planes*2)
|
| 351 |
+
# self.conv_half2 = self._half_conv(64*2)
|
| 352 |
+
# self.conv_half3 = self._half_conv(96*2)
|
| 353 |
+
# self.conv_half4 = self._half_conv(128*2)
|
| 354 |
+
self.layer1_i = self._make_layer(self.in_planes*2, 64, stride=1)
|
| 355 |
+
self.layer2_i = self._make_layer(64*2, 96, stride=2)
|
| 356 |
+
self.layer3_i = self._make_layer(96*2, 128, stride=2)
|
| 357 |
+
self.layer4_i = self._make_layer(128*2, 128, stride=2)
|
| 358 |
+
self.conv2 = nn.Conv2d(128 + 128 + 96 + 64, output_dim, kernel_size=3, padding=1, padding_mode="zeros")
|
| 359 |
+
self.relu2 = nn.ReLU(inplace=True)
|
| 360 |
+
self.conv3 = nn.Conv2d(output_dim , output_dim, kernel_size=1)
|
| 361 |
+
|
| 362 |
+
self.dropout = None
|
| 363 |
+
if dropout > 0:
|
| 364 |
+
self.dropout = nn.Dropout2d(p=dropout)
|
| 365 |
+
|
| 366 |
+
def _half_conv(self, in_planes):
|
| 367 |
+
conv = nn.Conv2d(in_planes, in_planes // 2, kernel_size=1)
|
| 368 |
+
norm = nn.BatchNorm2d(in_planes // 2)
|
| 369 |
+
relu = nn.ReLU(inplace=True)
|
| 370 |
+
layers = (conv, norm, relu)
|
| 371 |
+
return nn.Sequential(*layers)
|
| 372 |
+
|
| 373 |
+
def _make_layer(self, input_dim, output_dim, stride=1):
|
| 374 |
+
layer1 = ResidualBlock(input_dim, output_dim, self.norm_fn, stride=stride)
|
| 375 |
+
layer2 = ResidualBlock(output_dim, output_dim, self.norm_fn, stride=1)
|
| 376 |
+
layers = (layer1, layer2)
|
| 377 |
+
|
| 378 |
+
return nn.Sequential(*layers)
|
| 379 |
+
|
| 380 |
+
def forward(self, x_i, x_e):
|
| 381 |
+
_, _, H, W = x_i.size()
|
| 382 |
+
|
| 383 |
+
x_i = self.conv1_i(x_i)
|
| 384 |
+
x_e = self.conv1_e(x_e)
|
| 385 |
+
x_i = self.relu1(self.norm1_i(x_i))
|
| 386 |
+
x_e = self.relu1(self.norm1_e(x_e))
|
| 387 |
+
|
| 388 |
+
if self.shallow:
|
| 389 |
+
x_e_a = self.layer1_e(x_e)
|
| 390 |
+
x_e_b = self.layer2_e(x_e_a)
|
| 391 |
+
x_e_c = self.layer3_e(x_e_b)
|
| 392 |
+
# x_i_a = self.layer1_i(self.conv_half1(torch.cat((x_i, x_e), dim=1)))
|
| 393 |
+
# x_i_b = self.layer2_i(self.conv_half2(torch.cat((x_i_a, x_e_a), dim=1)))
|
| 394 |
+
# x_i_c = self.layer3_i(self.conv_half3(torch.cat((x_i_b, x_e_b), dim=1)))
|
| 395 |
+
x_i_a = self.layer1_i(torch.cat((x_i, x_e), dim=1))
|
| 396 |
+
x_i_b = self.layer2_i(torch.cat((x_i_a, x_e_a), dim=1))
|
| 397 |
+
x_i_c = self.layer3_i(torch.cat((x_i_b, x_e_b), dim=1))
|
| 398 |
+
x_i_a = F.interpolate(x_i_a, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 399 |
+
x_i_b = F.interpolate(x_i_b, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 400 |
+
x_i_c = F.interpolate(x_i_c, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 401 |
+
x = self.conv2(torch.cat((x_i_a, x_i_b, x_i_c), dim=1))
|
| 402 |
+
else:
|
| 403 |
+
x_e_a = self.layer1_e(x_e)
|
| 404 |
+
x_e_b = self.layer2_e(x_e_a)
|
| 405 |
+
x_e_c = self.layer3_e(x_e_b)
|
| 406 |
+
# x_e_d = self.layer4_e(x_e_c)
|
| 407 |
+
# x_i_a = self.layer1_i(self.conv_half1(torch.cat((x_i, x_e), dim=1)))
|
| 408 |
+
# x_i_b = self.layer2_i(self.conv_half2(torch.cat((x_i_a, x_e_a), dim=1)))
|
| 409 |
+
# x_i_c = self.layer3_i(self.conv_half3(torch.cat((x_i_b, x_e_b), dim=1)))
|
| 410 |
+
# x_i_d = self.layer4_i(self.conv_half4(torch.cat((x_i_c, x_e_c), dim=1)))
|
| 411 |
+
x_i_a = self.layer1_i(torch.cat((x_i, x_e), dim=1))
|
| 412 |
+
x_i_b = self.layer2_i(torch.cat((x_i_a, x_e_a), dim=1))
|
| 413 |
+
x_i_c = self.layer3_i(torch.cat((x_i_b, x_e_b), dim=1))
|
| 414 |
+
x_i_d = self.layer4_i(torch.cat((x_i_c, x_e_c), dim=1))
|
| 415 |
+
x_i_a = F.interpolate(x_i_a, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 416 |
+
x_i_b = F.interpolate(x_i_b, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 417 |
+
x_i_c = F.interpolate(x_i_c, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 418 |
+
x_i_d = F.interpolate(x_i_d, size=(H // self.stride, W // self.stride), mode='bilinear', align_corners=True)
|
| 419 |
+
x = self.conv2(torch.cat((x_i_a, x_i_b, x_i_c, x_i_d), dim=1))
|
| 420 |
+
x = self.relu2(self.norm2(x))
|
| 421 |
+
x = self.conv3(x)
|
| 422 |
+
|
| 423 |
+
if self.training and self.dropout is not None:
|
| 424 |
+
x = self.dropout(x)
|
| 425 |
+
|
| 426 |
+
return x
|
| 427 |
+
|
| 428 |
+
|
| 429 |
+
class AttnBlock(nn.Module):
|
| 430 |
+
"""
|
| 431 |
+
A DiT block with adaptive layer norm zero (adaLN-Zero) conditioning.
|
| 432 |
+
"""
|
| 433 |
+
|
| 434 |
+
def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, **block_kwargs):
|
| 435 |
+
super().__init__()
|
| 436 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 437 |
+
self.attn = Attention(
|
| 438 |
+
hidden_size, num_heads=num_heads, qkv_bias=True, **block_kwargs
|
| 439 |
+
)
|
| 440 |
+
|
| 441 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 442 |
+
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
| 443 |
+
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
| 444 |
+
self.mlp = Mlp(
|
| 445 |
+
in_features=hidden_size,
|
| 446 |
+
hidden_features=mlp_hidden_dim,
|
| 447 |
+
act_layer=approx_gelu,
|
| 448 |
+
drop=0,
|
| 449 |
+
)
|
| 450 |
+
|
| 451 |
+
def forward(self, x):
|
| 452 |
+
x = x + self.attn(self.norm1(x))
|
| 453 |
+
x = x + self.mlp(self.norm2(x))
|
| 454 |
+
return x
|
| 455 |
+
|
| 456 |
+
|
| 457 |
+
class UpdateFormer(nn.Module):
|
| 458 |
+
def __init__(self, space_depth=12, time_depth=12, input_dim=320, hidden_size=384, num_heads=8, output_dim=130, mlp_ratio=4.0):
|
| 459 |
+
super(UpdateFormer, self).__init__()
|
| 460 |
+
self.hidden_size = hidden_size
|
| 461 |
+
self.input_transform = nn.Linear(input_dim, hidden_size, bias=True)
|
| 462 |
+
self.flow_head = nn.Linear(hidden_size, output_dim, bias=True)
|
| 463 |
+
|
| 464 |
+
self.time_blocks = nn.ModuleList(
|
| 465 |
+
[
|
| 466 |
+
AttnBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio)
|
| 467 |
+
for _ in range(time_depth)
|
| 468 |
+
]
|
| 469 |
+
)
|
| 470 |
+
|
| 471 |
+
self.space_blocks = nn.ModuleList(
|
| 472 |
+
[
|
| 473 |
+
AttnBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio)
|
| 474 |
+
for _ in range(space_depth)
|
| 475 |
+
]
|
| 476 |
+
)
|
| 477 |
+
assert len(self.time_blocks) >= len(self.space_blocks)
|
| 478 |
+
self.initialize_weights()
|
| 479 |
+
|
| 480 |
+
def initialize_weights(self):
|
| 481 |
+
def _basic_init(module):
|
| 482 |
+
if isinstance(module, nn.Linear):
|
| 483 |
+
torch.nn.init.xavier_uniform_(module.weight)
|
| 484 |
+
if module.bias is not None:
|
| 485 |
+
nn.init.constant_(module.bias, 0)
|
| 486 |
+
|
| 487 |
+
self.apply(_basic_init)
|
| 488 |
+
|
| 489 |
+
def forward(self, x):
|
| 490 |
+
x = self.input_transform(x)
|
| 491 |
+
j = 0
|
| 492 |
+
for i in range(len(self.time_blocks)):
|
| 493 |
+
B, N, T, _ = x.shape
|
| 494 |
+
x_time = rearrange(x, "b n t c -> (b n) t c", b=B, t=T, n=N)
|
| 495 |
+
x_time = self.time_blocks[i](x_time)
|
| 496 |
+
|
| 497 |
+
x = rearrange(x_time, "(b n) t c -> b n t c", b=B, t=T, n=N)
|
| 498 |
+
if self.add_space_attn and (i % (len(self.time_blocks) // len(self.space_blocks)) == 0):
|
| 499 |
+
x_space = rearrange(x, "b n t c -> (b n) t c", b=B, t=T, n=N)
|
| 500 |
+
x_space = self.space_blocks[j](x_space)
|
| 501 |
+
x = rearrange(x_space, "(b n) t c -> b n t c", b=B, t=T, n=N)
|
| 502 |
+
j += 1
|
| 503 |
+
|
| 504 |
+
flow = self.flow_head(x)
|
| 505 |
+
return flow
|
| 506 |
+
|
| 507 |
+
|
| 508 |
+
class Attention2(nn.Module):
|
| 509 |
+
def __init__(
|
| 510 |
+
self, query_dim, context_dim=None, num_heads=8, dim_head=48, qkv_bias=False
|
| 511 |
+
):
|
| 512 |
+
super().__init__()
|
| 513 |
+
inner_dim = dim_head * num_heads
|
| 514 |
+
context_dim = default(context_dim, query_dim)
|
| 515 |
+
self.scale = dim_head**-0.5
|
| 516 |
+
self.heads = num_heads
|
| 517 |
+
|
| 518 |
+
self.to_q = nn.Linear(query_dim, inner_dim, bias=qkv_bias)
|
| 519 |
+
self.to_kv = nn.Linear(context_dim, inner_dim * 2, bias=qkv_bias)
|
| 520 |
+
self.to_out = nn.Linear(inner_dim, query_dim)
|
| 521 |
+
|
| 522 |
+
def forward(self, x, context=None, attn_bias=None):
|
| 523 |
+
B, N1, C = x.shape
|
| 524 |
+
h = self.heads
|
| 525 |
+
|
| 526 |
+
q = self.to_q(x).reshape(B, N1, h, C // h).permute(0, 2, 1, 3)
|
| 527 |
+
context = default(context, x)
|
| 528 |
+
k, v = self.to_kv(context).chunk(2, dim=-1)
|
| 529 |
+
|
| 530 |
+
N2 = context.shape[1]
|
| 531 |
+
k = k.reshape(B, N2, h, C // h).permute(0, 2, 1, 3)
|
| 532 |
+
v = v.reshape(B, N2, h, C // h).permute(0, 2, 1, 3)
|
| 533 |
+
|
| 534 |
+
sim = (q @ k.transpose(-2, -1)) * self.scale
|
| 535 |
+
|
| 536 |
+
if attn_bias is not None:
|
| 537 |
+
sim = sim + attn_bias
|
| 538 |
+
attn = sim.softmax(dim=-1)
|
| 539 |
+
|
| 540 |
+
x = (attn @ v).transpose(1, 2).reshape(B, N1, C)
|
| 541 |
+
return self.to_out(x)
|
| 542 |
+
|
| 543 |
+
|
| 544 |
+
class CrossAttnBlock(nn.Module):
|
| 545 |
+
def __init__(
|
| 546 |
+
self, hidden_size, context_dim, num_heads=1, mlp_ratio=4.0, **block_kwargs
|
| 547 |
+
):
|
| 548 |
+
super().__init__()
|
| 549 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 550 |
+
self.norm_context = nn.LayerNorm(hidden_size)
|
| 551 |
+
self.cross_attn = Attention2(
|
| 552 |
+
hidden_size,
|
| 553 |
+
context_dim=context_dim,
|
| 554 |
+
num_heads=num_heads,
|
| 555 |
+
qkv_bias=True,
|
| 556 |
+
**block_kwargs
|
| 557 |
+
)
|
| 558 |
+
|
| 559 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 560 |
+
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
| 561 |
+
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
| 562 |
+
self.mlp = Mlp(
|
| 563 |
+
in_features=hidden_size,
|
| 564 |
+
hidden_features=mlp_hidden_dim,
|
| 565 |
+
act_layer=approx_gelu,
|
| 566 |
+
drop=0,
|
| 567 |
+
)
|
| 568 |
+
|
| 569 |
+
def forward(self, x, context, mask=None):
|
| 570 |
+
attn_bias = None
|
| 571 |
+
if mask is not None:
|
| 572 |
+
if mask.shape[1] == x.shape[1]:
|
| 573 |
+
mask = mask[:, None, :, None].expand(
|
| 574 |
+
-1, self.cross_attn.heads, -1, context.shape[1]
|
| 575 |
+
)
|
| 576 |
+
else:
|
| 577 |
+
mask = mask[:, None, None].expand(
|
| 578 |
+
-1, self.cross_attn.heads, x.shape[1], -1
|
| 579 |
+
)
|
| 580 |
+
|
| 581 |
+
max_neg_value = -torch.finfo(x.dtype).max
|
| 582 |
+
attn_bias = (~mask) * max_neg_value
|
| 583 |
+
x = x + self.cross_attn(
|
| 584 |
+
self.norm1(x), context=self.norm_context(context), attn_bias=attn_bias
|
| 585 |
+
)
|
| 586 |
+
x = x + self.mlp(self.norm2(x))
|
| 587 |
+
return x
|
| 588 |
+
|
| 589 |
+
|
| 590 |
+
class AttnBlock2(nn.Module):
|
| 591 |
+
def __init__(
|
| 592 |
+
self,
|
| 593 |
+
hidden_size,
|
| 594 |
+
num_heads,
|
| 595 |
+
attn_class: Callable[..., nn.Module] = Attention2,
|
| 596 |
+
mlp_ratio=4.0,
|
| 597 |
+
**block_kwargs
|
| 598 |
+
):
|
| 599 |
+
super().__init__()
|
| 600 |
+
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 601 |
+
self.attn = attn_class(
|
| 602 |
+
hidden_size, num_heads=num_heads, qkv_bias=True, **block_kwargs
|
| 603 |
+
)
|
| 604 |
+
|
| 605 |
+
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
| 606 |
+
mlp_hidden_dim = int(hidden_size * mlp_ratio)
|
| 607 |
+
approx_gelu = lambda: nn.GELU(approximate="tanh")
|
| 608 |
+
self.mlp = Mlp(
|
| 609 |
+
in_features=hidden_size,
|
| 610 |
+
hidden_features=mlp_hidden_dim,
|
| 611 |
+
act_layer=approx_gelu,
|
| 612 |
+
drop=0,
|
| 613 |
+
)
|
| 614 |
+
|
| 615 |
+
def forward(self, x, mask=None):
|
| 616 |
+
attn_bias = mask
|
| 617 |
+
if mask is not None:
|
| 618 |
+
mask = (
|
| 619 |
+
(mask[:, None] * mask[:, :, None])
|
| 620 |
+
.unsqueeze(1)
|
| 621 |
+
.expand(-1, self.attn.num_heads, -1, -1)
|
| 622 |
+
)
|
| 623 |
+
max_neg_value = -torch.finfo(x.dtype).max
|
| 624 |
+
attn_bias = (~mask) * max_neg_value
|
| 625 |
+
x = x + self.attn(self.norm1(x), attn_bias=attn_bias)
|
| 626 |
+
x = x + self.mlp(self.norm2(x))
|
| 627 |
+
return x
|
| 628 |
+
|
| 629 |
+
|
| 630 |
+
class EfficientUpdateFormer(nn.Module):
|
| 631 |
+
"""
|
| 632 |
+
Transformer model that updates track estimates.
|
| 633 |
+
"""
|
| 634 |
+
|
| 635 |
+
def __init__(
|
| 636 |
+
self,
|
| 637 |
+
space_depth=6,
|
| 638 |
+
time_depth=6,
|
| 639 |
+
input_dim=320,
|
| 640 |
+
hidden_size=384,
|
| 641 |
+
num_heads=8,
|
| 642 |
+
output_dim=130,
|
| 643 |
+
mlp_ratio=4.0,
|
| 644 |
+
num_virtual_tracks=32,
|
| 645 |
+
add_space_attn=True,
|
| 646 |
+
linear_layer_for_vis_conf=False,
|
| 647 |
+
):
|
| 648 |
+
super().__init__()
|
| 649 |
+
self.out_channels = 2
|
| 650 |
+
self.num_heads = num_heads
|
| 651 |
+
self.hidden_size = hidden_size
|
| 652 |
+
self.input_transform = torch.nn.Linear(input_dim, hidden_size, bias=True)
|
| 653 |
+
if linear_layer_for_vis_conf:
|
| 654 |
+
self.flow_head = torch.nn.Linear(hidden_size, output_dim - 2, bias=True)
|
| 655 |
+
self.vis_conf_head = torch.nn.Linear(hidden_size, 2, bias=True)
|
| 656 |
+
else:
|
| 657 |
+
self.flow_head = torch.nn.Linear(hidden_size, output_dim, bias=True)
|
| 658 |
+
self.num_virtual_tracks = num_virtual_tracks
|
| 659 |
+
self.virual_tracks = nn.Parameter(
|
| 660 |
+
torch.randn(1, num_virtual_tracks, 1, hidden_size)
|
| 661 |
+
)
|
| 662 |
+
self.add_space_attn = add_space_attn
|
| 663 |
+
self.linear_layer_for_vis_conf = linear_layer_for_vis_conf
|
| 664 |
+
self.time_blocks = nn.ModuleList(
|
| 665 |
+
[
|
| 666 |
+
AttnBlock2(
|
| 667 |
+
hidden_size,
|
| 668 |
+
num_heads,
|
| 669 |
+
mlp_ratio=mlp_ratio,
|
| 670 |
+
attn_class=Attention2,
|
| 671 |
+
)
|
| 672 |
+
for _ in range(time_depth)
|
| 673 |
+
]
|
| 674 |
+
)
|
| 675 |
+
|
| 676 |
+
if add_space_attn:
|
| 677 |
+
self.space_virtual_blocks = nn.ModuleList(
|
| 678 |
+
[
|
| 679 |
+
AttnBlock2(
|
| 680 |
+
hidden_size,
|
| 681 |
+
num_heads,
|
| 682 |
+
mlp_ratio=mlp_ratio,
|
| 683 |
+
attn_class=Attention2,
|
| 684 |
+
)
|
| 685 |
+
for _ in range(space_depth)
|
| 686 |
+
]
|
| 687 |
+
)
|
| 688 |
+
self.space_point2virtual_blocks = nn.ModuleList(
|
| 689 |
+
[
|
| 690 |
+
CrossAttnBlock(
|
| 691 |
+
hidden_size, hidden_size, num_heads, mlp_ratio=mlp_ratio
|
| 692 |
+
)
|
| 693 |
+
for _ in range(space_depth)
|
| 694 |
+
]
|
| 695 |
+
)
|
| 696 |
+
self.space_virtual2point_blocks = nn.ModuleList(
|
| 697 |
+
[
|
| 698 |
+
CrossAttnBlock(
|
| 699 |
+
hidden_size, hidden_size, num_heads, mlp_ratio=mlp_ratio
|
| 700 |
+
)
|
| 701 |
+
for _ in range(space_depth)
|
| 702 |
+
]
|
| 703 |
+
)
|
| 704 |
+
assert len(self.time_blocks) >= len(self.space_virtual2point_blocks)
|
| 705 |
+
self.initialize_weights()
|
| 706 |
+
|
| 707 |
+
def initialize_weights(self):
|
| 708 |
+
def _basic_init(module):
|
| 709 |
+
if isinstance(module, nn.Linear):
|
| 710 |
+
torch.nn.init.xavier_uniform_(module.weight)
|
| 711 |
+
if module.bias is not None:
|
| 712 |
+
nn.init.constant_(module.bias, 0)
|
| 713 |
+
torch.nn.init.trunc_normal_(self.flow_head.weight, std=0.001)
|
| 714 |
+
if self.linear_layer_for_vis_conf:
|
| 715 |
+
torch.nn.init.trunc_normal_(self.vis_conf_head.weight, std=0.001)
|
| 716 |
+
|
| 717 |
+
def _trunc_init(module):
|
| 718 |
+
"""ViT weight initialization, original timm impl (for reproducibility)"""
|
| 719 |
+
if isinstance(module, nn.Linear):
|
| 720 |
+
torch.nn.init.trunc_normal_(module.weight, std=0.02)
|
| 721 |
+
if module.bias is not None:
|
| 722 |
+
nn.init.zeros_(module.bias)
|
| 723 |
+
|
| 724 |
+
self.apply(_basic_init)
|
| 725 |
+
|
| 726 |
+
def forward(self, input_tensor, mask=None, add_space_attn=True):
|
| 727 |
+
tokens = self.input_transform(input_tensor)
|
| 728 |
+
|
| 729 |
+
B, _, T, _ = tokens.shape
|
| 730 |
+
virtual_tokens = self.virual_tracks.repeat(B, 1, T, 1)
|
| 731 |
+
tokens = torch.cat([tokens, virtual_tokens], dim=1)
|
| 732 |
+
|
| 733 |
+
_, N, _, _ = tokens.shape
|
| 734 |
+
j = 0
|
| 735 |
+
layers = []
|
| 736 |
+
for i in range(len(self.time_blocks)):
|
| 737 |
+
time_tokens = tokens.contiguous().view(B * N, T, -1) # B N T C -> (B N) T C
|
| 738 |
+
time_tokens = self.time_blocks[i](time_tokens)
|
| 739 |
+
|
| 740 |
+
tokens = time_tokens.view(B, N, T, -1) # (B N) T C -> B N T C
|
| 741 |
+
if (
|
| 742 |
+
add_space_attn
|
| 743 |
+
and hasattr(self, "space_virtual_blocks")
|
| 744 |
+
and (i % (len(self.time_blocks) // len(self.space_virtual_blocks)) == 0)
|
| 745 |
+
):
|
| 746 |
+
space_tokens = (
|
| 747 |
+
tokens.permute(0, 2, 1, 3).contiguous().view(B * T, N, -1)
|
| 748 |
+
) # B N T C -> (B T) N C
|
| 749 |
+
|
| 750 |
+
point_tokens = space_tokens[:, : N - self.num_virtual_tracks]
|
| 751 |
+
virtual_tokens = space_tokens[:, N - self.num_virtual_tracks :]
|
| 752 |
+
|
| 753 |
+
virtual_tokens = self.space_virtual2point_blocks[j](
|
| 754 |
+
virtual_tokens, point_tokens, mask=mask
|
| 755 |
+
)
|
| 756 |
+
|
| 757 |
+
virtual_tokens = self.space_virtual_blocks[j](virtual_tokens)
|
| 758 |
+
point_tokens = self.space_point2virtual_blocks[j](
|
| 759 |
+
point_tokens, virtual_tokens, mask=mask
|
| 760 |
+
)
|
| 761 |
+
|
| 762 |
+
space_tokens = torch.cat([point_tokens, virtual_tokens], dim=1)
|
| 763 |
+
tokens = space_tokens.view(B, T, N, -1).permute(
|
| 764 |
+
0, 2, 1, 3
|
| 765 |
+
) # (B T) N C -> B N T C
|
| 766 |
+
j += 1
|
| 767 |
+
tokens = tokens[:, : N - self.num_virtual_tracks]
|
| 768 |
+
|
| 769 |
+
flow = self.flow_head(tokens)
|
| 770 |
+
if self.linear_layer_for_vis_conf:
|
| 771 |
+
vis_conf = self.vis_conf_head(tokens)
|
| 772 |
+
flow = torch.cat([flow, vis_conf], dim=-1)
|
| 773 |
+
|
| 774 |
+
return flow
|
| 775 |
+
|
| 776 |
+
|
| 777 |
+
class BaseBackbone(nn.Module):
|
| 778 |
+
def __init__(self):
|
| 779 |
+
super().__init__()
|
| 780 |
+
self.patch_size = 16
|
| 781 |
+
|
| 782 |
+
self.pos_embed_x = None
|
| 783 |
+
|
| 784 |
+
self.return_inter = False
|
| 785 |
+
|
| 786 |
+
def finetune_track(self, img_size):
|
| 787 |
+
new_patch_size = 16
|
| 788 |
+
|
| 789 |
+
patch_pos_embed = self.absolute_pos_embed
|
| 790 |
+
patch_pos_embed = patch_pos_embed.transpose(1, 2)
|
| 791 |
+
B, E, Q = patch_pos_embed.shape
|
| 792 |
+
P_H, P_W = img_size[0] // self.patch_size, img_size[1] // self.patch_size
|
| 793 |
+
patch_pos_embed = patch_pos_embed.view(B, E, P_H, P_W)
|
| 794 |
+
|
| 795 |
+
# for search region
|
| 796 |
+
H, W = img_size
|
| 797 |
+
new_P_H, new_P_W = H // new_patch_size, W // new_patch_size
|
| 798 |
+
img_patch_pos_embed = nn.functional.interpolate(patch_pos_embed, size=(new_P_H, new_P_W), mode='bicubic',
|
| 799 |
+
align_corners=False)
|
| 800 |
+
img_patch_pos_embed = img_patch_pos_embed.flatten(2).transpose(1, 2)
|
| 801 |
+
|
| 802 |
+
self.pos_embed_x = nn.Parameter(img_patch_pos_embed)
|
| 803 |
+
|
| 804 |
+
# if self.return_inter:
|
| 805 |
+
# for i_layer in self.fpn_stage:
|
| 806 |
+
# if i_layer != 11:
|
| 807 |
+
# norm_layer = partial(nn.LayerNorm, eps=1e-6)
|
| 808 |
+
# layer = norm_layer(self.embed_dim)
|
| 809 |
+
# layer_name = f'norm{i_layer}'
|
| 810 |
+
# self.add_module(layer_name, layer)
|
| 811 |
+
|
| 812 |
+
|
| 813 |
+
def forward(self, x, **kwargs):
|
| 814 |
+
"""
|
| 815 |
+
Joint feature extraction and relation modeling for the basic HiViT backbone.
|
| 816 |
+
Args:
|
| 817 |
+
z (torch.Tensor): template feature, [B, C, H_z, W_z]
|
| 818 |
+
x (torch.Tensor): search region feature, [B, C, H_x, W_x]
|
| 819 |
+
|
| 820 |
+
Returns:
|
| 821 |
+
x (torch.Tensor): merged template and search region feature, [B, L_z+L_x, C]
|
| 822 |
+
attn : None
|
| 823 |
+
"""
|
| 824 |
+
B = x.shape[0]
|
| 825 |
+
|
| 826 |
+
x = self.patch_embed(x)
|
| 827 |
+
|
| 828 |
+
for blk in self.blocks[:-self.num_main_blocks]:
|
| 829 |
+
x = blk(x)
|
| 830 |
+
|
| 831 |
+
x = x[..., 0, 0, :]
|
| 832 |
+
|
| 833 |
+
x += self.pos_embed_x # 添加位置编码
|
| 834 |
+
|
| 835 |
+
x = self.pos_drop(x)
|
| 836 |
+
|
| 837 |
+
for blk in self.blocks[-self.num_main_blocks:]:
|
| 838 |
+
x = blk(x)
|
| 839 |
+
|
| 840 |
+
aux_dict = {"attn": None}
|
| 841 |
+
x = self.norm_(x)
|
| 842 |
+
|
| 843 |
+
return x, aux_dict
|
| 844 |
+
|
| 845 |
+
|
| 846 |
+
class CorrBlock:
|
| 847 |
+
def __init__(
|
| 848 |
+
self,
|
| 849 |
+
fmaps,
|
| 850 |
+
num_levels=4,
|
| 851 |
+
radius=4,
|
| 852 |
+
multiple_track_feats=False,
|
| 853 |
+
padding_mode="zeros",
|
| 854 |
+
appearance_fact_flow_dim=None,
|
| 855 |
+
):
|
| 856 |
+
B, S, C, H, W = fmaps.shape
|
| 857 |
+
self.S, self.C, self.H, self.W = S, C, H, W
|
| 858 |
+
self.padding_mode = padding_mode
|
| 859 |
+
self.num_levels = num_levels
|
| 860 |
+
self.radius = radius
|
| 861 |
+
self.fmaps_pyramid = []
|
| 862 |
+
self.multiple_track_feats = multiple_track_feats
|
| 863 |
+
self.appearance_fact_flow_dim = appearance_fact_flow_dim
|
| 864 |
+
|
| 865 |
+
self.fmaps_pyramid.append(fmaps)
|
| 866 |
+
for i in range(self.num_levels - 1):
|
| 867 |
+
fmaps_ = fmaps.reshape(B * S, C, H, W)
|
| 868 |
+
fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2)
|
| 869 |
+
_, _, H, W = fmaps_.shape
|
| 870 |
+
fmaps = fmaps_.reshape(B, S, C, H, W)
|
| 871 |
+
self.fmaps_pyramid.append(fmaps)
|
| 872 |
+
|
| 873 |
+
def sample(self, coords):
|
| 874 |
+
r = self.radius
|
| 875 |
+
B, S, N, D = coords.shape
|
| 876 |
+
assert D == 2
|
| 877 |
+
|
| 878 |
+
H, W = self.H, self.W
|
| 879 |
+
out_pyramid = []
|
| 880 |
+
for i in range(self.num_levels):
|
| 881 |
+
corrs = self.corrs_pyramid[i] # B, S, N, H, W
|
| 882 |
+
*_, H, W = corrs.shape
|
| 883 |
+
|
| 884 |
+
dx = torch.linspace(-r, r, 2 * r + 1)
|
| 885 |
+
dy = torch.linspace(-r, r, 2 * r + 1)
|
| 886 |
+
delta = torch.stack(torch.meshgrid(dy, dx, indexing="ij"), axis=-1).to(coords.device)
|
| 887 |
+
|
| 888 |
+
centroid_lvl = coords.reshape(B * S * N, 1, 1, 2) / 2**i
|
| 889 |
+
delta_lvl = delta.view(1, 2 * r + 1, 2 * r + 1, 2)
|
| 890 |
+
coords_lvl = centroid_lvl + delta_lvl
|
| 891 |
+
|
| 892 |
+
corrs = bilinear_sampler(
|
| 893 |
+
corrs.reshape(B * S * N, 1, H, W),
|
| 894 |
+
coords_lvl,
|
| 895 |
+
padding_mode=self.padding_mode,
|
| 896 |
+
)
|
| 897 |
+
corrs = corrs.view(B, S, N, -1)
|
| 898 |
+
out_pyramid.append(corrs)
|
| 899 |
+
|
| 900 |
+
out = torch.cat(out_pyramid, dim=-1) # B, S, N, LRR*2
|
| 901 |
+
out = out.permute(0, 2, 1, 3).contiguous().view(B * N, S, -1).float()
|
| 902 |
+
return out
|
| 903 |
+
|
| 904 |
+
def corr(self, targets, coords=None, gt_flow=None, use_gt_flow=False, use_flow_tokens=False,
|
| 905 |
+
use_af_high_dim=False, interaction_network=None):
|
| 906 |
+
assert sum([coords is not None,
|
| 907 |
+
use_gt_flow,
|
| 908 |
+
use_flow_tokens,
|
| 909 |
+
use_af_high_dim,
|
| 910 |
+
interaction_network is not None]) <= 1, \
|
| 911 |
+
"Exactly one of coords, use_gt_flow, use_flow_tokens, use_af_high_dim, or interaction_network must be specified."
|
| 912 |
+
assert not use_flow_tokens or self.appearance_fact_flow_dim is not None
|
| 913 |
+
|
| 914 |
+
# Appearance factorization
|
| 915 |
+
if coords is not None:
|
| 916 |
+
B, S, N, D2 = targets.shape
|
| 917 |
+
targets = targets.reshape(B, S, N, 2, D2 // 2)
|
| 918 |
+
if use_gt_flow:
|
| 919 |
+
flow = gt_flow
|
| 920 |
+
else:
|
| 921 |
+
flow = coords[:, 1:] - coords[:, :-1]
|
| 922 |
+
flow = torch.cat([flow[:, 0:1], flow], dim=1)
|
| 923 |
+
flow = gt_flow if use_gt_flow else flow
|
| 924 |
+
targets = flow[..., 0:1] * targets[..., 0, :] + flow[..., 1:2] * targets[..., 1, :]
|
| 925 |
+
|
| 926 |
+
###### DEBUG ########
|
| 927 |
+
# flow: [B, S, N, 2]
|
| 928 |
+
# targets: [B, S, N, 2, feat_dim], remember feat_dim == latent_dim / 2
|
| 929 |
+
|
| 930 |
+
# Appearane factorization with flow tokens
|
| 931 |
+
if use_flow_tokens:
|
| 932 |
+
fdim = self.appearance_fact_flow_dim
|
| 933 |
+
flow = targets[..., -fdim:]
|
| 934 |
+
targets = targets[..., :-fdim]
|
| 935 |
+
B, S, N, D2 = targets.shape
|
| 936 |
+
targets = targets.reshape(B, S, N, fdim, D2 // fdim)
|
| 937 |
+
#targets = flow[..., 0:1] * targets[..., 0, :] + flow[..., 1:2] * targets[..., 1, :]
|
| 938 |
+
targets = torch.einsum('btnji,btnj->btni', targets, flow)
|
| 939 |
+
|
| 940 |
+
if use_af_high_dim:
|
| 941 |
+
fdim = self.appearance_fact_flow_dim
|
| 942 |
+
flow = targets[..., -fdim:]
|
| 943 |
+
targets = targets[..., :-fdim]
|
| 944 |
+
targets = flow * targets
|
| 945 |
+
|
| 946 |
+
#if interaction_network is not None:
|
| 947 |
+
# flow_feat = targets[..., :self.appearance_fact_flow_dim]
|
| 948 |
+
# targets = targets[..., self.appearance_fact_flow_dim:]
|
| 949 |
+
|
| 950 |
+
# Correlation
|
| 951 |
+
B, S, N, C = targets.shape
|
| 952 |
+
#assert C == self.C
|
| 953 |
+
assert S == self.S
|
| 954 |
+
|
| 955 |
+
fmap1 = targets
|
| 956 |
+
|
| 957 |
+
self.corrs_pyramid = []
|
| 958 |
+
for _, fmaps in enumerate(self.fmaps_pyramid):
|
| 959 |
+
*_, H, W = fmaps.shape
|
| 960 |
+
fmap2s = fmaps.view(B, S, C, H * W) # B S C H W -> B S C (H W)
|
| 961 |
+
corrs = torch.matmul(fmap1, fmap2s)
|
| 962 |
+
corrs = corrs.view(B, S, N, H, W) # B S N (H W) -> B S N H W
|
| 963 |
+
corrs = corrs / torch.sqrt(torch.tensor(C).float())
|
| 964 |
+
self.corrs_pyramid.append(corrs)
|
| 965 |
+
|
| 966 |
+
def sample_fmap(self, coords):
|
| 967 |
+
'''Sample at coords directly from scaled feature map.
|
| 968 |
+
'''
|
| 969 |
+
B, S, N, D = coords.shape
|
| 970 |
+
assert D == 2
|
| 971 |
+
|
| 972 |
+
H, W = self.H, self.W
|
| 973 |
+
out_pyramid = []
|
| 974 |
+
for i in range(self.num_levels):
|
| 975 |
+
|
| 976 |
+
fmap = self.fmaps_pyramid[i] # B, S, C, H, W
|
| 977 |
+
B, S, C, H, W = fmap.shape
|
| 978 |
+
|
| 979 |
+
coords_lvl = coords / 2**i
|
| 980 |
+
|
| 981 |
+
fmap = fmap.reshape(B * S, C, H, W)
|
| 982 |
+
coords_lvl = coords_lvl.reshape(B * S, N, 2)
|
| 983 |
+
|
| 984 |
+
coords_normalized = coords_lvl.clone()
|
| 985 |
+
coords_normalized[..., 0] = coords_normalized[..., 0] / (W - 1) * 2 - 1
|
| 986 |
+
coords_normalized[..., 1] = coords_normalized[..., 1] / (H - 1) * 2 - 1
|
| 987 |
+
coords_normalized = coords_normalized.unsqueeze(1)
|
| 988 |
+
feature_at_coords = F.grid_sample(fmap, coords_normalized,
|
| 989 |
+
mode='bilinear', align_corners=True)
|
| 990 |
+
feature_at_coords = feature_at_coords.permute(0, 2, 3, 1).view(B, S, N, C)
|
| 991 |
+
out_pyramid.append(feature_at_coords)
|
| 992 |
+
|
| 993 |
+
out = torch.cat(out_pyramid, dim=-1) # B, S, N, LRR*2
|
| 994 |
+
return out
|
LFE_TAP/models/embeddings.py
ADDED
|
@@ -0,0 +1,110 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
from typing import Tuple, Union
|
| 3 |
+
|
| 4 |
+
def get_1d_sincos_pos_embed_from_grid(
|
| 5 |
+
embed_dim: int, pos: torch.Tensor
|
| 6 |
+
) -> torch.Tensor:
|
| 7 |
+
"""
|
| 8 |
+
This function generates a 1D positional embedding from a given grid using sine and cosine functions.
|
| 9 |
+
|
| 10 |
+
Args:
|
| 11 |
+
- embed_dim: The embedding dimension.
|
| 12 |
+
- pos: The position to generate the embedding from.
|
| 13 |
+
|
| 14 |
+
Returns:
|
| 15 |
+
- emb: The generated 1D positional embedding.
|
| 16 |
+
"""
|
| 17 |
+
assert embed_dim % 2 == 0
|
| 18 |
+
omega = torch.arange(embed_dim // 2, dtype=torch.double)
|
| 19 |
+
omega /= embed_dim / 2.0
|
| 20 |
+
omega = 1.0 / 10000**omega # (D/2,)
|
| 21 |
+
|
| 22 |
+
pos = pos.reshape(-1) # (M,)
|
| 23 |
+
out = torch.einsum("m,d->md", pos, omega) # (M, D/2), outer product
|
| 24 |
+
|
| 25 |
+
emb_sin = torch.sin(out) # (M, D/2)
|
| 26 |
+
emb_cos = torch.cos(out) # (M, D/2)
|
| 27 |
+
|
| 28 |
+
emb = torch.cat([emb_sin, emb_cos], dim=1) # (M, D)
|
| 29 |
+
return emb[None].float()
|
| 30 |
+
|
| 31 |
+
def get_2d_embedding(xy: torch.Tensor, C: int, cat_coords: bool = True) -> torch.Tensor:
|
| 32 |
+
"""
|
| 33 |
+
This function generates a 2D positional embedding from given coordinates using sine and cosine functions.
|
| 34 |
+
|
| 35 |
+
Args:
|
| 36 |
+
- xy: The coordinates to generate the embedding from.
|
| 37 |
+
- C: The size of the embedding.
|
| 38 |
+
- cat_coords: A flag to indicate whether to concatenate the original coordinates to the embedding.
|
| 39 |
+
|
| 40 |
+
Returns:
|
| 41 |
+
- pe: The generated 2D positional embedding.
|
| 42 |
+
"""
|
| 43 |
+
B, N, D = xy.shape
|
| 44 |
+
assert D == 2
|
| 45 |
+
|
| 46 |
+
x = xy[:, :, 0:1]
|
| 47 |
+
y = xy[:, :, 1:2]
|
| 48 |
+
div_term = (
|
| 49 |
+
torch.arange(0, C, 2, device=xy.device, dtype=torch.float32) * (1000.0 / C)
|
| 50 |
+
).reshape(1, 1, int(C / 2))
|
| 51 |
+
|
| 52 |
+
pe_x = torch.zeros(B, N, C, device=xy.device, dtype=torch.float32)
|
| 53 |
+
pe_y = torch.zeros(B, N, C, device=xy.device, dtype=torch.float32)
|
| 54 |
+
|
| 55 |
+
pe_x[:, :, 0::2] = torch.sin(x * div_term)
|
| 56 |
+
pe_x[:, :, 1::2] = torch.cos(x * div_term)
|
| 57 |
+
|
| 58 |
+
pe_y[:, :, 0::2] = torch.sin(y * div_term)
|
| 59 |
+
pe_y[:, :, 1::2] = torch.cos(y * div_term)
|
| 60 |
+
|
| 61 |
+
pe = torch.cat([pe_x, pe_y], dim=2) # (B, N, C*3)
|
| 62 |
+
if cat_coords:
|
| 63 |
+
pe = torch.cat([xy, pe], dim=2) # (B, N, C*3+3)
|
| 64 |
+
return pe
|
| 65 |
+
|
| 66 |
+
def get_2d_sincos_pos_embed_from_grid(
|
| 67 |
+
embed_dim: int, grid: torch.Tensor
|
| 68 |
+
) -> torch.Tensor:
|
| 69 |
+
"""
|
| 70 |
+
This function generates a 2D positional embedding from a given grid using sine and cosine functions.
|
| 71 |
+
|
| 72 |
+
Args:
|
| 73 |
+
- embed_dim: The embedding dimension.
|
| 74 |
+
- grid: The grid to generate the embedding from.
|
| 75 |
+
|
| 76 |
+
Returns:
|
| 77 |
+
- emb: The generated 2D positional embedding.
|
| 78 |
+
"""
|
| 79 |
+
assert embed_dim % 2 == 0
|
| 80 |
+
|
| 81 |
+
# use half of dimensions to encode grid_h
|
| 82 |
+
emb_h = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[0]) # (H*W, D/2)
|
| 83 |
+
emb_w = get_1d_sincos_pos_embed_from_grid(embed_dim // 2, grid[1]) # (H*W, D/2)
|
| 84 |
+
|
| 85 |
+
emb = torch.cat([emb_h, emb_w], dim=2) # (H*W, D)
|
| 86 |
+
return emb
|
| 87 |
+
|
| 88 |
+
def get_2d_sincos_pos_embed(
|
| 89 |
+
embed_dim: int, grid_size: Union[int, Tuple[int, int]]
|
| 90 |
+
) -> torch.Tensor:
|
| 91 |
+
"""
|
| 92 |
+
This function initializes a grid and generates a 2D positional embedding using sine and cosine functions.
|
| 93 |
+
It is a wrapper of get_2d_sincos_pos_embed_from_grid.
|
| 94 |
+
Args:
|
| 95 |
+
- embed_dim: The embedding dimension.
|
| 96 |
+
- grid_size: The grid size.
|
| 97 |
+
Returns:
|
| 98 |
+
- pos_embed: The generated 2D positional embedding.
|
| 99 |
+
"""
|
| 100 |
+
if isinstance(grid_size, tuple):
|
| 101 |
+
grid_size_h, grid_size_w = grid_size
|
| 102 |
+
else:
|
| 103 |
+
grid_size_h = grid_size_w = grid_size
|
| 104 |
+
grid_h = torch.arange(grid_size_h, dtype=torch.float)
|
| 105 |
+
grid_w = torch.arange(grid_size_w, dtype=torch.float)
|
| 106 |
+
grid = torch.meshgrid(grid_w, grid_h, indexing="xy")
|
| 107 |
+
grid = torch.stack(grid, dim=0)
|
| 108 |
+
grid = grid.reshape([2, 1, grid_size_h, grid_size_w])
|
| 109 |
+
pos_embed = get_2d_sincos_pos_embed_from_grid(embed_dim, grid)
|
| 110 |
+
return pos_embed.reshape(1, grid_size_h, grid_size_w, -1).permute(0, 3, 1, 2)
|
LFE_TAP/models/fusionFormer.py
ADDED
|
@@ -0,0 +1,253 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
from einops import rearrange
|
| 5 |
+
from einops.layers.torch import Rearrange
|
| 6 |
+
from LFE_TAP.models.blocks import ResidualBlock, CrossAttnBlock, AttnBlock2, BasicEncoder
|
| 7 |
+
from timm.models.vision_transformer import Mlp
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class ST_Transformer(nn.Module):
|
| 11 |
+
def __init__(self, dim, heads, mlp_dim=512, dropout=0., mlp_ratio=4.0):
|
| 12 |
+
super().__init__()
|
| 13 |
+
self.event_t = CrossAttnBlock(128, 128, num_heads=8, mlp_ratio=4.0, dim_head=16)
|
| 14 |
+
self.fe_space = CrossAttnBlock(128, 128, num_heads=8, mlp_ratio=4.0, dim_head=16)
|
| 15 |
+
|
| 16 |
+
def forward(self, x_i, x_e, x_e_pre):
|
| 17 |
+
x_q = self.event_t(x_e, x_e_pre)
|
| 18 |
+
x = self.fe_space(x_q, x_i)
|
| 19 |
+
|
| 20 |
+
return x, x_q
|
| 21 |
+
|
| 22 |
+
class downsample(nn.Module):
|
| 23 |
+
def __init__(self, in_channels, out_channels):
|
| 24 |
+
super().__init__()
|
| 25 |
+
self.layer = nn.Sequential(
|
| 26 |
+
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=2, padding=1),
|
| 27 |
+
nn.BatchNorm2d(out_channels),
|
| 28 |
+
nn.ReLU(inplace=True),
|
| 29 |
+
)
|
| 30 |
+
def forward(self, x):
|
| 31 |
+
return self.layer(x)
|
| 32 |
+
|
| 33 |
+
|
| 34 |
+
class upsample(nn.Module):
|
| 35 |
+
def __init__(self, in_channels, out_channels):
|
| 36 |
+
super().__init__()
|
| 37 |
+
self.layer = nn.Sequential(
|
| 38 |
+
nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1),
|
| 39 |
+
nn.BatchNorm2d(out_channels),
|
| 40 |
+
nn.ReLU(inplace=True),
|
| 41 |
+
)
|
| 42 |
+
self.cov = nn.Sequential(
|
| 43 |
+
nn.Conv2d(out_channels * 2, out_channels, kernel_size=1, stride=1, padding=0),
|
| 44 |
+
nn.BatchNorm2d(out_channels),
|
| 45 |
+
nn.ReLU(inplace=True),
|
| 46 |
+
)
|
| 47 |
+
def forward(self, x1, x2):
|
| 48 |
+
x1 = F.interpolate(x1, scale_factor=2, mode="bilinear", align_corners=True)
|
| 49 |
+
x_ = self.layer(x1)
|
| 50 |
+
if x2 is not None:
|
| 51 |
+
x_ = torch.cat([x2, x_], dim=1)
|
| 52 |
+
x_ = self.cov(x_)
|
| 53 |
+
else:
|
| 54 |
+
x_ = self.cov(torch.cat([x_, x_], dim=1))
|
| 55 |
+
return x_
|
| 56 |
+
|
| 57 |
+
|
| 58 |
+
class Fusionformer(nn.Module):
|
| 59 |
+
def __init__(self, image_size=(384, 512), out_dim=128, mlp_dim=512, depth=6, stride=8, dropout=0.):
|
| 60 |
+
super().__init__()
|
| 61 |
+
img_h, img_w = image_size
|
| 62 |
+
self.stride = stride
|
| 63 |
+
self.in_planes = 32
|
| 64 |
+
|
| 65 |
+
self.fnet_img = BasicEncoder(
|
| 66 |
+
output_dim=128, norm_fn="instance", dropout=0, stride=stride, shallow=True, in_planes=32
|
| 67 |
+
)
|
| 68 |
+
self.fnet_event = BasicEncoder(
|
| 69 |
+
input_dim=10, output_dim=128, norm_fn="instance", dropout=0, stride=stride, shallow=True, in_planes=32, dilation=1
|
| 70 |
+
)
|
| 71 |
+
|
| 72 |
+
self.transunet = CLWF(128, out_dim, image_size, stride, mlp_dim, depth, dropout)
|
| 73 |
+
|
| 74 |
+
# self.resnet = ResidualBlock(128, out_dim, stride=1)
|
| 75 |
+
|
| 76 |
+
def forward(self, x_i, x_e, img_ifnew=None, feature_teacher=None):
|
| 77 |
+
_, _, H, W = x_i.size()
|
| 78 |
+
|
| 79 |
+
x_i = self.fnet_img(x_i)
|
| 80 |
+
x_e = self.fnet_event(x_e)
|
| 81 |
+
|
| 82 |
+
x_out = self.transunet(x_i, x_e, img_ifnew, feature_teacher)
|
| 83 |
+
return x_out
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
class CLWF(nn.Module):
|
| 87 |
+
def __init__(self, in_channels, out_channels, image_size=(384, 512), stride=8, mlp_dim=512, depth=3, dropout=0.):
|
| 88 |
+
super().__init__()
|
| 89 |
+
img_h, img_w = image_size
|
| 90 |
+
self.patches_resolution = (img_h // (stride * 4), img_w // (stride * 4))
|
| 91 |
+
num_patches = self.patches_resolution[0] * self.patches_resolution[1]
|
| 92 |
+
|
| 93 |
+
self.down1 = downsample(in_channels, 192)
|
| 94 |
+
self.down2 = downsample(192, 256)
|
| 95 |
+
|
| 96 |
+
self.up1 = upsample(256, 192)
|
| 97 |
+
self.up2 = upsample(192, out_channels)
|
| 98 |
+
|
| 99 |
+
self.pos_embedding = nn.Parameter(torch.randn(1, num_patches, 256))
|
| 100 |
+
self.dropout_i = nn.Dropout(dropout)
|
| 101 |
+
self.dropout_e = nn.Dropout(dropout)
|
| 102 |
+
|
| 103 |
+
self.xe_history = []
|
| 104 |
+
self.x_e_pre = None
|
| 105 |
+
self.x_out_ = None
|
| 106 |
+
# self.x_i_last = None
|
| 107 |
+
|
| 108 |
+
self.cov_out1 = nn.Conv2d(256, 128, kernel_size=1)
|
| 109 |
+
self.cov_out2 = nn.Conv2d(192, 128, kernel_size=1)
|
| 110 |
+
|
| 111 |
+
self.layers = nn.ModuleList(
|
| 112 |
+
[
|
| 113 |
+
CrossAttnBlock(256, 256, num_heads=8, mlp_ratio=4.0, dim_head=32)
|
| 114 |
+
for _ in range(depth)
|
| 115 |
+
]
|
| 116 |
+
)
|
| 117 |
+
self.TemporalAdapter = nn.ModuleList(
|
| 118 |
+
[
|
| 119 |
+
AttnBlock2(256, num_heads=8, mlp_ratio=4.0, dim_head=32)
|
| 120 |
+
for _ in range(depth)
|
| 121 |
+
]
|
| 122 |
+
)
|
| 123 |
+
|
| 124 |
+
def forward(self, x_i, x_e, img_ifnew=None, feature_teacher=None):
|
| 125 |
+
x1_i = self.down1(x_i)
|
| 126 |
+
x2_i = self.down2(x1_i)
|
| 127 |
+
|
| 128 |
+
x1_e = self.down1(x_e)
|
| 129 |
+
x2_e = self.down2(x1_e)
|
| 130 |
+
|
| 131 |
+
x2_i_ = x2_i.clone()
|
| 132 |
+
x2_e_ = x2_e.clone()
|
| 133 |
+
x2_i_ = rearrange(x2_i_, 'b c h w -> b (h w) c')
|
| 134 |
+
x2_e_ = rearrange(x2_e_, 'b c h w -> b (h w) c')
|
| 135 |
+
|
| 136 |
+
b, n, _ = x2_i_.shape
|
| 137 |
+
|
| 138 |
+
x2_i_ = self.dropout_i(x2_i_)
|
| 139 |
+
x2_e_ = self.dropout_e(x2_e_)
|
| 140 |
+
|
| 141 |
+
x_out = []
|
| 142 |
+
for time in range(len(x_i)):
|
| 143 |
+
x_i_t = x2_i_[time:time+1]
|
| 144 |
+
x_e_t = x2_e_[time:time+1]
|
| 145 |
+
|
| 146 |
+
if img_ifnew[time] == 1:
|
| 147 |
+
for st_attn in self.layers:
|
| 148 |
+
x_i_t = st_attn(x_i_t, x_e_t)
|
| 149 |
+
self.x_out_ = x_i_t
|
| 150 |
+
else:
|
| 151 |
+
x_i_t = self.x_out_
|
| 152 |
+
for st_attn in self.layers:
|
| 153 |
+
x_i_t = st_attn(x_i_t, x_e_t)
|
| 154 |
+
self.x_out_ = x_i_t
|
| 155 |
+
|
| 156 |
+
x_out.append(rearrange(self.x_out_, 'b (h w) out_dim -> b out_dim h w', h=self.patches_resolution[0], w=self.patches_resolution[1]))
|
| 157 |
+
|
| 158 |
+
x_out = torch.cat(x_out, dim=0)
|
| 159 |
+
x_out = rearrange(x_out, 't c h w -> (h w) t c')
|
| 160 |
+
for time_attn in self.TemporalAdapter:
|
| 161 |
+
x_out = time_attn(x_out)
|
| 162 |
+
x_out1 = rearrange(x_out, '(h w) t c -> t c h w', h=self.patches_resolution[0], w=self.patches_resolution[1])
|
| 163 |
+
x_out2 = self.up1(x_out1, torch.tensor(img_ifnew)[:,None,None,None].to(x1_e.device).float() * (x1_i) + torch.tensor(1-img_ifnew)[:,None,None,None].to(x1_e.device).float() * (x1_e))
|
| 164 |
+
x_out3 = self.up2(x_out2, torch.tensor(img_ifnew)[:,None,None,None].to(x1_e.device).float() * (x_i) + torch.tensor(1-img_ifnew)[:,None,None,None].to(x1_e.device).float() * (x_e))
|
| 165 |
+
x_out1 = self.cov_out1(x_out1)
|
| 166 |
+
x_out2 = self.cov_out2(x_out2)
|
| 167 |
+
x_out_pyramid = [x_out3, x_out2, x_out1]
|
| 168 |
+
return x_out_pyramid
|
| 169 |
+
|
| 170 |
+
|
| 171 |
+
class Unet_Transformer(nn.Module):
|
| 172 |
+
def __init__(self, input_dim=3, image_size=(384, 512), out_dim=128, mlp_dim=512, depth=6, stride=8, dropout=0.):
|
| 173 |
+
super().__init__()
|
| 174 |
+
img_h, img_w = image_size
|
| 175 |
+
self.stride = stride
|
| 176 |
+
self.in_planes = 32
|
| 177 |
+
|
| 178 |
+
self.fnet = BasicEncoder(
|
| 179 |
+
input_dim=input_dim, output_dim=128, norm_fn="instance", dropout=0, stride=stride, shallow=True, in_planes=32
|
| 180 |
+
)
|
| 181 |
+
|
| 182 |
+
self.transunet = TransUnet_pyramid_onemod(128, out_dim, image_size, stride, mlp_dim, depth, dropout)
|
| 183 |
+
|
| 184 |
+
# self.resnet = ResidualBlock(128, out_dim, stride=1)
|
| 185 |
+
|
| 186 |
+
def forward(self, x, feature_teacher=None):
|
| 187 |
+
_, _, H, W = x.size()
|
| 188 |
+
|
| 189 |
+
x = self.fnet(x)
|
| 190 |
+
|
| 191 |
+
x_out = self.transunet(x, feature_teacher)
|
| 192 |
+
|
| 193 |
+
# x_out = self.resnet(x_out)
|
| 194 |
+
|
| 195 |
+
return x_out
|
| 196 |
+
|
| 197 |
+
|
| 198 |
+
class TransUnet_pyramid_onemod(nn.Module):
|
| 199 |
+
def __init__(self, in_channels, out_channels, image_size=(384, 512), stride=8, mlp_dim=512, depth=3, dropout=0.):
|
| 200 |
+
super().__init__()
|
| 201 |
+
img_h, img_w = image_size
|
| 202 |
+
self.patches_resolution = (img_h // (stride * 4), img_w // (stride * 4))
|
| 203 |
+
num_patches = self.patches_resolution[0] * self.patches_resolution[1]
|
| 204 |
+
|
| 205 |
+
self.down1 = downsample(in_channels, 192)
|
| 206 |
+
self.down2 = downsample(192, 256)
|
| 207 |
+
|
| 208 |
+
self.up1 = upsample(256, 192)
|
| 209 |
+
self.up2 = upsample(192, out_channels)
|
| 210 |
+
|
| 211 |
+
self.pos_embedding = nn.Parameter(torch.randn(1, num_patches, 256))
|
| 212 |
+
self.dropout_i = nn.Dropout(dropout)
|
| 213 |
+
self.dropout_e = nn.Dropout(dropout)
|
| 214 |
+
|
| 215 |
+
self.xe_history = []
|
| 216 |
+
self.x_e_pre = None
|
| 217 |
+
self.x_out_ = None
|
| 218 |
+
|
| 219 |
+
self.cov_out1 = nn.Conv2d(256, 128, kernel_size=1)
|
| 220 |
+
self.cov_out2 = nn.Conv2d(192, 128, kernel_size=1)
|
| 221 |
+
|
| 222 |
+
|
| 223 |
+
self.TemporalAdapter = nn.ModuleList(
|
| 224 |
+
[
|
| 225 |
+
AttnBlock2(256, num_heads=8, mlp_ratio=4.0, dim_head=32)
|
| 226 |
+
for _ in range(depth)
|
| 227 |
+
]
|
| 228 |
+
)
|
| 229 |
+
|
| 230 |
+
def forward(self, x, feature_teacher=None):
|
| 231 |
+
x1 = self.down1(x)
|
| 232 |
+
x2 = self.down2(x1)
|
| 233 |
+
|
| 234 |
+
x2_ = x2.clone()
|
| 235 |
+
x2_ = rearrange(x2_, 'b c h w -> b (h w) c')
|
| 236 |
+
|
| 237 |
+
b, n, _ = x2_.shape
|
| 238 |
+
|
| 239 |
+
x2_ = self.dropout_i(x2_)
|
| 240 |
+
|
| 241 |
+
x2_ = x2_.permute(1,0,2)
|
| 242 |
+
|
| 243 |
+
for time_attn in self.TemporalAdapter:
|
| 244 |
+
x_out = time_attn(x2_)
|
| 245 |
+
x_out1 = rearrange(x_out, '(h w) t c -> t c h w', h=self.patches_resolution[0], w=self.patches_resolution[1])
|
| 246 |
+
x_out2 = self.up1(x_out1, x1)
|
| 247 |
+
x_out3 = self.up2(x_out2, x)
|
| 248 |
+
x_out1 = self.cov_out1(x_out1)
|
| 249 |
+
x_out2 = self.cov_out2(x_out2)
|
| 250 |
+
x_out_pyramid = [x_out3, x_out2, x_out1]
|
| 251 |
+
|
| 252 |
+
|
| 253 |
+
return x_out_pyramid
|
LFE_TAP/models/tapformer.py
ADDED
|
@@ -0,0 +1,308 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import torch.nn as nn
|
| 3 |
+
import torch.nn.functional as F
|
| 4 |
+
import time
|
| 5 |
+
|
| 6 |
+
from LFE_TAP.models.blocks import (
|
| 7 |
+
BasicEncoder,
|
| 8 |
+
FusionBlock,
|
| 9 |
+
FusionBlock_basic,
|
| 10 |
+
EfficientUpdateFormer,
|
| 11 |
+
UpdateFormer,
|
| 12 |
+
Mlp,
|
| 13 |
+
)
|
| 14 |
+
from LFE_TAP.models.fusionFormer import Fusionformer
|
| 15 |
+
from LFE_TAP.utils.model_utils import get_track_feat, bilinear_sampler, get_support_points
|
| 16 |
+
from LFE_TAP.models.embeddings import get_1d_sincos_pos_embed_from_grid
|
| 17 |
+
|
| 18 |
+
torch.manual_seed(0)
|
| 19 |
+
|
| 20 |
+
|
| 21 |
+
def posenc(x, min_deg, max_deg):
|
| 22 |
+
"""Cat x with a positional encoding of x with scales 2^[min_deg, max_deg-1].
|
| 23 |
+
Instead of computing [sin(x), cos(x)], we use the trig identity
|
| 24 |
+
cos(x) = sin(x + pi/2) and do one vectorized call to sin([x, x+pi/2]).
|
| 25 |
+
Args:
|
| 26 |
+
x: torch.Tensor, variables to be encoded. Note that x should be in [-pi, pi].
|
| 27 |
+
min_deg: int, the minimum (inclusive) degree of the encoding.
|
| 28 |
+
max_deg: int, the maximum (exclusive) degree of the encoding.
|
| 29 |
+
legacy_posenc_order: bool, keep the same ordering as the original tf code.
|
| 30 |
+
Returns:
|
| 31 |
+
encoded: torch.Tensor, encoded variables.
|
| 32 |
+
"""
|
| 33 |
+
if min_deg == max_deg:
|
| 34 |
+
return x
|
| 35 |
+
scales = torch.tensor(
|
| 36 |
+
[2**i for i in range(min_deg, max_deg)], dtype=x.dtype, device=x.device
|
| 37 |
+
)
|
| 38 |
+
|
| 39 |
+
xb = (x[..., None, :] * scales[:, None]).reshape(list(x.shape[:-1]) + [-1])
|
| 40 |
+
four_feat = torch.sin(torch.cat([xb, xb + 0.5 * torch.pi], dim=-1))
|
| 41 |
+
return torch.cat([x] + [four_feat], dim=-1)
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
class TAPFormer(nn.Module):
|
| 45 |
+
def __init__(self, window_size=16, stride=8, corr_radius=3, corr_levels=3, backbone="basic", num_heads=8, hidden_size=384, space_depth=3, time_depth=3):
|
| 46 |
+
super(TAPFormer, self).__init__()
|
| 47 |
+
self.window_size = window_size
|
| 48 |
+
self.stride = stride
|
| 49 |
+
self.corr_radius = corr_radius
|
| 50 |
+
self.corr_levels = corr_levels
|
| 51 |
+
self.hidden_size = hidden_size
|
| 52 |
+
self.space_depth = space_depth
|
| 53 |
+
self.time_depth = time_depth
|
| 54 |
+
self.latent_dim = 128
|
| 55 |
+
self.backbone = backbone
|
| 56 |
+
self.model_resolution = (384, 512)
|
| 57 |
+
self.mlp_output_dim = 256
|
| 58 |
+
self.input_dim = 2 + 84 + self.mlp_output_dim * self.corr_levels
|
| 59 |
+
num_virtual_tracks = 32
|
| 60 |
+
|
| 61 |
+
|
| 62 |
+
self.fusion_block = Fusionformer(image_size=self.model_resolution, out_dim=self.latent_dim, mlp_dim=512, stride=self.stride, depth=2)
|
| 63 |
+
self.updateformer2 = EfficientUpdateFormer(
|
| 64 |
+
space_depth=space_depth,
|
| 65 |
+
time_depth=time_depth,
|
| 66 |
+
input_dim=self.input_dim,
|
| 67 |
+
hidden_size=hidden_size,
|
| 68 |
+
output_dim=4,
|
| 69 |
+
mlp_ratio=4.0,
|
| 70 |
+
num_virtual_tracks=num_virtual_tracks,
|
| 71 |
+
linear_layer_for_vis_conf=True,
|
| 72 |
+
)
|
| 73 |
+
|
| 74 |
+
# self.norm = nn.GroupNorm(1, self.latent_dim)
|
| 75 |
+
# self.ffeat_updater = nn.Sequential(
|
| 76 |
+
# nn.Linear(self.latent_dim, self.latent_dim),
|
| 77 |
+
# nn.ReLU(),
|
| 78 |
+
# )
|
| 79 |
+
self.corr_mlp = Mlp(in_features=(2*corr_radius+1) ** 4, hidden_features=384, out_features=256)
|
| 80 |
+
|
| 81 |
+
time_grid = torch.linspace(0, window_size - 1, window_size).reshape(1, window_size, 1)
|
| 82 |
+
self.register_buffer(
|
| 83 |
+
"time_emb", get_1d_sincos_pos_embed_from_grid(self.input_dim, time_grid[0])
|
| 84 |
+
)
|
| 85 |
+
|
| 86 |
+
def interpolate_time_embed(self, x, t):
|
| 87 |
+
previous_dtype = x.dtype
|
| 88 |
+
T = self.time_emb.shape[1]
|
| 89 |
+
|
| 90 |
+
if t == T:
|
| 91 |
+
return self.time_emb
|
| 92 |
+
|
| 93 |
+
time_emb = self.time_emb.float()
|
| 94 |
+
time_emb = F.interpolate(
|
| 95 |
+
time_emb.permute(0, 2, 1), size=t, mode="linear"
|
| 96 |
+
).permute(0, 2, 1)
|
| 97 |
+
return time_emb.to(previous_dtype)
|
| 98 |
+
|
| 99 |
+
def get_correlation_feat(self, fmaps, queried_coords):
|
| 100 |
+
B, T, D, H_, W_ = fmaps.shape
|
| 101 |
+
N = queried_coords.shape[1]
|
| 102 |
+
r = self.corr_radius
|
| 103 |
+
sample_coords = torch.cat(
|
| 104 |
+
[torch.zeros_like(queried_coords[..., :1]), queried_coords], dim=-1
|
| 105 |
+
)[:, None]
|
| 106 |
+
support_points = get_support_points(sample_coords, r, reshape_back=False)
|
| 107 |
+
correlation_feat = bilinear_sampler(
|
| 108 |
+
fmaps.reshape(B * T, D, 1, H_, W_), support_points
|
| 109 |
+
)
|
| 110 |
+
return correlation_feat.view(B, T, D, N, (2 * r + 1), (2 * r + 1)).permute(
|
| 111 |
+
0, 1, 3, 4, 5, 2
|
| 112 |
+
)
|
| 113 |
+
|
| 114 |
+
def forward_window(self, fmaps_pyramid, coords, track_feat_support_pyramid, vis, conf, attenstion_mask, iters=4):
|
| 115 |
+
B, S, D, *_ = fmaps_pyramid[0].shape
|
| 116 |
+
N = coords.shape[2]
|
| 117 |
+
r = 2 * self.corr_radius + 1
|
| 118 |
+
|
| 119 |
+
coord_preds, vis_preds, conf_preds = [], [], []
|
| 120 |
+
for it in range(iters):
|
| 121 |
+
coords = coords.detach()
|
| 122 |
+
coord_init = coords.view(B * S, N, 2)
|
| 123 |
+
corr_embs = []
|
| 124 |
+
for i in range(self.corr_levels):
|
| 125 |
+
corr_feat = self.get_correlation_feat(fmaps_pyramid[i], coord_init / 2**i)
|
| 126 |
+
track_feat_support = (
|
| 127 |
+
track_feat_support_pyramid[i]
|
| 128 |
+
.view(B, 1, r, r, N, self.latent_dim)
|
| 129 |
+
.squeeze(1)
|
| 130 |
+
.permute(0, 3, 1, 2, 4)
|
| 131 |
+
)
|
| 132 |
+
corr_volume = torch.einsum("btnhwc,bnijc->btnhwij", corr_feat, track_feat_support)
|
| 133 |
+
corr_emb = self.corr_mlp(corr_volume.reshape(B * S * N, r * r * r *r))
|
| 134 |
+
# del corr_volume, corr_feat
|
| 135 |
+
# torch.cuda.empty_cache()
|
| 136 |
+
corr_embs.append(corr_emb)
|
| 137 |
+
|
| 138 |
+
corr_embs = torch.cat(corr_embs, dim=1)
|
| 139 |
+
corr_embs = corr_embs.view(B, S, N, corr_embs.shape[-1])
|
| 140 |
+
|
| 141 |
+
transformer_input = [vis, conf, corr_embs]
|
| 142 |
+
|
| 143 |
+
rel_coords_forward = coords[:, :-1] - coords[:, 1:]
|
| 144 |
+
rel_coords_backward = coords[:, 1:] - coords[:, :-1]
|
| 145 |
+
|
| 146 |
+
rel_coords_forward = torch.nn.functional.pad(rel_coords_forward, (0, 0, 0, 0, 0, 1))
|
| 147 |
+
rel_coords_backward = torch.nn.functional.pad(rel_coords_backward, (0, 0, 0, 0, 1, 0))
|
| 148 |
+
|
| 149 |
+
scale = (torch.tensor([self.model_resolution[1], self.model_resolution[0]], device=coords.device,) / self.stride)
|
| 150 |
+
rel_coords_forward = rel_coords_forward / scale # 归一化到[-1, 1]
|
| 151 |
+
rel_coords_backward = rel_coords_backward / scale
|
| 152 |
+
|
| 153 |
+
rel_pos_emb_input = posenc(torch.cat([rel_coords_forward, rel_coords_backward], dim=-1), min_deg=0, max_deg=10,)
|
| 154 |
+
transformer_input.append(rel_pos_emb_input)
|
| 155 |
+
|
| 156 |
+
x = (torch.cat(transformer_input, dim=-1).permute(0, 2, 1, 3).reshape(B*N, S, -1))
|
| 157 |
+
|
| 158 |
+
x = x + self.interpolate_time_embed(x, S)
|
| 159 |
+
x = x.view(B, N, S, -1)
|
| 160 |
+
|
| 161 |
+
delta = self.updateformer2(x)
|
| 162 |
+
|
| 163 |
+
delta_coords = delta[..., :2].permute(0, 2, 1, 3)
|
| 164 |
+
delta_vis = delta[..., 2:3].permute(0, 2, 1, 3)
|
| 165 |
+
delta_conf = delta[..., 3:].permute(0, 2, 1, 3)
|
| 166 |
+
|
| 167 |
+
vis = vis + delta_vis
|
| 168 |
+
conf = conf + delta_conf
|
| 169 |
+
|
| 170 |
+
coords = coords + delta_coords
|
| 171 |
+
coord_preds.append(coords[..., :2] * float(self.stride))
|
| 172 |
+
|
| 173 |
+
vis_preds.append(vis[..., 0])
|
| 174 |
+
conf_preds.append(conf[..., 0])
|
| 175 |
+
return coord_preds, vis_preds, conf_preds
|
| 176 |
+
|
| 177 |
+
def forward(self, rgbs, events, queries, iters=4, img_ifnew=None, feat_init=None, is_train=False):
|
| 178 |
+
B, T, C, H, W = events.shape
|
| 179 |
+
B, N, _ = queries.shape
|
| 180 |
+
_, _, C_img, _, _ = rgbs.shape
|
| 181 |
+
S = self.window_size
|
| 182 |
+
step = S // 2
|
| 183 |
+
device = events.device
|
| 184 |
+
assert H % self.stride == 0 and W % self.stride == 0
|
| 185 |
+
assert B == 1, "batch size should be 1"
|
| 186 |
+
|
| 187 |
+
queried_frames = queries[:, :, 0].long()
|
| 188 |
+
queried_coords = queries[..., 1:3]
|
| 189 |
+
queried_coords = queried_coords / self.stride
|
| 190 |
+
|
| 191 |
+
coords_predicted = torch.zeros((B, T, N, 2), device=device)
|
| 192 |
+
vis_predicted= torch.zeros((B, T, N), device=device)
|
| 193 |
+
conf_predicted = torch.zeros((B, T, N), device=device)
|
| 194 |
+
|
| 195 |
+
all_coords_predictions, all_vis_predictions, all_confidence_predictions = ([], [], [])
|
| 196 |
+
H_stride, W_stride = H // self.stride, W // self.stride
|
| 197 |
+
|
| 198 |
+
rgbs = 2 * (rgbs / 255.0) - 1.0
|
| 199 |
+
# events = torch.sigmoid(events)
|
| 200 |
+
events = 2 * events - 1.0
|
| 201 |
+
dtype = rgbs.dtype
|
| 202 |
+
|
| 203 |
+
# fusion event and image to get fusion feature
|
| 204 |
+
# start = time.time()
|
| 205 |
+
fmaps = self.fusion_block(rgbs.reshape(-1, C_img, H, W), events.reshape(-1, C, H, W), img_ifnew)
|
| 206 |
+
# print(f"fusion block time: {time.time() - start}")
|
| 207 |
+
fmaps = fmaps.permute(0, 2, 3, 1)
|
| 208 |
+
fmaps = fmaps / torch.sqrt(
|
| 209 |
+
torch.maximum(
|
| 210 |
+
torch.sum(torch.square(fmaps), axis=-1, keepdims=True),
|
| 211 |
+
torch.tensor(1e-12, device=fmaps.device),
|
| 212 |
+
)
|
| 213 |
+
)
|
| 214 |
+
fmaps = fmaps.permute(0, 3, 1, 2).reshape(
|
| 215 |
+
B, -1, self.latent_dim, H_stride, W_stride
|
| 216 |
+
)
|
| 217 |
+
fmaps = fmaps.to(dtype)
|
| 218 |
+
|
| 219 |
+
# compute queries point feature
|
| 220 |
+
fmaps_pyramid = []
|
| 221 |
+
track_feat_pyramid = []
|
| 222 |
+
track_feat_support_pyramid = []
|
| 223 |
+
fmaps_pyramid.append(fmaps)
|
| 224 |
+
for i in range(self.corr_levels - 1):
|
| 225 |
+
fmaps_ = fmaps.reshape(B * T, self.latent_dim, fmaps.shape[-2], fmaps.shape[-1])
|
| 226 |
+
fmaps_ = F.avg_pool2d(fmaps_, 2, stride=2)
|
| 227 |
+
fmaps = fmaps_.reshape(B, T, self.latent_dim, fmaps_.shape[-2], fmaps_.shape[-1])
|
| 228 |
+
fmaps_pyramid.append(fmaps)
|
| 229 |
+
|
| 230 |
+
for i in range(self.corr_levels):
|
| 231 |
+
track_feat, track_feat_support = get_track_feat(fmaps_pyramid[i], queried_frames, queried_coords/2**i, support_radius=self.corr_radius)
|
| 232 |
+
track_feat_pyramid.append(track_feat.repeat(1, T, 1, 1))
|
| 233 |
+
track_feat_support_pyramid.append(track_feat_support.unsqueeze(1))
|
| 234 |
+
|
| 235 |
+
vis_init = torch.zeros((B, S, N, 1), device=device).float()
|
| 236 |
+
conf_init = torch.zeros((B, S, N, 1), device=device).float()
|
| 237 |
+
coords_init = queried_coords.reshape(B, 1, N, 2).expand(B, S, N, 2).float()
|
| 238 |
+
|
| 239 |
+
num_windows = (T - S + step - 1) // step + 1
|
| 240 |
+
indices = range(0, step * num_windows, step)
|
| 241 |
+
|
| 242 |
+
for ind in indices:
|
| 243 |
+
if ind > 0:
|
| 244 |
+
overlap = S - step
|
| 245 |
+
copy_over = (queried_frames < ind + overlap)[:, None, :, None] # B, 1, N, 1
|
| 246 |
+
coords_prev = coords_predicted[:, ind : ind + overlap] / self.stride
|
| 247 |
+
padding_tensor = coords_prev[:, -1:, :, :].expand(-1, step, -1, -1) # 将上一个时刻的坐标作为待优化坐标的初始值
|
| 248 |
+
coords_prev = torch.cat([coords_prev, padding_tensor], dim=1)
|
| 249 |
+
|
| 250 |
+
vis_prev = vis_predicted[:, ind : ind + overlap, :, None].clone()
|
| 251 |
+
padding_tensor = vis_prev[:, -1:, :, :].expand(-1, step, -1, -1)
|
| 252 |
+
vis_prev = torch.cat([vis_prev, padding_tensor], dim=1)
|
| 253 |
+
|
| 254 |
+
conf_prev = conf_predicted[:, ind : ind + overlap, :, None].clone()
|
| 255 |
+
padding_tensor = conf_prev[:, -1:, :, :].expand(-1, step, -1, -1)
|
| 256 |
+
conf_prev = torch.cat([conf_prev, padding_tensor], dim=1)
|
| 257 |
+
|
| 258 |
+
coords_init = torch.where(copy_over.expand_as(coords_init), coords_prev, coords_init)
|
| 259 |
+
vis_init = torch.where(copy_over.expand_as(vis_init), vis_prev, vis_init)
|
| 260 |
+
conf_init = torch.where(copy_over.expand_as(conf_init), conf_prev, conf_init)
|
| 261 |
+
|
| 262 |
+
attenstion_mask = (queried_frames < ind + S).reshape(B, 1, N) # B, 1, N
|
| 263 |
+
# start = time.time()
|
| 264 |
+
coords, viss, confs = self.forward_window(
|
| 265 |
+
fmaps_pyramid=[fmap[:, ind : ind +S] for fmap in fmaps_pyramid],
|
| 266 |
+
coords=coords_init,
|
| 267 |
+
track_feat_support_pyramid=[attenstion_mask[:, None, :, :, None]*tfeat for tfeat in track_feat_support_pyramid],
|
| 268 |
+
vis=vis_init,
|
| 269 |
+
conf=conf_init,
|
| 270 |
+
attenstion_mask=attenstion_mask.repeat(1, S, 1),
|
| 271 |
+
iters=iters,
|
| 272 |
+
)
|
| 273 |
+
# print(f"forward window time: {time.time() - start}")
|
| 274 |
+
S_trimmed = min(T - ind, S) # accounts for last window duration
|
| 275 |
+
coords_predicted[:, ind : ind + S] = coords[-1][:, :S_trimmed]
|
| 276 |
+
vis_predicted[:, ind : ind + S] = viss[-1][:, :S_trimmed]
|
| 277 |
+
conf_predicted[:, ind : ind + S] = confs[-1][:, :S_trimmed]
|
| 278 |
+
if is_train:
|
| 279 |
+
all_coords_predictions.append(
|
| 280 |
+
[coord[:, :S_trimmed] for coord in coords]
|
| 281 |
+
)
|
| 282 |
+
all_vis_predictions.append(
|
| 283 |
+
[torch.sigmoid(vis[:, :S_trimmed]) for vis in viss]
|
| 284 |
+
)
|
| 285 |
+
all_confidence_predictions.append(
|
| 286 |
+
[torch.sigmoid(conf[:, :S_trimmed]) for conf in confs]
|
| 287 |
+
)
|
| 288 |
+
|
| 289 |
+
vis_predicted = torch.sigmoid(vis_predicted)
|
| 290 |
+
conf_predicted = torch.sigmoid(conf_predicted)
|
| 291 |
+
|
| 292 |
+
if is_train:
|
| 293 |
+
valid_mask = (
|
| 294 |
+
queried_frames[:, None]
|
| 295 |
+
<= torch.arange(0, T, device=device)[None, :, None]
|
| 296 |
+
)
|
| 297 |
+
train_data = (
|
| 298 |
+
all_coords_predictions,
|
| 299 |
+
all_vis_predictions,
|
| 300 |
+
all_confidence_predictions,
|
| 301 |
+
valid_mask,
|
| 302 |
+
)
|
| 303 |
+
else:
|
| 304 |
+
train_data = None
|
| 305 |
+
|
| 306 |
+
return coords_predicted, vis_predicted, conf_predicted, train_data
|
| 307 |
+
|
| 308 |
+
|
LFE_TAP/utils/__pycache__/dataset_utils.cpython-38.pyc
ADDED
|
Binary file (5.86 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/dataset_utils.cpython-39.pyc
ADDED
|
Binary file (5.84 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/feature_map_vis.cpython-38.pyc
ADDED
|
Binary file (1.01 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/feature_map_vis.cpython-39.pyc
ADDED
|
Binary file (1.01 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/model_utils.cpython-38.pyc
ADDED
|
Binary file (12.5 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/model_utils.cpython-39.pyc
ADDED
|
Binary file (12.4 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/predictor.cpython-39.pyc
ADDED
|
Binary file (6.16 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/train_utils.cpython-39.pyc
ADDED
|
Binary file (3.62 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/visualizer.cpython-38.pyc
ADDED
|
Binary file (7.25 kB). View file
|
|
|
LFE_TAP/utils/__pycache__/visualizer.cpython-39.pyc
ADDED
|
Binary file (7.22 kB). View file
|
|
|
LFE_TAP/utils/dataset_utils.py
ADDED
|
@@ -0,0 +1,115 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import dataclasses
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
from typing import Optional, Any, Union
|
| 6 |
+
|
| 7 |
+
@dataclasses.dataclass(eq=False)
|
| 8 |
+
class FrameEventData:
|
| 9 |
+
""" Data class for frame, event, tracks data. """
|
| 10 |
+
video: torch.Tensor # (B, S, C_i, H, W)
|
| 11 |
+
events: torch.Tensor # (B, S, C_e, H, W)
|
| 12 |
+
segmentation: torch.Tensor # (B, S, 1, H, W)
|
| 13 |
+
trajectory: torch.Tensor # (B, S, N, 2)
|
| 14 |
+
visibility: torch.Tensor # (B, S, N)
|
| 15 |
+
img_ifnew: torch.Tensor = None
|
| 16 |
+
# optional daa
|
| 17 |
+
clear_video: Optional[torch.Tensor] = None # (B, S, C_i, H, W)
|
| 18 |
+
valid: Optional[torch.Tensor] = None # (B, S, N)
|
| 19 |
+
seq_name: Optional[torch.Tensor] = None
|
| 20 |
+
|
| 21 |
+
|
| 22 |
+
@dataclasses.dataclass(eq=False)
|
| 23 |
+
class FrameEventData_test:
|
| 24 |
+
video: np.array
|
| 25 |
+
events: Union[np.array, list]
|
| 26 |
+
segmentation: np.array
|
| 27 |
+
trajectory: np.array
|
| 28 |
+
query_points: torch.Tensor
|
| 29 |
+
# optional daa
|
| 30 |
+
visibility: Optional[torch.Tensor] = None
|
| 31 |
+
valid: Optional[torch.Tensor] = None # (B, S, N)
|
| 32 |
+
seq_name: Optional[torch.Tensor] = None
|
| 33 |
+
img_ifnew: Optional[np.array] = None
|
| 34 |
+
img_ifnew_full: Optional[np.array] = None
|
| 35 |
+
rgb_timestamp: Optional[np.array] = None
|
| 36 |
+
|
| 37 |
+
|
| 38 |
+
def collate_fn(batch):
|
| 39 |
+
""" Collate function for frame, event, tracks data. """
|
| 40 |
+
video = torch.stack([b.video for b, in batch], dim=0)
|
| 41 |
+
events = torch.stack([b.events for b in batch], dim=0)
|
| 42 |
+
segmentation = torch.stack([b.segmentation for b in batch], dim=0)
|
| 43 |
+
trajectory = torch.stack([b.trajectory for b in batch], dim=0)
|
| 44 |
+
visibility = torch.stack([b.visibility for b in batch], dim=0)
|
| 45 |
+
|
| 46 |
+
seq_name = [b.seq_name for b in batch]
|
| 47 |
+
|
| 48 |
+
return FrameEventData(video, events, segmentation, trajectory, visibility, seq_name=seq_name)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def collate_fn_train(batch):
|
| 52 |
+
""" Collate function for training data. """
|
| 53 |
+
gotit = [gotit for _, gotit in batch]
|
| 54 |
+
video = torch.stack([b.video for b, _ in batch], dim=0)
|
| 55 |
+
events = torch.stack([b.events for b, _ in batch], dim=0)
|
| 56 |
+
segmentation = torch.stack([b.segmentation for b, _ in batch], dim=0)
|
| 57 |
+
trajectory = torch.stack([b.trajectory for b, _ in batch], dim=0)
|
| 58 |
+
visibility = torch.stack([b.visibility for b, _ in batch], dim=0)
|
| 59 |
+
clear_video = torch.stack([b.clear_video for b, _ in batch], dim=0)
|
| 60 |
+
valid = torch.stack([b.valid for b, _ in batch], dim=0)
|
| 61 |
+
seq_name = [b.seq_name for b, _ in batch]
|
| 62 |
+
for b, _ in batch:
|
| 63 |
+
if b.img_ifnew is None:
|
| 64 |
+
gotit = [False]
|
| 65 |
+
img_ifnew = []
|
| 66 |
+
else:
|
| 67 |
+
img_ifnew = torch.stack([b.img_ifnew for b, _ in batch], dim=0)
|
| 68 |
+
return (FrameEventData(video, events, segmentation, trajectory, visibility, clear_video=clear_video, valid=valid, seq_name=seq_name, img_ifnew=img_ifnew), gotit)
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
def collate_fn_EDS(batch):
|
| 72 |
+
video = np.stack([b.video for b, _ in batch], axis=0)
|
| 73 |
+
events = np.stack([b.events for b, _ in batch], axis=0)
|
| 74 |
+
segmentation = torch.stack([b.segmentation for b, _ in batch], axis=0)
|
| 75 |
+
trajectory = np.stack([b.trajectory for b, _ in batch], axis=0)
|
| 76 |
+
query_points = None
|
| 77 |
+
if batch[0][0].query_points is not None:
|
| 78 |
+
query_points = torch.stack([b.query_points for b, _ in batch], dim=0)
|
| 79 |
+
|
| 80 |
+
seq_name = [b.seq_name for b, _ in batch]
|
| 81 |
+
rgb_timestamp = [b.rgb_timestamp for b, _ in batch]
|
| 82 |
+
|
| 83 |
+
return FrameEventData_test(video, events, segmentation, trajectory, query_points, seq_name=seq_name, rgb_timestamp=rgb_timestamp)
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def try_to_cuda(t: Any) -> Any:
|
| 87 |
+
"""
|
| 88 |
+
Try to move the input variable `t` to a cuda device.
|
| 89 |
+
|
| 90 |
+
Args:
|
| 91 |
+
t: Input.
|
| 92 |
+
|
| 93 |
+
Returns:
|
| 94 |
+
t_cuda: `t` moved to a cuda device, if supported.
|
| 95 |
+
"""
|
| 96 |
+
try:
|
| 97 |
+
t = t.float().cuda()
|
| 98 |
+
except AttributeError:
|
| 99 |
+
pass
|
| 100 |
+
return t
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
def dataclass_to_cuda_(obj):
|
| 104 |
+
"""
|
| 105 |
+
Move all contents of a dataclass to cuda inplace if supported.
|
| 106 |
+
|
| 107 |
+
Args:
|
| 108 |
+
batch: Input dataclass.
|
| 109 |
+
|
| 110 |
+
Returns:
|
| 111 |
+
batch_cuda: `batch` moved to a cuda device, if supported.
|
| 112 |
+
"""
|
| 113 |
+
for f in dataclasses.fields(obj):
|
| 114 |
+
setattr(obj, f.name, try_to_cuda(getattr(obj, f.name)))
|
| 115 |
+
return obj
|
LFE_TAP/utils/event/__pycache__/representations.cpython-39.pyc
ADDED
|
Binary file (8.58 kB). View file
|
|
|
LFE_TAP/utils/event/__pycache__/utils.cpython-39.pyc
ADDED
|
Binary file (7.44 kB). View file
|
|
|
LFE_TAP/utils/event/representations.py
ADDED
|
@@ -0,0 +1,312 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import numpy as np
|
| 3 |
+
import cv2
|
| 4 |
+
from enum import Enum, auto
|
| 5 |
+
|
| 6 |
+
import hdf5plugin
|
| 7 |
+
import h5py
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class EventRepresentationTypes(Enum):
|
| 11 |
+
time_surface = 0
|
| 12 |
+
voxel_grid = 1
|
| 13 |
+
event_stack = 2
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
class EventRepresentation:
|
| 17 |
+
def __init__(self):
|
| 18 |
+
pass
|
| 19 |
+
|
| 20 |
+
def convert(self, events):
|
| 21 |
+
raise NotImplementedError
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
class TimeSurface(EventRepresentation):
|
| 25 |
+
def __init__(self, input_size: tuple, p, t, x, y):
|
| 26 |
+
assert len(input_size) == 3
|
| 27 |
+
t = t.astype('float32')
|
| 28 |
+
x = x.astype('float32')
|
| 29 |
+
y = y.astype('float32')
|
| 30 |
+
pol = p.astype('float32')
|
| 31 |
+
self.input_size = input_size
|
| 32 |
+
self.time_surface = torch.zeros(input_size, dtype=torch.float, requires_grad=False)
|
| 33 |
+
self.n_bins = input_size[0] // 2
|
| 34 |
+
|
| 35 |
+
def convert(self, events):
|
| 36 |
+
_, H, W = self.time_surface.shape
|
| 37 |
+
with torch.no_grad():
|
| 38 |
+
self.time_surface = torch.zeros(self.input_size, dtype=torch.float, requires_grad=False,
|
| 39 |
+
device=events['p'].device)
|
| 40 |
+
time_surface = self.time_surface.clone()
|
| 41 |
+
|
| 42 |
+
t = events['t'].cpu().numpy()
|
| 43 |
+
dt_bin = 1. / self.n_bins
|
| 44 |
+
x0 = events['x'].int()
|
| 45 |
+
y0 = events['y'].int()
|
| 46 |
+
p0 = events['p'].int()
|
| 47 |
+
t0 = events['t']
|
| 48 |
+
|
| 49 |
+
# iterate over bins
|
| 50 |
+
for i_bin in range(self.n_bins):
|
| 51 |
+
t0_bin = i_bin * dt_bin
|
| 52 |
+
t1_bin = t0_bin + dt_bin
|
| 53 |
+
|
| 54 |
+
# mask_t = np.logical_and(time > t0_bin, time <= t1_bin)
|
| 55 |
+
# x_bin, y_bin, p_bin, t_bin = x[mask_t], y[mask_t], p[mask_t], time[mask_t]
|
| 56 |
+
idx0 = np.searchsorted(t, t0_bin, side='left')
|
| 57 |
+
idx1 = np.searchsorted(t, t1_bin, side='right')
|
| 58 |
+
x_bin = x0[idx0:idx1]
|
| 59 |
+
y_bin = y0[idx0:idx1]
|
| 60 |
+
p_bin = p0[idx0:idx1]
|
| 61 |
+
t_bin = t0[idx0:idx1]
|
| 62 |
+
|
| 63 |
+
n_events = len(x_bin)
|
| 64 |
+
for i in range(n_events):
|
| 65 |
+
if 0 <= x_bin[i] < W and 0 <= y_bin[i] < H:
|
| 66 |
+
time_surface[2*i_bin+p_bin[i], y_bin[i], x_bin[i]] = t_bin[i]
|
| 67 |
+
|
| 68 |
+
return time_surface
|
| 69 |
+
|
| 70 |
+
|
| 71 |
+
class TimeOrderSurface():
|
| 72 |
+
def __init__(self, input_size: tuple, x, y, p, t):
|
| 73 |
+
assert len(input_size) == 3
|
| 74 |
+
H, W, C = input_size
|
| 75 |
+
self.t = torch.from_numpy(t.astype('int32'))
|
| 76 |
+
self.x = torch.from_numpy(x.astype('int32'))
|
| 77 |
+
self.y = torch.from_numpy(y.astype('int32'))
|
| 78 |
+
self.pol = torch.from_numpy(p.astype('int32'))
|
| 79 |
+
|
| 80 |
+
mask = (self.x >= 3) & (self.x < W - 3) & (self.y >= 3) & (self.y < H - 3) & (self.t >= 370000) & (self.t <= 900000)
|
| 81 |
+
self.x = self.x[mask]
|
| 82 |
+
self.y = self.y[mask]
|
| 83 |
+
self.pol = self.pol[mask]
|
| 84 |
+
self.t = self.t[mask]
|
| 85 |
+
|
| 86 |
+
self.index = torch.tensor(0 , device=self.t.device)
|
| 87 |
+
self.input_size = input_size
|
| 88 |
+
self.n_bins = input_size[2] // 2
|
| 89 |
+
self.tos = torch.zeros((input_size[0], input_size[1], 2), dtype=torch.float, requires_grad=False, device=self.t.device)
|
| 90 |
+
self.sae = torch.zeros((input_size[0], input_size[1], 2), dtype=torch.float, requires_grad=False, device=self.t.device)
|
| 91 |
+
self.sae_latest = torch.zeros((input_size[0], input_size[1], 2), dtype=torch.float, requires_grad=False, device=self.t.device)
|
| 92 |
+
self.TOS_bins_001 = torch.zeros(self.input_size, dtype=torch.float, requires_grad=False,
|
| 93 |
+
device=self.t.device)
|
| 94 |
+
self.TOS_bins_002 = torch.zeros(self.input_size, dtype=torch.float, requires_grad=False,
|
| 95 |
+
device=self.t.device)
|
| 96 |
+
|
| 97 |
+
def convert(self, time, n_bins, type):
|
| 98 |
+
assert n_bins < self.n_bins
|
| 99 |
+
with torch.no_grad():
|
| 100 |
+
# TOS_bins_001 = self.TOS_bins_001.clone()
|
| 101 |
+
# TOS_bins_002 = self.TOS_bins_002.clone()
|
| 102 |
+
|
| 103 |
+
while self.index < len(self.t) and self.t[self.index] <= time:
|
| 104 |
+
pol = 1 if self.pol[self.index] else 0
|
| 105 |
+
pol_inv = 0 if self.pol[self.index] else 1
|
| 106 |
+
if ((self.t[self.index] > self.sae_latest[self.y[self.index]][self.x[self.index]][pol] + 20000) or
|
| 107 |
+
(self.sae_latest[self.y[self.index]][self.x[self.index]][pol_inv] > self.sae_latest[self.y[self.index]][self.x[self.index]][pol])):
|
| 108 |
+
self.sae_latest[self.y[self.index]][self.x[self.index]][pol] = self.t[self.index]
|
| 109 |
+
self.sae[self.y[self.index]][self.x[self.index]][pol] = self.t[self.index]
|
| 110 |
+
self.tos[self.y[self.index] - 3:self.y[self.index] + 4,
|
| 111 |
+
self.x[self.index] - 3:self.x[self.index] + 4, self.pol[self.index]] -= 1
|
| 112 |
+
self.tos[self.y[self.index] - 3:self.y[self.index] + 4, self.x[self.index] - 3:self.x[self.index] + 4, self.pol[self.index]][
|
| 113 |
+
self.tos[self.y[self.index] - 3:self.y[self.index] + 4, self.x[self.index] - 3:self.x[self.index] + 4, self.pol[self.index]] < 241] = 0
|
| 114 |
+
self.tos[self.y[self.index], self.x[self.index], self.pol[self.index]] = 255
|
| 115 |
+
else:
|
| 116 |
+
self.sae_latest[self.y[self.index]][self.x[self.index]][pol] = self.t[self.index]
|
| 117 |
+
|
| 118 |
+
self.index += 1
|
| 119 |
+
|
| 120 |
+
# while self.t[self.index] <= 900000 and self.t[self.index] <= time:
|
| 121 |
+
# for x0 in range(self.x[self.index]-3, self.x[self.index]+4):
|
| 122 |
+
# for y0 in range(self.y[self.index]-3, self.y[self.index]+4):
|
| 123 |
+
# if self.tos[y0, x0, self.pol[self.index]] != 0:
|
| 124 |
+
# self.tos[y0, x0, self.pol[self.index]] -= 1
|
| 125 |
+
# if self.tos[y0, x0, self.pol[self.index]] < 241:
|
| 126 |
+
# self.tos[y0, x0, self.pol[self.index]] = 0
|
| 127 |
+
# # self.tos[y0, x0, self.pol[self.index]].clamp_(min=0, max=240) # 使用clamp函数限制值的范围
|
| 128 |
+
# self.tos[self.y[self.index], self.x[self.index], self.pol[self.index]] = 255
|
| 129 |
+
# self.index += 1
|
| 130 |
+
|
| 131 |
+
# time_threshold = 900000
|
| 132 |
+
#
|
| 133 |
+
# # 使用torch.where函数实现条件操作
|
| 134 |
+
# while torch.any(torch.logical_and(self.t[self.index] <= time_threshold, self.t[self.index] <= time)):
|
| 135 |
+
# y_slice = slice(self.y[self.index]-3, self.y[self.index]+4)
|
| 136 |
+
# x_slice = slice(self.x[self.index]-3, self.x[self.index]+4)
|
| 137 |
+
# mask = (self.tos[y_slice, x_slice, self.pol[self.index]] != 0)
|
| 138 |
+
# self.tos[y_slice, x_slice, self.pol[self.index]].masked_scatter_(mask, self.tos[y_slice, x_slice, self.pol[self.index]] - 1)
|
| 139 |
+
# self.tos[y_slice, x_slice, self.pol[self.index]] = torch.where(self.tos[y_slice, x_slice, self.pol[self.index]] < 241,
|
| 140 |
+
# torch.tensor(0), self.tos[y_slice, x_slice, self.pol[self.index]])
|
| 141 |
+
# self.tos[self.y[self.index], self.x[self.index], self.pol[self.index]] = 255
|
| 142 |
+
# self.index += 1
|
| 143 |
+
|
| 144 |
+
if type == 0:
|
| 145 |
+
self.TOS_bins_001[:, :, 2*n_bins] = self.tos[:, :, 0]
|
| 146 |
+
self.TOS_bins_001[:, :, 2*n_bins+1] = self.tos[:, :, 1]
|
| 147 |
+
elif type == 1:
|
| 148 |
+
self.TOS_bins_002[:, :, 2*n_bins] = self.tos[:, :, 0]
|
| 149 |
+
self.TOS_bins_002[:, :, 2*n_bins+1] = self.tos[:, :, 1]
|
| 150 |
+
|
| 151 |
+
def get_time_order_surface(self, type):
|
| 152 |
+
# tos = self.tos.numpy()
|
| 153 |
+
# cv2.imshow("p_tos", tos[:, :, 0])
|
| 154 |
+
# cv2.imshow("n_tos", tos[:, :, 1])
|
| 155 |
+
# cv2.waitKey(0)
|
| 156 |
+
if type == 0:
|
| 157 |
+
return self.TOS_bins_001
|
| 158 |
+
elif type == 1:
|
| 159 |
+
return self.TOS_bins_002
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
class VoxelGrid(EventRepresentation):
|
| 163 |
+
def __init__(self, input_size: tuple, normalize: bool):
|
| 164 |
+
assert len(input_size) == 3
|
| 165 |
+
self.voxel_grid = torch.zeros((input_size), dtype=torch.float, requires_grad=False)
|
| 166 |
+
self.input_size = input_size
|
| 167 |
+
self.nb_channels = input_size[0]
|
| 168 |
+
self.normalize = normalize
|
| 169 |
+
|
| 170 |
+
def convert(self, events):
|
| 171 |
+
C, H, W = self.voxel_grid.shape
|
| 172 |
+
with torch.no_grad():
|
| 173 |
+
self.voxel_grid = torch.zeros((self.input_size), dtype=torch.float, requires_grad=False,
|
| 174 |
+
device=events['p'].device)
|
| 175 |
+
voxel_grid = self.voxel_grid.clone()
|
| 176 |
+
|
| 177 |
+
t_norm = events['t']
|
| 178 |
+
t_norm = (C - 1) * (t_norm-t_norm[0]) / (t_norm[-1]-t_norm[0])
|
| 179 |
+
|
| 180 |
+
x0 = events['x'].int()
|
| 181 |
+
y0 = events['y'].int()
|
| 182 |
+
t0 = t_norm.int()
|
| 183 |
+
|
| 184 |
+
value = 2*events['p']-1
|
| 185 |
+
|
| 186 |
+
for xlim in [x0,x0+1]:
|
| 187 |
+
for ylim in [y0,y0+1]:
|
| 188 |
+
for tlim in [t0,t0+1]:
|
| 189 |
+
|
| 190 |
+
mask = (xlim < W) & (xlim >= 0) & (ylim < H) & (ylim >= 0) & (tlim >= 0) & (tlim < self.nb_channels)
|
| 191 |
+
interp_weights = value * (1 - (xlim-events['x']).abs()) * (1 - (ylim-events['y']).abs()) * (1 - (tlim - t_norm).abs())
|
| 192 |
+
|
| 193 |
+
index = H * W * tlim.long() + \
|
| 194 |
+
W * ylim.long() + \
|
| 195 |
+
xlim.long()
|
| 196 |
+
|
| 197 |
+
voxel_grid.put_(index[mask], interp_weights[mask], accumulate=True)
|
| 198 |
+
|
| 199 |
+
if self.normalize:
|
| 200 |
+
mask = torch.nonzero(voxel_grid, as_tuple=True)
|
| 201 |
+
if mask[0].size()[0] > 0:
|
| 202 |
+
mean = voxel_grid[mask].mean()
|
| 203 |
+
std = voxel_grid[mask].std()
|
| 204 |
+
if std > 0:
|
| 205 |
+
voxel_grid[mask] = (voxel_grid[mask] - mean) / std
|
| 206 |
+
else:
|
| 207 |
+
voxel_grid[mask] = voxel_grid[mask] - mean
|
| 208 |
+
|
| 209 |
+
return voxel_grid
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
class EventStack(EventRepresentation):
|
| 213 |
+
def __init__(self, input_size: tuple):
|
| 214 |
+
"""
|
| 215 |
+
:param input_size: (C, H, W)
|
| 216 |
+
"""
|
| 217 |
+
assert len(input_size) == 3
|
| 218 |
+
self.input_size = input_size
|
| 219 |
+
self.event_stack = torch.zeros((input_size), dtype=torch.float, requires_grad=False)
|
| 220 |
+
self.nb_channels = input_size[0]
|
| 221 |
+
|
| 222 |
+
def convert(self, events):
|
| 223 |
+
C, H, W = self.event_stack.shape
|
| 224 |
+
with torch.no_grad():
|
| 225 |
+
self.event_stack = torch.zeros((self.input_size), dtype=torch.float, requires_grad=False,
|
| 226 |
+
device=events['p'].device)
|
| 227 |
+
event_stack = self.event_stack.clone()
|
| 228 |
+
|
| 229 |
+
t = events['t'].cpu().numpy()
|
| 230 |
+
dt_bin = 1. / self.nb_channels
|
| 231 |
+
x0 = events['x'].int()
|
| 232 |
+
y0 = events['y'].int()
|
| 233 |
+
p0 = 2*events['p'].int()-1
|
| 234 |
+
t0 = events['t']
|
| 235 |
+
|
| 236 |
+
# iterate over bins
|
| 237 |
+
for i_bin in range(self.nb_channels):
|
| 238 |
+
t0_bin = i_bin * dt_bin
|
| 239 |
+
t1_bin = t0_bin + dt_bin
|
| 240 |
+
|
| 241 |
+
# mask_t = np.logical_and(time > t0_bin, time <= t1_bin)
|
| 242 |
+
# x_bin, y_bin, p_bin, t_bin = x[mask_t], y[mask_t], p[mask_t], time[mask_t]
|
| 243 |
+
idx0 = np.searchsorted(t, t0_bin, side='left')
|
| 244 |
+
idx1 = np.searchsorted(t, t1_bin, side='right')
|
| 245 |
+
x_bin = x0[idx0:idx1]
|
| 246 |
+
y_bin = y0[idx0:idx1]
|
| 247 |
+
p_bin = p0[idx0:idx1]
|
| 248 |
+
|
| 249 |
+
n_events = len(x_bin)
|
| 250 |
+
for i in range(n_events):
|
| 251 |
+
if 0 <= x_bin[i] < W and 0 <= y_bin[i] < H:
|
| 252 |
+
event_stack[i_bin, y_bin[i], x_bin[i]] += p_bin[i]
|
| 253 |
+
|
| 254 |
+
return event_stack
|
| 255 |
+
|
| 256 |
+
|
| 257 |
+
def events_to_time_surface(time_surface, p, t, x, y):
|
| 258 |
+
t = (t - t[0]).astype('float32')
|
| 259 |
+
t = (t/t[-1])
|
| 260 |
+
x = x.astype('float32')
|
| 261 |
+
y = y.astype('float32')
|
| 262 |
+
pol = p.astype('float32')
|
| 263 |
+
event_data_torch = {
|
| 264 |
+
'p': torch.from_numpy(pol),
|
| 265 |
+
't': torch.from_numpy(t),
|
| 266 |
+
'x': torch.from_numpy(x),
|
| 267 |
+
'y': torch.from_numpy(y),
|
| 268 |
+
}
|
| 269 |
+
return time_surface.convert(event_data_torch)
|
| 270 |
+
|
| 271 |
+
def events_to_time_order_surface(time_order_surface, p, t, x, y, dt):
|
| 272 |
+
t = t.astype('float32')
|
| 273 |
+
x = x.astype('float32')
|
| 274 |
+
y = y.astype('float32')
|
| 275 |
+
pol = p.astype('float32')
|
| 276 |
+
event_data_torch = {
|
| 277 |
+
'p': torch.from_numpy(pol),
|
| 278 |
+
't': torch.from_numpy(t),
|
| 279 |
+
'x': torch.from_numpy(x),
|
| 280 |
+
'y': torch.from_numpy(y),
|
| 281 |
+
}
|
| 282 |
+
|
| 283 |
+
time_order_surface.convert(event_data_torch, dt)
|
| 284 |
+
|
| 285 |
+
def events_to_event_stack(event_stack, p, t, x, y, dt):
|
| 286 |
+
t = (t - t[0]).astype('float32')
|
| 287 |
+
t = (t/t[-1])
|
| 288 |
+
x = x.astype('float32')
|
| 289 |
+
y = y.astype('float32')
|
| 290 |
+
pol = p.astype('float32')
|
| 291 |
+
event_data_torch = {
|
| 292 |
+
'p': torch.from_numpy(pol),
|
| 293 |
+
't': torch.from_numpy(t),
|
| 294 |
+
'x': torch.from_numpy(x),
|
| 295 |
+
'y': torch.from_numpy(y),
|
| 296 |
+
}
|
| 297 |
+
return event_stack.convert(event_data_torch)
|
| 298 |
+
|
| 299 |
+
|
| 300 |
+
def events_to_voxel_grid(voxel_grid, p, t, x, y):
|
| 301 |
+
t = (t - t[0]).astype('float32')
|
| 302 |
+
t = (t/t[-1])
|
| 303 |
+
x = x.astype('float32')
|
| 304 |
+
y = y.astype('float32')
|
| 305 |
+
pol = p.astype('float32')
|
| 306 |
+
event_data_torch = {
|
| 307 |
+
'p': torch.from_numpy(pol),
|
| 308 |
+
't': torch.from_numpy(t),
|
| 309 |
+
'x': torch.from_numpy(x),
|
| 310 |
+
'y': torch.from_numpy(y),
|
| 311 |
+
}
|
| 312 |
+
return voxel_grid.convert(event_data_torch)
|