File size: 4,411 Bytes
fbd9366
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""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")

# Half-open slices in the canonical concatenation order.
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")