multimodalart HF Staff commited on
Commit
6d14c19
·
verified ·
1 Parent(s): 161a7b2

Fix: re-upload correct other_tools_hf.py

Browse files
Files changed (1) hide show
  1. utils/other_tools_hf.py +959 -6
utils/other_tools_hf.py CHANGED
@@ -1,6 +1,959 @@
1
- libegl1
2
- libgles2
3
- libgl1-mesa-glx
4
- libgl1-mesa-dri
5
- libegl1-mesa
6
- libwayland-egl1-mesa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import numpy as np
3
+ import random
4
+ import torch
5
+ import shutil
6
+ import csv
7
+ import pprint
8
+ import pandas as pd
9
+ from loguru import logger
10
+ from collections import OrderedDict
11
+ import matplotlib.pyplot as plt
12
+ import pickle
13
+ import time
14
+ import hashlib
15
+ from scipy.spatial.transform import Rotation as R
16
+ from scipy.spatial.transform import Slerp
17
+ import cv2
18
+ # Defer pyrender-dependent imports to avoid EGL issues at module scope
19
+ # import utils.media
20
+ # import utils.fast_render
21
+
22
+ def write_wav_names_to_csv(folder_path, csv_path):
23
+ """
24
+ Traverse a folder and write the base names of all .wav files to a CSV file.
25
+
26
+ :param folder_path: Path to the folder to traverse.
27
+ :param csv_path: Path to the CSV file to write.
28
+ """
29
+ # Open the CSV file for writing
30
+ with open(csv_path, mode='w', newline='') as file:
31
+ writer = csv.writer(file)
32
+ # Write the header
33
+ writer.writerow(['id', 'type'])
34
+
35
+ # Walk through the folder
36
+ for root, dirs, files in os.walk(folder_path):
37
+ for file in files:
38
+ # Check if the file ends with .wav
39
+ if file.endswith('.wav'):
40
+ # Extract the base name without the extension
41
+ base_name = os.path.splitext(file)[0]
42
+ # Write the base name and type to the CSV
43
+ writer.writerow([base_name, 'test'])
44
+
45
+ def resize_motion_sequence_tensor(sequence, target_frames):
46
+ """
47
+ Resize a batch of 8-frame motion sequences to a specified number of frames using interpolation.
48
+
49
+ :param sequence: A (bs, 8, 165) tensor representing a batch of 8-frame motion sequences
50
+ :param target_frames: An integer representing the desired number of frames in the output sequences
51
+ :return: A (bs, target_frames, 165) tensor representing the resized motion sequences
52
+ """
53
+ bs, _, _ = sequence.shape
54
+
55
+ # Create a time vector for the original and target sequences
56
+ original_time = torch.linspace(0, 1, 8, device=sequence.device).view(1, -1, 1)
57
+ target_time = torch.linspace(0, 1, target_frames, device=sequence.device).view(1, -1, 1)
58
+
59
+ # Permute the dimensions to (bs, 165, 8) for interpolation
60
+ sequence = sequence.permute(0, 2, 1)
61
+
62
+ # Interpolate each joint's motion to the target number of frames
63
+ resized_sequence = torch.nn.functional.interpolate(sequence, size=target_frames, mode='linear', align_corners=True)
64
+
65
+ # Permute the dimensions back to (bs, target_frames, 165)
66
+ resized_sequence = resized_sequence.permute(0, 2, 1)
67
+
68
+ return resized_sequence
69
+
70
+ def adjust_speed_according_to_ratio_tensor(chunks):
71
+ """
72
+ Adjust the playback speed within a batch of 32-frame chunks according to random intervals.
73
+
74
+ :param chunks: A (bs, 32, 165) tensor representing a batch of motion chunks
75
+ :return: A (bs, 32, 165) tensor representing the motion chunks after speed adjustment
76
+ """
77
+ bs, _, _ = chunks.shape
78
+
79
+ # Step 1: Divide the chunk into 4 equal intervals of 8 frames
80
+ equal_intervals = torch.chunk(chunks, 4, dim=1)
81
+
82
+ # Step 2: Randomly sample 3 points within the chunk to determine new intervals
83
+ success = 0
84
+ all_success = []
85
+ #sample_points = torch.sort(torch.randint(1, 32, (bs, 3), device=chunks.device), dim=1).values
86
+ # new_intervals_boundaries = torch.cat([torch.zeros((bs, 1), device=chunks.device, dtype=torch.long), sample_points, 32*torch.ones((bs, 1), device=chunks.device, dtype=torch.long)], dim=1)
87
+ while success != 1:
88
+ sample_points = sorted(random.sample(range(1, 32), 3))
89
+ new_intervals_boundaries = [0] + sample_points + [32]
90
+ new_intervals = [chunks[0][new_intervals_boundaries[i]:new_intervals_boundaries[i+1]] for i in range(4)]
91
+ speed_ratios = [8 / len(new_interval) for new_interval in new_intervals]
92
+ # if any of the speed ratios is greater than 3 or less than 0.33, resample
93
+ if all([0.33 <= speed_ratio <= 3 for speed_ratio in speed_ratios]):
94
+ success += 1
95
+ all_success.append(new_intervals_boundaries)
96
+ new_intervals_boundaries = torch.from_numpy(np.array(all_success))
97
+ # print(new_intervals_boundaries)
98
+ all_shapes = new_intervals_boundaries[:, 1:] - new_intervals_boundaries[:, :-1]
99
+ # Step 4: Adjust the speed of each new interval
100
+ adjusted_intervals = []
101
+ # print(equal_intervals[0].shape)
102
+ for i in range(4):
103
+ adjusted_interval = resize_motion_sequence_tensor(equal_intervals[i], all_shapes[0, i])
104
+ adjusted_intervals.append(adjusted_interval)
105
+
106
+ # Step 5: Concatenate the adjusted intervals
107
+ adjusted_chunk = torch.cat(adjusted_intervals, dim=1)
108
+
109
+ return adjusted_chunk
110
+
111
+ def compute_exact_iou(bbox1, bbox2):
112
+ x1 = max(bbox1[0], bbox2[0])
113
+ y1 = max(bbox1[1], bbox2[1])
114
+ x2 = min(bbox1[0] + bbox1[2], bbox2[0] + bbox2[2])
115
+ y2 = min(bbox1[1] + bbox1[3], bbox2[1] + bbox2[3])
116
+
117
+ intersection_area = max(0, x2 - x1) * max(0, y2 - y1)
118
+ bbox1_area = bbox1[2] * bbox1[3]
119
+ bbox2_area = bbox2[2] * bbox2[3]
120
+ union_area = bbox1_area + bbox2_area - intersection_area
121
+
122
+ if union_area == 0:
123
+ return 0
124
+
125
+ return intersection_area / union_area
126
+
127
+ def compute_iou(mask1, mask2):
128
+ # Compute the intersection
129
+ intersection = np.logical_and(mask1, mask2).sum()
130
+
131
+ # Compute the union
132
+ union = np.logical_or(mask1, mask2).sum()
133
+
134
+ # Compute the IoU
135
+ iou = intersection / union
136
+
137
+ return iou
138
+
139
+ def blankblending(all_frames, x, n):
140
+ return all_frames[x:x+n+1]
141
+
142
+
143
+ def load_video_as_numpy_array(video_path):
144
+ cap = cv2.VideoCapture(video_path)
145
+
146
+ # Using list comprehension to read frames and store in a list
147
+ frames = [frame for ret, frame in iter(lambda: cap.read(), (False, None)) if ret]
148
+
149
+ cap.release()
150
+
151
+ return np.array(frames)
152
+
153
+ def synthesize_intermediate_frames_bidirectional(all_frames, x, n):
154
+ frame1 = all_frames[x]
155
+ frame2 = all_frames[x + n]
156
+
157
+ # Convert the frames to grayscale
158
+ gray1 = cv2.cvtColor(frame1, cv2.COLOR_BGR2GRAY)
159
+ gray2 = cv2.cvtColor(frame2, cv2.COLOR_BGR2GRAY)
160
+
161
+ # Calculate the forward and backward optical flow
162
+ forward_flow = cv2.calcOpticalFlowFarneback(gray1, gray2, None, 0.5, 3, 15, 3, 5, 1.2, 0)
163
+ backward_flow = cv2.calcOpticalFlowFarneback(gray2, gray1, None, 0.5, 3, 15, 3, 5, 1.2, 0)
164
+
165
+ synthesized_frames = []
166
+ for i in range(1, n): # For each intermediate frame between x and x + n
167
+ alpha = i / n # Interpolation factor
168
+
169
+ # Compute the intermediate forward and backward flow
170
+ intermediate_forward_flow = forward_flow * alpha
171
+ intermediate_backward_flow = backward_flow * (1 - alpha)
172
+
173
+ # Warp the frames based on the intermediate flow
174
+ h, w = frame1.shape[:2]
175
+ flow_map = np.column_stack((np.repeat(np.arange(h), w), np.tile(np.arange(w), h)))
176
+ forward_displacement = flow_map + intermediate_forward_flow.reshape(-1, 2)
177
+ backward_displacement = flow_map - intermediate_backward_flow.reshape(-1, 2)
178
+
179
+ # Use cv2.remap for efficient warping
180
+ remap_x_forward, remap_y_forward = np.clip(forward_displacement[:, 1], 0, w - 1), np.clip(forward_displacement[:, 0], 0, h - 1)
181
+ remap_x_backward, remap_y_backward = np.clip(backward_displacement[:, 1], 0, w - 1), np.clip(backward_displacement[:, 0], 0, h - 1)
182
+
183
+ warped_forward = cv2.remap(frame1, remap_x_forward.reshape(h, w).astype(np.float32), remap_y_forward.reshape(h, w).astype(np.float32), interpolation=cv2.INTER_LINEAR)
184
+ warped_backward = cv2.remap(frame2, remap_x_backward.reshape(h, w).astype(np.float32), remap_y_backward.reshape(h, w).astype(np.float32), interpolation=cv2.INTER_LINEAR)
185
+
186
+ # Blend the warped frames to generate the intermediate frame
187
+ intermediate_frame = cv2.addWeighted(warped_forward, 1 - alpha, warped_backward, alpha, 0)
188
+ synthesized_frames.append(intermediate_frame)
189
+
190
+ return synthesized_frames # Return n-2 synthesized intermediate frames
191
+
192
+
193
+ def linear_interpolate_frames(all_frames, x, n):
194
+ frame1 = all_frames[x]
195
+ frame2 = all_frames[x + n]
196
+
197
+ synthesized_frames = []
198
+ for i in range(1, n): # For each intermediate frame between x and x + n
199
+ alpha = i / (n) # Correct interpolation factor
200
+ inter_frame = cv2.addWeighted(frame1, 1 - alpha, frame2, alpha, 0)
201
+ synthesized_frames.append(inter_frame)
202
+ return synthesized_frames[:-1]
203
+
204
+ def warp_frame(src_frame, flow):
205
+ h, w = flow.shape[:2]
206
+ flow_map = np.column_stack((np.repeat(np.arange(h), w), np.tile(np.arange(w), h)))
207
+ displacement = flow_map + flow.reshape(-1, 2)
208
+
209
+ # Extract x and y coordinates of the displacement
210
+ x_coords = np.clip(displacement[:, 1], 0, w - 1).reshape(h, w).astype(np.float32)
211
+ y_coords = np.clip(displacement[:, 0], 0, h - 1).reshape(h, w).astype(np.float32)
212
+
213
+ # Use cv2.remap for efficient warping
214
+ warped_frame = cv2.remap(src_frame, x_coords, y_coords, interpolation=cv2.INTER_LINEAR)
215
+
216
+ return warped_frame
217
+
218
+ def synthesize_intermediate_frames(all_frames, x, n):
219
+ # Calculate Optical Flow between the first and last frame
220
+ frame1 = cv2.cvtColor(all_frames[x], cv2.COLOR_BGR2GRAY)
221
+ frame2 = cv2.cvtColor(all_frames[x + n], cv2.COLOR_BGR2GRAY)
222
+ flow = cv2.calcOpticalFlowFarneback(frame1, frame2, None, 0.5, 3, 15, 3, 5, 1.2, 0)
223
+
224
+ synthesized_frames = []
225
+ for i in range(1, n): # For each intermediate frame
226
+ alpha = i / (n) # Interpolation factor
227
+ intermediate_flow = flow * alpha # Interpolate the flow
228
+ intermediate_frame = warp_frame(all_frames[x], intermediate_flow) # Warp the first frame
229
+ synthesized_frames.append(intermediate_frame)
230
+
231
+ return synthesized_frames
232
+
233
+
234
+ def map2color(s):
235
+ m = hashlib.md5()
236
+ m.update(s.encode('utf-8'))
237
+ color_code = m.hexdigest()[:6]
238
+ return '#' + color_code
239
+
240
+ def euclidean_distance(a, b):
241
+ return np.sqrt(np.sum((a - b)**2))
242
+
243
+ def adjust_array(x, k):
244
+ len_x = len(x)
245
+ len_k = len(k)
246
+
247
+ # If x is shorter than k, pad with zeros
248
+ if len_x < len_k:
249
+ return np.pad(x, (0, len_k - len_x), 'constant')
250
+
251
+ # If x is longer than k, truncate x
252
+ elif len_x > len_k:
253
+ return x[:len_k]
254
+
255
+ # If both are of same length
256
+ else:
257
+ return x
258
+
259
+ def onset_to_frame(onset_times, audio_length, fps):
260
+ # Calculate total number of frames for the given audio length
261
+ total_frames = int(audio_length * fps)
262
+
263
+ # Create an array of zeros of shape (total_frames,)
264
+ frame_array = np.zeros(total_frames, dtype=np.int32)
265
+
266
+ # For each onset time, calculate the frame number and set it to 1
267
+ for onset in onset_times:
268
+ frame_num = int(onset * fps)
269
+ # Check if the frame number is within the array bounds
270
+ if 0 <= frame_num < total_frames:
271
+ frame_array[frame_num] = 1
272
+
273
+ return frame_array
274
+
275
+ # def np_slerp(q1, q2, t):
276
+ # dot_product = np.sum(q1 * q2, axis=-1)
277
+ # q2_flip = np.where(dot_product[:, None] < 0, -q2, q2) # Flip quaternions where dot_product is negative
278
+ # dot_product = np.abs(dot_product)
279
+
280
+ # angle = np.arccos(np.clip(dot_product, -1, 1))
281
+ # sin_angle = np.sin(angle)
282
+
283
+ # t1 = np.sin((1.0 - t) * angle) / sin_angle
284
+ # t2 = np.sin(t * angle) / sin_angle
285
+
286
+ # return t1 * q1 + t2 * q2_flip
287
+
288
+
289
+ def smooth_rotvec_animations(animation1, animation2, blend_frames):
290
+ """
291
+ Smoothly transition between two animation clips using SLERP.
292
+
293
+ Parameters:
294
+ - animation1: The first animation clip, a numpy array of shape [n, k].
295
+ - animation2: The second animation clip, a numpy array of shape [n, k].
296
+ - blend_frames: Number of frames over which to blend the two animations.
297
+
298
+ Returns:
299
+ - A smoothly blended animation clip of shape [2n, k].
300
+ """
301
+
302
+ # Ensure blend_frames doesn't exceed the length of either animation
303
+ n1, k1 = animation1.shape
304
+ n2, k2 = animation2.shape
305
+ animation1 = animation1.reshape(n1, k1//3, 3)
306
+ animation2 = animation2.reshape(n2, k2//3, 3)
307
+ blend_frames = min(blend_frames, len(animation1), len(animation2))
308
+ all_int = []
309
+ for i in range(k1//3):
310
+ # Convert rotation vectors to quaternion for the overlapping part
311
+ q = R.from_rotvec(np.concatenate([animation1[0:1, i], animation2[-2:-1, i]], axis=0))#.as_quat()
312
+ # q2 = R.from_rotvec()#.as_quat()
313
+ times = [0, blend_frames * 2 - 1]
314
+ slerp = Slerp(times, q)
315
+ interpolated = slerp(np.arange(blend_frames * 2))
316
+ interpolated_rotvecs = interpolated.as_rotvec()
317
+ all_int.append(interpolated_rotvecs)
318
+ interpolated_rotvecs = np.concatenate(all_int, axis=1)
319
+ # result = np.vstack((animation1[:-blend_frames], interpolated_rotvecs, animation2[blend_frames:]))
320
+ result = interpolated_rotvecs.reshape(2*n1, k1)
321
+ return result
322
+
323
+ def smooth_animations(animation1, animation2, blend_frames):
324
+ """
325
+ Smoothly transition between two animation clips using linear interpolation.
326
+
327
+ Parameters:
328
+ - animation1: The first animation clip, a numpy array of shape [n, k].
329
+ - animation2: The second animation clip, a numpy array of shape [n, k].
330
+ - blend_frames: Number of frames over which to blend the two animations.
331
+
332
+ Returns:
333
+ - A smoothly blended animation clip of shape [2n, k].
334
+ """
335
+
336
+ # Ensure blend_frames doesn't exceed the length of either animation
337
+ blend_frames = min(blend_frames, len(animation1), len(animation2))
338
+
339
+ # Extract overlapping sections
340
+ overlap_a1 = animation1[-blend_frames:-blend_frames+1, :]
341
+ overlap_a2 = animation2[blend_frames-1:blend_frames, :]
342
+
343
+ # Create blend weights for linear interpolation
344
+ alpha = np.linspace(0, 1, 2 * blend_frames).reshape(-1, 1)
345
+
346
+ # Linearly interpolate between overlapping sections
347
+ blended_overlap = overlap_a1 * (1 - alpha) + overlap_a2 * alpha
348
+
349
+ # Extend the animations to form the result with 2n frames
350
+ if blend_frames == len(animation1) and blend_frames == len(animation2):
351
+ result = blended_overlap
352
+ else:
353
+ before_blend = animation1[:-blend_frames]
354
+ after_blend = animation2[blend_frames:]
355
+ result = np.vstack((before_blend, blended_overlap, after_blend))
356
+ return result
357
+
358
+ def interpolate_sequence(quaternions):
359
+ bs, n, j, _ = quaternions.shape
360
+ new_n = 2 * n
361
+ new_quaternions = torch.zeros((bs, new_n, j, 4), device=quaternions.device, dtype=quaternions.dtype)
362
+
363
+ for i in range(n):
364
+ q1 = quaternions[:, i, :, :]
365
+ new_quaternions[:, 2*i, :, :] = q1
366
+
367
+ if i < n - 1:
368
+ q2 = quaternions[:, i + 1, :, :]
369
+ new_quaternions[:, 2*i + 1, :, :] = slerp(q1, q2, 0.5)
370
+ else:
371
+ # For the last point, duplicate the value
372
+ new_quaternions[:, 2*i + 1, :, :] = q1
373
+
374
+ return new_quaternions
375
+
376
+ def quaternion_multiply(q1, q2):
377
+ w1, x1, y1, z1 = q1
378
+ w2, x2, y2, z2 = q2
379
+ w = w1 * w2 - x1 * x2 - y1 * y2 - z1 * z2
380
+ x = w1 * x2 + x1 * w2 + y1 * z2 - z1 * y2
381
+ y = w1 * y2 + y1 * w2 + z1 * x2 - x1 * z2
382
+ z = w1 * z2 + z1 * w2 + x1 * y2 - y1 * x2
383
+ return w, x, y, z
384
+
385
+ def quaternion_conjugate(q):
386
+ w, x, y, z = q
387
+ return (w, -x, -y, -z)
388
+
389
+ def slerp(q1, q2, t):
390
+ dot = torch.sum(q1 * q2, dim=-1, keepdim=True)
391
+
392
+ flip = (dot < 0).float()
393
+ q2 = (1 - flip * 2) * q2
394
+ dot = dot * (1 - flip * 2)
395
+
396
+ DOT_THRESHOLD = 0.9995
397
+ mask = (dot > DOT_THRESHOLD).float()
398
+
399
+ theta_0 = torch.acos(dot)
400
+ theta = theta_0 * t
401
+
402
+ q3 = q2 - q1 * dot
403
+ q3 = q3 / torch.norm(q3, dim=-1, keepdim=True)
404
+
405
+ interpolated = (torch.cos(theta) * q1 + torch.sin(theta) * q3)
406
+
407
+ return mask * (q1 + t * (q2 - q1)) + (1 - mask) * interpolated
408
+
409
+ def estimate_linear_velocity(data_seq, dt):
410
+ '''
411
+ Given some batched data sequences of T timesteps in the shape (B, T, ...), estimates
412
+ the velocity for the middle T-2 steps using a second order central difference scheme.
413
+ The first and last frames are with forward and backward first-order
414
+ differences, respectively
415
+ - h : step size
416
+ '''
417
+ # first steps is forward diff (t+1 - t) / dt
418
+ init_vel = (data_seq[:, 1:2] - data_seq[:, :1]) / dt
419
+ # middle steps are second order (t+1 - t-1) / 2dt
420
+ middle_vel = (data_seq[:, 2:] - data_seq[:, 0:-2]) / (2 * dt)
421
+ # last step is backward diff (t - t-1) / dt
422
+ final_vel = (data_seq[:, -1:] - data_seq[:, -2:-1]) / dt
423
+
424
+ vel_seq = torch.cat([init_vel, middle_vel, final_vel], dim=1)
425
+ return vel_seq
426
+
427
+ def velocity2position(data_seq, dt, init_pos):
428
+ res_trans = []
429
+ for i in range(data_seq.shape[1]):
430
+ if i == 0:
431
+ res_trans.append(init_pos.unsqueeze(1))
432
+ else:
433
+ res = data_seq[:, i-1:i] * dt + res_trans[-1]
434
+ res_trans.append(res)
435
+ return torch.cat(res_trans, dim=1)
436
+
437
+ def estimate_angular_velocity(rot_seq, dt):
438
+ '''
439
+ Given a batch of sequences of T rotation matrices, estimates angular velocity at T-2 steps.
440
+ Input sequence should be of shape (B, T, ..., 3, 3)
441
+ '''
442
+ # see https://en.wikipedia.org/wiki/Angular_velocity#Calculation_from_the_orientation_matrix
443
+ dRdt = estimate_linear_velocity(rot_seq, dt)
444
+ R = rot_seq
445
+ RT = R.transpose(-1, -2)
446
+ # compute skew-symmetric angular velocity tensor
447
+ w_mat = torch.matmul(dRdt, RT)
448
+ # pull out angular velocity vector by averaging symmetric entries
449
+ w_x = (-w_mat[..., 1, 2] + w_mat[..., 2, 1]) / 2.0
450
+ w_y = (w_mat[..., 0, 2] - w_mat[..., 2, 0]) / 2.0
451
+ w_z = (-w_mat[..., 0, 1] + w_mat[..., 1, 0]) / 2.0
452
+ w = torch.stack([w_x, w_y, w_z], axis=-1)
453
+ return w
454
+
455
+ def image_from_bytes(image_bytes):
456
+ import matplotlib.image as mpimg
457
+ from io import BytesIO
458
+ return mpimg.imread(BytesIO(image_bytes), format='PNG')
459
+
460
+ def process_frame(i, vertices_all, vertices1_all, faces, output_dir, filenames):
461
+ import matplotlib
462
+ matplotlib.use('Agg')
463
+ import matplotlib.pyplot as plt
464
+ import trimesh
465
+ import pyrender
466
+
467
+ def deg_to_rad(degrees):
468
+ return degrees * np.pi / 180
469
+
470
+ uniform_color = [220, 220, 220, 255]
471
+ resolution = (1000, 1000)
472
+ figsize = (10, 10)
473
+
474
+ fig, axs = plt.subplots(
475
+ nrows=1,
476
+ ncols=2,
477
+ figsize=(figsize[0] * 2, figsize[1] * 1)
478
+ )
479
+ axs = axs.flatten()
480
+
481
+ vertices = vertices_all[i]
482
+ vertices1 = vertices1_all[i]
483
+ filename = f"{output_dir}frame_{i}.png"
484
+ filenames.append(filename)
485
+ if i%100 == 0:
486
+ print('processed', i, 'frames')
487
+ #time_s = time.time()
488
+ #print(vertices.shape)
489
+ angle_rad = deg_to_rad(-2)
490
+ pose_camera = np.array([
491
+ [1.0, 0.0, 0.0, 0.0],
492
+ [0.0, np.cos(angle_rad), -np.sin(angle_rad), 1.0],
493
+ [0.0, np.sin(angle_rad), np.cos(angle_rad), 5.0],
494
+ [0.0, 0.0, 0.0, 1.0]
495
+ ])
496
+ angle_rad = deg_to_rad(-30)
497
+ pose_light = np.array([
498
+ [1.0, 0.0, 0.0, 0.0],
499
+ [0.0, np.cos(angle_rad), -np.sin(angle_rad), 0.0],
500
+ [0.0, np.sin(angle_rad), np.cos(angle_rad), 3.0],
501
+ [0.0, 0.0, 0.0, 1.0]
502
+ ])
503
+
504
+ for vtx_idx, vtx in enumerate([vertices, vertices1]):
505
+ trimesh_mesh = trimesh.Trimesh(
506
+ vertices=vtx,
507
+ faces=faces,
508
+ vertex_colors=uniform_color
509
+ )
510
+ mesh = pyrender.Mesh.from_trimesh(
511
+ trimesh_mesh, smooth=True
512
+ )
513
+ scene = pyrender.Scene()
514
+ scene.add(mesh)
515
+ camera = pyrender.OrthographicCamera(xmag=1.0, ymag=1.0)
516
+ scene.add(camera, pose=pose_camera)
517
+ light = pyrender.DirectionalLight(color=[1.0, 1.0, 1.0], intensity=4.0)
518
+ scene.add(light, pose=pose_light)
519
+ renderer = pyrender.OffscreenRenderer(*resolution)
520
+ color, _ = renderer.render(scene)
521
+ axs[vtx_idx].imshow(color)
522
+ axs[vtx_idx].axis('off')
523
+ renderer.delete()
524
+
525
+ plt.savefig(filename, bbox_inches='tight')
526
+ plt.close(fig)
527
+
528
+ def generate_images(frames, vertices_all, vertices1_all, faces, output_dir, filenames):
529
+ import multiprocessing
530
+ # import trimesh
531
+ num_cores = multiprocessing.cpu_count() - 1 # This will get the number of cores on your machine.
532
+ # mesh = trimesh.Trimesh(vertices_all[0], faces)
533
+ # scene = mesh.scene()
534
+ # fov = scene.camera.fov.copy()
535
+ # fov[0] = 80.0
536
+ # fov[1] = 60.0
537
+ # camera_params = {
538
+ # 'fov': fov,
539
+ # 'resolution': scene.camera.resolution,
540
+ # 'focal': scene.camera.focal,
541
+ # 'z_near': scene.camera.z_near,
542
+ # "z_far": scene.camera.z_far,
543
+ # 'transform': scene.graph[scene.camera.name][0]
544
+ # }
545
+ # mesh1 = trimesh.Trimesh(vertices1_all[0], faces)
546
+ # scene1 = mesh1.scene()
547
+ # camera_params1 = {
548
+ # 'fov': fov,
549
+ # 'resolution': scene1.camera.resolution,
550
+ # 'focal': scene1.camera.focal,
551
+ # 'z_near': scene1.camera.z_near,
552
+ # "z_far": scene1.camera.z_far,
553
+ # 'transform': scene1.graph[scene1.camera.name][0]
554
+ # }
555
+ # Use a Pool to manage the processes
556
+ # print(num_cores)
557
+ # for i in range(frames):
558
+ # process_frame(i, vertices_all, vertices1_all, faces, output_dir, use_matplotlib, filenames, camera_params, camera_params1)
559
+ for i in range(frames):
560
+ process_frame(i*3, vertices_all, vertices1_all, faces, output_dir, filenames)
561
+
562
+ # progress = multiprocessing.Value('i', 0)
563
+ # lock = multiprocessing.Lock()
564
+ # with multiprocessing.Pool(num_cores) as pool:
565
+ # # pool.starmap(process_frame, [(i, vertices_all, vertices1_all, faces, output_dir, use_matplotlib, filenames, camera_params, camera_params1) for i in range(frames)])
566
+ # pool.starmap(
567
+ # process_frame,
568
+ # [
569
+ # (i, vertices_all, vertices1_all, faces, output_dir, filenames)
570
+ # for i in range(frames)
571
+ # ]
572
+ # )
573
+
574
+ # progress = multiprocessing.Value('i', 0)
575
+ # lock = multiprocessing.Lock()
576
+ # with multiprocessing.Pool(num_cores) as pool:
577
+ # # pool.starmap(process_frame, [(i, vertices_all, vertices1_all, faces, output_dir, use_matplotlib, filenames, camera_params, camera_params1) for i in range(frames)])
578
+ # pool.starmap(
579
+ # process_frame,
580
+ # [
581
+ # (i, vertices_all, vertices1_all, faces, output_dir, filenames)
582
+ # for i in range(frames)
583
+ # ]
584
+ # )
585
+
586
+ def render_one_sequence(
587
+ res_npz_path,
588
+ gt_npz_path,
589
+ output_dir,
590
+ audio_path,
591
+ model_folder="/data/datasets/smplx_models/",
592
+ model_type='smplx',
593
+ gender='NEUTRAL_2020',
594
+ ext='npz',
595
+ num_betas=300,
596
+ num_expression_coeffs=100,
597
+ use_face_contour=False,
598
+ use_matplotlib=False,
599
+ args=None):
600
+ import smplx
601
+ import matplotlib.pyplot as plt
602
+ import imageio
603
+ from tqdm import tqdm
604
+ import os
605
+ import numpy as np
606
+ import torch
607
+ import moviepy.editor as mp
608
+ import librosa
609
+ import utils.media
610
+ import utils.fast_render
611
+
612
+ model = smplx.create(model_folder, model_type=model_type,
613
+ gender=gender, use_face_contour=use_face_contour,
614
+ num_betas=num_betas,
615
+ num_expression_coeffs=num_expression_coeffs,
616
+ ext=ext, use_pca=False).cuda()
617
+
618
+ #data_npz = np.load(f"{output_dir}{res_npz_path}.npz")
619
+ data_np_body = np.load(res_npz_path, allow_pickle=True)
620
+ gt_np_body = np.load(gt_npz_path, allow_pickle=True)
621
+ # if not use_matplotlib:
622
+ # import trimesh
623
+ #import pyrender
624
+ from pyvirtualdisplay import Display
625
+ #'''
626
+ #display = Display(visible=0, size=(1000, 1000))
627
+ #display.start()
628
+ faces = np.load(f"{model_folder}/smplx/SMPLX_NEUTRAL_2020.npz", allow_pickle=True)["f"]
629
+ seconds = 1
630
+ #data_npz["jaw_pose"].shape[0]
631
+ n = data_np_body["poses"].shape[0]
632
+ beta = torch.from_numpy(data_np_body["betas"]).to(torch.float32).unsqueeze(0).cuda()
633
+ beta = beta.repeat(n, 1)
634
+ expression = torch.from_numpy(data_np_body["expressions"][:n]).to(torch.float32).cuda()
635
+ jaw_pose = torch.from_numpy(data_np_body["poses"][:n, 66:69]).to(torch.float32).cuda()
636
+ pose = torch.from_numpy(data_np_body["poses"][:n]).to(torch.float32).cuda()
637
+ transl = torch.from_numpy(data_np_body["trans"][:n]).to(torch.float32).cuda()
638
+ # print(beta.shape, expression.shape, jaw_pose.shape, pose.shape, transl.shape, pose[:,:3].shape)
639
+ output = model(betas=beta, transl=transl, expression=expression, jaw_pose=jaw_pose,
640
+ global_orient=pose[:,:3], body_pose=pose[:,3:21*3+3], left_hand_pose=pose[:,25*3:40*3], right_hand_pose=pose[:,40*3:55*3],
641
+ leye_pose=pose[:, 69:72],
642
+ reye_pose=pose[:, 72:75],
643
+ return_verts=True)
644
+ vertices_all = output["vertices"].cpu().detach().numpy()
645
+
646
+ beta1 = torch.from_numpy(gt_np_body["betas"]).to(torch.float32).unsqueeze(0).cuda()
647
+ expression1 = torch.from_numpy(gt_np_body["expressions"][:n]).to(torch.float32).cuda()
648
+ jaw_pose1 = torch.from_numpy(gt_np_body["poses"][:n,66:69]).to(torch.float32).cuda()
649
+ pose1 = torch.from_numpy(gt_np_body["poses"][:n]).to(torch.float32).cuda()
650
+ transl1 = torch.from_numpy(gt_np_body["trans"][:n]).to(torch.float32).cuda()
651
+ output1 = model(betas=beta1, transl=transl1, expression=expression1, jaw_pose=jaw_pose1, global_orient=pose1[:,:3], body_pose=pose1[:,3:21*3+3], left_hand_pose=pose1[:,25*3:40*3], right_hand_pose=pose1[:,40*3:55*3],
652
+ leye_pose=pose1[:, 69:72],
653
+ reye_pose=pose1[:, 72:75],return_verts=True)
654
+ vertices1_all = output1["vertices"].cpu().detach().numpy()
655
+ if args.debug:
656
+ seconds = 1
657
+ else:
658
+ seconds = vertices_all.shape[0]//30
659
+ silent_video_file_path = utils.fast_render.generate_silent_videos(args.render_video_fps,
660
+ args.render_video_width,
661
+ args.render_video_height,
662
+ args.render_concurrent_num,
663
+ args.render_tmp_img_filetype,
664
+ int(seconds*args.render_video_fps),
665
+ vertices_all,
666
+ vertices1_all,
667
+ faces,
668
+ output_dir)
669
+ base_filename_without_ext = os.path.splitext(os.path.basename(res_npz_path))[0]
670
+ final_clip = os.path.join(output_dir, f"{base_filename_without_ext}.mp4")
671
+ utils.media.add_audio_to_video(silent_video_file_path, audio_path, final_clip)
672
+ os.remove(silent_video_file_path)
673
+ return final_clip
674
+
675
+ def render_one_sequence_no_gt(
676
+ res_npz_path,
677
+ output_dir,
678
+ audio_path,
679
+ model_folder="/data/datasets/smplx_models/",
680
+ model_type='smplx',
681
+ gender='NEUTRAL_2020',
682
+ ext='npz',
683
+ num_betas=300,
684
+ num_expression_coeffs=100,
685
+ use_face_contour=False,
686
+ use_matplotlib=False,
687
+ args=None):
688
+ import smplx
689
+ import matplotlib.pyplot as plt
690
+ import imageio
691
+ from tqdm import tqdm
692
+ import os
693
+ import numpy as np
694
+ import torch
695
+ import moviepy.editor as mp
696
+ import librosa
697
+ import utils.media
698
+ import utils.fast_render
699
+
700
+ model = smplx.create(model_folder, model_type=model_type,
701
+ gender=gender, use_face_contour=use_face_contour,
702
+ num_betas=num_betas,
703
+ num_expression_coeffs=num_expression_coeffs,
704
+ ext=ext, use_pca=False).cuda()
705
+
706
+ #data_npz = np.load(f"{output_dir}{res_npz_path}.npz")
707
+ data_np_body = np.load(res_npz_path, allow_pickle=True)
708
+ # gt_np_body = np.load(gt_npz_path, allow_pickle=True)
709
+
710
+ if not os.path.exists(output_dir): os.makedirs(output_dir)
711
+ # if not use_matplotlib:
712
+ # import trimesh
713
+ #import pyrender
714
+ #'''
715
+ #display = Display(visible=0, size=(1000, 1000))
716
+ #display.start()
717
+ faces = np.load(f"{model_folder}/smplx/SMPLX_NEUTRAL_2020.npz", allow_pickle=True)["f"]
718
+ seconds = 1
719
+ #data_npz["jaw_pose"].shape[0]
720
+ n = data_np_body["poses"].shape[0]
721
+ beta = torch.from_numpy(data_np_body["betas"]).to(torch.float32).unsqueeze(0).cuda()
722
+ beta = beta.repeat(n, 1)
723
+ expression = torch.from_numpy(data_np_body["expressions"][:n]).to(torch.float32).cuda()
724
+ jaw_pose = torch.from_numpy(data_np_body["poses"][:n, 66:69]).to(torch.float32).cuda()
725
+ pose = torch.from_numpy(data_np_body["poses"][:n]).to(torch.float32).cuda()
726
+ transl = torch.from_numpy(data_np_body["trans"][:n]).to(torch.float32).cuda()
727
+ # print(beta.shape, expression.shape, jaw_pose.shape, pose.shape, transl.shape, pose[:,:3].shape)
728
+ output = model(betas=beta, transl=transl, expression=expression, jaw_pose=jaw_pose,
729
+ global_orient=pose[:,:3], body_pose=pose[:,3:21*3+3], left_hand_pose=pose[:,25*3:40*3], right_hand_pose=pose[:,40*3:55*3],
730
+ leye_pose=pose[:, 69:72],
731
+ reye_pose=pose[:, 72:75],
732
+ return_verts=True)
733
+ vertices_all = output["vertices"].cpu().detach().numpy()
734
+
735
+ # beta1 = torch.from_numpy(gt_np_body["betas"]).to(torch.float32).unsqueeze(0).cuda()
736
+ # expression1 = torch.from_numpy(gt_np_body["expressions"][:n]).to(torch.float32).cuda()
737
+ # jaw_pose1 = torch.from_numpy(gt_np_body["poses"][:n,66:69]).to(torch.float32).cuda()
738
+ # pose1 = torch.from_numpy(gt_np_body["poses"][:n]).to(torch.float32).cuda()
739
+ # transl1 = torch.from_numpy(gt_np_body["trans"][:n]).to(torch.float32).cuda()
740
+ # output1 = model(betas=beta1, transl=transl1, expression=expression1, jaw_pose=jaw_pose1, global_orient=pose1[:,:3], body_pose=pose1[:,3:21*3+3], left_hand_pose=pose1[:,25*3:40*3], right_hand_pose=pose1[:,40*3:55*3],
741
+ # leye_pose=pose1[:, 69:72],
742
+ # reye_pose=pose1[:, 72:75],return_verts=True)
743
+ # vertices1_all = output1["vertices"].cpu().detach().numpy()
744
+ if args.debug:
745
+ seconds = 1
746
+ else:
747
+ seconds = vertices_all.shape[0]//30
748
+ silent_video_file_path = utils.fast_render.generate_silent_videos_no_gt(args.render_video_fps,
749
+ args.render_video_width,
750
+ args.render_video_height,
751
+ args.render_concurrent_num,
752
+ args.render_tmp_img_filetype,
753
+ int(seconds*args.render_video_fps),
754
+ vertices_all,
755
+ faces,
756
+ output_dir)
757
+ base_filename_without_ext = os.path.splitext(os.path.basename(res_npz_path))[0]
758
+ final_clip = os.path.join(output_dir, f"{base_filename_without_ext}.mp4")
759
+ utils.media.add_audio_to_video(silent_video_file_path, audio_path, final_clip)
760
+ os.remove(silent_video_file_path)
761
+ return final_clip
762
+
763
+ def print_exp_info(args):
764
+ logger.info(pprint.pformat(vars(args)))
765
+ logger.info(f"# ------------ {args.name} ----------- #")
766
+ logger.info("PyTorch version: {}".format(torch.__version__))
767
+ logger.info("CUDA version: {}".format(torch.version.cuda))
768
+ logger.info("{} GPUs".format(torch.cuda.device_count()))
769
+ logger.info(f"Random Seed: {args.random_seed}")
770
+
771
+ def args2csv(args, get_head=False, list4print=[]):
772
+ for k, v in args.items():
773
+ if isinstance(args[k], dict):
774
+ args2csv(args[k], get_head, list4print)
775
+ else: list4print.append(k) if get_head else list4print.append(v)
776
+ return list4print
777
+
778
+ class EpochTracker:
779
+ def __init__(self, metric_names, metric_directions):
780
+ assert len(metric_names) == len(metric_directions), "Metric names and directions should have the same length"
781
+
782
+
783
+ self.metric_names = metric_names
784
+ self.states = ['train', 'val', 'test']
785
+ self.types = ['last', 'best']
786
+
787
+
788
+ self.values = {name: {state: {type_: {'value': np.inf if not is_higher_better else -np.inf, 'epoch': 0}
789
+ for type_ in self.types}
790
+ for state in self.states}
791
+ for name, is_higher_better in zip(metric_names, metric_directions)}
792
+
793
+ self.loss_meters = {name: {state: AverageMeter(f"{name}_{state}")
794
+ for state in self.states}
795
+ for name in metric_names}
796
+
797
+
798
+ self.is_higher_better = {name: direction for name, direction in zip(metric_names, metric_directions)}
799
+ self.train_history = {name: [] for name in metric_names}
800
+ self.val_history = {name: [] for name in metric_names}
801
+
802
+
803
+ def update_meter(self, name, state, value):
804
+ self.loss_meters[name][state].update(value)
805
+
806
+
807
+ def update_values(self, name, state, epoch):
808
+ value_avg = self.loss_meters[name][state].avg
809
+ new_best = False
810
+
811
+
812
+ if ((value_avg < self.values[name][state]['best']['value'] and not self.is_higher_better[name]) or
813
+ (value_avg > self.values[name][state]['best']['value'] and self.is_higher_better[name])):
814
+ self.values[name][state]['best']['value'] = value_avg
815
+ self.values[name][state]['best']['epoch'] = epoch
816
+ new_best = True
817
+ self.values[name][state]['last']['value'] = value_avg
818
+ self.values[name][state]['last']['epoch'] = epoch
819
+ return new_best
820
+
821
+
822
+ def get(self, name, state, type_):
823
+ return self.values[name][state][type_]
824
+
825
+
826
+ def reset(self):
827
+ for name in self.metric_names:
828
+ for state in self.states:
829
+ self.loss_meters[name][state].reset()
830
+
831
+
832
+ def flatten_values(self):
833
+ flat_dict = {}
834
+ for name in self.metric_names:
835
+ for state in self.states:
836
+ for type_ in self.types:
837
+ value_key = f"{name}_{state}_{type_}"
838
+ epoch_key = f"{name}_{state}_{type_}_epoch"
839
+ flat_dict[value_key] = self.values[name][state][type_]['value']
840
+ flat_dict[epoch_key] = self.values[name][state][type_]['epoch']
841
+ return flat_dict
842
+
843
+ def update_and_plot(self, name, epoch, save_path):
844
+ new_best_train = self.update_values(name, 'train', epoch)
845
+ new_best_val = self.update_values(name, 'val', epoch)
846
+
847
+
848
+ self.train_history[name].append(self.loss_meters[name]['train'].avg)
849
+ self.val_history[name].append(self.loss_meters[name]['val'].avg)
850
+
851
+
852
+ train_values = self.train_history[name]
853
+ val_values = self.val_history[name]
854
+ epochs = list(range(1, len(train_values) + 1))
855
+
856
+
857
+ plt.figure(figsize=(10, 6))
858
+ plt.plot(epochs, train_values, label='Train')
859
+ plt.plot(epochs, val_values, label='Val')
860
+ plt.title(f'Train vs Val {name} over epochs')
861
+ plt.xlabel('Epochs')
862
+ plt.ylabel(name)
863
+ plt.legend()
864
+ plt.savefig(save_path)
865
+ plt.close()
866
+
867
+
868
+ return new_best_train, new_best_val
869
+
870
+ def record_trial(args, tracker):
871
+ """
872
+ 1. record notes, score, env_name, experments_path,
873
+ """
874
+ csv_path = args.out_path + "custom/" +args.csv_name+".csv"
875
+ all_print_dict = vars(args)
876
+ all_print_dict.update(tracker.flatten_values())
877
+ if not os.path.exists(csv_path):
878
+ pd.DataFrame([all_print_dict]).to_csv(csv_path, index=False)
879
+ else:
880
+ df_existing = pd.read_csv(csv_path)
881
+ df_new = pd.DataFrame([all_print_dict])
882
+ df_aligned = df_existing.append(df_new).fillna("")
883
+ df_aligned.to_csv(csv_path, index=False)
884
+
885
+ def set_random_seed(args):
886
+ os.environ['PYTHONHASHSEED'] = str(args.random_seed)
887
+ random.seed(args.random_seed)
888
+ np.random.seed(args.random_seed)
889
+ torch.manual_seed(args.random_seed)
890
+ torch.cuda.manual_seed_all(args.random_seed)
891
+ torch.cuda.manual_seed(args.random_seed)
892
+ torch.backends.cudnn.deterministic = args.deterministic #args.CUDNN_DETERMINISTIC
893
+ torch.backends.cudnn.benchmark = args.benchmark
894
+ torch.backends.cudnn.enabled = args.cudnn_enabled
895
+
896
+ def save_checkpoints(save_path, model, opt=None, epoch=None, lrs=None):
897
+ if lrs is not None:
898
+ states = { 'model_state': model.state_dict(),
899
+ 'epoch': epoch + 1,
900
+ 'opt_state': opt.state_dict(),
901
+ 'lrs':lrs.state_dict(),}
902
+ elif opt is not None:
903
+ states = { 'model_state': model.state_dict(),
904
+ 'epoch': epoch + 1,
905
+ 'opt_state': opt.state_dict(),}
906
+ else:
907
+ states = { 'model_state': model.state_dict(),}
908
+ torch.save(states, save_path)
909
+
910
+ def load_checkpoints(model, save_path, load_name='model'):
911
+ states = torch.load(save_path)
912
+ new_weights = OrderedDict()
913
+ flag=False
914
+ for k, v in states['model_state'].items():
915
+ #print(k)
916
+ if "module" not in k:
917
+ break
918
+ else:
919
+ new_weights[k[7:]]=v
920
+ flag=True
921
+ if flag:
922
+ try:
923
+ model.load_state_dict(new_weights)
924
+ except:
925
+ #print(states['model_state'])
926
+ model.load_state_dict(states['model_state'])
927
+ else:
928
+ model.load_state_dict(states['model_state'])
929
+ logger.info(f"load self-pretrained checkpoints for {load_name}")
930
+
931
+ def model_complexity(model, args):
932
+ from ptflops import get_model_complexity_info
933
+ flops, params = get_model_complexity_info(model, (args.T_GLOBAL._DIM, args.TRAIN.CROP, args.TRAIN),
934
+ as_strings=False, print_per_layer_stat=False)
935
+ logging.info('{:<30} {:<8} BFlops'.format('Computational complexity: ', flops / 1e9))
936
+ logging.info('{:<30} {:<8} MParams'.format('Number of parameters: ', params / 1e6))
937
+
938
+ class AverageMeter(object):
939
+ """Computes and stores the average and current value"""
940
+ def __init__(self, name, fmt=':f'):
941
+ self.name = name
942
+ self.fmt = fmt
943
+ self.reset()
944
+
945
+ def reset(self):
946
+ self.val = 0
947
+ self.avg = 0
948
+ self.sum = 0
949
+ self.count = 0
950
+
951
+ def update(self, val, n=1):
952
+ self.val = val
953
+ self.sum += val * n
954
+ self.count += n
955
+ self.avg = self.sum / self.count
956
+
957
+ def __str__(self):
958
+ fmtstr = '{name} {val' + self.fmt + '} ({avg' + self.fmt + '})'
959
+ return fmtstr.format(**self.__dict__)