File size: 8,896 Bytes
5261696
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
"""Raw-IMU preprocessing and prompt construction for AnyMo."""

from __future__ import annotations

from typing import Any, Sequence

import numpy as np
import torch
from scipy.signal import resample
from transformers import ProcessorMixin
from transformers.feature_extraction_utils import BatchFeature

from .modeling_components import SEGMENT_NAMES


IMU_BOS_TOKEN = "<imu_bos>"
IMU_EOS_TOKEN = "<imu_eos>"

IMU_CONTRASTIVE_TEMPLATE = (
    "Represent the human motion from the wearable IMU motion tokens.\n\n"
    "The IMU tokens are from IMU sensors attached to the user's {sensor_context}.\n\n"
    "Input IMU token:\n{imu_token}\n\nReturn a compact embedding of the motion."
)
TEXT_CONTRASTIVE_PREFIX = "Represent the human motion described by the text.\n\nMotion description:\n"
TEXT_CONTRASTIVE_SUFFIX = "\n\nReturn a compact embedding of the motion."
CAPTION_TEMPLATE = (
    "Describe the human motion represented by the wearable IMU motion tokens.\n\n"
    "The IMU tokens are from IMU sensors attached to the user's {sensor_context}.\n\n"
    "Input IMU token:\n{imu_token}"
)
HAR_TEMPLATE = (
    "Recognize the activity represented by the wearable IMU motion tokens.\n\n"
    "The IMU tokens are from IMU sensors attached to the user's {sensor_context}.\n\n"
    "Input IMU token:\n{imu_token}\n\nCHOICES:\n{choices}\n\n"
    "Choose the best matching option. Output the option key followed by the selected activity label."
)


DEFAULT_LOCATION_ALIASES = {
    "head": "Head",
    "neck": "Neck",
    "pelvis": "Pelvis",
    "waist": "Pelvis",
    "lower back": "L5",
    "chest": "T8",
    "left shoulder": "L_Shoulder",
    "right shoulder": "R_Shoulder",
    "left upper arm": "L_UpperArm",
    "right upper arm": "R_UpperArm",
    "left forearm": "L_Forearm",
    "right forearm": "R_Forearm",
    "left wrist": "L_Forearm",
    "right wrist": "R_Forearm",
    "left hand": "L_Hand",
    "right hand": "R_Hand",
    "left thigh": "L_UpperLeg",
    "right thigh": "R_UpperLeg",
    "left lower leg": "L_LowerLeg",
    "right lower leg": "R_LowerLeg",
    "left ankle": "L_LowerLeg",
    "right ankle": "R_LowerLeg",
    "left foot": "L_Foot",
    "right foot": "R_Foot",
}


class AnyMoProcessor(ProcessorMixin):
    """Maps raw wearable IMU streams to the AnyMo graph and prompt formats."""

    tokenizer_class = "AutoTokenizer"
    valid_kwargs = ["target_sample_rate_hz", "location_aliases", "channel_order"]

    def __init__(
        self,
        tokenizer,
        target_sample_rate_hz: int = 60,
        location_aliases: dict[str, str] | None = None,
        channel_order: Sequence[str] = ("acc_x", "acc_y", "acc_z", "gyro_x", "gyro_y", "gyro_z"),
    ):
        super().__init__(tokenizer=tokenizer)
        self.target_sample_rate_hz = int(target_sample_rate_hz)
        self.location_aliases = dict(DEFAULT_LOCATION_ALIASES)
        if location_aliases:
            self.location_aliases.update(
                {str(key).strip().lower(): str(value) for key, value in location_aliases.items()}
            )
        self.channel_order = list(channel_order)

    def _canonical_segment(self, location: str) -> str:
        raw = str(location).strip()
        if raw in SEGMENT_NAMES:
            return raw
        normalized = raw.replace("_", " ").strip().lower()
        if normalized in self.location_aliases:
            return self.location_aliases[normalized]
        matches = [name for name in SEGMENT_NAMES if name.replace("_", " ").lower() == normalized]
        if matches:
            return matches[0]
        raise ValueError(
            f"Unknown sensor location {location!r}. Use one of the canonical 23 segments "
            f"or provide an explicit location alias."
        )

    def prepare_imu(
        self,
        imu: np.ndarray | torch.Tensor,
        sensor_locations: Sequence[str] | Sequence[Sequence[str]],
        sampling_rate: float | Sequence[float],
        *,
        return_tensors: str = "pt",
    ) -> BatchFeature:
        values = imu.detach().cpu().numpy() if isinstance(imu, torch.Tensor) else np.asarray(imu)
        if values.ndim == 3:
            values = values[None]
        if values.ndim != 4 or values.shape[-1] != 6:
            raise ValueError(f"Expected IMU shape [T, S, 6] or [B, T, S, 6], got {values.shape}")
        batch_size, _, sensor_count, _ = values.shape
        if sensor_locations and isinstance(sensor_locations[0], str):
            locations = [list(sensor_locations)] * batch_size
        else:
            locations = [list(row) for row in sensor_locations]
        if len(locations) != batch_size or any(len(row) != sensor_count for row in locations):
            raise ValueError("sensor_locations must contain one location for every sensor in every sample")
        rates = [float(sampling_rate)] * batch_size if np.isscalar(sampling_rate) else [float(x) for x in sampling_rate]
        if len(rates) != batch_size:
            raise ValueError("sampling_rate must be a scalar or contain one value per sample")

        graph_rows, mask_rows, context_rows = [], [], []
        for sample, sample_locations, rate in zip(values, locations, rates):
            if rate <= 0:
                raise ValueError(f"sampling_rate must be positive, got {rate}")
            target_length = max(1, int(round(sample.shape[0] * self.target_sample_rate_hz / rate)))
            if target_length != sample.shape[0]:
                sample = np.asarray(resample(sample, target_length, axis=0), dtype=np.float32)
            else:
                sample = np.asarray(sample, dtype=np.float32)
            graph = np.zeros((6, sample.shape[0], len(SEGMENT_NAMES), 1), dtype=np.float32)
            mask = np.zeros(len(SEGMENT_NAMES), dtype=bool)
            canonical = [self._canonical_segment(name) for name in sample_locations]
            if len(set(canonical)) != len(canonical):
                raise ValueError("Multiple sensors map to the same graph segment; provide distinct canonical segments")
            for sensor_index, segment in enumerate(canonical):
                node_index = SEGMENT_NAMES.index(segment)
                graph[:, :, node_index, 0] = sample[:, sensor_index, :].T
                mask[node_index] = True
            graph_rows.append(graph)
            mask_rows.append(mask)
            context_rows.append(self.format_sensor_context(canonical))

        if len({row.shape[1] for row in graph_rows}) != 1:
            raise ValueError("A batch must have equal resampled sequence lengths")
        data = {
            "imu_values": torch.from_numpy(np.stack(graph_rows)),
            "visible_node_mask": torch.from_numpy(np.stack(mask_rows)),
            "sensor_context": context_rows,
        }
        if return_tensors != "pt":
            data["imu_values"] = data["imu_values"].numpy()
            data["visible_node_mask"] = data["visible_node_mask"].numpy()
        return BatchFeature(data=data, tensor_type=None)

    @staticmethod
    def format_sensor_context(segments: Sequence[str]) -> str:
        names = []
        for segment in segments:
            if segment.startswith("L_"):
                segment = "Left " + segment[2:]
            elif segment.startswith("R_"):
                segment = "Right " + segment[2:]
            names.append(segment.replace("_", " "))
        return ", ".join(names) if names else "visible body segments"

    @staticmethod
    def local_ids_to_text(local_ids: Sequence[int]) -> str:
        tokens = "".join(f"<imu_{int(index):04d}>" for index in local_ids)
        return f"{IMU_BOS_TOKEN}{tokens}{IMU_EOS_TOKEN}"

    def build_motion_prompt(self, local_ids: Sequence[int], sensor_context: str) -> str:
        return IMU_CONTRASTIVE_TEMPLATE.format(
            imu_token=self.local_ids_to_text(local_ids), sensor_context=sensor_context
        )

    @staticmethod
    def build_text_prompt(text: str) -> str:
        return f"{TEXT_CONTRASTIVE_PREFIX}{text}{TEXT_CONTRASTIVE_SUFFIX}"

    def build_caption_prompt(self, local_ids: Sequence[int], sensor_context: str) -> str:
        content = CAPTION_TEMPLATE.format(
            imu_token=self.local_ids_to_text(local_ids), sensor_context=sensor_context
        )
        return f"<|im_start|>user\n{content}<|im_end|>\n<|im_start|>assistant\n"

    def build_har_prompt(
        self, local_ids: Sequence[int], sensor_context: str, candidate_labels: Sequence[str]
    ) -> str:
        choices = "\n".join(f"{chr(65 + index)}: {label}" for index, label in enumerate(candidate_labels))
        content = HAR_TEMPLATE.format(
            imu_token=self.local_ids_to_text(local_ids),
            sensor_context=sensor_context,
            choices=choices,
        )
        return f"<|im_start|>user\n{content}<|im_end|>\n<|im_start|>assistant\n"

    def __call__(self, *args, **kwargs) -> BatchFeature:
        return self.prepare_imu(*args, **kwargs)