File size: 13,210 Bytes
b9c996e 0782212 b9c996e 0782212 b9c996e 0782212 b9c996e c9de639 b9c996e c9de639 0782212 b9c996e 84332ee b9c996e c9de639 0782212 b9c996e 84332ee c9de639 b9c996e 84332ee b9c996e c9de639 b9c996e c9de639 39ef827 c9de639 b9c996e 0782212 b9c996e 0782212 b9c996e 0782212 b9c996e 0782212 b9c996e 0782212 b9c996e 0782212 b9c996e 0782212 b9c996e | 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 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 | """Tremor v0.1 inference — search body-worn motion in natural language.
Tremor maps a short window of accelerometer motion into the Qwen3-VL-Embedding-2B
text space, so a stream of inertial motion becomes retrievable with plain-language
activity queries ("walking upstairs", "sitting", "picking something up").
Pipeline: a raw 3-axis accelerometer window (in m/s^2, gravity ~9.8) is resampled to a
200-sample, 20 Hz window and placed at one fixed joint of a frozen UniMTS ST-GCN encoder,
which outputs a 512-d motion feature; a small trained projector (this repository's only
weights) maps that to the frozen Qwen base's 2048-d text space. Motion and text embeddings
are L2-normalized and compared by cosine.
Give the encoder the window the way it was trained: pass your sensor's true sample rate so
the window is resampled to 20 Hz and cropped/wrapped to 200 samples (10 s), and supply the
accelerometer in m/s^2 (pass unit="g" if your sensor reports g). Both are optional and the
API stays backward-compatible: with no sample rate the window is linearly resampled to 200
samples, as before.
from inference import TremorEmbedder
tr = TremorEmbedder.from_pretrained("EximiusLabs/fusion-embedding-2-tremor")
m = tr.embed_motion(accel, sample_rate_hz=50) # accel: np.ndarray [3, T] in m/s^2
scores = tr.rank(accel, ["walking", "sitting", "running", "climbing stairs"],
sample_rate_hz=50)
Requirements:
- torch (CUDA recommended), numpy, transformers>=4.46, huggingface_hub, scipy
- The frozen UniMTS encoder code and weights (Apache-2.0):
git clone https://github.com/xiyuanzh/UniMTS # provides model.py::ST_GCN_18
Point UNIMTS_REPO at the clone, or pass unimts_repo= to from_pretrained.
The UniMTS weights (checkpoint/UniMTS.pth) download automatically from the HF hub.
The frozen Qwen3-VL-Embedding-2B base downloads from its original repository. Embedding
quality is sensitive to the base's chat-template formatting; use the template provided
here rather than constructing your own.
"""
from __future__ import annotations
import json
import os
import sys
from typing import List, Optional, Sequence, Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
BASE_MODEL = "Qwen/Qwen3-VL-Embedding-2B"
UNIMTS_REPO = os.environ.get("UNIMTS_REPO", "UniMTS") # local clone of xiyuanzh/UniMTS
CKPT_FILE = "tremor_projector.pt"
SAFETENSORS_FILE = "model.safetensors" # projector weights; config in the file's metadata
GRAVITY = 9.80665 # m/s^2 per g
DEFAULT_JOINT = 5 # fixed single joint for a general single-IMU mount,
# selected on in-domain held-out data disjoint from
# the reported zero-shot sets; supersedes the
# checkpoint's legacy 'joint' field
def _chat(text: str) -> str:
"""The Qwen base's embedding chat-template. Motion is matched against text embedded this way."""
return ("<|im_start|>system\nYou are a helpful assistant.<|im_end|>\n"
f"<|im_start|>user\n{text}<|im_end|>\n<|im_start|>assistant\n")
def _build_projector(in_dim: int, out_dim: int, arch: str = "ln") -> nn.Sequential:
"""Projector head. 'ln' = the general base (LayerNorm front); 'plain' = the per-fleet
variants (e.g. Tremor-G1), which omit the input LayerNorm. Selected by the checkpoint config."""
if arch == "plain":
return nn.Sequential(nn.Linear(in_dim, 1024), nn.GELU(),
nn.Dropout(0.3), nn.Linear(1024, out_dim))
return nn.Sequential(nn.LayerNorm(in_dim), nn.Linear(in_dim, 1024), nn.GELU(),
nn.Dropout(0.3), nn.Linear(1024, out_dim))
class TremorEmbedder:
def __init__(self, ckpt_path: str, device: str = "cuda", dtype=torch.bfloat16,
unimts_repo: str = UNIMTS_REPO):
from huggingface_hub import hf_hub_download
from transformers import AutoModel, AutoTokenizer
self.device = device
if ckpt_path.endswith(".safetensors"):
import safetensors.torch as st
from safetensors import safe_open
proj_sd = st.load_file(ckpt_path)
with safe_open(ckpt_path, framework="pt") as f:
self.cfg = json.loads((f.metadata() or {})["config"])
else: # legacy .pt checkpoint
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
self.cfg = ck["config"]; proj_sd = ck["proj"]
# Placement joint. The general base ('ln' arch) mounts a single IMU at DEFAULT_JOINT,
# a fixed skeleton joint selected on in-domain data disjoint from the reported zero-shot
# sets. Per-fleet heads ('plain' arch, e.g. Tremor-G1) are trained at, and keep, their
# own joint read from the checkpoint.
self.joint = self.cfg["joint"] if self.cfg.get("arch") == "plain" else DEFAULT_JOINT
self.n_joints = self.cfg["n_joints"]
self.win = self.cfg["window_samples"]
# frozen UniMTS motion encoder (accelerometer-only ST-GCN)
if not os.path.isdir(unimts_repo):
raise FileNotFoundError(
f"UniMTS repo not found at '{unimts_repo}'. Clone it and set unimts_repo=/UNIMTS_REPO:\n"
" git clone https://github.com/xiyuanzh/UniMTS")
sys.path.insert(0, unimts_repo)
from model import ST_GCN_18
w = hf_hub_download(self.cfg["unimts_repo"], self.cfg["unimts_file"])
sd = torch.load(w, map_location="cpu", weights_only=False)
sd = sd.get("state_dict", sd)
acc = {k[len("acc."):]: v for k, v in sd.items() if k.startswith("acc.")}
enc = ST_GCN_18(in_channels=3)
enc.load_state_dict(acc, strict=False)
self.enc = enc.eval().to(device)
for p in self.enc.parameters():
p.requires_grad_(False)
# trained projector (this repository)
self.proj = _build_projector(self.cfg["in_dim"], self.cfg["out_dim"],
self.cfg.get("arch", "ln")).to(device)
self.proj.load_state_dict({k: v.float() for k, v in proj_sd.items()})
self.proj.eval()
# frozen Qwen text space
self.base = AutoModel.from_pretrained(BASE_MODEL, trust_remote_code=True,
dtype=dtype).to(device).eval()
for p in self.base.parameters():
p.requires_grad_(False)
self.tok = AutoTokenizer.from_pretrained(BASE_MODEL, trust_remote_code=True)
self.tok.padding_side = "right"
@classmethod
def from_pretrained(cls, repo_or_path: str, device: str = "cuda",
revision: Optional[str] = None, unimts_repo: str = UNIMTS_REPO,
subfolder: str = "", **kw) -> "TremorEmbedder":
"""Load from a local checkpoint (tremor_projector.pt) or an HF repo.
subfolder selects a variant that ships alongside the general base in the same repo, e.g.
the Unitree-specialized head: from_pretrained("EximiusLabs/fusion-embedding-2-tremor",
subfolder="g1")."""
if os.path.isfile(repo_or_path):
path = repo_or_path
elif os.path.isdir(repo_or_path):
cand = os.path.join(repo_or_path, subfolder, SAFETENSORS_FILE)
path = cand if os.path.exists(cand) else os.path.join(repo_or_path, subfolder, CKPT_FILE)
else:
from huggingface_hub import hf_hub_download
from huggingface_hub.utils import EntryNotFoundError
cfg_fn = f"{subfolder}/config.json" if subfolder else "config.json"
try: # fetch config.json so a real load registers a Hub download
hf_hub_download(repo_or_path, cfg_fn, revision=revision)
except Exception:
pass
st_fn = f"{subfolder}/{SAFETENSORS_FILE}" if subfolder else SAFETENSORS_FILE
try: # prefer safetensors, fall back to the legacy pickle
path = hf_hub_download(repo_or_path, st_fn, revision=revision)
except EntryNotFoundError:
pt_fn = f"{subfolder}/{CKPT_FILE}" if subfolder else CKPT_FILE
path = hf_hub_download(repo_or_path, pt_fn, revision=revision)
return cls(path, device=device, unimts_repo=unimts_repo, **kw)
# --------------------------------------------------------------- motion
def _prep(self, accel: np.ndarray, sample_rate_hz: Optional[float] = None,
unit: str = "m/s2") -> torch.Tensor:
"""Raw accelerometer [3, T] -> encoder input [1, 3, win, n_joints, 1] at one joint.
Units: the encoder expects m/s^2 (gravity reads ~9.8 at rest), the convention UniMTS
was pretrained on. Pass unit="g" if your sensor reports g (gravity ~1.0) and the window
is scaled to m/s^2. The default assumes m/s^2 and applies no scaling; unknown input is
never silently rescaled.
Sampling rate: pass the sensor's true rate (Hz) and the window is resampled to the
model's 20 Hz, then wrap-padded or truncated to exactly `win` samples (a 10 s window),
matching how the model was trained and evaluated. With no rate the window is linearly
resampled to `win` samples (rate-agnostic fallback, backward-compatible)."""
a = np.asarray(accel, dtype="float64")
if a.ndim != 2 or a.shape[0] != 3:
raise ValueError(f"expected accelerometer of shape [3, T], got {list(a.shape)}")
if unit == "g":
a = a * GRAVITY
elif unit not in ("m/s2", "m/s^2"):
raise ValueError("unit must be 'm/s2' (default) or 'g'")
if sample_rate_hz:
from scipy.signal import resample
n = max(1, int((a.shape[1] / float(sample_rate_hz)) * 20)) # -> 20 Hz length
a = resample(a, n, axis=1) # [3, n]
if n < self.win: # wrap-pad short windows
a = np.pad(a, ((0, 0), (0, self.win - n)), "wrap")
a = a[:, :self.win] # truncate to the 10 s window
t = torch.as_tensor(a, dtype=torch.float32).unsqueeze(0) # [1,3,win]
else:
t = torch.as_tensor(a, dtype=torch.float32).unsqueeze(0) # [1,3,T]
if t.shape[-1] != self.win:
t = F.interpolate(t, size=self.win, mode="linear", align_corners=False)
g = torch.zeros(1, 3, self.win, self.n_joints, 1)
g[:, :, :, self.joint, 0] = t
return g.to(self.device)
@torch.no_grad()
def embed_motion(self, accel: np.ndarray, sample_rate_hz: Optional[float] = None,
unit: str = "m/s2") -> torch.Tensor:
"""Embed a raw 3-axis accelerometer window (np.ndarray [3, T], any length).
Pass sample_rate_hz (the sensor's true rate) for the trained-time 20 Hz / 200-sample
handling, and unit="g" if the sensor reports g rather than m/s^2. See _prep."""
feat = self.enc(self._prep(accel, sample_rate_hz, unit)).squeeze(-1).squeeze(-1).float()
return F.normalize(self.proj(feat), dim=-1).squeeze(0).cpu()
# --------------------------------------------------------------- text
@torch.no_grad()
def embed_text(self, text: Union[str, Sequence[str]]) -> torch.Tensor:
one = isinstance(text, str)
texts = [text] if one else list(text)
out = []
for t in texts: # padding-free (base is sensitive to padding)
enc = self.tok(_chat(t), return_tensors="pt", truncation=True, max_length=64).to(self.device)
h = self.base(**enc).last_hidden_state
idx = int(enc["attention_mask"].sum().item()) - 1
out.append(F.normalize(h[0, idx].float(), dim=-1).cpu())
e = torch.stack(out)
return e.squeeze(0) if one else e
# --------------------------------------------------------------- readout
@torch.no_grad()
def rank(self, accel: np.ndarray, texts: Sequence[str],
sample_rate_hz: Optional[float] = None, unit: str = "m/s2") -> List[tuple]:
"""Rank candidate activity texts by cosine similarity to a motion window.
Pass sample_rate_hz / unit through to embed_motion (see _prep).
Returns [(text, score), ...] sorted high to low."""
m = self.embed_motion(accel, sample_rate_hz, unit)
te = self.embed_text(list(texts))
scores = (te @ m).tolist()
return sorted(zip(texts, scores), key=lambda x: -x[1])
if __name__ == "__main__":
# smoke: random motion, four candidate activities (needs the models + UniMTS clone to run)
tr = TremorEmbedder.from_pretrained(os.path.join(os.path.dirname(__file__), "out"))
demo = np.random.randn(3, 300).astype("float32")
for text, score in tr.rank(demo, ["walking", "sitting", "running", "climbing stairs"]):
print(f"{score:+.3f} {text}")
|