"""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_EOS_TOKEN = "" 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"" 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)