AnyMo / processing_anymo.py
Breezelled's picture
Release AnyMo model and raw-IMU pipeline
5261696 verified
Raw History Blame Contribute Delete
8.9 kB
"""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)