plice13 commited on
Commit
efaa107
·
verified ·
1 Parent(s): 988f575

final demo #2

Browse files
.gitattributes CHANGED
@@ -33,3 +33,6 @@ 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
+ checkpoints/pose/face_landmarker.task filter=lfs diff=lfs merge=lfs -text
37
+ checkpoints/pose/hand_landmarker.task filter=lfs diff=lfs merge=lfs -text
38
+ checkpoints/pose/pose_landmarker_full.task filter=lfs diff=lfs merge=lfs -text
Uni_Sign/__pycache__/datasets.cpython-311.pyc ADDED
Binary file (54.1 kB). View file
 
Uni_Sign/__pycache__/deformable_attention_2d.cpython-311.pyc ADDED
Binary file (18.2 kB). View file
 
Uni_Sign/__pycache__/models.cpython-311.pyc ADDED
Binary file (21.3 kB). View file
 
Uni_Sign/__pycache__/normalization.cpython-311.pyc ADDED
Binary file (7 kB). View file
 
Uni_Sign/datasets.py ADDED
@@ -0,0 +1,1178 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import torch.utils.data.dataset as Dataset
3
+ from torch.nn.utils.rnn import pad_sequence
4
+ from PIL import Image
5
+ import os
6
+ import random
7
+ import numpy as np
8
+ import copy
9
+ import pickle
10
+ from decord import VideoReader, cpu
11
+ import json
12
+ import pathlib
13
+ import re
14
+ from torchvision import transforms
15
+ # from config import rgb_dirs, pose_dirs
16
+ from Uni_Sign.normalization import (local_keypoint_normalization, global_keypoint_normalization)
17
+
18
+
19
+ def all_same(keypoints):
20
+ return np.sum(keypoints == keypoints[0, 0]) == keypoints.size
21
+
22
+
23
+ def sign_space_normalization(raw_keypoints, missing_values=None, layout='default'):
24
+ local_landmarks = {}
25
+ global_landmarks = {}
26
+ kp_normalization = ('global-body', 'local-right', 'local-left', 'local-face_all')
27
+ part_order = [i.removeprefix('local-').removeprefix('global-') for i in kp_normalization]
28
+ part_order = {k: v for v, k in enumerate(part_order)}
29
+
30
+ for idx, landmarks in enumerate(kp_normalization):
31
+ prefix, landmarks = landmarks.split("-")
32
+ if prefix == "local":
33
+ local_landmarks[idx] = landmarks
34
+ elif prefix == "global":
35
+ global_landmarks[idx] = landmarks
36
+
37
+ # local normalization
38
+ for idx, landmarks in local_landmarks.items():
39
+ normalized_keypoints = local_keypoint_normalization(raw_keypoints, landmarks, padding=0.2)
40
+ local_landmarks[idx] = normalized_keypoints
41
+
42
+ # global normalization
43
+ additional_landmarks = list(global_landmarks.values())
44
+ if "body" in additional_landmarks:
45
+ additional_landmarks.remove("body")
46
+
47
+ if layout == 'default':
48
+ l_shoulder_idx, r_shoulder_idx = 11, 12
49
+ else:
50
+ l_shoulder_idx, r_shoulder_idx = 3, 4
51
+ keypoints, additional_keypoints = global_keypoint_normalization(
52
+ raw_keypoints,
53
+ "body",
54
+ additional_landmarks,
55
+ l_shoulder_idx=l_shoulder_idx,
56
+ r_shoulder_idx=r_shoulder_idx,
57
+ )
58
+
59
+ for k, landmark in global_landmarks.items():
60
+ if landmark == "body":
61
+ global_landmarks[k] = keypoints
62
+ else:
63
+ global_landmarks[k] = additional_keypoints[landmark]
64
+
65
+ all_landmarks = {**local_landmarks, **global_landmarks}
66
+ all_landmarks_per_part = {k: all_landmarks[v] for k, v in part_order.items()}
67
+
68
+ if missing_values is not None:
69
+ for part, data in all_landmarks_per_part.items():
70
+ for fidx in range(len(data)):
71
+ if not all_same(data[fidx]):
72
+ continue
73
+ all_landmarks_per_part[part][fidx] = np.zeros_like(data[fidx]) + missing_values
74
+
75
+ return all_landmarks_per_part
76
+
77
+
78
+ def load_part_kp_YTASL(skeletons, confs, normalization, layout):
79
+ thr = 0.3
80
+ # kps_with_scores = {}
81
+ kps_all_parts = {}
82
+ confs_all_parts = {}
83
+ scale = None
84
+
85
+ for part in ['body', 'left', 'right', 'face_all']:
86
+ kps = []
87
+ confidences = []
88
+ for i, (skeleton, conf) in enumerate(zip(skeletons, confs)):
89
+
90
+ if part == 'body':
91
+ if layout == 'default':
92
+ hand_kp2d = np.stack(skeleton['pose_landmarks'][:25])
93
+ confidence = np.stack(conf['pose_landmarks'][:25])
94
+ elif layout == 'pruned':
95
+ pose_landmarks = [0, 7, 8, 11, 12, 13, 14, 15, 16]
96
+ hand_kp2d = np.stack([skeleton['pose_landmarks'][i] for i in pose_landmarks])
97
+ confidence = np.stack([conf['pose_landmarks'][i] for i in pose_landmarks])
98
+ elif layout == 'isharah':
99
+ pose_landmarks = [0, 7, 8, 11, 12, 13, 14, 15, 16]
100
+ hand_kp2d = np.stack([skeleton['pose_landmarks'][i] for i in pose_landmarks])
101
+ confidence = np.stack([conf['pose_landmarks'][i] for i in pose_landmarks])
102
+
103
+ elif part == 'left':
104
+ if layout in ['default', 'pruned', 'isharah']:
105
+ hand_kp2d = np.stack(skeleton['left_hand_landmarks'])
106
+ confidence = np.stack(conf['left_hand_landmarks'])
107
+
108
+ elif part == 'right':
109
+ if layout in ['default', 'pruned', 'isharah']:
110
+ hand_kp2d = np.stack(skeleton['right_hand_landmarks'])
111
+ confidence = np.stack(conf['right_hand_landmarks'])
112
+
113
+ elif part == 'face_all':
114
+ if layout == 'default':
115
+ face_landmarks = [
116
+ 0, 4, 13, 14, 17, 33, 39, 46, 52, 55, 61, 64, 81,
117
+ 93, 133, 151, 152, 159, 172, 178, 181, 263, 269, 276,
118
+ 282, 285, 291, 294, 311, 323, 362, 386, 397, 402, 405, 468, 473
119
+ ]
120
+ hand_kp2d = np.stack([skeleton['face_landmarks'][i] for i in face_landmarks])
121
+ confidence = np.stack([conf['face_landmarks'][i] for i in face_landmarks])
122
+ elif layout == 'pruned':
123
+ face_landmarks = [4, 13, 14, 61, 81, 93, 152, 159, 172, 178, 291, 311, 323, 386, 397, 402, 472, 477]
124
+ hand_kp2d = np.stack([skeleton['face_landmarks'][i] for i in face_landmarks])
125
+ confidence = np.stack([conf['face_landmarks'][i] for i in face_landmarks])
126
+ elif layout == 'isharah':
127
+ face_landmarks = [0, 17, 37, 39, 40, 61, 84, 91, 146, 181, 185, 267, 269, 270, 291, 314, 321, 375, 405]
128
+ hand_kp2d = np.stack([skeleton['face_landmarks'][i] for i in face_landmarks])
129
+ confidence = np.stack([conf['face_landmarks'][i] for i in face_landmarks])
130
+
131
+ else:
132
+ raise NotImplementedError
133
+ kps.append(hand_kp2d)
134
+ confidences.append(confidence)
135
+
136
+ kps = np.stack(kps, axis=0)
137
+ confidences = np.stack(confidences, axis=0)
138
+
139
+ kps_all_parts[part] = kps
140
+ confs_all_parts[part] = confidences[..., None]
141
+
142
+ if normalization == 'signspace':
143
+ normalized_kps = sign_space_normalization(kps_all_parts.copy(), layout=layout)
144
+ else:
145
+ normalized_kps = kps_all_parts
146
+
147
+ kps_with_scores = {}
148
+ for part in normalized_kps.keys():
149
+ kps_with_scores[part] = np.concatenate([normalized_kps[part], confs_all_parts[part]], axis=-1)
150
+
151
+ kps_with_scores = {k: torch.as_tensor(v, dtype=torch.float32) for k, v in kps_with_scores.items()}
152
+ return kps_with_scores
153
+
154
+
155
+ def load_part_kp_Isharah(skeletons, confs, normalization, layout):
156
+ # kps_with_scores = {}
157
+ kps_all_parts = {}
158
+ confs_all_parts = {}
159
+
160
+ for part in ['body', 'left', 'right', 'face_all']:
161
+ kps = []
162
+ confidences = []
163
+ for i, (skeleton, conf) in enumerate(zip(skeletons, confs)):
164
+
165
+ if part == 'body':
166
+ pose_landmarks = [0, 7, 8, 11, 12, 13, 14, 15, 16]
167
+ hand_kp2d = np.stack([skeleton['pose_landmarks'][i] for i in pose_landmarks])
168
+ confidence = np.stack([conf['pose_landmarks'][i] for i in pose_landmarks])
169
+
170
+ elif part == 'left':
171
+ hand_kp2d = np.stack(skeleton['left_hand_landmarks'])
172
+ confidence = np.stack(conf['left_hand_landmarks'])
173
+
174
+ elif part == 'right':
175
+ hand_kp2d = np.stack(skeleton['right_hand_landmarks'])
176
+ confidence = np.stack(conf['right_hand_landmarks'])
177
+
178
+ elif part == 'face_all':
179
+ hand_kp2d = np.stack(skeleton['face_landmarks'])
180
+ confidence = np.stack(conf['face_landmarks'])
181
+
182
+ else:
183
+ raise NotImplementedError
184
+ kps.append(hand_kp2d)
185
+ confidences.append(confidence)
186
+
187
+ kps = np.stack(kps, axis=0)
188
+ confidences = np.stack(confidences, axis=0)
189
+
190
+ kps_all_parts[part] = kps
191
+ confs_all_parts[part] = confidences[..., None]
192
+
193
+ if normalization == 'signspace':
194
+ normalized_kps = sign_space_normalization(kps_all_parts.copy(), layout=layout)
195
+ else:
196
+ normalized_kps = kps_all_parts
197
+
198
+ kps_with_scores = {}
199
+ for part in normalized_kps.keys():
200
+ kps_with_scores[part] = np.concatenate([normalized_kps[part], confs_all_parts[part]], axis=-1)
201
+
202
+ kps_with_scores = {k: torch.as_tensor(v, dtype=torch.float32) for k, v in kps_with_scores.items()}
203
+ return kps_with_scores
204
+
205
+
206
+ YTASL_GROUP_SIZES = {
207
+ 'pose_landmarks': 33,
208
+ 'right_hand_landmarks': 21,
209
+ 'left_hand_landmarks': 21,
210
+ 'face_landmarks': 478,
211
+ }
212
+
213
+ YTASL_GROUP_ERROR_LABELS = {
214
+ 'pose_landmarks': 'a pose group',
215
+ 'right_hand_landmarks': 'a Rhand group',
216
+ 'left_hand_landmarks': 'a Lhand group',
217
+ 'face_landmarks': 'a face group',
218
+ }
219
+
220
+ ISHARAH_GROUP_SIZES = {
221
+ 'pose_landmarks': 25,
222
+ 'right_hand_landmarks': 21,
223
+ 'left_hand_landmarks': 21,
224
+ 'face_landmarks': 19,
225
+ }
226
+
227
+
228
+ def _fill_missing_landmarks(
229
+ skeleton,
230
+ conf,
231
+ group_name,
232
+ expected_size,
233
+ clip_name,
234
+ frame_idx,
235
+ error_group_label=None,
236
+ include_size_details=True,
237
+ strict_key_access=False,
238
+ ):
239
+ points = skeleton[group_name] if strict_key_access else skeleton.get(group_name, [])
240
+ if len(points) == 0:
241
+ conf[group_name] = [0] * expected_size
242
+ skeleton[group_name] = [[0.0, 0.0]] * expected_size
243
+ elif len(points) != expected_size:
244
+ group_label = error_group_label or f"group '{group_name}'"
245
+ if include_size_details:
246
+ raise NotImplementedError(
247
+ f"Unexpected number of keypoints in {group_label}: {clip_name}, frame {frame_idx}, "
248
+ f"expected {expected_size}, got {len(points)}"
249
+ )
250
+ raise NotImplementedError(f"Unexpected number of keypoints in {group_label}: {clip_name}, {frame_idx}")
251
+ else:
252
+ conf[group_name] = [1] * expected_size
253
+
254
+
255
+ # load sub-pose
256
+ def load_part_kp(skeletons, confs, force_ok=False):
257
+ thr = 0.3
258
+ kps_with_scores = {}
259
+ scale = None
260
+
261
+ for part in ['body', 'left', 'right', 'face_all']:
262
+ kps = []
263
+ confidences = []
264
+
265
+ for skeleton, conf in zip(skeletons, confs):
266
+ if skeleton.ndim == 4: # if (1,133,2) - wrapped in list
267
+ skeleton = skeleton[0]
268
+ conf = conf[0]
269
+
270
+ if part == 'body': # [0, 3, 4, 5, 6, 7, 8, 9, 10]
271
+ hand_kp2d = skeleton[[0] + [i for i in range(3, 11)], :]
272
+ confidence = conf[[0] + [i for i in range(3, 11)]]
273
+ elif part == 'left': # [91, 92, 93, 94, 95, 96, 97, 98, 99, 100, 101, 102, 103, 104, 105, 106, 107, 108, 109, 110, 111]
274
+ hand_kp2d = skeleton[91:112, :]
275
+ hand_kp2d = hand_kp2d - hand_kp2d[0, :]
276
+ confidence = conf[91:112]
277
+ elif part == 'right': # [112, 113, 114, 115, 116, 117, 118, 119, 120, 121, 122, 123, 124, 125, 126, 127, 128, 129, 130, 131, 132]
278
+ hand_kp2d = skeleton[112:133, :]
279
+ hand_kp2d = hand_kp2d - hand_kp2d[0, :]
280
+ confidence = conf[112:133]
281
+ elif part == 'face_all': # [23, 25, 27, 29, 31, 33, 35, 37, 39, 83, 84, 85, 86, 87, 88, 89, 90, 53]
282
+ hand_kp2d = skeleton[[i for i in list(range(23, 23 + 17))[::2]] + [i for i in range(83, 83 + 8)] + [53], :]
283
+ hand_kp2d = hand_kp2d - hand_kp2d[-1, :]
284
+ confidence = conf[[i for i in list(range(23, 23 + 17))[::2]] + [i for i in range(83, 83 + 8)] + [53]]
285
+
286
+ else:
287
+ raise NotImplementedError
288
+
289
+ kps.append(hand_kp2d)
290
+ confidences.append(confidence)
291
+
292
+ kps = np.stack(kps, axis=0)
293
+ confidences = np.stack(confidences, axis=0)
294
+
295
+ if part == 'body':
296
+ if force_ok:
297
+ result, scale, _ = crop_scale(np.concatenate([kps, confidences[..., None]], axis=-1), thr)
298
+
299
+ else:
300
+ result, scale, _ = crop_scale(np.concatenate([kps, confidences[..., None]], axis=-1), thr)
301
+ else:
302
+ assert not scale is None
303
+ result = np.concatenate([kps, confidences[..., None]], axis=-1)
304
+ if scale == 0:
305
+ result = np.zeros(result.shape)
306
+ else:
307
+ result[..., :2] = (result[..., :2]) / scale
308
+ result = np.clip(result, -1, 1)
309
+ # mask useless kp
310
+ result[result[..., 2] <= thr] = 0
311
+
312
+ kps_with_scores[part] = torch.tensor(result)
313
+
314
+ return kps_with_scores
315
+
316
+
317
+ # input: T, N, 3
318
+ # input is un-normed joints
319
+ def crop_scale(motion, thr):
320
+ '''
321
+ Motion: [(M), T, 17, 3].
322
+ Normalize to [-1, 1]
323
+ '''
324
+ result = copy.deepcopy(motion)
325
+ valid_coords = motion[motion[..., 2] > thr][:, :2]
326
+ if len(valid_coords) < 4:
327
+ return np.zeros(motion.shape), 0, None
328
+ xmin = min(valid_coords[:, 0])
329
+ xmax = max(valid_coords[:, 0])
330
+ ymin = min(valid_coords[:, 1])
331
+ ymax = max(valid_coords[:, 1])
332
+ # ratio = np.random.uniform(low=scale_range[0], high=scale_range[1], size=1)[0]
333
+ ratio = 1
334
+ scale = max(xmax - xmin, ymax - ymin) * ratio
335
+ if scale == 0:
336
+ return np.zeros(motion.shape), 0, None
337
+ xs = (xmin + xmax - scale) / 2
338
+ ys = (ymin + ymax - scale) / 2
339
+ result[..., :2] = (motion[..., :2] - [xs, ys]) / scale
340
+ result[..., :2] = (result[..., :2] - 0.5) * 2
341
+ result = np.clip(result, -1, 1)
342
+ # mask useless kp
343
+ result[result[..., 2] <= thr] = 0
344
+ return result, scale, [xs, ys]
345
+
346
+
347
+ # bbox of hands
348
+ def bbox_4hands(left_keypoints, right_keypoints, hw):
349
+ # keypoints --> T,21,2
350
+ # keypoints --> T,21,2
351
+
352
+ def compute_bbox(keypoints):
353
+ min_x = np.min(keypoints[..., 0], axis=1)
354
+ min_y = np.min(keypoints[..., 1], axis=1)
355
+ max_x = np.max(keypoints[..., 0], axis=1)
356
+ max_y = np.max(keypoints[..., 1], axis=1)
357
+
358
+ return (max_x + min_x) / 2, (max_y + min_y) / 2, (max_x - min_x), (max_y - min_y)
359
+
360
+ H, W = hw
361
+
362
+ if left_keypoints is None:
363
+ left_keypoints = np.zeros([1, 21, 2])
364
+
365
+ if right_keypoints is None:
366
+ right_keypoints = np.zeros([1, 21, 2])
367
+ # [T, 21, 2]
368
+ left_mean_x, left_mean_y, left_diff_x, left_diff_y = compute_bbox(left_keypoints)
369
+ left_mean_x = W * left_mean_x
370
+ left_mean_y = H * left_mean_y
371
+
372
+ left_diff_x = W * left_diff_x
373
+ left_diff_y = H * left_diff_y
374
+
375
+ left_diff_x = max(left_diff_x)
376
+ left_diff_y = max(left_diff_y)
377
+ left_box_hw = max(left_diff_x, left_diff_y)
378
+
379
+ right_mean_x, right_mean_y, right_diff_x, right_diff_y = compute_bbox(right_keypoints)
380
+ right_mean_x = W * right_mean_x
381
+ right_mean_y = H * right_mean_y
382
+
383
+ right_diff_x = W * right_diff_x
384
+ right_diff_y = H * right_diff_y
385
+
386
+ right_diff_x = max(right_diff_x)
387
+ right_diff_y = max(right_diff_y)
388
+ right_box_hw = max(right_diff_x, right_diff_y)
389
+
390
+ box_hw = int(max(left_box_hw, right_box_hw) * 1.2 / 2) * 2
391
+ box_hw = max(box_hw, 0)
392
+
393
+ left_new_box = np.stack([left_mean_x - box_hw / 2, left_mean_y - box_hw / 2, left_mean_x + box_hw / 2,
394
+ left_mean_y + box_hw / 2]).astype(np.int16)
395
+ right_new_box = np.stack([right_mean_x - box_hw / 2, right_mean_y - box_hw / 2, right_mean_x + box_hw / 2,
396
+ right_mean_y + box_hw / 2]).astype(np.int16)
397
+
398
+ return left_new_box.transpose(1, 0), right_new_box.transpose(1, 0), box_hw
399
+
400
+
401
+ def load_support_rgb_dict(tmp, skeletons, confs, full_path, data_transform):
402
+ support_rgb_dict = {}
403
+
404
+ confs = np.array(confs)
405
+ skeletons = np.array(skeletons)
406
+
407
+ # sample index of low scores
408
+ left_confs_filter = confs[:, 0, 91:112].mean(-1)
409
+ left_confs_filter_indices = np.where(left_confs_filter > 0.3)[0]
410
+
411
+ if len(left_confs_filter_indices) == 0:
412
+ left_sampled_indices = None
413
+ left_skeletons = None
414
+ else:
415
+
416
+ left_confs = confs[left_confs_filter_indices]
417
+ left_confs = left_confs[:, 0, [95, 99, 103, 107, 111]].min(-1)
418
+
419
+ left_weights = np.max(left_confs) - left_confs + 1e-5
420
+ left_probabilities = left_weights / np.sum(left_weights)
421
+
422
+ left_sample_size = int(np.ceil(0.1 * len(left_confs_filter_indices)))
423
+
424
+ left_sampled_indices = np.random.choice(left_confs_filter_indices.tolist(),
425
+ size=left_sample_size,
426
+ replace=False,
427
+ p=left_probabilities)
428
+ # left_sampled_indices: values: 0-255(0,max_len)
429
+ # tmp: values: 0-(end-start)
430
+ left_sampled_indices = np.sort(left_sampled_indices)
431
+
432
+ left_skeletons = skeletons[left_sampled_indices, 0, 91:112]
433
+
434
+ right_confs_filter = confs[:, 0, 112:].mean(-1)
435
+ right_confs_filter_indices = np.where(right_confs_filter > 0.3)[0]
436
+ if len(right_confs_filter_indices) == 0:
437
+ right_sampled_indices = None
438
+ right_skeletons = None
439
+
440
+ else:
441
+ right_confs = confs[right_confs_filter_indices]
442
+ right_confs = right_confs[:, 0, [95 + 21, 99 + 21, 103 + 21, 107 + 21, 111 + 21]].min(-1)
443
+
444
+ right_weights = np.max(right_confs) - right_confs + 1e-5
445
+ right_probabilities = right_weights / np.sum(right_weights)
446
+
447
+ right_sample_size = int(np.ceil(0.1 * len(right_confs_filter_indices)))
448
+
449
+ right_sampled_indices = np.random.choice(right_confs_filter_indices.tolist(),
450
+ size=right_sample_size,
451
+ replace=False,
452
+ p=right_probabilities)
453
+ right_sampled_indices = np.sort(right_sampled_indices)
454
+
455
+ right_skeletons = skeletons[right_sampled_indices, 0, 112:133]
456
+
457
+ image_size = 112
458
+ all_indices = []
459
+ if not left_sampled_indices is None:
460
+ all_indices.append(left_sampled_indices)
461
+ if not right_sampled_indices is None:
462
+ all_indices.append(right_sampled_indices)
463
+ if len(all_indices) == 0:
464
+ support_rgb_dict['left_sampled_indices'] = torch.tensor([-1])
465
+ support_rgb_dict['left_hands'] = torch.zeros(1, 3, image_size, image_size)
466
+ support_rgb_dict['left_skeletons_norm'] = torch.zeros(1, 21, 2)
467
+
468
+ support_rgb_dict['right_sampled_indices'] = torch.tensor([-1])
469
+ support_rgb_dict['right_hands'] = torch.zeros(1, 3, image_size, image_size)
470
+ support_rgb_dict['right_skeletons_norm'] = torch.zeros(1, 21, 2)
471
+
472
+ return support_rgb_dict
473
+
474
+ sampled_indices = np.concatenate(all_indices)
475
+ sampled_indices = np.unique(sampled_indices)
476
+ sampled_indices_real = tmp[sampled_indices]
477
+
478
+ # load image sample
479
+ imgs = load_video_support_rgb(full_path, sampled_indices_real)
480
+
481
+ # get hand bbox
482
+ left_new_box, right_new_box, box_hw = bbox_4hands(left_skeletons,
483
+ right_skeletons,
484
+ imgs[0].shape[:2])
485
+
486
+ # crop left and right hand
487
+ image_size = 112
488
+ if box_hw == 0:
489
+ support_rgb_dict['left_sampled_indices'] = torch.tensor([-1])
490
+ support_rgb_dict['left_hands'] = torch.zeros(1, 3, image_size, image_size)
491
+ support_rgb_dict['left_skeletons_norm'] = torch.zeros(1, 21, 2)
492
+
493
+ support_rgb_dict['right_sampled_indices'] = torch.tensor([-1])
494
+ support_rgb_dict['right_hands'] = torch.zeros(1, 3, image_size, image_size)
495
+ support_rgb_dict['right_skeletons_norm'] = torch.zeros(1, 21, 2)
496
+
497
+ return support_rgb_dict
498
+
499
+ factor = image_size / box_hw
500
+
501
+ if left_sampled_indices is None:
502
+ left_hands = torch.zeros(1, 3, image_size, image_size)
503
+ left_skeletons_norm = torch.zeros(1, 21, 2)
504
+
505
+ else:
506
+ left_hands = torch.zeros(len(left_sampled_indices), 3, image_size, image_size)
507
+
508
+ left_skeletons_norm = left_skeletons * imgs[0].shape[:2][::-1] - left_new_box[:, None, [0, 1]]
509
+ left_skeletons_norm = left_skeletons_norm / box_hw
510
+ left_skeletons_norm = left_skeletons_norm.clip(0, 1)
511
+
512
+ if right_sampled_indices is None:
513
+ right_hands = torch.zeros(1, 3, image_size, image_size)
514
+ right_skeletons_norm = torch.zeros(1, 21, 2)
515
+
516
+ else:
517
+ right_hands = torch.zeros(len(right_sampled_indices), 3, image_size, image_size)
518
+
519
+ right_skeletons_norm = right_skeletons * imgs[0].shape[:2][::-1] - right_new_box[:, None, [0, 1]]
520
+ right_skeletons_norm = right_skeletons_norm / box_hw
521
+ right_skeletons_norm = right_skeletons_norm.clip(0, 1)
522
+ left_idx = 0
523
+ right_idx = 0
524
+
525
+ for idx, img in enumerate(imgs):
526
+ mapping_idx = sampled_indices[idx]
527
+ if not left_sampled_indices is None and left_idx < len(left_sampled_indices) and mapping_idx == \
528
+ left_sampled_indices[left_idx]:
529
+ box = left_new_box[left_idx]
530
+
531
+ img_draw = np.uint8(copy.deepcopy(img))[box[1]:box[3], box[0]:box[2], :]
532
+ img_draw = np.pad(img_draw,
533
+ ((0, max(0, box_hw - img_draw.shape[0])), (0, max(0, box_hw - img_draw.shape[1])),
534
+ (0, 0)), mode='constant', constant_values=0)
535
+
536
+ f_img = Image.fromarray(img_draw).convert('RGB').resize((image_size, image_size))
537
+ f_img = data_transform(f_img).unsqueeze(0)
538
+ left_hands[left_idx] = f_img
539
+ left_idx += 1
540
+
541
+ if not right_sampled_indices is None and right_idx < len(right_sampled_indices) and mapping_idx == \
542
+ right_sampled_indices[right_idx]:
543
+ box = right_new_box[right_idx]
544
+
545
+ img_draw = np.uint8(copy.deepcopy(img))[box[1]:box[3], box[0]:box[2], :]
546
+ img_draw = np.pad(img_draw,
547
+ ((0, max(0, box_hw - img_draw.shape[0])), (0, max(0, box_hw - img_draw.shape[1])),
548
+ (0, 0)), mode='constant', constant_values=0)
549
+
550
+ f_img = Image.fromarray(img_draw).convert('RGB').resize((image_size, image_size))
551
+ f_img = data_transform(f_img).unsqueeze(0)
552
+ right_hands[right_idx] = f_img
553
+ right_idx += 1
554
+
555
+ if left_sampled_indices is None:
556
+ left_sampled_indices = np.array([-1])
557
+
558
+ if right_sampled_indices is None:
559
+ right_sampled_indices = np.array([-1])
560
+
561
+ # get index, images and keypoints priors
562
+ support_rgb_dict['left_sampled_indices'] = torch.tensor(left_sampled_indices)
563
+ support_rgb_dict['left_hands'] = left_hands
564
+ support_rgb_dict['left_skeletons_norm'] = torch.tensor(left_skeletons_norm)
565
+
566
+ support_rgb_dict['right_sampled_indices'] = torch.tensor(right_sampled_indices)
567
+ support_rgb_dict['right_hands'] = right_hands
568
+ support_rgb_dict['right_skeletons_norm'] = torch.tensor(right_skeletons_norm)
569
+
570
+ return support_rgb_dict
571
+
572
+
573
+ # use split rgb video for save time
574
+ def load_video_support_rgb(path, tmp):
575
+ vr = VideoReader(path, num_threads=1, ctx=cpu(0))
576
+
577
+ vr.seek(0)
578
+ buffer = vr.get_batch(tmp).asnumpy()
579
+ batch_image = buffer
580
+ del vr
581
+
582
+ return batch_image
583
+
584
+
585
+ def load_json(path):
586
+ with open(path, "r", encoding="utf-8") as f:
587
+ return json.load(f)
588
+
589
+ def is_valid_metric_label(text):
590
+ if text is None:
591
+ return False
592
+ text = " ".join(str(text).split()).strip()
593
+ if not text:
594
+ return False
595
+ # Require at least one alnum/letter token after punctuation/symbol stripping.
596
+ return re.search(r"\w", text, flags=re.UNICODE) is not None
597
+
598
+ def select_frame_indices(duration, max_length, phase):
599
+ if duration <= max_length:
600
+ return list(range(duration))
601
+ if phase == 'train':
602
+ return sorted(random.sample(range(duration), k=max_length))
603
+ # Deterministic, near-uniform coverage for dev/test.
604
+ return ((np.arange(max_length) * duration) // max_length).tolist()
605
+
606
+
607
+ # build base dataset
608
+ class Base_Dataset(Dataset.Dataset):
609
+ def collate_fn(self, batch):
610
+ tgt_batch, src_length_batch, name_batch, pose_tmp, gloss_batch = [], [], [], [], []
611
+
612
+ for name_sample, pose_sample, text, gloss, _ in batch:
613
+ name_batch.append(name_sample)
614
+ pose_tmp.append(pose_sample)
615
+ tgt_batch.append(text)
616
+ gloss_batch.append(gloss)
617
+
618
+ src_input = {}
619
+
620
+ keys = pose_tmp[0].keys()
621
+ for key in keys:
622
+ max_len = max([len(vid[key]) for vid in pose_tmp])
623
+ video_length = torch.LongTensor([len(vid[key]) for vid in pose_tmp])
624
+
625
+ padded_video = [torch.cat(
626
+ (
627
+ vid[key],
628
+ vid[key][-1][None].expand(max_len - len(vid[key]), -1, -1),
629
+ )
630
+ , dim=0)
631
+ for vid in pose_tmp]
632
+
633
+ img_batch = torch.stack(padded_video, 0)
634
+
635
+ src_input[key] = img_batch
636
+ if 'attention_mask' not in src_input.keys():
637
+ src_length_batch = video_length
638
+
639
+ mask_gen = []
640
+ for i in src_length_batch:
641
+ tmp = torch.ones([i]) + 7
642
+ mask_gen.append(tmp)
643
+ mask_gen = pad_sequence(mask_gen, padding_value=0, batch_first=True)
644
+ img_padding_mask = (mask_gen != 0).long()
645
+ src_input['attention_mask'] = img_padding_mask
646
+
647
+ src_input['name_batch'] = name_batch
648
+ src_input['src_length_batch'] = src_length_batch
649
+
650
+ if self.rgb_support:
651
+ support_rgb_dicts = {key: [] for key in batch[0][-1].keys()}
652
+ for _, _, _, _, support_rgb_dict in batch:
653
+ for key in support_rgb_dict.keys():
654
+ support_rgb_dicts[key].append(support_rgb_dict[key])
655
+
656
+ for part in ['left', 'right']:
657
+ index_key = f'{part}_sampled_indices'
658
+ skeletons_key = f'{part}_skeletons_norm'
659
+ rgb_key = f'{part}_hands'
660
+ len_key = f'{part}_rgb_len'
661
+
662
+ index_batch = torch.cat(support_rgb_dicts[index_key], 0)
663
+ skeletons_batch = torch.cat(support_rgb_dicts[skeletons_key], 0)
664
+ img_batch = torch.cat(support_rgb_dicts[rgb_key], 0)
665
+
666
+ src_input[index_key] = index_batch
667
+ src_input[skeletons_key] = skeletons_batch
668
+ src_input[rgb_key] = img_batch
669
+ src_input[len_key] = [len(index) for index in support_rgb_dicts[index_key]]
670
+
671
+ tgt_input = {}
672
+ tgt_input['gt_sentence'] = tgt_batch
673
+ tgt_input['gt_gloss'] = gloss_batch
674
+
675
+ return src_input, tgt_input
676
+
677
+ import gzip
678
+ def load_dataset_file(filename):
679
+ with gzip.open(filename, "rb") as f:
680
+ loaded_object = pickle.load(f)
681
+ return loaded_object
682
+
683
+ class S2T_Dataset(Base_Dataset):
684
+ def __init__(self, path, args, phase):
685
+ super(S2T_Dataset, self).__init__()
686
+ self.args = args
687
+ self.rgb_support = self.args.rgb_support
688
+ self.max_length = args.max_length
689
+ self.raw_data = load_dataset_file(path)
690
+ self.phase = phase
691
+
692
+ if self.args.dataset == "CSL_Daily":
693
+ self.pose_dir = pose_dirs[args.dataset]
694
+ self.rgb_dir = rgb_dirs[args.dataset]
695
+
696
+ elif "WLASL" in self.args.dataset:
697
+ self.pose_dir = os.path.join(pose_dirs[args.dataset], phase)
698
+ self.rgb_dir = os.path.join(rgb_dirs[args.dataset], phase)
699
+
700
+ elif self.args.dataset == "Isharah":
701
+ self.pose_dir = pose_dirs[args.dataset] # pose only
702
+ self.rgb_dir = rgb_dirs[args.dataset] # ""
703
+
704
+ else:
705
+ raise NotImplementedError(f"dataset {self.args.dataset} not supported")
706
+
707
+ self.list = list(self.raw_data.keys())
708
+
709
+ self.data_transform = transforms.Compose([
710
+ transforms.ToTensor(),
711
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
712
+ ])
713
+
714
+ def __len__(self):
715
+ return len(self.list)
716
+
717
+ def __getitem__(self, index):
718
+ key = self.list[index]
719
+ sample = self.raw_data[key]
720
+
721
+ text = sample['text']
722
+ if "gloss" in sample.keys():
723
+ gloss = " ".join(sample['gloss'])
724
+ else:
725
+ gloss = ''
726
+
727
+ name_sample = sample['name']
728
+ pose_sample, support_rgb_dict = self.load_pose(sample['video_path'])
729
+
730
+ return name_sample, pose_sample, text, gloss, support_rgb_dict
731
+
732
+ def load_pose(self, path):
733
+ pose = pickle.load(open(os.path.join(self.pose_dir, path.replace(".mp4", '.pkl')), 'rb'))
734
+
735
+ if 'start' in pose.keys():
736
+ assert pose['start'] < pose['end']
737
+ duration = pose['end'] - pose['start']
738
+ start = pose['start']
739
+ else:
740
+ duration = len(pose['scores'])
741
+ start = 0
742
+
743
+ tmp = select_frame_indices(duration, self.max_length, self.phase)
744
+
745
+ tmp = np.array(tmp) + start
746
+
747
+ skeletons = pose['keypoints']
748
+ confs = pose['scores']
749
+ skeletons_tmp = []
750
+ confs_tmp = []
751
+ for index in tmp:
752
+ skeletons_tmp.append(skeletons[index])
753
+ confs_tmp.append(confs[index])
754
+
755
+ skeletons = skeletons_tmp
756
+ confs = confs_tmp
757
+
758
+ kps_with_scores = load_part_kp(skeletons, confs, force_ok=True)
759
+
760
+ support_rgb_dict = {}
761
+ if self.rgb_support:
762
+ full_path = os.path.join(self.rgb_dir, path)
763
+ support_rgb_dict = load_support_rgb_dict(tmp, skeletons, confs, full_path, self.data_transform)
764
+
765
+ return kps_with_scores, support_rgb_dict
766
+
767
+ def __str__(self):
768
+ return f'#total {len(self)}'
769
+
770
+
771
+ class S2T_Dataset_YTASL(Base_Dataset):
772
+ def __init__(self, path, args, phase):
773
+ super(S2T_Dataset_YTASL, self).__init__()
774
+ self.args = args
775
+ self.max_length = args.max_length
776
+ self.phase = phase
777
+ self.annotation = load_json(path)
778
+ self.rgb_support = self.args.rgb_support
779
+ self.normalization = self.args.normalization
780
+ self.layout = args.layout
781
+
782
+ self.pose_dir = pose_dirs[args.dataset]
783
+ self.rgb_dir = rgb_dirs[args.dataset]
784
+
785
+ self.list_data = [] # [(video_id, clip_id), ...]
786
+ self.clip_order_to_int = {}
787
+ self.clip_order_from_int = {}
788
+
789
+ for video_id in self.annotation.keys():
790
+ co = self.annotation[video_id]['clip_order']
791
+ self.clip_order_from_int[video_id] = dict(zip(range(len(co)), co))
792
+ self.clip_order_to_int[video_id] = dict(zip(co, range(len(co))))
793
+
794
+ for video_id, clip_dict in self.annotation.items():
795
+ for clip_name in clip_dict['clip_order']:
796
+ self.list_data.append((video_id, self.clip_order_to_int[video_id][clip_name]))
797
+
798
+ available_clip_names = {
799
+ pathlib.Path(clip).stem # remove suffix
800
+ for clip in os.listdir(self.pose_dir)
801
+ if clip.endswith(".json")
802
+ }
803
+ video_clips = set()
804
+ for video_id, clip_dict in self.annotation.items():
805
+ for clip_name in clip_dict['clip_order']:
806
+ if clip_name in available_clip_names:
807
+ video_clips.add((video_id, self.clip_order_to_int[video_id][clip_name]))
808
+
809
+ self.remove_missing_annotation(video_clips) # Remove data in annotations that are missing in h5 file
810
+
811
+ def remove_missing_annotation(self, h5_video_clip):
812
+ annotations_to_delete = set(self.list_data) - h5_video_clip
813
+ for a in annotations_to_delete:
814
+ self.list_data.remove(a)
815
+
816
+ def __getitem__(self, index):
817
+ video_id, clip_id = self.list_data[index]
818
+ clip_name = self.clip_order_from_int[video_id][clip_id]
819
+
820
+ # Get translation
821
+ clip_dict = self.annotation[video_id][clip_name]
822
+ text = clip_dict['translation']
823
+
824
+ # Get the pose features
825
+ pose_sample = self.load_pose(clip_name)
826
+
827
+ # TODO: rgb support
828
+ video_path = ""
829
+ support_rgb_dict = {}
830
+
831
+ # sample = {"name": clip_name,
832
+ # "video_path": video_path,
833
+ # "pose_features": pose_features,
834
+ # "text": translation}
835
+ # Crop long sequences to desired max length. Random sample
836
+
837
+ # skeletons = pose['keypoints']
838
+ # confs = pose['scores']
839
+ # skeletons_tmp = []
840
+ # confs_tmp = []
841
+ # for index in tmp:
842
+ # skeletons_tmp.append(skeletons[index])
843
+ # confs_tmp.append(confs[index])
844
+ #
845
+ # skeletons = skeletons_tmp
846
+ # confs = confs_tmp
847
+
848
+ # confs = [np.ones(int(pose_features.shape[1]/2)) for _ in range(pose_features.shape[0])]
849
+ # confs = [np.ones(pose_features.shape[0])] * pose_features.shape[1]
850
+ # skeletons = [] # List of ndarrays (133,2) - full keypoints
851
+ # kps_with_scores = load_part_kp(skeletons, confs, force_ok=True)
852
+
853
+ name_sample = clip_name
854
+ gloss = ''
855
+
856
+ return name_sample, pose_sample, text, gloss, support_rgb_dict
857
+
858
+ def load_pose(self, clip_name):
859
+ path = os.path.join(self.pose_dir, f"{clip_name}.json")
860
+ pose_data = load_json(path)
861
+ pose = pose_data['cropped_keypoints']
862
+
863
+ duration = len(pose)
864
+ start = 0
865
+
866
+ tmp = select_frame_indices(duration, self.max_length, self.phase)
867
+ tmp = np.array(tmp) + start
868
+ skeletons = [pose[i] for i in tmp]
869
+
870
+ confs = []
871
+ for i, skeleton in enumerate(skeletons):
872
+ conf = {}
873
+ for group_name, expected_size in YTASL_GROUP_SIZES.items():
874
+ _fill_missing_landmarks(
875
+ skeleton=skeleton,
876
+ conf=conf,
877
+ group_name=group_name,
878
+ expected_size=expected_size,
879
+ clip_name=clip_name,
880
+ frame_idx=i,
881
+ error_group_label=YTASL_GROUP_ERROR_LABELS[group_name],
882
+ include_size_details=False,
883
+ strict_key_access=True,
884
+ )
885
+
886
+ confs.append(conf)
887
+
888
+ kps_with_scores = load_part_kp_YTASL(skeletons, confs, self.normalization, self.layout)
889
+ return kps_with_scores
890
+
891
+ def __len__(self):
892
+ return len(self.list_data)
893
+
894
+ def __str__(self):
895
+ return f'#total {len(self)}'
896
+
897
+
898
+ class S2T_Dataset_Isharah(S2T_Dataset_YTASL):
899
+ def __init__(self, path, args, phase):
900
+ super(S2T_Dataset_Isharah, self).__init__(path=path, args=args, phase=phase)
901
+
902
+ def load_pose(self, clip_name):
903
+ path = os.path.join(self.pose_dir, f"{clip_name}.json")
904
+ pose_data = load_json(path)
905
+ pose = pose_data['cropped_keypoints']
906
+
907
+ duration = len(pose)
908
+ tmp = select_frame_indices(duration, self.max_length, self.phase)
909
+ tmp = np.array(tmp)
910
+ skeletons = [pose[i] for i in tmp]
911
+
912
+ confs = []
913
+ for i, skeleton in enumerate(skeletons):
914
+ conf = {}
915
+ for group_name, expected_size in ISHARAH_GROUP_SIZES.items():
916
+ _fill_missing_landmarks(
917
+ skeleton=skeleton,
918
+ conf=conf,
919
+ group_name=group_name,
920
+ expected_size=expected_size,
921
+ clip_name=clip_name,
922
+ frame_idx=i,
923
+ )
924
+ confs.append(conf)
925
+
926
+ kps_with_scores = load_part_kp_Isharah(skeletons, confs, self.normalization, self.layout)
927
+ return kps_with_scores
928
+
929
+
930
+ # class S2T_Dataset_YTASL_h5(Base_Dataset):
931
+ # def __init__(self, path, args, phase):
932
+ # super(S2T_Dataset_YTASL_h5, self).__init__()
933
+ # self.args = args
934
+ # self.max_length = args.max_length
935
+ # self.phase = phase
936
+ # self.annotation = load_json(path)
937
+ # self.rgb_support = self.args.rgb_support
938
+ #
939
+ # # Load poses
940
+ # self.list_data = [] # [(video_id, clip_id), ...]
941
+ # self.h5_data = {}
942
+ # self.h5shard = defaultdict(lambda: defaultdict(dict))
943
+ # self.clip_order_to_int = {}
944
+ # self.clip_order_from_int = {}
945
+ #
946
+ # for video_id in self.annotation.keys():
947
+ # co = self.annotation[video_id]['clip_order']
948
+ # self.clip_order_from_int[video_id] = dict(zip(range(len(co)), co))
949
+ # self.clip_order_to_int[video_id] = dict(zip(co, range(len(co))))
950
+ #
951
+ # for video_id, clip_dict in self.annotation.items():
952
+ # for clip_name in clip_dict:
953
+ # if clip_name != "clip_order":
954
+ # self.list_data.append((video_id, self.clip_order_to_int[video_id][clip_name]))
955
+ #
956
+ # self.vf_path = os.path.join(pose_dirs["YTASL"], "YouTubeASL.keypoints.{}.json".format(phase))
957
+ #
958
+ # h5_video_clip = self.read_multih5_json(self.vf_path)
959
+ # self.remove_missing_annotation(h5_video_clip) # Remove data in annotations that are missing in h5 file
960
+ #
961
+ # def read_multih5_json(self, json_filename):
962
+ # """Helper function for reading json specifications of multiple H5 files for visual features"""
963
+ # h5_video_clip = set()
964
+ # with open(json_filename, 'r') as F:
965
+ # self.h5shard = json.load(F)
966
+ # self.h5_data = {}
967
+ # print(f"Pose {self.phase} data are loaded from: ")
968
+ # for k in set(self.h5shard.values()):
969
+ # h5file = json_filename.replace('metadata_', '').replace('.json', ".%s.h5" % k)
970
+ # print("--" + h5file) # ,k,json_filename,data_dir)
971
+ # self.h5_data[k] = h5py.File(h5file, 'r')
972
+ #
973
+ # for vi in self.h5_data[k].keys():
974
+ # for ci in self.h5_data[k][vi].keys():
975
+ # if vi in self.clip_order_to_int:
976
+ # if ci in self.clip_order_to_int[vi]:
977
+ # clip_id = self.clip_order_to_int[vi][ci]
978
+ # h5_video_clip.add((vi, clip_id))
979
+ # return h5_video_clip
980
+ #
981
+ # def remove_missing_annotation(self, h5_video_clip):
982
+ # annotations_to_delete = set(self.list_data) - h5_video_clip
983
+ # for a in annotations_to_delete:
984
+ # self.list_data.remove(a)
985
+ #
986
+ # def __getitem__(self, index):
987
+ # video_id, clip_id = self.list_data[index]
988
+ # clip_name = self.clip_order_from_int[video_id][clip_id]
989
+ #
990
+ # # Get the pose features
991
+ # shard = self.h5shard[video_id]
992
+ # pose_features = torch.tensor(np.array(self.h5_data[shard][video_id][clip_name]))
993
+ #
994
+ # # TODO: rgb support
995
+ # video_path = ""
996
+ #
997
+ # # Get translation
998
+ # clip_dict = self.annotation[video_id][clip_name]
999
+ # translation = clip_dict['translation']
1000
+ #
1001
+ # # sample = {"name": clip_name,
1002
+ # # "video_path": video_path,
1003
+ # # "pose_features": pose_features,
1004
+ # # "text": translation}
1005
+ # #
1006
+ # # # Crop long sequences to desired max length. Random sample
1007
+ # # duration = len(pose_features) # TODO: works?
1008
+ # # if duration > self.max_length:
1009
+ # # tmp = sorted(random.sample(range(duration), k=self.max_length))
1010
+ # # else:
1011
+ # # tmp = list(range(duration))
1012
+ # #
1013
+ # # tmp = np.array(tmp)
1014
+ #
1015
+ # # skeletons = pose['keypoints']
1016
+ # # confs = pose['scores']
1017
+ # # skeletons_tmp = []
1018
+ # # confs_tmp = []
1019
+ # # for index in tmp:
1020
+ # # skeletons_tmp.append(skeletons[index])
1021
+ # # confs_tmp.append(confs[index])
1022
+ # #
1023
+ # # skeletons = skeletons_tmp
1024
+ # # confs = confs_tmp
1025
+ #
1026
+ # # confs = [np.ones(int(pose_features.shape[1]/2)) for _ in range(pose_features.shape[0])]
1027
+ # # confs = [np.ones(pose_features.shape[0])] * pose_features.shape[1]
1028
+ # # skeletons = [] # List of ndarrays (133,2) - full keypoints
1029
+ # # kps_with_scores = load_part_kp(skeletons, confs, force_ok=True)
1030
+ #
1031
+ # # decoded = self.tokenizer(
1032
+ # # translation,
1033
+ # # max_length=self.max_token_length,
1034
+ # # padding="max_length",
1035
+ # # truncation=True,
1036
+ # # return_tensors="pt",
1037
+ # # )
1038
+ # # labels = decoded.input_ids
1039
+ #
1040
+ # # Skip frames for the keypoints
1041
+ # # if self.skip_frames:
1042
+ # # if type(self.skip_frames) == bool:
1043
+ # # for input_type in INPUT_TYPES:
1044
+ # # if visual_features[input_type] is not None:
1045
+ # # visual_features[input_type] = visual_features[input_type][::2]
1046
+ # # elif type(self.skip_frames) == int:
1047
+ # # for input_type in INPUT_TYPES:
1048
+ # # if visual_features[input_type] is not None:
1049
+ # # visual_features[input_type] = visual_features[input_type][::self.skip_frames]
1050
+ # #
1051
+ # # # Trim the keypoints to the max sequence length
1052
+ # # if self.max_sequence_length:
1053
+ # # for input_type in INPUT_TYPES:
1054
+ # # if visual_features[input_type] is not None:
1055
+ # # visual_features[input_type] = visual_features[input_type][: self.max_sequence_length]
1056
+ # # seq_len = len(visual_features[input_type])
1057
+ # #
1058
+ # # assert seq_len, "No modality provided or clip has no length!"
1059
+ # # attention_mask = torch.ones(seq_len)
1060
+ # #
1061
+ # # return {
1062
+ # # "sign_inputs": {'pose': visual_features['pose'],
1063
+ # # 'mae': visual_features['mae'],
1064
+ # # 'dino': visual_features['dino'],
1065
+ # # 'sign2vec': visual_features['sign2vec']},
1066
+ # # "sentence": translation,
1067
+ # # "labels": labels,
1068
+ # # "attention_mask": attention_mask,
1069
+ # # }
1070
+ # return kps
1071
+ #
1072
+ # def __len__(self):
1073
+ # return len(self.list_data)
1074
+ #
1075
+ # def __str__(self):
1076
+ # return f'#total {len(self)}'
1077
+
1078
+
1079
+ class S2T_Dataset_news(Base_Dataset):
1080
+ def __init__(self, path, args, phase):
1081
+ super(S2T_Dataset_news, self).__init__()
1082
+ self.args = args
1083
+ self.rgb_support = self.args.rgb_support
1084
+ self.phase = phase
1085
+ self.max_length = args.max_length
1086
+
1087
+ path = pathlib.Path(path)
1088
+
1089
+ with path.open(encoding='utf-8') as f:
1090
+ self.annotation = json.load(f)
1091
+
1092
+ if self.args.dataset == "CSL_News":
1093
+ self.pose_dir = pose_dirs[args.dataset]
1094
+ self.rgb_dir = rgb_dirs[args.dataset]
1095
+
1096
+ else:
1097
+ raise NotImplementedError
1098
+ sum_sample = len(self.annotation)
1099
+ self.data_transform = transforms.Compose([
1100
+ transforms.ToTensor(),
1101
+ transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),
1102
+ ])
1103
+
1104
+ if phase == 'train':
1105
+ self.start_idx = int(sum_sample * 0.0)
1106
+ self.end_idx = int(sum_sample * 0.99)
1107
+ else:
1108
+ self.start_idx = int(sum_sample * 0.99)
1109
+ self.end_idx = int(sum_sample)
1110
+
1111
+ def __len__(self):
1112
+ return self.end_idx - self.start_idx
1113
+
1114
+ def __getitem__(self, index):
1115
+ num_retries = 10
1116
+
1117
+ # skip some invalid video sample
1118
+ for _ in range(num_retries):
1119
+ sample = self.annotation[self.start_idx:self.end_idx][index]
1120
+
1121
+ text = sample['text']
1122
+ name_sample = sample['video']
1123
+
1124
+ try:
1125
+ pose_sample, support_rgb_dict = self.load_pose(sample['pose'], sample['video'])
1126
+
1127
+ except:
1128
+ import traceback
1129
+
1130
+ traceback.print_exc()
1131
+ print(f"Failed to load examples with video: {name_sample}. "
1132
+ f"Will randomly sample an example as a replacement.")
1133
+ index = random.randint(0, len(self) - 1)
1134
+ continue
1135
+
1136
+ break
1137
+
1138
+ else:
1139
+ raise RuntimeError(f"Failed to fetch video after {num_retries} retries.")
1140
+
1141
+ return name_sample, pose_sample, text, _, support_rgb_dict
1142
+
1143
+ def load_pose(self, pose_name, rgb_name):
1144
+ pose = pickle.load(open(os.path.join(self.pose_dir, pose_name), 'rb'))
1145
+ full_path = os.path.join(self.rgb_dir, rgb_name)
1146
+
1147
+ duration = len(pose['scores'])
1148
+
1149
+ tmp = select_frame_indices(duration, self.max_length, self.phase)
1150
+
1151
+ tmp = np.array(tmp)
1152
+
1153
+ # dict_keys(['keypoints', 'scores'])
1154
+ # keypoints (1, 133, 2)
1155
+ # scores (1, 133)
1156
+
1157
+ skeletons = pose['keypoints']
1158
+ confs = pose['scores']
1159
+ skeletons_tmp = []
1160
+ confs_tmp = []
1161
+
1162
+ for index in tmp:
1163
+ skeletons_tmp.append(skeletons[index])
1164
+ confs_tmp.append(confs[index])
1165
+
1166
+ skeletons = skeletons_tmp
1167
+ confs = confs_tmp
1168
+
1169
+ kps_with_scores = load_part_kp(skeletons, confs)
1170
+
1171
+ support_rgb_dict = {}
1172
+ if self.rgb_support:
1173
+ support_rgb_dict = load_support_rgb_dict(tmp, skeletons, confs, full_path, self.data_transform)
1174
+
1175
+ return kps_with_scores, support_rgb_dict
1176
+
1177
+ def __str__(self):
1178
+ return f'#total {len(self)}'
Uni_Sign/deformable_attention_2d.py ADDED
@@ -0,0 +1,311 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Clone from https://github.com/lucidrains/deformable-attention
2
+ import torch
3
+ import torch.nn.functional as F
4
+ from torch import nn, einsum
5
+
6
+ from einops import rearrange, repeat
7
+ from einops.layers.torch import Rearrange
8
+ # helper functions
9
+ from typing import Any, Optional, Tuple, Type
10
+ import numpy as np
11
+
12
+ def exists(val):
13
+ return val is not None
14
+
15
+ def default(val, d):
16
+ return val if exists(val) else d
17
+
18
+ def divisible_by(numer, denom):
19
+ return (numer % denom) == 0
20
+
21
+ # tensor helpers
22
+
23
+ def create_grid_like(t, dim = 0):
24
+ h, w, device = *t.shape[-2:], t.device
25
+
26
+ grid = torch.stack(torch.meshgrid(
27
+ torch.arange(w, device = device),
28
+ torch.arange(h, device = device),
29
+ indexing = 'xy'), dim = dim)
30
+
31
+ grid.requires_grad = False
32
+ grid = grid.type_as(t)
33
+ return grid
34
+
35
+ def normalize_grid(grid, dim = 1, out_dim = -1):
36
+ # normalizes a grid to range from -1 to 1
37
+ h, w = grid.shape[-2:]
38
+ grid_h, grid_w = grid.unbind(dim = dim)
39
+
40
+ grid_h = 2.0 * grid_h / max(h - 1, 1) - 1.0
41
+ grid_w = 2.0 * grid_w / max(w - 1, 1) - 1.0
42
+
43
+ return torch.stack((grid_h, grid_w), dim = out_dim)
44
+
45
+ def reshape_grid_1d(grid, dim = 1, out_dim = -1):
46
+ # normalizes a grid to range from -1 to 1
47
+ n = grid.shape[-1]
48
+ grid_h, grid_w = grid.unbind(dim = dim)
49
+
50
+ return torch.stack((grid_h, grid_w), dim = out_dim)
51
+
52
+ class Scale(nn.Module):
53
+ def __init__(self, scale):
54
+ super().__init__()
55
+ self.scale = scale
56
+
57
+ def forward(self, x):
58
+ return x * self.scale
59
+
60
+ # continuous positional bias from SwinV2
61
+
62
+ class CPB(nn.Module):
63
+ """ https://arxiv.org/abs/2111.09883v1 """
64
+
65
+ def __init__(self, dim, *, heads, offset_groups, depth):
66
+ super().__init__()
67
+ self.heads = heads
68
+ self.offset_groups = offset_groups
69
+
70
+ self.mlp = nn.ModuleList([])
71
+
72
+ self.mlp.append(nn.Sequential(
73
+ nn.Linear(2, dim),
74
+ nn.ReLU()
75
+ ))
76
+
77
+ for _ in range(depth - 1):
78
+ self.mlp.append(nn.Sequential(
79
+ nn.Linear(dim, dim),
80
+ nn.ReLU()
81
+ ))
82
+
83
+ self.mlp.append(nn.Linear(dim, heads // offset_groups))
84
+
85
+ def forward(self, grid_q, grid_kv):
86
+ device, dtype = grid_q.device, grid_kv.dtype
87
+
88
+ if grid_q.ndim == 3:
89
+ raise AttributeError
90
+ grid_q = rearrange(grid_q, 'h w c -> 1 (h w) c')
91
+ elif grid_q.ndim == 4:
92
+ grid_q = rearrange(grid_q, 'b h w c -> b (h w) c')
93
+ else:
94
+ raise AttributeError
95
+ grid_kv = rearrange(grid_kv, 'b h w c -> b (h w) c')
96
+
97
+ pos = rearrange(grid_q, 'b i c -> b i 1 c') - rearrange(grid_kv, 'b j c -> b 1 j c')
98
+ bias = torch.sign(pos) * torch.log(pos.abs() + 1) # log of distance is sign(rel_pos) * log(abs(rel_pos) + 1)
99
+
100
+ for layer in self.mlp:
101
+ bias = layer(bias)
102
+
103
+ bias = rearrange(bias, '(b g) i j o -> b (g o) i j', g = self.offset_groups)
104
+
105
+ return bias
106
+
107
+
108
+ class PositionEmbeddingRandom(nn.Module):
109
+ """
110
+ Positional encoding using random spatial frequencies.
111
+ """
112
+
113
+ def __init__(self, num_pos_feats: int = 64, scale: Optional[float] = None) -> None:
114
+ super().__init__()
115
+ if scale is None or scale <= 0.0:
116
+ scale = 1.0
117
+ self.register_buffer(
118
+ "positional_encoding_gaussian_matrix",
119
+ scale * torch.randn((2, num_pos_feats)),
120
+ )
121
+
122
+ def _pe_encoding(self, coords: torch.Tensor) -> torch.Tensor:
123
+ """Positionally encode points that are normalized to [0,1]."""
124
+ # assuming coords are in [0, 1]^2 square and have d_1 x ... x d_n x 2 shape
125
+ coords = 2 * coords - 1
126
+ coords = coords @ self.positional_encoding_gaussian_matrix
127
+ coords = 2 * np.pi * coords
128
+ # outputs d_1 x ... x d_n x C shape
129
+ return torch.cat([torch.sin(coords), torch.cos(coords)], dim=-1)
130
+
131
+ def forward(self, size: Tuple[int, int]) -> torch.Tensor:
132
+ """Generate positional encoding for a grid of the specified size."""
133
+ h, w = size
134
+ device: Any = self.positional_encoding_gaussian_matrix.device
135
+ grid = torch.ones((h, w), device=device, dtype=torch.float32)
136
+ y_embed = grid.cumsum(dim=0) - 0.5
137
+ x_embed = grid.cumsum(dim=1) - 0.5
138
+ y_embed = y_embed / h
139
+ x_embed = x_embed / w
140
+
141
+ pe = self._pe_encoding(torch.stack([x_embed, y_embed], dim=-1))
142
+ return pe.permute(2, 0, 1) # C x H x W
143
+
144
+ def forward_with_coords(
145
+ self, coords_input: torch.Tensor, image_size: Tuple[int, int]
146
+ ) -> torch.Tensor:
147
+ """Positionally encode points that are not normalized to [0,1]."""
148
+ coords = coords_input.clone()
149
+ coords[:, :, 0] = coords[:, :, 0] / image_size[1]
150
+ coords[:, :, 1] = coords[:, :, 1] / image_size[0]
151
+ return self._pe_encoding(coords.to(torch.float)) # B x N x C
152
+
153
+ def get_sinusoid_encoding_table(n_position, d_hid):
154
+ ''' Sinusoid position encoding table '''
155
+ # TODO: make it with torch instead of numpy
156
+ def get_position_angle_vec(position):
157
+ return [position / np.power(10000, 2 * (hid_j // 2) / d_hid) for hid_j in range(d_hid)]
158
+
159
+ sinusoid_table = np.array([get_position_angle_vec(pos_i) for pos_i in range(n_position)])
160
+ sinusoid_table[:, 0::2] = np.sin(sinusoid_table[:, 0::2]) # dim 2i
161
+ sinusoid_table[:, 1::2] = np.cos(sinusoid_table[:, 1::2]) # dim 2i+1
162
+
163
+ return torch.tensor(sinusoid_table,dtype=torch.float, requires_grad=False).unsqueeze(0)
164
+
165
+ # main class
166
+
167
+ class DeformableAttention2D(nn.Module):
168
+ def __init__(
169
+ self,
170
+ *,
171
+ dim,
172
+ dim_head = 64,
173
+ heads = 8,
174
+ dropout = 0.,
175
+ downsample_factor = 1,
176
+ offset_scale = None,
177
+ offset_groups = None,
178
+ offset_kernel_size = 1,
179
+ group_queries = True,
180
+ group_key_values = True
181
+ ):
182
+ super().__init__()
183
+ offset_scale = default(offset_scale, downsample_factor)
184
+ assert offset_kernel_size >= downsample_factor, 'offset kernel size must be greater than or equal to the downsample factor'
185
+ assert divisible_by(offset_kernel_size - downsample_factor, 2)
186
+
187
+ offset_groups = default(offset_groups, heads)
188
+ assert divisible_by(heads, offset_groups)
189
+
190
+ inner_dim = dim_head * heads
191
+ self.scale = dim_head ** -0.5
192
+ self.heads = heads
193
+ self.offset_groups = offset_groups
194
+
195
+ offset_dims = inner_dim // offset_groups
196
+
197
+ self.downsample_factor = downsample_factor
198
+
199
+ self.to_offsets = nn.Sequential(
200
+ nn.Conv1d(offset_dims, offset_dims, kernel_size=1, groups = offset_dims, stride = 1, padding = 0),
201
+ nn.GELU(),
202
+ nn.Conv1d(offset_dims, 2, 1, bias = False),
203
+ # Rearrange('b 1 n -> b n'),
204
+ nn.Tanh(),
205
+ Scale(4/6)
206
+ )
207
+
208
+ self.rel_pos_bias = CPB(dim // 4, offset_groups = offset_groups, heads = heads, depth = 2)
209
+
210
+ self.dropout = nn.Dropout(dropout)
211
+ self.to_q = nn.Conv1d(dim, inner_dim, 1, groups = offset_groups if group_queries else 1, bias = False)
212
+ self.to_k = nn.Conv2d(dim, inner_dim, 1, groups = offset_groups if group_key_values else 1, bias = False)
213
+ self.to_v = nn.Conv2d(dim, inner_dim, 1, groups = offset_groups if group_key_values else 1, bias = False)
214
+ self.to_out = nn.Conv1d(inner_dim, dim, 1)
215
+
216
+ self.cross_attn = nn.MultiheadAttention(embed_dim=dim, num_heads=heads, batch_first=True)
217
+ # for pose
218
+ self.pe_layer = PositionEmbeddingRandom(dim//2)
219
+ # for rgb
220
+ self.pos_embed = get_sinusoid_encoding_table(4*4, dim)
221
+
222
+ def forward(self, pose_feat, rgb_feat, pose_init, return_vgrid = False):
223
+ """
224
+ b - batch
225
+ h - heads
226
+ x - height
227
+ y - width
228
+ d - dimension
229
+ g - offset groups
230
+ """
231
+
232
+ heads, b, h, w, downsample_factor, device = self.heads, rgb_feat.shape[0], *rgb_feat.shape[-2:], self.downsample_factor, rgb_feat.device
233
+
234
+ # queries
235
+ # pose_feat: bt, c, n
236
+ pose_feat_cross = rearrange(pose_feat, 'b d n -> b n d')
237
+ rgb_feat_cross = rearrange(rgb_feat, 'b d h w -> b (h w) d')
238
+ pose_init_cross = rearrange(pose_init, 'b d n -> b n d')
239
+
240
+ point_embedding = self.pe_layer._pe_encoding(pose_init_cross.detach())
241
+
242
+ kv = rgb_feat_cross + self.pos_embed.expand(b, -1, -1).type_as(rgb_feat_cross).to(rgb_feat_cross.device).clone().detach()
243
+
244
+ pose_feat_cross, _ = self.cross_attn(pose_feat_cross + point_embedding,
245
+ kv, kv)
246
+
247
+ pose_feat_cross = rearrange(pose_feat_cross, 'b n d -> b d n')
248
+
249
+ q = self.to_q(pose_feat + pose_feat_cross)
250
+
251
+ # calculate offsets - offset MLP shared across all groups
252
+
253
+ group_1d = lambda t: rearrange(t, 'b (g d) n -> (b g) d n', g = self.offset_groups)
254
+ group_2d = lambda t: rearrange(t, 'b (g d) ... -> (b g) d ...', g = self.offset_groups)
255
+
256
+ grouped_queries = group_1d(q)
257
+
258
+ offsets = self.to_offsets(grouped_queries)
259
+
260
+ # calculate grid + offsets
261
+
262
+ # pose_init --> [0, 1]
263
+ grid = pose_init[:,None].repeat(1, self.offset_groups, 1, 1) * 2 - 1
264
+ # grid --> [-1, 1]
265
+ grid = grid.reshape(-1, *pose_init.shape[-2:])
266
+
267
+ vgrid = grid + offsets
268
+ vgrid_scaled = reshape_grid_1d(vgrid)[:,None]
269
+ # vgrid_scaled = normalize_grid(vgrid)
270
+
271
+ kv_feats = F.grid_sample(
272
+ group_2d(rgb_feat),
273
+ vgrid_scaled,
274
+ mode = 'bilinear', padding_mode = 'zeros', align_corners = False)
275
+
276
+ kv_feats = rearrange(kv_feats, '(b g) d ... -> b (g d) ...', b = b)
277
+
278
+ # derive key / values
279
+ k, v = self.to_k(kv_feats), self.to_v(kv_feats)
280
+
281
+ # scale queries
282
+ q = q * self.scale
283
+
284
+ # split out heads
285
+ q, k, v = map(lambda t: rearrange(t, 'b (h d) ... -> b h (...) d', h = heads), (q, k, v))
286
+
287
+ # query / key similarity
288
+ sim = einsum('b h i d, b h j d -> b h i j', q, k)
289
+
290
+ # relative positional bias
291
+ grid = reshape_grid_1d(grid)[:,None]
292
+
293
+ rel_pos_bias = self.rel_pos_bias(grid, vgrid_scaled)
294
+ sim = sim + rel_pos_bias
295
+
296
+ # numerical stability
297
+ sim = sim - sim.amax(dim = -1, keepdim = True).detach()
298
+
299
+ # attention
300
+ attn = sim.softmax(dim = -1)
301
+ attn = self.dropout(attn)
302
+
303
+ # aggregate and combine heads
304
+ out = einsum('b h i j, b h j d -> b h i d', attn, v)
305
+ out = rearrange(out, 'b h n d -> b (h d) n')
306
+ out = self.to_out(out)
307
+
308
+ if return_vgrid:
309
+ return out, vgrid
310
+
311
+ return out
Uni_Sign/models.py ADDED
@@ -0,0 +1,415 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from torch import Tensor
2
+ import torch
3
+ from torch import nn
4
+ import torch.utils.checkpoint
5
+ import contextlib
6
+ import torchvision
7
+ from einops import rearrange
8
+
9
+ import math
10
+ from Uni_Sign.stgcn_layers import Graph, get_stgcn_chain
11
+ from Uni_Sign.deformable_attention_2d import DeformableAttention2D
12
+ from transformers import MT5ForConditionalGeneration, T5Tokenizer, MT5Config
13
+ import warnings
14
+
15
+ mt5_path = r"./Uni_Sign/unisign_model"
16
+
17
+ def _no_grad_trunc_normal_(tensor, mean, std, a, b):
18
+ # Cut & paste from PyTorch official master until it's in a few official releases - RW
19
+ # Method based on https://people.sc.fsu.edu/~jburkardt/presentations/truncated_normal.pdf
20
+ def norm_cdf(x):
21
+ # Computes standard normal cumulative distribution function
22
+ return (1. + math.erf(x / math.sqrt(2.))) / 2.
23
+
24
+ if (mean < a - 2 * std) or (mean > b + 2 * std):
25
+ warnings.warn("mean is more than 2 std from [a, b] in nn.init.trunc_normal_. "
26
+ "The distribution of values may be incorrect.",
27
+ stacklevel=2)
28
+
29
+ with torch.no_grad():
30
+ # Values are generated by using a truncated uniform distribution and
31
+ # then using the inverse CDF for the normal distribution.
32
+ # Get upper and lower cdf values
33
+ l = norm_cdf((a - mean) / std)
34
+ u = norm_cdf((b - mean) / std)
35
+
36
+ # Uniformly fill tensor with values from [l, u], then translate to
37
+ # [2l-1, 2u-1].
38
+ tensor.uniform_(2 * l - 1, 2 * u - 1)
39
+
40
+ # Use inverse cdf transform for normal distribution to get truncated
41
+ # standard normal
42
+ tensor.erfinv_()
43
+
44
+ # Transform to proper mean, std
45
+ tensor.mul_(std * math.sqrt(2.))
46
+ tensor.add_(mean)
47
+
48
+ # Clamp to ensure it's in the proper range
49
+ tensor.clamp_(min=a, max=b)
50
+ return tensor
51
+
52
+
53
+ def trunc_normal_(tensor, mean=0., std=1., a=-2., b=2.):
54
+ # type: (Tensor, float, float, float, float) -> Tensor
55
+ r"""Fills the input Tensor with values drawn from a truncated
56
+ normal distribution. The values are effectively drawn from the
57
+ normal distribution :math:`\mathcal{N}(\text{mean}, \text{std}^2)`
58
+ with values outside :math:`[a, b]` redrawn until they are within
59
+ the bounds. The method used for generating the random values works
60
+ best when :math:`a \leq \text{mean} \leq b`.
61
+ Args:
62
+ tensor: an n-dimensional `torch.Tensor`
63
+ mean: the mean of the normal distribution
64
+ std: the standard deviation of the normal distribution
65
+ a: the minimum cutoff value
66
+ b: the maximum cutoff value
67
+ Examples:
68
+ >>> w = torch.empty(3, 5)
69
+ >>> nn.init.trunc_normal_(w)
70
+ """
71
+ return _no_grad_trunc_normal_(tensor, mean, std, a, b)
72
+
73
+ class Uni_Sign(nn.Module):
74
+ def __init__(self, args):
75
+ super(Uni_Sign, self).__init__()
76
+ self.args = args
77
+
78
+ self.modes = ['body', 'left', 'right', 'face_all']
79
+
80
+ self.graph, A = {}, []
81
+ # project (x,y,score) to hidden dim
82
+ hidden_dim = args.hidden_dim
83
+ self.proj_linear = nn.ModuleDict()
84
+ for mode in self.modes:
85
+ graph_layout = f'{args.layout}_ytasl_{mode}' if self.args.dataset in ["YTASL", "Isharah"] else f'{args.layout}_{mode}'
86
+ self.graph[mode] = Graph(layout=graph_layout, strategy='distance', max_hop=1)
87
+ A.append(torch.tensor(self.graph[mode].A, dtype=torch.float32, requires_grad=False))
88
+ self.proj_linear[mode] = nn.Linear(3, 64)
89
+
90
+ self.gcn_modules = nn.ModuleDict()
91
+ self.fusion_gcn_modules = nn.ModuleDict()
92
+ spatial_kernel_size = A[0].size(0)
93
+ for index, mode in enumerate(self.modes):
94
+ self.gcn_modules[mode], final_dim = get_stgcn_chain(64, 'spatial', (1, spatial_kernel_size), A[index].clone(), adaptive=not self.args.no_adaptive_gcn)
95
+ self.fusion_gcn_modules[mode], _ = get_stgcn_chain(final_dim, 'temporal', (5, spatial_kernel_size), A[index].clone(), adaptive=not self.args.no_adaptive_gcn)
96
+
97
+ self.gcn_modules['left'] = self.gcn_modules['right']
98
+ self.fusion_gcn_modules['left'] = self.fusion_gcn_modules['right']
99
+ self.proj_linear['left'] = self.proj_linear['right']
100
+
101
+ self.part_para = nn.Parameter(torch.zeros(hidden_dim*len(self.modes)))
102
+ self.pose_proj = nn.Linear(256*4, 768)
103
+
104
+ self.apply(self._init_weights)
105
+
106
+ if self.args.dataset == "Isharah":
107
+ self.lang = 'Arabic'
108
+ elif "CSL" in self.args.dataset:
109
+ self.lang = 'Chinese'
110
+ else:
111
+ self.lang = 'English'
112
+
113
+ if self.args.rgb_support:
114
+ self.rgb_support_backbone = torch.nn.Sequential(*list(torchvision.models.efficientnet_b0(pretrained=True).children())[:-2])
115
+ self.rgb_proj = nn.Conv2d(1280, hidden_dim, kernel_size=1)
116
+
117
+ self.fusion_pose_rgb_linear = nn.Linear(hidden_dim, hidden_dim)
118
+
119
+ # PGF
120
+ self.fusion_pose_rgb_DA = DeformableAttention2D(
121
+ dim = hidden_dim, # feature dimensions
122
+ dim_head = 32, # dimension per head
123
+ heads = 8, # attention heads
124
+ dropout = 0., # dropout
125
+ downsample_factor = 1, # downsample factor (r in paper)
126
+ offset_scale = None, # scale of offset, maximum offset
127
+ offset_groups = None, # number of offset groups, should be multiple of heads
128
+ offset_kernel_size = 1, # offset kernel size
129
+ )
130
+
131
+ self.fusion_gate = nn.Sequential(nn.Conv1d(hidden_dim*2, hidden_dim, 1),
132
+ nn.GELU(),
133
+ nn.Conv1d(hidden_dim, 1, 1),
134
+ nn.Tanh(),
135
+ nn.ReLU(),
136
+ )
137
+
138
+ for layer in self.fusion_gate:
139
+ try:
140
+ if isinDataLoaderance(layer, nn.Conv1d):
141
+ nn.init.constant_(layer.weight, 0)
142
+ nn.init.constant_(layer.bias, 0)
143
+ except:
144
+ print("NOT IMPLEMENTED...")
145
+
146
+ # Načte pouze strukturu architektury z config.json
147
+ mt5_config = MT5Config.from_pretrained(mt5_path)
148
+ # Vytvoří model s prázdnými vahami, které hned v dalším kroku přepíšeme
149
+ self.mt5_model = MT5ForConditionalGeneration(mt5_config)
150
+
151
+ self.mt5_tokenizer = T5Tokenizer.from_pretrained(mt5_path, legacy=False)
152
+
153
+ self.n_registers = args.n_registers
154
+ self.register_position = args.register_position
155
+ self.d_model = self.mt5_model.config.d_model # should be 768
156
+
157
+ if self.n_registers > 0:
158
+ self.register_tokens = nn.Parameter(torch.zeros(self.n_registers, self.d_model))
159
+ # init like other embeddings
160
+ trunc_normal_(self.register_tokens, std=0.02)
161
+ else:
162
+ self.register_tokens = None
163
+
164
+
165
+ def _init_weights(self, m):
166
+ if isinstance(m, nn.Linear):
167
+ trunc_normal_(m.weight, std=.02)
168
+ if isinstance(m, nn.Linear) and m.bias is not None:
169
+ nn.init.constant_(m.bias, 0)
170
+ elif isinstance(m, nn.LayerNorm):
171
+ nn.init.constant_(m.bias, 0)
172
+ nn.init.constant_(m.weight, 1.0)
173
+
174
+ def maybe_autocast(self, dtype=torch.float32):
175
+ # if on cpu, don't use autocast
176
+ # if on gpu, use autocast with dtype if provided, otherwise use torch.float16
177
+ # enable_autocast = self.device != torch.device("cpu")
178
+ enable_autocast = True
179
+
180
+ if enable_autocast:
181
+ return torch.cuda.amp.autocast(dtype=dtype)
182
+ else:
183
+ return contextlib.nullcontext()
184
+
185
+ def gather_feat_pose_rgb(self, gcn_feat, rgb_feat, indices, rgb_len, pose_init):
186
+ b, c, T, n = gcn_feat.shape
187
+ assert rgb_feat.shape[0] == indices.shape[0]
188
+ rgb_feat = self.rgb_proj(rgb_feat)
189
+
190
+ assert len(rgb_len) == b
191
+ start = 0
192
+ for batch in range(b):
193
+ index = indices[start:start + rgb_len[batch]].to(torch.long)
194
+ # ignore some invalid rgb clip
195
+ if rgb_len[batch] == 1 and -1 in index:
196
+ start = start + rgb_len[batch]
197
+ continue
198
+
199
+ # index selection
200
+ gcn_feat_selected = gcn_feat[batch, :, index]
201
+ rgb_feat_selected = rgb_feat[start:start + rgb_len[batch]]
202
+ pose_init_selected = pose_init[start:start + rgb_len[batch]]
203
+
204
+ gcn_feat_selected = rearrange(gcn_feat_selected, 'c t n -> t c n')
205
+ pose_init_selected = rearrange(pose_init_selected, 't n c -> t c n')
206
+
207
+ # PGF forward
208
+ with self.maybe_autocast():
209
+ fused_transposed = self.fusion_pose_rgb_DA(pose_feat=gcn_feat_selected,
210
+ rgb_feat=rgb_feat_selected,
211
+ pose_init=pose_init_selected, )
212
+
213
+ fused_transposed = fused_transposed.to(gcn_feat.dtype)
214
+ gate_feature = torch.concat([fused_transposed, gcn_feat_selected,], dim=-2)
215
+ gate_score = self.fusion_gate(gate_feature)
216
+ fused_transposed_post = (gate_score) * fused_transposed + (1 - gate_score) * gcn_feat_selected
217
+
218
+ gcn_feat = gcn_feat.clone()
219
+ fused_transposed_post = rearrange(fused_transposed_post, 't c n -> c t n')
220
+
221
+ # replace gcn feature
222
+ gcn_feat[batch, :, index] = fused_transposed_post
223
+ start = start + rgb_len[batch]
224
+
225
+ assert start == rgb_feat.shape[0]
226
+ return gcn_feat
227
+
228
+ def forward(self, src_input, tgt_input):
229
+ # RGB branch forward
230
+ if self.args.rgb_support:
231
+ rgb_support_dict = {}
232
+ for index_key, rgb_key in zip(['left_sampled_indices', 'right_sampled_indices'], ['left_hands', 'right_hands']):
233
+ rgb_feat = self.rgb_support_backbone(src_input[rgb_key])
234
+
235
+ rgb_support_dict[index_key] = src_input[index_key]
236
+ rgb_support_dict[rgb_key] = rgb_feat
237
+
238
+ # Pose branch forward
239
+ features = []
240
+
241
+ body_feat = None
242
+ for part in self.modes:
243
+ # project position to hidden dim
244
+ proj_feat = self.proj_linear[part](src_input[part]).permute(0,3,1,2) #B,C,T,V
245
+ # spatial gcn forward
246
+ gcn_feat = self.gcn_modules[part](proj_feat)
247
+ if part == 'body':
248
+ body_feat = gcn_feat
249
+
250
+ else:
251
+ assert not body_feat is None
252
+ if part == 'left':
253
+ # Pose RGB fusion
254
+ if self.args.rgb_support:
255
+ gcn_feat = self.gather_feat_pose_rgb(gcn_feat,
256
+ rgb_support_dict[f'{part}_hands'],
257
+ rgb_support_dict[f'{part}_sampled_indices'],
258
+ src_input[f'{part}_rgb_len'],
259
+ src_input[f'{part}_skeletons_norm'],
260
+ )
261
+
262
+ gcn_feat = gcn_feat + body_feat[..., -2][...,None].detach()
263
+
264
+ elif part == 'right':
265
+ # Pose RGB fusion
266
+ if self.args.rgb_support:
267
+ gcn_feat = self.gather_feat_pose_rgb(gcn_feat,
268
+ rgb_support_dict[f'{part}_hands'],
269
+ rgb_support_dict[f'{part}_sampled_indices'],
270
+ src_input[f'{part}_rgb_len'],
271
+ src_input[f'{part}_skeletons_norm'],
272
+ )
273
+
274
+ gcn_feat = gcn_feat + body_feat[..., -1][...,None].detach()
275
+
276
+ elif part == 'face_all':
277
+ gcn_feat = gcn_feat + body_feat[..., 0][...,None].detach()
278
+
279
+ else:
280
+ raise NotImplementedError
281
+
282
+ # temporal gcn forward
283
+ gcn_feat = self.fusion_gcn_modules[part](gcn_feat) #B,C,T,V
284
+ pool_feat = gcn_feat.mean(-1).transpose(1,2) #B,T,C
285
+ features.append(pool_feat)
286
+
287
+ # concat sub-pose feature across token dimension
288
+ inputs_embeds = torch.cat(features, dim=-1) + self.part_para
289
+ inputs_embeds = self.pose_proj(inputs_embeds)
290
+
291
+ prefix_token = self.mt5_tokenizer(
292
+ [f"Translate sign language video to {self.lang}: "] * len(tgt_input["gt_sentence"]),
293
+ padding="longest",
294
+ truncation=True,
295
+ return_tensors="pt",
296
+ ).to(inputs_embeds.device)
297
+
298
+ prefix_embeds = self.mt5_model.encoder.embed_tokens(prefix_token['input_ids'])
299
+
300
+ if self.n_registers > 0:
301
+ B = inputs_embeds.size(0)
302
+
303
+ # expand registers for batch
304
+ register_embeds = self.register_tokens.unsqueeze(0).expand(B, -1, -1)
305
+ # shape: (B, 4, 768)
306
+
307
+ register_mask = torch.ones((B, self.n_registers), device=inputs_embeds.device,dtype=prefix_token['attention_mask'].dtype)
308
+
309
+ if self.register_position == 'before_all':
310
+ # prepend order: [registers | prefix | pose_tokens]
311
+ inputs_embeds = torch.cat([register_embeds, prefix_embeds, inputs_embeds], dim=1)
312
+ attention_mask = torch.cat([register_mask, prefix_token['attention_mask'], src_input['attention_mask']], dim=1)
313
+
314
+ elif self.register_position == 'after_prefix' or self.register_position == 'before_pose':
315
+ # prepend order: [prefix | registers | pose_tokens]
316
+ inputs_embeds = torch.cat([prefix_embeds, register_embeds, inputs_embeds], dim=1)
317
+ attention_mask = torch.cat([prefix_token['attention_mask'], register_mask, src_input['attention_mask']], dim=1)
318
+
319
+ elif self.register_position == 'after_valid_pose':
320
+ # prepend order: [prefix | valid_pose_tokens | registers | padded_pose_tokens]
321
+ inputs_list = []
322
+ mask_list = []
323
+
324
+ for b in range(B):
325
+ valid_len = int(src_input['attention_mask'][b].sum().item())
326
+
327
+ pose_valid = inputs_embeds[b, :valid_len]
328
+ pose_pad = inputs_embeds[b, valid_len:]
329
+
330
+ emb = torch.cat(
331
+ [prefix_embeds[b],
332
+ pose_valid,
333
+ register_embeds[b],
334
+ pose_pad],
335
+ dim=0
336
+ )
337
+
338
+ m = torch.cat(
339
+ [prefix_token['attention_mask'][b],
340
+ src_input['attention_mask'][b, :valid_len],
341
+ register_mask[b],
342
+ src_input['attention_mask'][b, valid_len:]],
343
+ dim=0
344
+ )
345
+
346
+ inputs_list.append(emb)
347
+ mask_list.append(m)
348
+
349
+ inputs_embeds = torch.stack(inputs_list, dim=0)
350
+ attention_mask = torch.stack(mask_list, dim=0)
351
+
352
+ elif self.register_position == 'after_all':
353
+ # prepend order: [prefix | pose_tokens | registers]
354
+ inputs_embeds = torch.cat([prefix_embeds, inputs_embeds, register_embeds], dim=1)
355
+ attention_mask = torch.cat([prefix_token['attention_mask'], src_input['attention_mask'], register_mask], dim=1)
356
+
357
+ else:
358
+ # prepend order: [prefix | pose_tokens]
359
+ inputs_embeds = torch.cat([prefix_embeds, inputs_embeds], dim=1)
360
+ attention_mask = torch.cat([prefix_token['attention_mask'], src_input['attention_mask']], dim=1)
361
+
362
+ tgt_input_tokenizer = self.mt5_tokenizer(tgt_input['gt_sentence'],
363
+ return_tensors="pt",
364
+ padding=True,
365
+ truncation=True,
366
+ max_length=50)
367
+
368
+ labels = tgt_input_tokenizer['input_ids']
369
+ labels[labels == self.mt5_tokenizer.pad_token_id] = -100
370
+
371
+ out = self.mt5_model(inputs_embeds = inputs_embeds,
372
+ attention_mask = attention_mask,
373
+ labels = labels.to(inputs_embeds.device),
374
+ return_dict = True,
375
+ )
376
+
377
+ label = labels.reshape(-1)
378
+ out_logits = out['logits']
379
+ logits = out_logits.reshape(-1,out_logits.shape[-1])
380
+ loss_fct = torch.nn.CrossEntropyLoss(label_smoothing=self.args.label_smoothing, ignore_index=-100)
381
+ loss = loss_fct(logits, label.to(out_logits.device, non_blocking=True))
382
+
383
+ stack_out = {
384
+ # use for inference
385
+ 'inputs_embeds':inputs_embeds,
386
+ 'attention_mask':attention_mask,
387
+ 'loss':loss,
388
+ }
389
+
390
+ return stack_out
391
+
392
+ @torch.no_grad()
393
+ def generate(self,pre_compute_item,max_new_tokens,num_beams):
394
+ inputs_embeds = pre_compute_item['inputs_embeds']
395
+ attention_mask = pre_compute_item['attention_mask']
396
+
397
+ out = self.mt5_model.generate(inputs_embeds = inputs_embeds,
398
+ attention_mask = attention_mask,
399
+ max_new_tokens=max_new_tokens,
400
+ num_beams = num_beams,
401
+ )
402
+
403
+ return out
404
+
405
+ def get_requires_grad_dict(model):
406
+ param_requires_grad = {name: True for name, param in model.named_parameters()}
407
+ param_requires_grad_right = {}
408
+ for key in param_requires_grad.keys():
409
+ if 'left' in key:
410
+ param_requires_grad_right[key.replace("left", 'right')] = param_requires_grad[key]
411
+ param_requires_grad = {**param_requires_grad,
412
+ **param_requires_grad_right}
413
+ params_to_update = {k: v for k, v in model.state_dict().items() if param_requires_grad.get(k, True)}
414
+
415
+ return params_to_update
Uni_Sign/normalization.py ADDED
@@ -0,0 +1,117 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import numpy as np
2
+ from typing import Tuple
3
+
4
+
5
+ def get_keypoints(joints, landmarks_name):
6
+ frames_keypoints = joints[landmarks_name]
7
+ frames_keypoints = frames_keypoints[:, :, :2]
8
+
9
+ return frames_keypoints
10
+
11
+
12
+ def output_keypoints(joints, valid_frames, frames_keypoints):
13
+ frames_names = np.array(list(joints.keys()))[valid_frames]
14
+ frames_keypoints = frames_keypoints[valid_frames]
15
+
16
+ return dict(zip(frames_names, frames_keypoints))
17
+
18
+
19
+ def safe_divide(a, b):
20
+ return np.divide(a, b, out=np.zeros_like(a), where=b != 0)
21
+
22
+
23
+ def local_keypoint_normalization(joints: dict, landmarks: str, select_idx: list = [], padding: float = 0.1) -> dict:
24
+ frames_keypoints = get_keypoints(joints, landmarks)
25
+
26
+ if select_idx:
27
+ frames_keypoints = frames_keypoints[:, select_idx, :]
28
+
29
+ # move to origin
30
+ xmin = np.min(frames_keypoints[:, :, 0], axis=1)
31
+ ymin = np.min(frames_keypoints[:, :, 1], axis=1)
32
+
33
+ frames_keypoints[:, :, 0] -= xmin[:, np.newaxis]
34
+ frames_keypoints[:, :, 1] -= ymin[:, np.newaxis]
35
+
36
+ # pad to square
37
+ xmax = np.max(frames_keypoints[:, :, 0], axis=1)
38
+ ymax = np.max(frames_keypoints[:, :, 1], axis=1)
39
+
40
+ dif_full = np.abs(xmax - ymax)
41
+ dif = np.floor(dif_full / 2)
42
+
43
+ for i in range(len(dif)):
44
+ if xmax[i] > ymax[i]:
45
+ ymax[i] += dif_full[i]
46
+ frames_keypoints[i, :, 1] += dif[i]
47
+ else:
48
+ xmax[i] += dif_full[i]
49
+ frames_keypoints[i, :, 0] += dif[i]
50
+
51
+ # add padding to all sides
52
+ side_size = np.max([xmax, ymax], axis=0)
53
+ padding = side_size * padding
54
+
55
+ frames_keypoints += padding[:, np.newaxis, np.newaxis]
56
+ xmax += padding * 2
57
+ ymax += padding * 2
58
+
59
+ # normalize to [-1, 1]
60
+ frames_keypoints = safe_divide(frames_keypoints, xmax[:, np.newaxis, np.newaxis])
61
+ # frames_keypoints /= xmax[:, np.newaxis, np.newaxis]
62
+ frames_keypoints = frames_keypoints * 2 - 1
63
+
64
+ return frames_keypoints
65
+
66
+
67
+ def global_keypoint_normalization(
68
+ joints: dict,
69
+ landmarks: str,
70
+ add_landmarks_names: list,
71
+ face_select_idx: list = [],
72
+ sign_area_size: tuple = (1.5, 1.5),
73
+ l_shoulder_idx: int = 11,
74
+ r_shoulder_idx: int = 12) -> Tuple[dict, dict]:
75
+ frames_keypoints = get_keypoints(joints, landmarks)
76
+
77
+ # get distance between right and left shoulder
78
+ l_shoulder_points = frames_keypoints[:, l_shoulder_idx, :]
79
+ r_shoulder_points = frames_keypoints[:, r_shoulder_idx, :]
80
+ distance = np.sqrt((l_shoulder_points[:, 0] - r_shoulder_points[:, 0]) ** 2 + (
81
+ l_shoulder_points[:, 1] - r_shoulder_points[:, 1]) ** 2)
82
+
83
+ # get center point between shoulders
84
+ center_x = np.abs(l_shoulder_points[:, 0] - r_shoulder_points[:, 0]) / 2 + np.min(
85
+ [l_shoulder_points[:, 0], r_shoulder_points[:, 0]], 0)
86
+ center_y = np.abs(l_shoulder_points[:, 1] - r_shoulder_points[:, 1]) / 2 + np.min(
87
+ [l_shoulder_points[:, 1], r_shoulder_points[:, 1]], 0)
88
+ sign_area_size = np.array(sign_area_size) * distance[:, np.newaxis]
89
+
90
+ # normalize
91
+ frames_keypoints[:, :, 0] -= center_x[:, np.newaxis]
92
+ frames_keypoints[:, :, 1] -= center_y[:, np.newaxis]
93
+
94
+ # frames_keypoints[:, :, 0] /= sign_area_size[:, 1, np.newaxis]
95
+ # frames_keypoints[:, :, 1] /= sign_area_size[:, 0, np.newaxis]
96
+ frames_keypoints[:, :, 0] = safe_divide(frames_keypoints[:, :, 0], sign_area_size[:, 0, np.newaxis])
97
+ frames_keypoints[:, :, 1] = safe_divide(frames_keypoints[:, :, 1], sign_area_size[:, 1, np.newaxis])
98
+
99
+ # normalize additional landmarks
100
+ add_landmarks = {}
101
+ for add_landmarks_name in add_landmarks_names:
102
+ add_frames_keypoints = get_keypoints(joints, add_landmarks_name)
103
+
104
+ if face_select_idx and add_landmarks_name == "face_landmarks":
105
+ add_frames_keypoints = add_frames_keypoints[:, face_select_idx, :]
106
+
107
+ add_frames_keypoints[:, :, 0] -= center_x[:, np.newaxis]
108
+ add_frames_keypoints[:, :, 1] -= center_y[:, np.newaxis]
109
+
110
+ # add_frames_keypoints[:, :, 0] /= sign_area_size[:, 1, np.newaxis]
111
+ # add_frames_keypoints[:, :, 1] /= sign_area_size[:, 0, np.newaxis]
112
+ add_frames_keypoints[:, :, 0] = safe_divide(add_frames_keypoints[:, :, 0], sign_area_size[:, 0, np.newaxis])
113
+ add_frames_keypoints[:, :, 1] = safe_divide(add_frames_keypoints[:, :, 1], sign_area_size[:, 1, np.newaxis])
114
+
115
+ add_landmarks[add_landmarks_name] = add_frames_keypoints
116
+
117
+ return frames_keypoints, add_landmarks
Uni_Sign/stgcn_layers/.DS_Store ADDED
Binary file (6.15 kB). View file
 
Uni_Sign/stgcn_layers/__init__.py ADDED
@@ -0,0 +1,2 @@
 
 
 
1
+ from .gcn_utils import Graph
2
+ from .stgcn_block import get_stgcn_chain
Uni_Sign/stgcn_layers/__pycache__/__init__.cpython-311.pyc ADDED
Binary file (294 Bytes). View file
 
Uni_Sign/stgcn_layers/__pycache__/__init__.cpython-39.pyc ADDED
Binary file (246 Bytes). View file
 
Uni_Sign/stgcn_layers/__pycache__/gcn_utils.cpython-311.pyc ADDED
Binary file (13.3 kB). View file
 
Uni_Sign/stgcn_layers/__pycache__/gcn_utils.cpython-39.pyc ADDED
Binary file (7.52 kB). View file
 
Uni_Sign/stgcn_layers/__pycache__/stgcn_block.cpython-311.pyc ADDED
Binary file (6.54 kB). View file
 
Uni_Sign/stgcn_layers/__pycache__/stgcn_block.cpython-39.pyc ADDED
Binary file (3.44 kB). View file
 
Uni_Sign/stgcn_layers/gcn_utils.py ADDED
@@ -0,0 +1,350 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import numpy as np
3
+ import torch.nn as nn
4
+ import pdb
5
+ import math
6
+ import copy
7
+
8
+
9
+ class Graph:
10
+ """The Graph to model the skeletons extracted by the openpose
11
+
12
+ Args:
13
+ strategy (string): must be one of the follow candidates
14
+ - uniform: Uniform Labeling
15
+ - distance: Distance Partitioning
16
+ - spatial: Spatial Configuration
17
+ For more information, please refer to the section 'Partition Strategies'
18
+ in our paper (https://arxiv.org/abs/1801.07455).
19
+
20
+ layout (string): must be one of the follow candidates
21
+ - openpose: Is consists of 18 joints. For more information, please
22
+ refer to https://github.com/CMU-Perceptual-Computing-Lab/openpose#output
23
+ - ntu-rgb+d: Is consists of 25 joints. For more information, please
24
+ refer to https://github.com/shahroudy/NTURGB-D
25
+
26
+ max_hop (int): the maximal distance between two connected nodes
27
+ dilation (int): controls the spacing between the kernel points
28
+
29
+ """
30
+
31
+ def __init__(self, layout='custom', strategy='uniform', max_hop=1, dilation=1):
32
+ self.max_hop = max_hop
33
+ self.dilation = dilation
34
+
35
+ self.get_edge(layout)
36
+ self.hop_dis = get_hop_distance(self.num_node, self.edge, max_hop=max_hop)
37
+ self.get_adjacency(strategy)
38
+
39
+ def __str__(self):
40
+ return self.A
41
+
42
+ def get_edge(self, layout):
43
+ # 'body', 'left', 'right', 'mouth', 'face'
44
+ # if layout == 'custom_hand21':
45
+
46
+ if layout == 'default_left' or layout == 'default_right':
47
+ self.num_node = 21
48
+ self_link = [(i, i) for i in range(self.num_node)]
49
+ neighbor_1base = [
50
+ [0, 1],
51
+ [1, 2],
52
+ [2, 3],
53
+ [3, 4],
54
+ [0, 5],
55
+ [5, 6],
56
+ [6, 7],
57
+ [7, 8],
58
+ [0, 9],
59
+ [9, 10],
60
+ [10, 11],
61
+ [11, 12],
62
+ [0, 13],
63
+ [13, 14],
64
+ [14, 15],
65
+ [15, 16],
66
+ [0, 17],
67
+ [17, 18],
68
+ [18, 19],
69
+ [19, 20],
70
+ ]
71
+ neighbor_link = neighbor_1base
72
+ self.edge = self_link + neighbor_link
73
+ self.center = 0
74
+
75
+ elif layout == 'default_body':
76
+ self.num_node = 9
77
+ self_link = [(i, i) for i in range(self.num_node)]
78
+ neighbor_1base = [
79
+ [0, 1],
80
+ [0, 2],
81
+ [0, 3],
82
+ [0, 4],
83
+ [3, 5],
84
+ [5, 7],
85
+ [4, 6],
86
+ [6, 8],
87
+ ]
88
+ neighbor_link = neighbor_1base
89
+ self.edge = self_link + neighbor_link
90
+ self.center = 0
91
+
92
+ elif layout == 'default_face_all':
93
+ self.num_node = 9 + 8 + 1
94
+ self_link = [(i, i) for i in range(self.num_node)]
95
+ neighbor_1base = [[i, i + 1] for i in range(9 - 1)] + \
96
+ [[i, i + 1] for i in range(9, 9 + 8 - 1)] + \
97
+ [[9 + 8 - 1, 9]] + \
98
+ [[17, i] for i in range(17)]
99
+ neighbor_link = neighbor_1base
100
+ self.edge = self_link + neighbor_link
101
+ self.center = self.num_node - 1
102
+
103
+ elif layout in ['default_ytasl_left', 'default_ytasl_right', 'pruned_ytasl_left', 'pruned_ytasl_right', 'isharah_ytasl_left', 'isharah_ytasl_right']:
104
+ self.num_node = 21
105
+ self_link = [(i, i) for i in range(self.num_node)]
106
+ neighbor_1base = [
107
+ [3, 4],
108
+ [0, 5],
109
+ [17, 18],
110
+ [0, 17],
111
+ [13, 14],
112
+ [13, 17],
113
+ [18, 19],
114
+ [5, 6],
115
+ [5, 9],
116
+ [14, 15],
117
+ [0, 1],
118
+ [9, 10],
119
+ [1, 2],
120
+ [9, 13],
121
+ [10, 11],
122
+ [19, 20],
123
+ [6, 7],
124
+ [15, 16],
125
+ [2, 3],
126
+ [11, 12],
127
+ [7, 8]
128
+ ]
129
+ neighbor_link = neighbor_1base
130
+ self.edge = self_link + neighbor_link
131
+ self.center = 0
132
+
133
+ elif layout == 'default_ytasl_body':
134
+ self.num_node = 25
135
+ self_link = [(i, i) for i in range(self.num_node)]
136
+ neighbor_1base = [
137
+ [15, 21],
138
+ [16, 20],
139
+ [18, 20],
140
+ [3, 7],
141
+ [14, 16],
142
+ [11, 23],
143
+ [6, 8],
144
+ [15, 17],
145
+ [16, 22],
146
+ [4, 5],
147
+ [5, 6],
148
+ [12, 24],
149
+ [23, 24],
150
+ [0, 1],
151
+ [9, 10],
152
+ [1, 2],
153
+ [0, 4],
154
+ [11, 13],
155
+ [15, 19],
156
+ [16, 18],
157
+ [12, 14],
158
+ [17, 19],
159
+ [2, 3],
160
+ [11, 12],
161
+ [13, 15]
162
+ ]
163
+ neighbor_link = neighbor_1base
164
+ self.edge = self_link + neighbor_link
165
+ self.center = 0
166
+
167
+ elif layout == 'default_ytasl_face_all':
168
+ self.num_node = 37
169
+ self_link = [(i, i) for i in range(self.num_node)]
170
+ neighbor_1base = [
171
+ [16, 18],
172
+ [18, 13],
173
+ [13, 7],
174
+ [7, 8],
175
+ [8, 15],
176
+ [15, 24],
177
+ [24, 23],
178
+ [23, 29],
179
+ [29, 32],
180
+ [32, 16],
181
+ [5, 17],
182
+ [17, 14],
183
+ [30, 31],
184
+ [31, 21],
185
+ [11, 1],
186
+ [1, 27],
187
+ [10, 6],
188
+ [6, 0],
189
+ [0, 22],
190
+ [22, 26],
191
+ [26, 34],
192
+ [34, 4],
193
+ [4, 20],
194
+ [20, 10],
195
+ [10, 12],
196
+ [12, 2],
197
+ [2, 28],
198
+ [28, 26],
199
+ [10, 19],
200
+ [19, 3],
201
+ [3, 33],
202
+ [33, 26]
203
+ ]
204
+ neighbor_link = neighbor_1base
205
+ self.edge = self_link + neighbor_link
206
+ self.center = self.num_node - 1
207
+
208
+ elif layout in ['pruned_ytasl_body', 'isharah_ytasl_body']:
209
+ self.num_node = 9
210
+ self_link = [(i, i) for i in range(self.num_node)]
211
+ neighbor_1base = [
212
+ [0, 1],
213
+ [0, 2],
214
+ [0, 3],
215
+ [0, 4],
216
+ [3, 5],
217
+ [4, 6],
218
+ [5, 7],
219
+ [6, 8],
220
+ ]
221
+ neighbor_link = neighbor_1base
222
+ self.edge = self_link + neighbor_link
223
+ self.center = 0
224
+
225
+ elif layout == 'pruned_ytasl_face_all':
226
+ self.num_node = 18
227
+ self_link = [(i, i) for i in range(self.num_node)]
228
+ neighbor_1base = [
229
+ [5, 8],
230
+ [8, 6],
231
+ [6, 14],
232
+ [14, 12],
233
+ [3, 4],
234
+ [4, 1],
235
+ [1, 11],
236
+ [11, 10],
237
+ [3, 9],
238
+ [9, 2],
239
+ [2, 15],
240
+ [15, 10],
241
+ [7, 16],
242
+ [13, 17],
243
+ ]
244
+ neighbor_link = neighbor_1base
245
+ self.edge = self_link + neighbor_link
246
+ self.center = self.num_node - 1
247
+
248
+ elif layout == 'isharah_ytasl_face_all':
249
+ self.num_node = 19
250
+ self_link = [(i, i) for i in range(self.num_node)]
251
+ neighbor_1base = [
252
+ [0, 2],
253
+ [2, 3],
254
+ [3, 4],
255
+ [4, 10],
256
+ [10, 5],
257
+ [5, 8],
258
+ [8, 7],
259
+ [7, 9],
260
+ [9, 6],
261
+ [6, 1],
262
+ [1, 15],
263
+ [15, 18],
264
+ [18, 16],
265
+ [16, 17],
266
+ [17, 14],
267
+ [14, 13],
268
+ [13, 12],
269
+ [12, 11],
270
+ [11, 0],
271
+ ]
272
+ neighbor_link = neighbor_1base
273
+ self.edge = self_link + neighbor_link
274
+ self.center = self.num_node - 1
275
+
276
+ else:
277
+ raise NotImplementedError(f"Layout not implemented for: {layout}")
278
+
279
+ def get_adjacency(self, strategy):
280
+ valid_hop = range(0, self.max_hop + 1, self.dilation)
281
+ adjacency = np.zeros((self.num_node, self.num_node))
282
+ for hop in valid_hop:
283
+ adjacency[self.hop_dis == hop] = 1
284
+ normalize_adjacency = normalize_digraph(adjacency)
285
+
286
+ if strategy == 'uniform':
287
+ A = np.zeros((1, self.num_node, self.num_node))
288
+ A[0] = normalize_adjacency
289
+ self.A = A
290
+ elif strategy == 'distance':
291
+ A = np.zeros((len(valid_hop), self.num_node, self.num_node))
292
+ for i, hop in enumerate(valid_hop):
293
+ A[i][self.hop_dis == hop] = normalize_adjacency[self.hop_dis == hop]
294
+ self.A = A
295
+ elif strategy == 'spatial':
296
+ A = []
297
+ for hop in valid_hop:
298
+ a_root = np.zeros((self.num_node, self.num_node))
299
+ a_close = np.zeros((self.num_node, self.num_node))
300
+ a_further = np.zeros((self.num_node, self.num_node))
301
+ for i in range(self.num_node):
302
+ for j in range(self.num_node):
303
+ if self.hop_dis[j, i] == hop:
304
+ if (
305
+ self.hop_dis[j, self.center]
306
+ == self.hop_dis[i, self.center]
307
+ ):
308
+ a_root[j, i] = normalize_adjacency[j, i]
309
+ elif (
310
+ self.hop_dis[j, self.center]
311
+ > self.hop_dis[i, self.center]
312
+ ):
313
+ a_close[j, i] = normalize_adjacency[j, i]
314
+ else:
315
+ a_further[j, i] = normalize_adjacency[j, i]
316
+ if hop == 0:
317
+ A.append(a_root)
318
+ else:
319
+ A.append(a_root + a_close)
320
+ A.append(a_further)
321
+ A = np.stack(A)
322
+ self.A = A
323
+ else:
324
+ raise ValueError("Do Not Exist This Strategy")
325
+
326
+
327
+ def get_hop_distance(num_node, edge, max_hop=1):
328
+ A = np.zeros((num_node, num_node))
329
+ for i, j in edge:
330
+ A[j, i] = 1
331
+ A[i, j] = 1
332
+
333
+ # compute hop steps
334
+ hop_dis = np.zeros((num_node, num_node)) + np.inf
335
+ transfer_mat = [np.linalg.matrix_power(A, d) for d in range(max_hop + 1)]
336
+ arrive_mat = np.stack(transfer_mat) > 0
337
+ for d in range(max_hop, -1, -1):
338
+ hop_dis[arrive_mat[d]] = d
339
+ return hop_dis
340
+
341
+
342
+ def normalize_digraph(A):
343
+ Dl = np.sum(A, 0)
344
+ num_node = A.shape[0]
345
+ Dn = np.zeros((num_node, num_node))
346
+ for i in range(num_node):
347
+ if Dl[i] > 0:
348
+ Dn[i, i] = Dl[i] ** (-1)
349
+ AD = np.dot(A, Dn)
350
+ return AD
Uni_Sign/stgcn_layers/stgcn_block.py ADDED
@@ -0,0 +1,128 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import torch
2
+ import numpy as np
3
+ import torch.nn as nn
4
+ import pdb
5
+ import math
6
+ import copy
7
+
8
+ class GCN_unit(nn.Module):
9
+ def __init__(
10
+ self,
11
+ in_channels,
12
+ out_channels,
13
+ kernel_size,
14
+ A,
15
+ adaptive=True,
16
+ t_kernel_size=1,
17
+ t_stride=1,
18
+ t_padding=0,
19
+ t_dilation=1,
20
+ bias=True,
21
+ ):
22
+ super().__init__()
23
+ self.kernel_size = kernel_size
24
+ assert A.size(0) == self.kernel_size
25
+ self.conv = nn.Conv2d(
26
+ in_channels,
27
+ out_channels * kernel_size,
28
+ kernel_size=(t_kernel_size, 1),
29
+ padding=(t_padding, 0),
30
+ stride=(t_stride, 1),
31
+ dilation=(t_dilation, 1),
32
+ bias=bias,
33
+ )
34
+ self.adaptive = adaptive
35
+ # print(self.adaptive)
36
+ if self.adaptive:
37
+ self.A = nn.Parameter(A.clone())
38
+ else:
39
+ self.register_buffer('A', A)
40
+ self.bn = nn.BatchNorm2d(out_channels)
41
+ self.relu = nn.ReLU(inplace=True)
42
+
43
+ def forward(self, x, len_x):
44
+ x = self.conv(x)
45
+
46
+ n, kc, t, v = x.size()
47
+ x = x.view(n, self.kernel_size, kc // self.kernel_size, t, v)
48
+ x = torch.einsum('nkctv,kvw->nctw', (x, self.A)).contiguous()
49
+ y = self.bn(x)
50
+ y = self.relu(y)
51
+ return y
52
+
53
+ class STGCN_block(nn.Module):
54
+ def __init__(
55
+ self,
56
+ in_channels,
57
+ out_channels,
58
+ kernel_size,
59
+ A,
60
+ adaptive=True,
61
+ stride=1,
62
+ dropout=0,
63
+ residual=True,
64
+ ):
65
+ super().__init__()
66
+
67
+ assert len(kernel_size) == 2
68
+ assert kernel_size[0] % 2 == 1
69
+ padding = ((kernel_size[0] - 1) // 2, 0)
70
+ self.gcn = GCN_unit(
71
+ in_channels,
72
+ out_channels,
73
+ kernel_size[1],
74
+ A,
75
+ adaptive=adaptive,
76
+ )
77
+ if kernel_size[0] > 1:
78
+ self.tcn = nn.Sequential(
79
+ nn.Conv2d(
80
+ out_channels,
81
+ out_channels,
82
+ (kernel_size[0], 1),
83
+ (stride, 1),
84
+ padding,
85
+ ),
86
+ nn.BatchNorm2d(out_channels),
87
+ nn.Dropout(dropout, inplace=True),
88
+ )
89
+ else:
90
+ self.tcn = nn.Identity()
91
+
92
+ if not residual:
93
+ self.residual = lambda x: 0
94
+
95
+ elif (in_channels == out_channels) and (stride == 1):
96
+ self.residual = lambda x: x
97
+
98
+ else:
99
+ self.residual = nn.Sequential(
100
+ nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=(stride, 1)),
101
+ nn.BatchNorm2d(out_channels),
102
+ )
103
+
104
+ self.relu = nn.ReLU(inplace=True)
105
+
106
+ def forward(self, x, len_x=None):
107
+ res = self.residual(x)
108
+ x = self.gcn(x, len_x)
109
+ x = self.tcn(x) + res
110
+ return self.relu(x)
111
+
112
+ class STGCNChain(nn.Sequential):
113
+ def __init__(self, in_dim, block_args, kernel_size, A, adaptive):
114
+ super(STGCNChain, self).__init__()
115
+ last_dim = in_dim
116
+ for i, [channel, depth] in enumerate(block_args):
117
+ for j in range(depth):
118
+ self.add_module(f'layer{i}_{j}', STGCN_block(last_dim, channel, kernel_size, A.clone(), adaptive))
119
+ last_dim = channel
120
+
121
+ def get_stgcn_chain(in_dim, level, kernel_size, A, adaptive):
122
+ if level == 'spatial':
123
+ block_args = [[64,1], [128,1], [256,1]]
124
+ elif level == 'temporal':
125
+ block_args = [[256,3]]
126
+ else:
127
+ raise NotImplementedError
128
+ return STGCNChain(in_dim, block_args, kernel_size, A, adaptive), block_args[-1][0]
Uni_Sign/unisign_model/config.json ADDED
@@ -0,0 +1,28 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_name_or_path": "/home/patrick/hugging_face/t5/mt5-base",
3
+ "architectures": [
4
+ "MT5ForConditionalGeneration"
5
+ ],
6
+ "d_ff": 2048,
7
+ "d_kv": 64,
8
+ "d_model": 768,
9
+ "decoder_start_token_id": 0,
10
+ "dropout_rate": 0.1,
11
+ "eos_token_id": 1,
12
+ "feed_forward_proj": "gated-gelu",
13
+ "initializer_factor": 1.0,
14
+ "is_encoder_decoder": true,
15
+ "layer_norm_epsilon": 1e-06,
16
+ "model_type": "mt5",
17
+ "num_decoder_layers": 12,
18
+ "num_heads": 12,
19
+ "num_layers": 12,
20
+ "output_past": true,
21
+ "pad_token_id": 0,
22
+ "relative_attention_num_buckets": 32,
23
+ "tie_word_embeddings": false,
24
+ "tokenizer_class": "T5Tokenizer",
25
+ "transformers_version": "4.10.0.dev0",
26
+ "use_cache": true,
27
+ "vocab_size": 250112
28
+ }
Uni_Sign/unisign_model/generation_config.json ADDED
@@ -0,0 +1,7 @@
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "decoder_start_token_id": 0,
4
+ "eos_token_id": 1,
5
+ "pad_token_id": 0,
6
+ "transformers_version": "4.27.0.dev0"
7
+ }
Uni_Sign/unisign_model/special_tokens_map.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"eos_token": "</s>", "unk_token": "<unk>", "pad_token": "<pad>"}
Uni_Sign/unisign_model/spiece.model ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:ef78f86560d809067d12bac6c09f19a462cb3af3f54d2b8acbba26e1433125d6
3
+ size 4309802
Uni_Sign/unisign_model/tokenizer_config.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"eos_token": "</s>", "unk_token": "<unk>", "pad_token": "<pad>", "extra_ids": 0, "additional_special_tokens": null, "special_tokens_map_file": "/home/patrick/.cache/torch/transformers/685ac0ca8568ec593a48b61b0a3c272beee9bc194a3c7241d15dcadb5f875e53.f76030f3ec1b96a8199b2593390c610e76ca8028ef3d24680000619ffb646276", "tokenizer_file": null, "name_or_path": "google/mt5-small"}
checkpoints/pose/face_landmarker.task ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:64184e229b263107bc2b804c6625db1341ff2bb731874b0bcc2fe6544e0bc9ff
3
+ size 3758596
checkpoints/pose/hand_landmarker.task ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:fbc2a30080c3c557093b5ddfc334698132eb341044ccee322ccf8bcf3607cde1
3
+ size 7819105
checkpoints/pose/pose_landmarker_full.task ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:4eaa5eb7a98365221087693fcc286334cf0858e2eb6e15b506aa4a7ecdcec4ad
3
+ size 9398198
checkpoints/pose/yolov8n-pose.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:c6fa93dd1ee4a2c18c900a45c1d864a1c6f7aba75d84f91648a30b7fb641d212
3
+ size 6832633