Process mmtracking results. Args: mmtracking_results ([list]): mmtracking_results. bbox_thr (float): threshold for bounding boxes. max_track_id (int): the maximum track id. Returns: person_results ([list]): a list of tracked bounding boxes max_track_i
(mmtracking_results,
max_track_id,
bbox_thr=None)
| 522 | |
| 523 | |
| 524 | def process_mmtracking_results(mmtracking_results, |
| 525 | max_track_id, |
| 526 | bbox_thr=None): |
| 527 | """Process mmtracking results. |
| 528 | |
| 529 | Args: |
| 530 | mmtracking_results ([list]): mmtracking_results. |
| 531 | bbox_thr (float): threshold for bounding boxes. |
| 532 | max_track_id (int): the maximum track id. |
| 533 | Returns: |
| 534 | person_results ([list]): a list of tracked bounding boxes |
| 535 | max_track_id (int): the maximum track id. |
| 536 | instance_num (int): the number of instance. |
| 537 | """ |
| 538 | person_results = [] |
| 539 | # 'track_results' is changed to 'track_bboxes' |
| 540 | # in https://github.com/open-mmlab/mmtracking/pull/300 |
| 541 | if 'track_bboxes' in mmtracking_results: |
| 542 | tracking_results = mmtracking_results['track_bboxes'][0] |
| 543 | elif 'track_results' in mmtracking_results: |
| 544 | tracking_results = mmtracking_results['track_results'][0] |
| 545 | |
| 546 | tracking_results = np.array(tracking_results) |
| 547 | |
| 548 | if bbox_thr is not None: |
| 549 | assert tracking_results.shape[-1] == 6 |
| 550 | valid_idx = np.where(tracking_results[:, 5] > bbox_thr)[0] |
| 551 | tracking_results = tracking_results[valid_idx] |
| 552 | |
| 553 | for track in tracking_results: |
| 554 | person = {} |
| 555 | person['track_id'] = int(track[0]) |
| 556 | if max_track_id < int(track[0]): |
| 557 | max_track_id = int(track[0]) |
| 558 | person['bbox'] = track[1:] |
| 559 | person_results.append(person) |
| 560 | person_results = sorted(person_results, key=lambda x: x.get('track_id', 0)) |
| 561 | instance_num = len(person_results) |
| 562 | return person_results, max_track_id, instance_num |
| 563 | |
| 564 | |
| 565 | def process_mmdet_results(mmdet_results, cat_id=1, bbox_thr=None): |