| """Canonical 250-point layout for T-Rex three-view tracks. |
| |
| The integer identities in this module are part of the on-disk dataset schema. |
| Do not reorder groups without also changing :data:`TRACK_LAYOUT_VERSION`. |
| """ |
|
|
| from __future__ import annotations |
|
|
| from typing import Final |
|
|
| TRACK_LAYOUT_VERSION: Final = "trex_track_250_v1" |
|
|
| VIEW_ORDER: Final = ("head_left", "left_wrist", "right_wrist") |
| VIEW_IDS: Final = {name: index for index, name in enumerate(VIEW_ORDER)} |
|
|
| NUM_HEAD_PER_HAND: Final = 50 |
| NUM_HEAD_LEFT: Final = NUM_HEAD_PER_HAND |
| NUM_HEAD_RIGHT: Final = NUM_HEAD_PER_HAND |
| NUM_HEAD_POINTS: Final = NUM_HEAD_LEFT + NUM_HEAD_RIGHT |
|
|
| NUM_WRIST_BACKGROUND: Final = 25 |
| NUM_WRIST_HAND: Final = 50 |
| NUM_WRIST_POINTS: Final = NUM_WRIST_BACKGROUND + NUM_WRIST_HAND |
|
|
| NUM_COMBINED_POINTS: Final = NUM_HEAD_POINTS + 2 * NUM_WRIST_POINTS |
| POINT_SLICES: Final = ( |
| 0, |
| NUM_HEAD_POINTS, |
| NUM_HEAD_POINTS + NUM_WRIST_POINTS, |
| NUM_COMBINED_POINTS, |
| ) |
|
|
| HAND_NONE: Final = 0 |
| HAND_LEFT: Final = 1 |
| HAND_RIGHT: Final = 2 |
| HAND_NAMES: Final = ("none", "left", "right") |
|
|
| ROLE_HEAD_HAND: Final = 0 |
| ROLE_WRIST_BACKGROUND: Final = 1 |
| ROLE_WRIST_HAND: Final = 2 |
| ROLE_NAMES: Final = ("head_hand", "wrist_background", "wrist_hand") |
|
|
| |
| COMPONENT_SLICES: Final = { |
| "head_left_hand": (0, 50), |
| "head_right_hand": (50, 100), |
| "left_wrist_background": (100, 125), |
| "left_wrist_hand": (125, 175), |
| "right_wrist_background": (175, 200), |
| "right_wrist_hand": (200, 250), |
| } |
| VIEW_SLICES: Final = { |
| "head_left": (0, 100), |
| "left_wrist": (100, 175), |
| "right_wrist": (175, 250), |
| } |
| VIEW_POINT_COUNTS: Final = { |
| view: end - start for view, (start, end) in VIEW_SLICES.items() |
| } |
|
|
|
|
| def identity_metadata() -> dict[str, list[int] | list[str]]: |
| """Return stable per-point IDs in canonical concatenation order.""" |
|
|
| view_ids: list[int] = [] |
| hand_ids: list[int] = [] |
| role_ids: list[int] = [] |
| local_ids: list[int] = [] |
| point_names: list[str] = [] |
|
|
| groups = ( |
| ("head_left", "left_hand", NUM_HEAD_LEFT, HAND_LEFT, ROLE_HEAD_HAND), |
| ("head_left", "right_hand", NUM_HEAD_RIGHT, HAND_RIGHT, ROLE_HEAD_HAND), |
| ( |
| "left_wrist", |
| "background", |
| NUM_WRIST_BACKGROUND, |
| HAND_NONE, |
| ROLE_WRIST_BACKGROUND, |
| ), |
| ("left_wrist", "left_hand", NUM_WRIST_HAND, HAND_LEFT, ROLE_WRIST_HAND), |
| ( |
| "right_wrist", |
| "background", |
| NUM_WRIST_BACKGROUND, |
| HAND_NONE, |
| ROLE_WRIST_BACKGROUND, |
| ), |
| ("right_wrist", "right_hand", NUM_WRIST_HAND, HAND_RIGHT, ROLE_WRIST_HAND), |
| ) |
| for view, component, count, hand_id, role_id in groups: |
| view_ids.extend([VIEW_IDS[view]] * count) |
| hand_ids.extend([hand_id] * count) |
| role_ids.extend([role_id] * count) |
| local_ids.extend(range(count)) |
| point_names.extend(f"{view}.{component}.{index:03d}" for index in range(count)) |
|
|
| lengths = { |
| len(view_ids), |
| len(hand_ids), |
| len(role_ids), |
| len(local_ids), |
| len(point_names), |
| } |
| if lengths != {NUM_COMBINED_POINTS}: |
| raise AssertionError(f"invalid identity lengths: {sorted(lengths)}") |
| return { |
| "view_ids": view_ids, |
| "hand_ids": hand_ids, |
| "role_ids": role_ids, |
| "local_ids": local_ids, |
| "global_ids": list(range(NUM_COMBINED_POINTS)), |
| "point_names": point_names, |
| } |
|
|
|
|
| def layout_metadata() -> dict[str, object]: |
| """Return the JSON-serializable schema metadata stored with the dataset.""" |
|
|
| return { |
| "version": TRACK_LAYOUT_VERSION, |
| "coordinate_space": "normalized_xy_div_wh", |
| "point_value_order": ["x", "y", "visibility"], |
| "view_order": list(VIEW_ORDER), |
| "point_slices": list(POINT_SLICES), |
| "view_slices": {key: list(value) for key, value in VIEW_SLICES.items()}, |
| "component_slices": { |
| key: list(value) for key, value in COMPONENT_SLICES.items() |
| }, |
| "view_point_counts": dict(VIEW_POINT_COUNTS), |
| "hand_id_names": list(HAND_NAMES), |
| "role_id_names": list(ROLE_NAMES), |
| **identity_metadata(), |
| } |
|
|
|
|
| if NUM_COMBINED_POINTS != 250 or POINT_SLICES != (0, 100, 175, 250): |
| raise AssertionError("the canonical T-Rex track layout must contain 250 points") |
|
|