ljx1002 commited on
Commit
315ffb3
·
verified ·
1 Parent(s): a77fb8f

Upload 96 files

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. .gitattributes +17 -0
  2. LFE_TAP/datasets/EC_dataset.py +111 -0
  3. LFE_TAP/datasets/EDS_dataset.py +115 -0
  4. LFE_TAP/datasets/TAPFormer_dataset.py +151 -0
  5. LFE_TAP/datasets/__pycache__/Aedat4_dataset.cpython-39.pyc +0 -0
  6. LFE_TAP/datasets/__pycache__/EC_dataset.cpython-39.pyc +0 -0
  7. LFE_TAP/datasets/__pycache__/EDS_dataset.cpython-39.pyc +0 -0
  8. LFE_TAP/datasets/__pycache__/MF_dataset.cpython-39.pyc +0 -0
  9. LFE_TAP/datasets/__pycache__/kubric_movif_dataset.cpython-39.pyc +0 -0
  10. LFE_TAP/datasets/__pycache__/prophesee_dataset.cpython-39.pyc +0 -0
  11. LFE_TAP/datasets/kubric_movif_dataset.py +778 -0
  12. LFE_TAP/evaluator/__pycache__/evaluation_pred.cpython-38.pyc +0 -0
  13. LFE_TAP/evaluator/__pycache__/evaluation_pred.cpython-39.pyc +0 -0
  14. LFE_TAP/evaluator/__pycache__/evaluator.cpython-38.pyc +0 -0
  15. LFE_TAP/evaluator/__pycache__/evaluator.cpython-39.pyc +0 -0
  16. LFE_TAP/evaluator/__pycache__/prediction_long.cpython-38.pyc +0 -0
  17. LFE_TAP/evaluator/__pycache__/prediction_long.cpython-39.pyc +0 -0
  18. LFE_TAP/evaluator/evaluation_pred.py +184 -0
  19. LFE_TAP/evaluator/evaluator.py +351 -0
  20. LFE_TAP/evaluator/prediction.py +311 -0
  21. LFE_TAP/models/__pycache__/blocks.cpython-38.pyc +0 -0
  22. LFE_TAP/models/__pycache__/blocks.cpython-39.pyc +0 -0
  23. LFE_TAP/models/__pycache__/embeddings.cpython-38.pyc +0 -0
  24. LFE_TAP/models/__pycache__/embeddings.cpython-39.pyc +0 -0
  25. LFE_TAP/models/__pycache__/etap.cpython-39.pyc +0 -0
  26. LFE_TAP/models/__pycache__/fusionFormer.cpython-38.pyc +0 -0
  27. LFE_TAP/models/__pycache__/fusionFormer.cpython-39.pyc +0 -0
  28. LFE_TAP/models/__pycache__/hivit.cpython-38.pyc +0 -0
  29. LFE_TAP/models/__pycache__/hivit.cpython-39.pyc +0 -0
  30. LFE_TAP/models/__pycache__/losses.cpython-39.pyc +0 -0
  31. LFE_TAP/models/__pycache__/tapfe.cpython-38.pyc +0 -0
  32. LFE_TAP/models/__pycache__/tapfe.cpython-39.pyc +0 -0
  33. LFE_TAP/models/blocks.py +994 -0
  34. LFE_TAP/models/embeddings.py +110 -0
  35. LFE_TAP/models/fusionFormer.py +253 -0
  36. LFE_TAP/models/tapformer.py +308 -0
  37. LFE_TAP/utils/__pycache__/dataset_utils.cpython-38.pyc +0 -0
  38. LFE_TAP/utils/__pycache__/dataset_utils.cpython-39.pyc +0 -0
  39. LFE_TAP/utils/__pycache__/feature_map_vis.cpython-38.pyc +0 -0
  40. LFE_TAP/utils/__pycache__/feature_map_vis.cpython-39.pyc +0 -0
  41. LFE_TAP/utils/__pycache__/model_utils.cpython-38.pyc +0 -0
  42. LFE_TAP/utils/__pycache__/model_utils.cpython-39.pyc +0 -0
  43. LFE_TAP/utils/__pycache__/predictor.cpython-39.pyc +0 -0
  44. LFE_TAP/utils/__pycache__/train_utils.cpython-39.pyc +0 -0
  45. LFE_TAP/utils/__pycache__/visualizer.cpython-38.pyc +0 -0
  46. LFE_TAP/utils/__pycache__/visualizer.cpython-39.pyc +0 -0
  47. LFE_TAP/utils/dataset_utils.py +115 -0
  48. LFE_TAP/utils/event/__pycache__/representations.cpython-39.pyc +0 -0
  49. LFE_TAP/utils/event/__pycache__/utils.cpython-39.pyc +0 -0
  50. 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)