Draw prediction for the given task. Args: task (TaskInfo object): task object that contain the necessary information for visualization. (e.g. frames, preds) All attributes must lie on CPU devices. video_vis (VideoVisualizer object): the video visualiz
(task, video_vis)
| 274 | |
| 275 | |
| 276 | def draw_predictions(task, video_vis): |
| 277 | """ |
| 278 | Draw prediction for the given task. |
| 279 | Args: |
| 280 | task (TaskInfo object): task object that contain |
| 281 | the necessary information for visualization. (e.g. frames, preds) |
| 282 | All attributes must lie on CPU devices. |
| 283 | video_vis (VideoVisualizer object): the video visualizer object. |
| 284 | """ |
| 285 | boxes = task.bboxes |
| 286 | frames = task.frames |
| 287 | preds = task.action_preds |
| 288 | if boxes is not None: |
| 289 | img_width = task.img_width |
| 290 | img_height = task.img_height |
| 291 | if boxes.device != torch.device("cpu"): |
| 292 | boxes = boxes.cpu() |
| 293 | boxes = cv2_transform.revert_scaled_boxes( |
| 294 | task.crop_size, boxes, img_height, img_width |
| 295 | ) |
| 296 | |
| 297 | keyframe_idx = len(frames) // 2 - task.num_buffer_frames |
| 298 | draw_range = [ |
| 299 | keyframe_idx - task.clip_vis_size, |
| 300 | keyframe_idx + task.clip_vis_size, |
| 301 | ] |
| 302 | buffer = frames[: task.num_buffer_frames] |
| 303 | frames = frames[task.num_buffer_frames :] |
| 304 | if boxes is not None: |
| 305 | if len(boxes) != 0: |
| 306 | frames = video_vis.draw_clip_range( |
| 307 | frames, |
| 308 | preds, |
| 309 | boxes, |
| 310 | keyframe_idx=keyframe_idx, |
| 311 | draw_range=draw_range, |
| 312 | ) |
| 313 | else: |
| 314 | frames = video_vis.draw_clip_range( |
| 315 | frames, preds, keyframe_idx=keyframe_idx, draw_range=draw_range |
| 316 | ) |
| 317 | del task |
| 318 | |
| 319 | return buffer + frames |