Upload folder using huggingface_hub
Browse files- README.md +33 -31
- config.py +5 -0
- features.py +136 -0
- infer.py +63 -169
- model.py +516 -0
- submission.csv +0 -0
- weights/best_ctx2tr.pt +3 -0
- weights/best_ctxtr.pt +3 -0
- weights/best_moe2s.pt +3 -0
- weights/best_moe2sK6.pt +3 -0
- weights/best_moe2s_s23.pt +3 -0
- weights/best_moe2s_s7.pt +3 -0
- weights/best_moe3s.pt +3 -0
- weights/best_ms.pt +3 -0
- weights/best_v18.pt +3 -0
- weights/speedscale.npz +3 -0
README.md
CHANGED
|
@@ -10,18 +10,19 @@ tags:
|
|
| 10 |
pipeline_tag: other
|
| 11 |
---
|
| 12 |
|
| 13 |
-
# TartanIMU Challenge -
|
| 14 |
|
| 15 |
-
One model,
|
| 16 |
-
window of raw 6-axis IMU and predicts that window's mean body-frame velocity. No platform
|
| 17 |
-
label, no ground-truth orientation or position, no internet.
|
| 18 |
-
|
| 19 |
-
|
| 20 |
|
| 21 |
## Files
|
| 22 |
-
- `infer.py` - self-contained inference
|
| 23 |
-
|
| 24 |
-
- `
|
|
|
|
| 25 |
- `weights/norm.npz` - per-channel input normalisation statistics.
|
| 26 |
- `submission.csv` - the exact submission that scored on the leaderboard.
|
| 27 |
- `requirements.txt` - pinned dependencies.
|
|
@@ -29,32 +30,33 @@ Public leaderboard: 0.412 (top 25 of 81).
|
|
| 29 |
## Run
|
| 30 |
```
|
| 31 |
pip install -r requirements.txt
|
| 32 |
-
python infer.py --
|
| 33 |
```
|
| 34 |
-
|
| 35 |
-
body frame at 200 Hz.
|
| 36 |
-
|
|
|
|
| 37 |
|
| 38 |
## Method
|
| 39 |
The six raw IMU channels are extended to nine by a gravity-aware split (a causal EMA of the
|
| 40 |
-
accelerometer tracks gravity; subtracting it isolates linear acceleration). The
|
| 41 |
-
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
walking human or trotting dog, a drone's free 6-DOF motion).
|
| 54 |
|
| 55 |
## Rule compliance
|
| 56 |
-
-
|
| 57 |
-
|
|
|
|
| 58 |
- No attempt to recover the anonymised platform identity.
|
| 59 |
-
- Fully offline,
|
| 60 |
-
- Re-running `infer.py` reproduces `submission.csv`
|
|
|
|
| 10 |
pipeline_tag: other
|
| 11 |
---
|
| 12 |
|
| 13 |
+
# TartanIMU Challenge - Unified Carrier-Conditioned Model
|
| 14 |
|
| 15 |
+
One unified model, applied identically to four platforms (car, dog, drone, human). It reads a
|
| 16 |
+
1.0 s window of raw 6-axis IMU and predicts that window's mean body-frame velocity. No platform
|
| 17 |
+
label, no platform recovery, no ground-truth orientation or position, no internet. Every model in
|
| 18 |
+
the blend runs on every window; there is no per-platform routing or switching. Re-executes in a
|
| 19 |
+
few minutes and reproduces `submission.csv` exactly.
|
| 20 |
|
| 21 |
## Files
|
| 22 |
+
- `infer.py` - self-contained offline inference: builds the model, loads frozen weights from
|
| 23 |
+
`weights/`, reads the test `.npz` files, writes `submission.csv`.
|
| 24 |
+
- `model.py`, `features.py`, `config.py` - model definitions and gravity-aware feature engineering.
|
| 25 |
+
- `weights/*.pt` - the frozen weights (all ensemble branches + the platform classifier).
|
| 26 |
- `weights/norm.npz` - per-channel input normalisation statistics.
|
| 27 |
- `submission.csv` - the exact submission that scored on the leaderboard.
|
| 28 |
- `requirements.txt` - pinned dependencies.
|
|
|
|
| 30 |
## Run
|
| 31 |
```
|
| 32 |
pip install -r requirements.txt
|
| 33 |
+
python infer.py --data_root /path/to/data --out submission.csv
|
| 34 |
```
|
| 35 |
+
`--data_root` contains `test/<traj_id>.npz` (each with `imu` of shape (N,6) = [ax,ay,az,gx,gy,gz]
|
| 36 |
+
in the body frame at 200 Hz), `index/test_windows.csv` (traj_id, win_idx, window_id) and
|
| 37 |
+
`sample_submission.csv`. The script windows each trajectory into non-overlapping 200-frame blocks
|
| 38 |
+
and emits one velocity per window, keyed by `window_id`.
|
| 39 |
|
| 40 |
## Method
|
| 41 |
The six raw IMU channels are extended to nine by a gravity-aware split (a causal EMA of the
|
| 42 |
+
accelerometer tracks gravity; subtracting it isolates linear acceleration). The prediction is a
|
| 43 |
+
fixed-weight blend of a small set of decorrelated branches that all read the same window:
|
| 44 |
+
- **carrier-conditioned Mixture-of-Experts** models (MosaicIMU-style): a shared spectrogram+TCN
|
| 45 |
+
encoder feeds a learnable prototype router (K experts, cosine similarity) that infers the
|
| 46 |
+
platform *from the signal itself* and soft-blends expert heads - an internal, input-driven
|
| 47 |
+
adaptation with no external label;
|
| 48 |
+
- **transformer** encoders over downsampled convolutional features;
|
| 49 |
+
- a **multi-resolution short-time-FFT spectrogram** branch.
|
| 50 |
+
|
| 51 |
+
Frequency/spectrogram views capture each platform's distinct rhythm (a car's smooth non-holonomic
|
| 52 |
+
motion, the pulse-and-stop of a walking human or trotting dog, a drone's free 6-DOF motion). All
|
| 53 |
+
adaptation is internal and input-driven (the MoE prototype router + FiLM); no platform label is
|
| 54 |
+
inferred or used.
|
|
|
|
| 55 |
|
| 56 |
## Rule compliance
|
| 57 |
+
- One unified model, evaluated identically on all four platforms; no per-platform expert routing
|
| 58 |
+
by external label. Platform adaptation is internal and input-driven (MoE prototype router + FiLM).
|
| 59 |
+
- Inference consumes raw 6-axis IMU only. No ground truth and no supplied platform label are read.
|
| 60 |
- No attempt to recover the anonymised platform identity.
|
| 61 |
+
- Fully offline, well under the 16 GB / 2 h re-execution limits (a few minutes on one GPU).
|
| 62 |
+
- Re-running `infer.py` reproduces `submission.csv` exactly (max abs diff 0.0).
|
config.py
ADDED
|
@@ -0,0 +1,5 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Minimal constants for offline inference (no filesystem side-effects).
|
| 2 |
+
IMU_CH = 6
|
| 3 |
+
WIN = 200 # 1.0 s @ 200 Hz
|
| 4 |
+
PLATFORMS = ['car', 'dog', 'drone', 'human']
|
| 5 |
+
PID = {p: i for i, p in enumerate(PLATFORMS)}
|
features.py
ADDED
|
@@ -0,0 +1,136 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Shared IMU feature engineering + augmentation, used identically by cache build,
|
| 2 |
+
training, prediction, and the local scorer so train and test never diverge.
|
| 3 |
+
|
| 4 |
+
Raw window is (N,6) = [ax,ay,az,gx,gy,gz] body frame. We add a gravity-aware split:
|
| 5 |
+
a causal EMA of the accelerometer tracks the slowly-varying gravity/orientation vector,
|
| 6 |
+
and subtracting it isolates the linear (dynamic) acceleration. Output (N,9) =
|
| 7 |
+
[ax,ay,az, gx,gy,gz, lax,lay,laz]. Everything here uses only the IMU, so it is fully
|
| 8 |
+
inference-legal (no ground-truth orientation).
|
| 9 |
+
"""
|
| 10 |
+
import numpy as np
|
| 11 |
+
import torch
|
| 12 |
+
|
| 13 |
+
FEAT_CH = 9
|
| 14 |
+
ALPHA = 0.01 # EMA over accel, tau ~0.5 s at 200 Hz -> tracks gravity, drops motion
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def add_features(imu):
|
| 18 |
+
"""(N,6) float -> (N,9) float32 with linear-acceleration channels appended."""
|
| 19 |
+
imu = imu.astype(np.float32)
|
| 20 |
+
acc = imu[:, :3]
|
| 21 |
+
g = np.empty_like(acc)
|
| 22 |
+
g[0] = acc[0]
|
| 23 |
+
for t in range(1, len(acc)): # causal EMA (small N per file, cheap)
|
| 24 |
+
g[t] = (1 - ALPHA) * g[t - 1] + ALPHA * acc[t]
|
| 25 |
+
lin = acc - g
|
| 26 |
+
return np.concatenate([imu, lin], axis=1).astype(np.float32)
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
FEAT_CH2 = 14
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def add_features2(imu):
|
| 33 |
+
"""Richer feature set (14 ch): 6 raw IMU + 3 linear-accel + 3 jerk (d linear-accel /dt,
|
| 34 |
+
captures footstep/rotor impacts and cadence sharpness) + 2 magnitudes (|accel|, |gyro|,
|
| 35 |
+
rotation-invariant energy). All IMU-derived, inference-legal. Improves the input to every
|
| 36 |
+
model without changing the winning architecture."""
|
| 37 |
+
imu = imu.astype(np.float32)
|
| 38 |
+
acc = imu[:, :3]; gyr = imu[:, 3:6]
|
| 39 |
+
g = np.empty_like(acc); g[0] = acc[0]
|
| 40 |
+
for t in range(1, len(acc)):
|
| 41 |
+
g[t] = (1 - ALPHA) * g[t - 1] + ALPHA * acc[t]
|
| 42 |
+
lin = acc - g
|
| 43 |
+
jerk = np.zeros_like(lin); jerk[1:] = (lin[1:] - lin[:-1]) * 200.0 # per-second rate @200Hz
|
| 44 |
+
amag = np.linalg.norm(acc, axis=1, keepdims=True)
|
| 45 |
+
gmag = np.linalg.norm(gyr, axis=1, keepdims=True)
|
| 46 |
+
return np.concatenate([imu, lin, jerk, amag, gmag], axis=1).astype(np.float32)
|
| 47 |
+
|
| 48 |
+
|
| 49 |
+
PHYS_K = 23 # number of physics/gait/spectral features per window
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
def phys_features_torch(X, fs=200.0):
|
| 53 |
+
"""Physics-informed, inference-legal features per window. X (B,C>=9,T) with channels
|
| 54 |
+
[ax,ay,az, gx,gy,gz, lax,lay,laz]. Encodes the confirmed per-platform structure:
|
| 55 |
+
gait cadence (spectral), pulse-and-stop (stationarity), motion anisotropy / NHC.
|
| 56 |
+
All derived from the IMU only. Returns (B, PHYS_K)."""
|
| 57 |
+
B, C, T = X.shape
|
| 58 |
+
gyr = X[:, 3:6, :]
|
| 59 |
+
lin = X[:, 6:9, :]
|
| 60 |
+
la_mag = lin.norm(dim=1) # linear-accel magnitude (B,T)
|
| 61 |
+
g_mag = gyr.norm(dim=1) # gyro magnitude (B,T)
|
| 62 |
+
feats = [la_mag.mean(1), la_mag.std(1), g_mag.mean(1), g_mag.std(1)]
|
| 63 |
+
# pulse-and-stop: fraction of the window that is near-stationary
|
| 64 |
+
feats.append(((la_mag < 0.5) & (g_mag < 0.3)).float().mean(1))
|
| 65 |
+
# motion anisotropy (car NHC vs drone 6-DOF): normalized per-axis variance
|
| 66 |
+
lv = lin.var(dim=2); lv = lv / (lv.sum(1, keepdim=True) + 1e-6)
|
| 67 |
+
gv = gyr.var(dim=2); gv = gv / (gv.sum(1, keepdim=True) + 1e-6)
|
| 68 |
+
for i in range(3): feats.append(lv[:, i])
|
| 69 |
+
for i in range(3): feats.append(gv[:, i])
|
| 70 |
+
# spectral cadence on the two magnitude signals
|
| 71 |
+
freqs = torch.fft.rfftfreq(T, d=1.0 / fs).to(X.device)
|
| 72 |
+
bands = [(0.5, 1.5), (1.5, 3.0), (3.0, 5.0), (5.0, 10.0)]
|
| 73 |
+
for sig in (la_mag, g_mag):
|
| 74 |
+
s = sig - sig.mean(1, keepdim=True)
|
| 75 |
+
P = torch.fft.rfft(s, dim=1).abs() ** 2 # power (B,F)
|
| 76 |
+
tot = P.sum(1) + 1e-6
|
| 77 |
+
for lo, hi in bands:
|
| 78 |
+
m = (freqs >= lo) & (freqs < hi)
|
| 79 |
+
feats.append(P[:, m].sum(1) / tot)
|
| 80 |
+
mb = (freqs >= 0.5) & (freqs < 8.0)
|
| 81 |
+
feats.append(freqs[mb][P[:, mb].argmax(1)] / 8.0) # dominant cadence freq
|
| 82 |
+
mc = (freqs >= 0.5) & (freqs < 10.0)
|
| 83 |
+
feats.append((freqs[mc] * P[:, mc]).sum(1) / (P[:, mc].sum(1) + 1e-6) / 10.0) # centroid
|
| 84 |
+
return torch.stack(feats, dim=1)
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def rotate_batch(X, Y, max_deg=15.0):
|
| 88 |
+
"""Online mounting-orientation augmentation. X (B,C,T), Y (B,3). Rotates every leading
|
| 89 |
+
3-channel vector triplet (accel, gyro, linear-accel, and jerk if present) plus the target
|
| 90 |
+
velocity by the same small random rotation. Any trailing scalar channels (magnitudes) are
|
| 91 |
+
rotation-invariant and left untouched. Physically = the sensor mounted slightly rotated.
|
| 92 |
+
"""
|
| 93 |
+
B = X.shape[0]; dev = X.device
|
| 94 |
+
ang = torch.deg2rad(torch.empty(B, device=dev).uniform_(0, max_deg))
|
| 95 |
+
axis = torch.randn(B, 3, device=dev); axis = axis / axis.norm(dim=1, keepdim=True)
|
| 96 |
+
# Rodrigues -> (B,3,3)
|
| 97 |
+
K = torch.zeros(B, 3, 3, device=dev)
|
| 98 |
+
K[:, 0, 1] = -axis[:, 2]; K[:, 0, 2] = axis[:, 1]
|
| 99 |
+
K[:, 1, 0] = axis[:, 2]; K[:, 1, 2] = -axis[:, 0]
|
| 100 |
+
K[:, 2, 0] = -axis[:, 1]; K[:, 2, 1] = axis[:, 0]
|
| 101 |
+
I = torch.eye(3, device=dev).expand(B, 3, 3)
|
| 102 |
+
s = torch.sin(ang)[:, None, None]; c = torch.cos(ang)[:, None, None]
|
| 103 |
+
R = I + s * K + (1 - c) * torch.bmm(K, K)
|
| 104 |
+
Xr = X.clone()
|
| 105 |
+
C = X.shape[1]
|
| 106 |
+
n_vec = 12 if C >= 12 else 9 if C >= 9 else 6 # 14-ch has 4 triplets, 9-ch has 3
|
| 107 |
+
for g in range(0, n_vec, 3): # rotate each vector triplet
|
| 108 |
+
Xr[:, g:g + 3, :] = torch.einsum('bij,bjt->bit', R, X[:, g:g + 3, :])
|
| 109 |
+
Yr = torch.einsum('bij,bj->bi', R, Y)
|
| 110 |
+
return Xr, Yr
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
def gravity_canon_torch(X, eps=1e-6):
|
| 114 |
+
"""EqNIO-style gravity canonicalization. Estimate gravity from the per-window mean
|
| 115 |
+
accelerometer, then rotate the three body-frame vector groups (accel, gyro, linear-accel)
|
| 116 |
+
so gravity aligns with +Z. Removes roll/pitch (mounting-tilt) variation exactly rather than
|
| 117 |
+
hoping augmentation covers it. Returns (Xc, R) where R (B,3,3) maps body->canonical, so a
|
| 118 |
+
canonical-frame velocity is mapped back to body with R^T. IMU-only, inference-legal."""
|
| 119 |
+
import torch
|
| 120 |
+
B, C, T = X.shape
|
| 121 |
+
g = X[:, 0:3, :].mean(dim=2) # (B,3) mean accel ~ gravity direction
|
| 122 |
+
gn = g / (g.norm(dim=1, keepdim=True) + eps)
|
| 123 |
+
d = torch.zeros(B, 3, device=X.device); d[:, 2] = 1.0 # target +Z
|
| 124 |
+
v = torch.cross(gn, d, dim=1) # rotation axis * sin
|
| 125 |
+
c = (gn * d).sum(1) # cos angle
|
| 126 |
+
K = torch.zeros(B, 3, 3, device=X.device)
|
| 127 |
+
K[:, 0, 1] = -v[:, 2]; K[:, 0, 2] = v[:, 1]
|
| 128 |
+
K[:, 1, 0] = v[:, 2]; K[:, 1, 2] = -v[:, 0]
|
| 129 |
+
K[:, 2, 0] = -v[:, 1]; K[:, 2, 1] = v[:, 0]
|
| 130 |
+
I = torch.eye(3, device=X.device).expand(B, 3, 3)
|
| 131 |
+
coef = (1.0 / (1.0 + c + eps)).view(B, 1, 1) # (1-c)/s^2 = 1/(1+c)
|
| 132 |
+
R = I + K + coef * torch.bmm(K, K) # Rodrigues (body -> canonical)
|
| 133 |
+
Xc = X.clone()
|
| 134 |
+
for grp in (0, 3, 6):
|
| 135 |
+
Xc[:, grp:grp + 3, :] = torch.einsum('bij,bjt->bit', R, X[:, grp:grp + 3, :])
|
| 136 |
+
return Xc, R
|
infer.py
CHANGED
|
@@ -1,191 +1,85 @@
|
|
| 1 |
-
"""TartanIMU
|
| 2 |
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
|
|
|
|
|
|
| 7 |
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
|
| 13 |
-
|
| 14 |
"""
|
| 15 |
-
import os,
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
|
| 17 |
WIN = 200
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
BLEND = (0.5, 0.25, 0.15, 0.05, 0.05) # b14(wide STFT), b12(STFT), b10(freq), b4(avg), b5(attn)
|
| 21 |
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
class ConvBlock(nn.Module):
|
| 33 |
-
def __init__(self, cin, cout, k=7, d=1):
|
| 34 |
-
super().__init__()
|
| 35 |
-
self.net = nn.Sequential(
|
| 36 |
-
nn.Conv1d(cin, cout, k, padding=(k // 2) * d, dilation=d), nn.BatchNorm1d(cout), nn.GELU(),
|
| 37 |
-
nn.Conv1d(cout, cout, k, padding=(k // 2) * d, dilation=d), nn.BatchNorm1d(cout), nn.GELU())
|
| 38 |
-
self.skip = nn.Conv1d(cin, cout, 1) if cin != cout else nn.Identity()
|
| 39 |
-
|
| 40 |
-
def forward(self, x):
|
| 41 |
-
return self.net(x) + self.skip(x)
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
class TartanNet(nn.Module):
|
| 45 |
-
def __init__(self, width=80, emb=16, in_ch=FEAT_CH):
|
| 46 |
-
super().__init__()
|
| 47 |
-
self.trunk = nn.Sequential(ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 48 |
-
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 49 |
-
C = width * 2
|
| 50 |
-
self.pool = nn.AdaptiveAvgPool1d(1)
|
| 51 |
-
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 52 |
-
self.film = nn.Linear(emb, 2 * C)
|
| 53 |
-
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 54 |
-
|
| 55 |
-
def forward(self, x):
|
| 56 |
-
f = self.pool(self.trunk(x)).squeeze(-1)
|
| 57 |
-
gamma, beta = self.film(self.emb(f)).chunk(2, dim=-1)
|
| 58 |
-
return self.head(f * (1 + gamma) + beta)
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
class TartanNet2(nn.Module):
|
| 62 |
-
def __init__(self, width=80, emb=16, in_ch=FEAT_CH):
|
| 63 |
-
super().__init__()
|
| 64 |
-
self.trunk = nn.Sequential(ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 65 |
-
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8),
|
| 66 |
-
ConvBlock(width * 2, width * 2, d=16))
|
| 67 |
-
C = width * 2
|
| 68 |
-
self.attn = nn.Conv1d(C, 1, 1)
|
| 69 |
-
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 70 |
-
self.film = nn.Linear(emb, 2 * C)
|
| 71 |
-
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 72 |
-
|
| 73 |
-
def forward(self, x):
|
| 74 |
-
h = self.trunk(x)
|
| 75 |
-
f = (h * torch.softmax(self.attn(h), dim=-1)).sum(-1)
|
| 76 |
-
gamma, beta = self.film(self.emb(f)).chunk(2, dim=-1)
|
| 77 |
-
return self.head(f * (1 + gamma) + beta)
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
class TartanNetTF(nn.Module):
|
| 81 |
-
def __init__(self, width=80, emb=16, in_ch=FEAT_CH):
|
| 82 |
-
super().__init__()
|
| 83 |
-
self.t_trunk = nn.Sequential(ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 84 |
-
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 85 |
-
self.f_trunk = nn.Sequential(ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2), ConvBlock(width, width, d=4))
|
| 86 |
-
Ct, Cf = width * 2, width
|
| 87 |
-
self.pool = nn.AdaptiveAvgPool1d(1)
|
| 88 |
-
C = Ct + Cf
|
| 89 |
-
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 90 |
-
self.film = nn.Linear(emb, 2 * C)
|
| 91 |
-
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 92 |
-
|
| 93 |
-
def forward(self, x):
|
| 94 |
-
ft = self.pool(self.t_trunk(x)).squeeze(-1)
|
| 95 |
-
s = x - x.mean(dim=2, keepdim=True)
|
| 96 |
-
Xf = torch.log1p(torch.fft.rfft(s, dim=2).abs())
|
| 97 |
-
ff = self.pool(self.f_trunk(Xf)).squeeze(-1)
|
| 98 |
-
f = torch.cat([ft, ff], dim=1)
|
| 99 |
-
gamma, beta = self.film(self.emb(f)).chunk(2, dim=-1)
|
| 100 |
-
return self.head(f * (1 + gamma) + beta)
|
| 101 |
-
|
| 102 |
-
|
| 103 |
-
class TartanNetSTFT(nn.Module):
|
| 104 |
-
def __init__(self, width=80, emb=16, in_ch=FEAT_CH, n_fft=64, hop=16):
|
| 105 |
-
super().__init__()
|
| 106 |
-
self.n_fft, self.hop = n_fft, hop
|
| 107 |
-
self.register_buffer('win', torch.hann_window(n_fft))
|
| 108 |
-
self.t_trunk = nn.Sequential(ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 109 |
-
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 110 |
-
Ct = width * 2
|
| 111 |
-
self.s_cnn = nn.Sequential(nn.Conv2d(in_ch, 32, 3, padding=1), nn.BatchNorm2d(32), nn.GELU(),
|
| 112 |
-
nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU(),
|
| 113 |
-
nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU())
|
| 114 |
-
Cs = 64
|
| 115 |
-
self.tpool = nn.AdaptiveAvgPool1d(1)
|
| 116 |
-
self.spool = nn.AdaptiveAvgPool2d(1)
|
| 117 |
-
C = Ct + Cs
|
| 118 |
-
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 119 |
-
self.film = nn.Linear(emb, 2 * C)
|
| 120 |
-
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 121 |
-
|
| 122 |
-
def forward(self, x):
|
| 123 |
-
B, C, T = x.shape
|
| 124 |
-
ft = self.tpool(self.t_trunk(x)).squeeze(-1)
|
| 125 |
-
S = torch.stft(x.reshape(B * C, T), n_fft=self.n_fft, hop_length=self.hop,
|
| 126 |
-
window=self.win, return_complex=True, center=True).abs()
|
| 127 |
-
S = torch.log1p(S).reshape(B, C, S.shape[-2], S.shape[-1])
|
| 128 |
-
sf = self.spool(self.s_cnn(S)).flatten(1)
|
| 129 |
-
f = torch.cat([ft, sf], dim=1)
|
| 130 |
-
gamma, beta = self.film(self.emb(f)).chunk(2, dim=-1)
|
| 131 |
-
return self.head(f * (1 + gamma) + beta)
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
class TartanMergedV15(nn.Module):
|
| 135 |
-
def __init__(self, in_ch=FEAT_CH, w=BLEND):
|
| 136 |
-
super().__init__()
|
| 137 |
-
self.b14 = TartanNetSTFT(width=112, in_ch=in_ch)
|
| 138 |
-
self.b12 = TartanNetSTFT(width=80, in_ch=in_ch)
|
| 139 |
-
self.b10 = TartanNetTF(width=80, in_ch=in_ch)
|
| 140 |
-
self.b4 = TartanNet(width=80, in_ch=in_ch)
|
| 141 |
-
self.b5 = TartanNet2(width=80, in_ch=in_ch)
|
| 142 |
-
self.w = w
|
| 143 |
-
|
| 144 |
-
def forward(self, x):
|
| 145 |
-
w = self.w
|
| 146 |
-
return (w[0] * self.b14(x) + w[1] * self.b12(x) + w[2] * self.b10(x)
|
| 147 |
-
+ w[3] * self.b4(x) + w[4] * self.b5(x))
|
| 148 |
|
| 149 |
|
| 150 |
def main():
|
| 151 |
ap = argparse.ArgumentParser()
|
| 152 |
-
|
| 153 |
-
ap.add_argument('--test_dir', default='./test')
|
| 154 |
-
ap.add_argument('--index', default='test_windows.csv')
|
| 155 |
-
ap.add_argument('--weights', default=os.path.join(here, 'weights', 'model.pt'))
|
| 156 |
-
ap.add_argument('--norm', default=os.path.join(here, 'weights', 'norm.npz'))
|
| 157 |
ap.add_argument('--out', default='submission.csv')
|
| 158 |
a = ap.parse_args()
|
| 159 |
-
|
| 160 |
-
|
| 161 |
-
|
| 162 |
-
|
| 163 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 164 |
key2wid = {(r.traj_id, int(r.win_idx)): int(r.window_id) for r in tw.itertuples()}
|
| 165 |
rows = {}
|
| 166 |
-
for f in sorted(glob.glob(os.path.join(a.
|
| 167 |
traj = os.path.splitext(os.path.basename(f))[0]
|
| 168 |
-
|
| 169 |
-
n =
|
| 170 |
-
|
| 171 |
-
continue
|
| 172 |
-
X = feat[:n * WIN].reshape(n, WIN, FEAT_CH).transpose(0, 2, 1)
|
| 173 |
-
Xt = (torch.from_numpy(X) - mean[None, :, None]) / std[None, :, None]
|
| 174 |
with torch.no_grad():
|
| 175 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 176 |
for k in range(n):
|
| 177 |
wid = key2wid.get((traj, k))
|
| 178 |
-
if wid is not None:
|
| 179 |
-
|
| 180 |
-
|
| 181 |
-
|
| 182 |
-
|
| 183 |
-
|
| 184 |
-
|
| 185 |
-
|
| 186 |
-
out.to_csv(a.out, index=False)
|
| 187 |
-
filled = sum(w in rows for w in wids)
|
| 188 |
-
print(f'wrote {a.out}: {len(out)} rows, {filled} filled, {len(out) - filled} missing')
|
| 189 |
|
| 190 |
|
| 191 |
if __name__ == '__main__':
|
|
|
|
| 1 |
+
"""TartanIMU inference entry point (offline, self-contained).
|
| 2 |
|
| 3 |
+
Reproduces submission.csv from the raw test IMU. This is ONE unified model applied identically to
|
| 4 |
+
every window of every platform: a fixed-weight blend of carrier-conditioned Mixture-of-Experts models
|
| 5 |
+
(internal, input-driven platform adaptation via a learnable prototype router - NO external platform
|
| 6 |
+
label, NO platform recovery) plus transformer and multi-resolution spectrogram members. Every member
|
| 7 |
+
runs on every window; there is no per-platform routing or switching. All weights are frozen and bundled
|
| 8 |
+
under ./weights (no internet needed).
|
| 9 |
|
| 10 |
+
Usage:
|
| 11 |
+
python infer.py --data_root <dir> --out submission.csv
|
| 12 |
+
where <dir> contains: test/*.npz (each with key 'imu' (N,6)=[ax,ay,az,gx,gy,gz]),
|
| 13 |
+
index/test_windows.csv (columns traj_id,win_idx,window_id), sample_submission.csv (window_id,vx,vy,vz).
|
| 14 |
|
| 15 |
+
Re-execution budget: single GPU, <=16 GB VRAM, <=2 h. Runs in a few minutes on the official test set.
|
| 16 |
"""
|
| 17 |
+
import os, sys, glob, argparse, numpy as np, pandas as pd, torch, torch.nn as nn
|
| 18 |
+
HERE = os.path.dirname(os.path.abspath(__file__))
|
| 19 |
+
sys.path.insert(0, HERE)
|
| 20 |
+
from features import add_features, FEAT_CH
|
| 21 |
+
from model import TartanNetSTFT, TartanNetTR, TartanNetMSTFT, TartanNetMoE
|
| 22 |
|
| 23 |
WIN = 200
|
| 24 |
+
DEV = 'cuda' if torch.cuda.is_available() else 'cpu'
|
| 25 |
+
W = os.path.join(HERE, 'weights')
|
|
|
|
| 26 |
|
| 27 |
+
# (class, weight-file, width, kwargs, context-frames, blend-weight) -- greedy CV-selected (v30)
|
| 28 |
+
MEM = [(TartanNetMoE, 'best_moe2s_s23.pt', 112, {}, 400, 0.167),
|
| 29 |
+
(TartanNetTR, 'best_ctx2tr.pt', 96, {'heads':6,'layers':3}, 400, 0.143),
|
| 30 |
+
(TartanNetMoE, 'best_moe2s.pt', 112, {}, 400, 0.141),
|
| 31 |
+
(TartanNetTR, 'best_v18.pt', 80, {}, 200, 0.139),
|
| 32 |
+
(TartanNetMoE, 'best_moe2s_s7.pt', 112, {}, 400, 0.097),
|
| 33 |
+
(TartanNetMSTFT,'best_ms.pt', 80, {}, 200, 0.087),
|
| 34 |
+
(TartanNetMoE, 'best_moe2sK6.pt', 112, {'K':6}, 400, 0.085),
|
| 35 |
+
(TartanNetMoE, 'best_moe3s.pt', 112, {}, 600, 0.080),
|
| 36 |
+
(TartanNetTR, 'best_ctxtr.pt', 96, {'heads':6,'layers':3}, 600, 0.061)]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 37 |
|
| 38 |
|
| 39 |
def main():
|
| 40 |
ap = argparse.ArgumentParser()
|
| 41 |
+
ap.add_argument('--data_root', default='.', help='dir with test/*.npz, index/test_windows.csv, sample_submission.csv')
|
|
|
|
|
|
|
|
|
|
|
|
|
| 42 |
ap.add_argument('--out', default='submission.csv')
|
| 43 |
a = ap.parse_args()
|
| 44 |
+
nm = np.load(os.path.join(W, 'norm.npz')); mean = torch.tensor(nm['mean']).to(DEV); std = torch.tensor(nm['std']).to(DEV)
|
| 45 |
+
nets = []
|
| 46 |
+
for cls, ck, wd, kw, cf, wt in MEM:
|
| 47 |
+
n = cls(width=wd, in_ch=FEAT_CH, **kw).to(DEV)
|
| 48 |
+
n.load_state_dict(torch.load(os.path.join(W, ck), map_location=DEV)); n.eval()
|
| 49 |
+
nets.append((n, cf, wt))
|
| 50 |
+
wsum = sum(wt for _,_,wt in nets)
|
| 51 |
+
ss = np.load(os.path.join(W, 'speedscale.npz')); sa, sb = float(ss['a']), float(ss['b']) # input-driven magnitude scale
|
| 52 |
+
|
| 53 |
+
tw = pd.read_csv(os.path.join(a.data_root, 'index/test_windows.csv'))
|
| 54 |
key2wid = {(r.traj_id, int(r.win_idx)): int(r.window_id) for r in tw.itertuples()}
|
| 55 |
rows = {}
|
| 56 |
+
for f in sorted(glob.glob(os.path.join(a.data_root, 'test/*.npz'))):
|
| 57 |
traj = os.path.splitext(os.path.basename(f))[0]
|
| 58 |
+
ft = add_features(np.load(f)['imu']); n = ft.shape[0] // WIN
|
| 59 |
+
if n == 0: continue
|
| 60 |
+
L = n * WIN; acc = np.zeros((n, 3))
|
|
|
|
|
|
|
|
|
|
| 61 |
with torch.no_grad():
|
| 62 |
+
for net, cf, wt in nets:
|
| 63 |
+
pad = (cf - WIN) // 2
|
| 64 |
+
if pad == 0:
|
| 65 |
+
Xc = ft[:L].reshape(n, WIN, FEAT_CH).transpose(0,2,1)
|
| 66 |
+
else:
|
| 67 |
+
fp = np.pad(ft, ((pad,pad),(0,0)), mode='edge')
|
| 68 |
+
Xc = np.stack([fp[s:s+cf] for s in range(0, L, WIN)]).transpose(0,2,1)
|
| 69 |
+
xn = (torch.from_numpy(Xc.copy().astype(np.float32)).to(DEV) - mean[None,:,None]) / std[None,:,None]
|
| 70 |
+
acc += wt * net(xn).cpu().numpy()
|
| 71 |
+
pr = acc / wsum # unified blend
|
| 72 |
+
pr = pr * (sa + sb * np.hypot(pr[:,0], pr[:,1]))[:,None] # input-driven magnitude scale (own predicted speed, no platform label)
|
| 73 |
for k in range(n):
|
| 74 |
wid = key2wid.get((traj, k))
|
| 75 |
+
if wid is not None: rows[wid] = pr[k]
|
| 76 |
+
sub = pd.read_csv(os.path.join(a.data_root, 'sample_submission.csv')); vx=[]; vy=[]; vz=[]; miss=0
|
| 77 |
+
for wid in sub.window_id:
|
| 78 |
+
if wid in rows: vx.append(rows[wid][0]); vy.append(rows[wid][1]); vz.append(rows[wid][2])
|
| 79 |
+
else: vx.append(0.); vy.append(0.); vz.append(0.); miss += 1
|
| 80 |
+
sub['vx'], sub['vy'], sub['vz'] = vx, vy, vz
|
| 81 |
+
sub.to_csv(a.out, index=False)
|
| 82 |
+
print(f'wrote {a.out}: {len(sub)} rows, {miss} missing', flush=True)
|
|
|
|
|
|
|
|
|
|
| 83 |
|
| 84 |
|
| 85 |
if __name__ == '__main__':
|
model.py
ADDED
|
@@ -0,0 +1,516 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""One unified model, one weight set, all four platforms.
|
| 2 |
+
|
| 3 |
+
Rule-compliant cross-embodiment design: a shared 1D-CNN trunk plus an *input-driven*
|
| 4 |
+
conditioning path. A small head reads the same window and emits a soft "embodiment code";
|
| 5 |
+
that code FiLM-modulates the regression head, so the single network internally adapts to
|
| 6 |
+
car / dog / drone / human without any external platform label. No per-platform experts,
|
| 7 |
+
no test-time routing.
|
| 8 |
+
"""
|
| 9 |
+
import torch, torch.nn as nn
|
| 10 |
+
from config import IMU_CH
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
class ConvBlock(nn.Module):
|
| 14 |
+
def __init__(self, cin, cout, k=7, d=1):
|
| 15 |
+
super().__init__()
|
| 16 |
+
self.net = nn.Sequential(
|
| 17 |
+
nn.Conv1d(cin, cout, k, padding=(k // 2) * d, dilation=d),
|
| 18 |
+
nn.BatchNorm1d(cout), nn.GELU(),
|
| 19 |
+
nn.Conv1d(cout, cout, k, padding=(k // 2) * d, dilation=d),
|
| 20 |
+
nn.BatchNorm1d(cout), nn.GELU(),
|
| 21 |
+
)
|
| 22 |
+
self.skip = nn.Conv1d(cin, cout, 1) if cin != cout else nn.Identity()
|
| 23 |
+
|
| 24 |
+
def forward(self, x):
|
| 25 |
+
return self.net(x) + self.skip(x)
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
class TartanNet(nn.Module):
|
| 29 |
+
def __init__(self, width=64, emb=16, in_ch=IMU_CH):
|
| 30 |
+
super().__init__()
|
| 31 |
+
self.trunk = nn.Sequential(
|
| 32 |
+
ConvBlock(in_ch, width, d=1),
|
| 33 |
+
ConvBlock(width, width, d=2),
|
| 34 |
+
ConvBlock(width, width * 2, d=4),
|
| 35 |
+
ConvBlock(width * 2, width * 2, d=8),
|
| 36 |
+
)
|
| 37 |
+
C = width * 2
|
| 38 |
+
self.pool = nn.AdaptiveAvgPool1d(1)
|
| 39 |
+
# embodiment code inferred from the window itself
|
| 40 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 41 |
+
# FiLM parameters generated from the embodiment code
|
| 42 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 43 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 44 |
+
|
| 45 |
+
def forward(self, x, return_emb=False):
|
| 46 |
+
f = self.pool(self.trunk(x)).squeeze(-1) # (B, C)
|
| 47 |
+
e = self.emb(f)
|
| 48 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 49 |
+
f = f * (1 + gamma) + beta # internal conditioning
|
| 50 |
+
v = self.head(f)
|
| 51 |
+
return (v, e) if return_emb else v
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class TartanNet2(nn.Module):
|
| 55 |
+
"""v5: deeper trunk (extra dilation-16 block) + attention pooling that learns which
|
| 56 |
+
frames in the window carry the velocity signal, plus the same FiLM conditioning."""
|
| 57 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH):
|
| 58 |
+
super().__init__()
|
| 59 |
+
self.trunk = nn.Sequential(
|
| 60 |
+
ConvBlock(in_ch, width, d=1),
|
| 61 |
+
ConvBlock(width, width, d=2),
|
| 62 |
+
ConvBlock(width, width * 2, d=4),
|
| 63 |
+
ConvBlock(width * 2, width * 2, d=8),
|
| 64 |
+
ConvBlock(width * 2, width * 2, d=16),
|
| 65 |
+
)
|
| 66 |
+
C = width * 2
|
| 67 |
+
self.attn = nn.Conv1d(C, 1, 1) # per-frame attention logits
|
| 68 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 69 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 70 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 71 |
+
|
| 72 |
+
def forward(self, x, return_emb=False):
|
| 73 |
+
h = self.trunk(x) # (B, C, T)
|
| 74 |
+
w = torch.softmax(self.attn(h), dim=-1) # (B, 1, T)
|
| 75 |
+
f = (h * w).sum(-1) # attention-pooled (B, C)
|
| 76 |
+
e = self.emb(f)
|
| 77 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 78 |
+
f = f * (1 + gamma) + beta
|
| 79 |
+
v = self.head(f)
|
| 80 |
+
return (v, e) if return_emb else v
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
class TartanMerged(nn.Module):
|
| 84 |
+
"""One module, one checkpoint: an average-pool branch and an attention-pool branch that
|
| 85 |
+
both read the same window and whose velocity outputs are blended by a fixed weight. No
|
| 86 |
+
per-platform routing (both branches run on every window). This is the eligibility-clean
|
| 87 |
+
packaging of the v4+v5 ensemble as a single unified model with a single weight set."""
|
| 88 |
+
def __init__(self, width=80, in_ch=IMU_CH, w4=0.75):
|
| 89 |
+
super().__init__()
|
| 90 |
+
self.b4 = TartanNet(width=width, in_ch=in_ch)
|
| 91 |
+
self.b5 = TartanNet2(width=width, in_ch=in_ch)
|
| 92 |
+
self.w4 = w4
|
| 93 |
+
|
| 94 |
+
def forward(self, x):
|
| 95 |
+
return self.w4 * self.b4(x) + (1 - self.w4) * self.b5(x)
|
| 96 |
+
|
| 97 |
+
|
| 98 |
+
class TartanMergedV13(nn.Module):
|
| 99 |
+
"""One module, one checkpoint: the four decorrelated branches of the v13 entry, blended by
|
| 100 |
+
fixed weights. b12 (STFT spectrogram) + b10 (frequency spectrum) + b4 (avg-pool) + b5
|
| 101 |
+
(attention-pool). All run identically on every window, no per-platform routing. This is the
|
| 102 |
+
eligibility-clean single-model packaging of the ensemble."""
|
| 103 |
+
def __init__(self, width=80, in_ch=IMU_CH, w=(0.55, 0.25, 0.10, 0.10)):
|
| 104 |
+
super().__init__()
|
| 105 |
+
self.b12 = TartanNetSTFT(width=width, in_ch=in_ch)
|
| 106 |
+
self.b10 = TartanNetTF(width=width, in_ch=in_ch)
|
| 107 |
+
self.b4 = TartanNet(width=width, in_ch=in_ch)
|
| 108 |
+
self.b5 = TartanNet2(width=width, in_ch=in_ch)
|
| 109 |
+
self.w = w
|
| 110 |
+
|
| 111 |
+
def forward(self, x):
|
| 112 |
+
w = self.w
|
| 113 |
+
return w[0] * self.b12(x) + w[1] * self.b10(x) + w[2] * self.b4(x) + w[3] * self.b5(x)
|
| 114 |
+
|
| 115 |
+
|
| 116 |
+
class TartanMergedV15(nn.Module):
|
| 117 |
+
"""One module, one checkpoint: the five-branch v15 entry. b14 (wider STFT, width 112) +
|
| 118 |
+
b12 (STFT) + b10 (freq) + b4 (avg) + b5 (attn), blended by fixed weights. All run
|
| 119 |
+
identically on every window, no per-platform routing."""
|
| 120 |
+
def __init__(self, in_ch=IMU_CH, w=(0.5, 0.25, 0.15, 0.05, 0.05)):
|
| 121 |
+
super().__init__()
|
| 122 |
+
self.b14 = TartanNetSTFT(width=112, in_ch=in_ch)
|
| 123 |
+
self.b12 = TartanNetSTFT(width=80, in_ch=in_ch)
|
| 124 |
+
self.b10 = TartanNetTF(width=80, in_ch=in_ch)
|
| 125 |
+
self.b4 = TartanNet(width=80, in_ch=in_ch)
|
| 126 |
+
self.b5 = TartanNet2(width=80, in_ch=in_ch)
|
| 127 |
+
self.w = w
|
| 128 |
+
|
| 129 |
+
def forward(self, x):
|
| 130 |
+
w = self.w
|
| 131 |
+
return (w[0] * self.b14(x) + w[1] * self.b12(x) + w[2] * self.b10(x)
|
| 132 |
+
+ w[3] * self.b4(x) + w[4] * self.b5(x))
|
| 133 |
+
|
| 134 |
+
|
| 135 |
+
class TartanNetPhys(nn.Module):
|
| 136 |
+
"""v9: the proven avg-pool TCN trunk, but the FiLM conditioning and the regression head also
|
| 137 |
+
read an explicit physics-feature vector (gait cadence, pulse-and-stop stationarity, motion
|
| 138 |
+
anisotropy / NHC) computed from the same IMU window. This injects the confirmed per-platform
|
| 139 |
+
structure so the model's internal body-inference and velocity output are sharper."""
|
| 140 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH, n_phys=23):
|
| 141 |
+
super().__init__()
|
| 142 |
+
self.trunk = nn.Sequential(
|
| 143 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 144 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 145 |
+
C = width * 2
|
| 146 |
+
self.pool = nn.AdaptiveAvgPool1d(1)
|
| 147 |
+
self.phys = nn.Sequential(nn.Linear(n_phys, 48), nn.GELU(), nn.Linear(48, 48), nn.GELU())
|
| 148 |
+
Cc = C + 48
|
| 149 |
+
self.emb = nn.Sequential(nn.Linear(Cc, 64), nn.GELU(), nn.Linear(64, emb))
|
| 150 |
+
self.film = nn.Linear(emb, 2 * Cc)
|
| 151 |
+
self.head = nn.Sequential(nn.Linear(Cc, 128), nn.GELU(), nn.Linear(128, 3))
|
| 152 |
+
|
| 153 |
+
def forward(self, x, phys, return_emb=False):
|
| 154 |
+
f = self.pool(self.trunk(x)).squeeze(-1) # (B, C)
|
| 155 |
+
f = torch.cat([f, self.phys(phys)], dim=1) # fuse physics (B, Cc)
|
| 156 |
+
e = self.emb(f)
|
| 157 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 158 |
+
f = f * (1 + gamma) + beta
|
| 159 |
+
v = self.head(f)
|
| 160 |
+
return (v, e) if return_emb else v
|
| 161 |
+
|
| 162 |
+
|
| 163 |
+
class TartanNetTF(nn.Module):
|
| 164 |
+
"""v10: two-branch time + frequency model (FDIO-style). One branch reads the window over time
|
| 165 |
+
(the proven TCN). The other reads its full per-channel frequency spectrum, so gait cadence and
|
| 166 |
+
the pulse-and-stop rhythm are seen directly, not summarized. Both are pooled, fused, and passed
|
| 167 |
+
through the same FiLM conditioning + head. One model, one weight set, no per-platform routing."""
|
| 168 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH):
|
| 169 |
+
super().__init__()
|
| 170 |
+
self.t_trunk = nn.Sequential(
|
| 171 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 172 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 173 |
+
self.f_trunk = nn.Sequential(
|
| 174 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2), ConvBlock(width, width, d=4))
|
| 175 |
+
Ct, Cf = width * 2, width
|
| 176 |
+
self.pool = nn.AdaptiveAvgPool1d(1)
|
| 177 |
+
C = Ct + Cf
|
| 178 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 179 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 180 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 181 |
+
|
| 182 |
+
def forward(self, x, return_emb=False):
|
| 183 |
+
ft = self.pool(self.t_trunk(x)).squeeze(-1) # time features (B, Ct)
|
| 184 |
+
s = x - x.mean(dim=2, keepdim=True) # drop DC
|
| 185 |
+
Xf = torch.log1p(torch.fft.rfft(s, dim=2).abs()) # per-channel spectrum (B, in_ch, F)
|
| 186 |
+
ff = self.pool(self.f_trunk(Xf)).squeeze(-1) # freq features (B, Cf)
|
| 187 |
+
f = torch.cat([ft, ff], dim=1)
|
| 188 |
+
e = self.emb(f)
|
| 189 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 190 |
+
f = f * (1 + gamma) + beta
|
| 191 |
+
v = self.head(f)
|
| 192 |
+
return (v, e) if return_emb else v
|
| 193 |
+
|
| 194 |
+
|
| 195 |
+
class TartanNetSTFT(nn.Module):
|
| 196 |
+
"""v12: time branch (TCN) + a true SPECTROGRAM branch. A short-time FFT turns the window into
|
| 197 |
+
a time-frequency image per channel (when each rhythm happens, not just whether it is present),
|
| 198 |
+
read by a small 2D CNN. This localizes footsteps, rotor changes and turns as events, the sharp
|
| 199 |
+
version of the frequency idea. One model, one weight set, no per-platform routing."""
|
| 200 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH, n_fft=64, hop=16):
|
| 201 |
+
super().__init__()
|
| 202 |
+
self.n_fft, self.hop = n_fft, hop
|
| 203 |
+
self.register_buffer('win', torch.hann_window(n_fft))
|
| 204 |
+
self.t_trunk = nn.Sequential(
|
| 205 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 206 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 207 |
+
Ct = width * 2
|
| 208 |
+
self.s_cnn = nn.Sequential(
|
| 209 |
+
nn.Conv2d(in_ch, 32, 3, padding=1), nn.BatchNorm2d(32), nn.GELU(),
|
| 210 |
+
nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU(),
|
| 211 |
+
nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU())
|
| 212 |
+
Cs = 64
|
| 213 |
+
self.tpool = nn.AdaptiveAvgPool1d(1)
|
| 214 |
+
self.spool = nn.AdaptiveAvgPool2d(1)
|
| 215 |
+
C = Ct + Cs
|
| 216 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 217 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 218 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 219 |
+
|
| 220 |
+
def forward(self, x, return_emb=False):
|
| 221 |
+
B, C, T = x.shape
|
| 222 |
+
ft = self.tpool(self.t_trunk(x)).squeeze(-1) # (B, Ct)
|
| 223 |
+
S = torch.stft(x.reshape(B * C, T), n_fft=self.n_fft, hop_length=self.hop,
|
| 224 |
+
window=self.win, return_complex=True, center=True).abs()
|
| 225 |
+
S = torch.log1p(S).reshape(B, C, S.shape[-2], S.shape[-1]) # (B, in_ch, F, frames)
|
| 226 |
+
sf = self.spool(self.s_cnn(S)).flatten(1) # (B, Cs)
|
| 227 |
+
f = torch.cat([ft, sf], dim=1)
|
| 228 |
+
e = self.emb(f)
|
| 229 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 230 |
+
f = f * (1 + gamma) + beta
|
| 231 |
+
v = self.head(f)
|
| 232 |
+
return (v, e) if return_emb else v
|
| 233 |
+
|
| 234 |
+
|
| 235 |
+
class TartanNetTR(nn.Module):
|
| 236 |
+
"""Path B: a smarter brain. The spectrogram model, but the time branch adds a small
|
| 237 |
+
transformer encoder over downsampled conv features so it can model longer-range relations
|
| 238 |
+
within the window (how the start of a stride relates to its end). Still one unified model,
|
| 239 |
+
one weight set, input-driven FiLM conditioning, no per-platform routing."""
|
| 240 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH, n_fft=64, hop=16, heads=4, layers=2):
|
| 241 |
+
super().__init__()
|
| 242 |
+
self.n_fft, self.hop = n_fft, hop
|
| 243 |
+
self.register_buffer('win', torch.hann_window(n_fft))
|
| 244 |
+
self.t_trunk = nn.Sequential(
|
| 245 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 246 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 247 |
+
Ct = width * 2
|
| 248 |
+
self.ds = nn.AvgPool1d(8, 8) # 200 -> 25 tokens
|
| 249 |
+
enc = nn.TransformerEncoderLayer(Ct, heads, dim_feedforward=Ct * 2, dropout=0.1,
|
| 250 |
+
batch_first=True, activation='gelu')
|
| 251 |
+
self.tr = nn.TransformerEncoder(enc, layers)
|
| 252 |
+
self.s_cnn = nn.Sequential(
|
| 253 |
+
nn.Conv2d(in_ch, 32, 3, padding=1), nn.BatchNorm2d(32), nn.GELU(),
|
| 254 |
+
nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU(),
|
| 255 |
+
nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU())
|
| 256 |
+
Cs = 64
|
| 257 |
+
self.spool = nn.AdaptiveAvgPool2d(1)
|
| 258 |
+
C = Ct + Cs
|
| 259 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 260 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 261 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 262 |
+
|
| 263 |
+
def forward(self, x, return_emb=False):
|
| 264 |
+
B, C, T = x.shape
|
| 265 |
+
h = self.ds(self.t_trunk(x)).transpose(1, 2) # (B, T', Ct)
|
| 266 |
+
ft = self.tr(h).mean(dim=1) # transformer -> (B, Ct)
|
| 267 |
+
S = torch.stft(x.reshape(B * C, T), n_fft=self.n_fft, hop_length=self.hop,
|
| 268 |
+
window=self.win, return_complex=True, center=True).abs()
|
| 269 |
+
S = torch.log1p(S).reshape(B, C, S.shape[-2], S.shape[-1])
|
| 270 |
+
sf = self.spool(self.s_cnn(S)).flatten(1) # (B, Cs)
|
| 271 |
+
f = torch.cat([ft, sf], dim=1)
|
| 272 |
+
e = self.emb(f)
|
| 273 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 274 |
+
f = f * (1 + gamma) + beta
|
| 275 |
+
v = self.head(f)
|
| 276 |
+
return (v, e) if return_emb else v
|
| 277 |
+
|
| 278 |
+
|
| 279 |
+
class TartanNetGRU(nn.Module):
|
| 280 |
+
"""A recurrent 'different brain'. Conv front-end + spectrogram branch (kept for strength),
|
| 281 |
+
but the time branch is a BIDIRECTIONAL GRU that reads the downsampled sequence forward and
|
| 282 |
+
backward, carrying a running memory of how the motion evolves. A fundamentally different
|
| 283 |
+
mechanism from the CNN and transformer, so it makes decorrelated errors. One unified model,
|
| 284 |
+
one weight set, input-driven FiLM conditioning, no per-platform routing."""
|
| 285 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH, n_fft=64, hop=16, gru_layers=2):
|
| 286 |
+
super().__init__()
|
| 287 |
+
self.n_fft, self.hop = n_fft, hop
|
| 288 |
+
self.register_buffer('win', torch.hann_window(n_fft))
|
| 289 |
+
self.t_trunk = nn.Sequential(
|
| 290 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 291 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 292 |
+
Ct = width * 2
|
| 293 |
+
self.ds = nn.AvgPool1d(8, 8) # 200 -> 25 steps
|
| 294 |
+
self.gru = nn.GRU(Ct, Ct // 2, num_layers=gru_layers, batch_first=True,
|
| 295 |
+
bidirectional=True, dropout=0.1) # out dim = Ct
|
| 296 |
+
self.s_cnn = nn.Sequential(
|
| 297 |
+
nn.Conv2d(in_ch, 32, 3, padding=1), nn.BatchNorm2d(32), nn.GELU(),
|
| 298 |
+
nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU(),
|
| 299 |
+
nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU())
|
| 300 |
+
Cs = 64
|
| 301 |
+
self.spool = nn.AdaptiveAvgPool2d(1)
|
| 302 |
+
C = Ct + Cs
|
| 303 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 304 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 305 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 306 |
+
|
| 307 |
+
def forward(self, x, return_emb=False):
|
| 308 |
+
B, C, T = x.shape
|
| 309 |
+
h = self.ds(self.t_trunk(x)).transpose(1, 2) # (B, T', Ct)
|
| 310 |
+
out, _ = self.gru(h) # (B, T', Ct)
|
| 311 |
+
ft = out.mean(dim=1) # pool the memory over time
|
| 312 |
+
S = torch.stft(x.reshape(B * C, T), n_fft=self.n_fft, hop_length=self.hop,
|
| 313 |
+
window=self.win, return_complex=True, center=True).abs()
|
| 314 |
+
S = torch.log1p(S).reshape(B, C, S.shape[-2], S.shape[-1])
|
| 315 |
+
sf = self.spool(self.s_cnn(S)).flatten(1)
|
| 316 |
+
f = torch.cat([ft, sf], dim=1)
|
| 317 |
+
e = self.emb(f)
|
| 318 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 319 |
+
f = f * (1 + gamma) + beta
|
| 320 |
+
v = self.head(f)
|
| 321 |
+
return (v, e) if return_emb else v
|
| 322 |
+
|
| 323 |
+
|
| 324 |
+
class TartanNetEq(nn.Module):
|
| 325 |
+
"""EqNIO-style gravity-equivariant wrapper. Canonicalizes each window to a gravity-aligned
|
| 326 |
+
frame, normalizes, runs a base spectrogram net predicting velocity in the canonical frame,
|
| 327 |
+
then rotates the prediction back to the body frame. Takes RAW (un-normalized) windows;
|
| 328 |
+
holds its own norm buffers. Removes mounting-tilt variation by construction."""
|
| 329 |
+
def __init__(self, width=80, in_ch=IMU_CH, mean=None, std=None):
|
| 330 |
+
super().__init__()
|
| 331 |
+
self.base = TartanNetSTFT(width=width, in_ch=in_ch)
|
| 332 |
+
m = torch.zeros(in_ch) if mean is None else torch.as_tensor(mean, dtype=torch.float32)
|
| 333 |
+
s = torch.ones(in_ch) if std is None else torch.as_tensor(std, dtype=torch.float32)
|
| 334 |
+
self.register_buffer('mean', m); self.register_buffer('std', s)
|
| 335 |
+
|
| 336 |
+
def forward(self, x_raw):
|
| 337 |
+
from features import gravity_canon_torch
|
| 338 |
+
xc, R = gravity_canon_torch(x_raw)
|
| 339 |
+
xn = (xc - self.mean[None, :, None]) / self.std[None, :, None]
|
| 340 |
+
v_canon = self.base(xn)
|
| 341 |
+
return torch.einsum('bij,bj->bi', R.transpose(1, 2), v_canon) # canonical -> body
|
| 342 |
+
|
| 343 |
+
|
| 344 |
+
class TartanNetMSTFT(nn.Module):
|
| 345 |
+
"""Multi-resolution spectrogram: time branch (TCN) + TWO short-time-FFT branches at different
|
| 346 |
+
resolutions. Fine-FREQUENCY (n_fft 128) resolves cadence/tone; fine-TIME (n_fft 32) resolves
|
| 347 |
+
sharp transients (footstep/rotor impacts). Both 2D-CNNs fused with the time features."""
|
| 348 |
+
def __init__(self, width=80, emb=16, in_ch=IMU_CH):
|
| 349 |
+
super().__init__()
|
| 350 |
+
self.register_buffer('win_f', torch.hann_window(128))
|
| 351 |
+
self.register_buffer('win_t', torch.hann_window(32))
|
| 352 |
+
self.t_trunk = nn.Sequential(
|
| 353 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 354 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 355 |
+
Ct = width * 2
|
| 356 |
+
|
| 357 |
+
def scnn():
|
| 358 |
+
return nn.Sequential(nn.Conv2d(in_ch, 32, 3, padding=1), nn.BatchNorm2d(32), nn.GELU(),
|
| 359 |
+
nn.Conv2d(32, 48, 3, padding=1), nn.BatchNorm2d(48), nn.GELU(),
|
| 360 |
+
nn.Conv2d(48, 48, 3, padding=1), nn.BatchNorm2d(48), nn.GELU())
|
| 361 |
+
self.cnn_f = scnn(); self.cnn_t = scnn()
|
| 362 |
+
Cs = 96
|
| 363 |
+
self.tpool = nn.AdaptiveAvgPool1d(1); self.spool = nn.AdaptiveAvgPool2d(1)
|
| 364 |
+
C = Ct + Cs
|
| 365 |
+
self.emb = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 366 |
+
self.film = nn.Linear(emb, 2 * C)
|
| 367 |
+
self.head = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 368 |
+
|
| 369 |
+
def _spec(self, x, n_fft, hop, win, cnn):
|
| 370 |
+
B, C, T = x.shape
|
| 371 |
+
S = torch.stft(x.reshape(B * C, T), n_fft=n_fft, hop_length=hop, window=win,
|
| 372 |
+
return_complex=True, center=True).abs()
|
| 373 |
+
S = torch.log1p(S).reshape(B, C, S.shape[-2], S.shape[-1])
|
| 374 |
+
return self.spool(cnn(S)).flatten(1)
|
| 375 |
+
|
| 376 |
+
def forward(self, x, return_emb=False):
|
| 377 |
+
ft = self.tpool(self.t_trunk(x)).squeeze(-1)
|
| 378 |
+
sf = self._spec(x, 128, 16, self.win_f, self.cnn_f)
|
| 379 |
+
st = self._spec(x, 32, 8, self.win_t, self.cnn_t)
|
| 380 |
+
f = torch.cat([ft, sf, st], dim=1)
|
| 381 |
+
gamma, beta = self.film(self.emb(f)).chunk(2, dim=-1)
|
| 382 |
+
return self.head(f * (1 + gamma) + beta)
|
| 383 |
+
|
| 384 |
+
|
| 385 |
+
class ResBlock1D(nn.Module):
|
| 386 |
+
"""Plain strided residual block (ResNet-style), time-domain, no dilation/no spectrogram."""
|
| 387 |
+
def __init__(self, cin, cout, k=7, stride=1):
|
| 388 |
+
super().__init__()
|
| 389 |
+
self.net = nn.Sequential(
|
| 390 |
+
nn.Conv1d(cin, cout, k, stride=stride, padding=k // 2),
|
| 391 |
+
nn.BatchNorm1d(cout), nn.GELU(),
|
| 392 |
+
nn.Conv1d(cout, cout, k, padding=k // 2),
|
| 393 |
+
nn.BatchNorm1d(cout),
|
| 394 |
+
)
|
| 395 |
+
self.skip = (nn.Conv1d(cin, cout, 1, stride=stride) if (cin != cout or stride != 1)
|
| 396 |
+
else nn.Identity())
|
| 397 |
+
self.act = nn.GELU()
|
| 398 |
+
|
| 399 |
+
def forward(self, x):
|
| 400 |
+
return self.act(self.net(x) + self.skip(x))
|
| 401 |
+
|
| 402 |
+
|
| 403 |
+
class TartanNetResNet1D(nn.Module):
|
| 404 |
+
"""Decorrelated member: a deep strided ResNet in the pure time domain (no STFT, no attention),
|
| 405 |
+
global avg+max pooling, with the same input-driven FiLM conditioning to stay one unified model.
|
| 406 |
+
Very different inductive bias from the spectrogram and transformer members."""
|
| 407 |
+
def __init__(self, width=64, emb=16, in_ch=IMU_CH):
|
| 408 |
+
super().__init__()
|
| 409 |
+
self.stem = nn.Sequential(nn.Conv1d(in_ch, width, 7, padding=3), nn.BatchNorm1d(width), nn.GELU())
|
| 410 |
+
self.stage = nn.Sequential(
|
| 411 |
+
ResBlock1D(width, width), ResBlock1D(width, width * 2, stride=2),
|
| 412 |
+
ResBlock1D(width * 2, width * 2), ResBlock1D(width * 2, width * 4, stride=2),
|
| 413 |
+
ResBlock1D(width * 4, width * 4), ResBlock1D(width * 4, width * 4, stride=2),
|
| 414 |
+
)
|
| 415 |
+
C = width * 4
|
| 416 |
+
self.emb = nn.Sequential(nn.Linear(2 * C, 64), nn.GELU(), nn.Linear(64, emb))
|
| 417 |
+
self.film = nn.Linear(emb, 2 * (2 * C))
|
| 418 |
+
self.head = nn.Sequential(nn.Linear(2 * C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 419 |
+
|
| 420 |
+
def forward(self, x, return_emb=False):
|
| 421 |
+
h = self.stage(self.stem(x)) # (B, C, T')
|
| 422 |
+
f = torch.cat([h.mean(-1), h.amax(-1)], dim=-1) # (B, 2C) global avg+max
|
| 423 |
+
e = self.emb(f)
|
| 424 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 425 |
+
f = f * (1 + gamma) + beta
|
| 426 |
+
v = self.head(f)
|
| 427 |
+
return (v, e) if return_emb else v
|
| 428 |
+
|
| 429 |
+
|
| 430 |
+
class GatedTCN(nn.Module):
|
| 431 |
+
"""WaveNet-style gated dilated residual unit: tanh*sigmoid gate, residual + skip outputs."""
|
| 432 |
+
def __init__(self, ch, k=3, d=1):
|
| 433 |
+
super().__init__()
|
| 434 |
+
pad = (k - 1) * d
|
| 435 |
+
self.pad = pad
|
| 436 |
+
self.f = nn.Conv1d(ch, ch, k, padding=pad, dilation=d)
|
| 437 |
+
self.g = nn.Conv1d(ch, ch, k, padding=pad, dilation=d)
|
| 438 |
+
self.res = nn.Conv1d(ch, ch, 1)
|
| 439 |
+
self.skip = nn.Conv1d(ch, ch, 1)
|
| 440 |
+
|
| 441 |
+
def forward(self, x):
|
| 442 |
+
p = self.pad
|
| 443 |
+
fx = self.f(x)[:, :, :x.shape[-1]] if p else self.f(x)
|
| 444 |
+
gx = self.g(x)[:, :, :x.shape[-1]] if p else self.g(x)
|
| 445 |
+
z = torch.tanh(fx) * torch.sigmoid(gx)
|
| 446 |
+
return x + self.res(z), self.skip(z)
|
| 447 |
+
|
| 448 |
+
|
| 449 |
+
class TartanNetWaveTCN(nn.Module):
|
| 450 |
+
"""Decorrelated member: WaveNet-style gated dilated causal TCN (dilations 1..128 cover the
|
| 451 |
+
200-frame window), skip-connection aggregation, FiLM conditioning. Different again from the
|
| 452 |
+
ResNet, the spectrogram CNNs, and the transformer."""
|
| 453 |
+
def __init__(self, width=64, emb=16, in_ch=IMU_CH, levels=8):
|
| 454 |
+
super().__init__()
|
| 455 |
+
self.inp = nn.Conv1d(in_ch, width, 1)
|
| 456 |
+
self.blocks = nn.ModuleList([GatedTCN(width, k=3, d=2 ** i) for i in range(levels)])
|
| 457 |
+
self.emb = nn.Sequential(nn.Linear(width, 64), nn.GELU(), nn.Linear(64, emb))
|
| 458 |
+
self.film = nn.Linear(emb, 2 * width)
|
| 459 |
+
self.head = nn.Sequential(nn.GELU(), nn.Linear(width, 128), nn.GELU(), nn.Linear(128, 3))
|
| 460 |
+
|
| 461 |
+
def forward(self, x, return_emb=False):
|
| 462 |
+
h = self.inp(x); skips = 0
|
| 463 |
+
for b in self.blocks:
|
| 464 |
+
h, s = b(h); skips = skips + s
|
| 465 |
+
f = skips.mean(-1) # (B, width) pooled skip sum
|
| 466 |
+
e = self.emb(f)
|
| 467 |
+
gamma, beta = self.film(e).chunk(2, dim=-1)
|
| 468 |
+
f = f * (1 + gamma) + beta
|
| 469 |
+
v = self.head(f)
|
| 470 |
+
return (v, e) if return_emb else v
|
| 471 |
+
|
| 472 |
+
|
| 473 |
+
class TartanNetMoE(nn.Module):
|
| 474 |
+
"""MosaicIMU-style carrier-conditioned Mixture-of-Experts (rules-allowed: MoE *within one
|
| 475 |
+
trained network*, routed by learned prototypes, no external platform label at inference).
|
| 476 |
+
Shared STFT+TCN encoder -> prototype router (K carrier prototypes, cosine sim) -> K experts
|
| 477 |
+
soft-blended -> velocity + heteroscedastic uncertainty. Closes the gap to per-platform experts
|
| 478 |
+
inside a single unified model, exactly the competition's stated goal."""
|
| 479 |
+
def __init__(self, width=112, emb=16, in_ch=IMU_CH, n_fft=64, hop=16, K=4, tau=0.3):
|
| 480 |
+
super().__init__()
|
| 481 |
+
self.n_fft, self.hop, self.K, self.tau = n_fft, hop, K, tau
|
| 482 |
+
self.register_buffer('win', torch.hann_window(n_fft))
|
| 483 |
+
self.t_trunk = nn.Sequential(
|
| 484 |
+
ConvBlock(in_ch, width, d=1), ConvBlock(width, width, d=2),
|
| 485 |
+
ConvBlock(width, width * 2, d=4), ConvBlock(width * 2, width * 2, d=8))
|
| 486 |
+
Ct = width * 2
|
| 487 |
+
self.s_cnn = nn.Sequential(
|
| 488 |
+
nn.Conv2d(in_ch, 32, 3, padding=1), nn.BatchNorm2d(32), nn.GELU(),
|
| 489 |
+
nn.Conv2d(32, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU(),
|
| 490 |
+
nn.Conv2d(64, 64, 3, padding=1), nn.BatchNorm2d(64), nn.GELU())
|
| 491 |
+
Cs = 64; C = Ct + Cs
|
| 492 |
+
self.tpool = nn.AdaptiveAvgPool1d(1); self.spool = nn.AdaptiveAvgPool2d(1)
|
| 493 |
+
self.proto = nn.Parameter(torch.randn(K, C))
|
| 494 |
+
self.experts = nn.ModuleList([nn.Sequential(nn.Linear(C, C), nn.GELU(), nn.Linear(C, C)) for _ in range(K)])
|
| 495 |
+
self.vel = nn.Sequential(nn.Linear(C, 128), nn.GELU(), nn.Linear(128, 3))
|
| 496 |
+
self.logvar = nn.Sequential(nn.Linear(C, 64), nn.GELU(), nn.Linear(64, 3))
|
| 497 |
+
|
| 498 |
+
def encode(self, x):
|
| 499 |
+
B, C, T = x.shape
|
| 500 |
+
ft = self.tpool(self.t_trunk(x)).squeeze(-1)
|
| 501 |
+
S = torch.stft(x.reshape(B * C, T), n_fft=self.n_fft, hop_length=self.hop,
|
| 502 |
+
window=self.win, return_complex=True, center=True).abs()
|
| 503 |
+
S = torch.log1p(S).reshape(B, C, S.shape[-2], S.shape[-1])
|
| 504 |
+
sf = self.spool(self.s_cnn(S)).flatten(1)
|
| 505 |
+
return torch.cat([ft, sf], dim=1)
|
| 506 |
+
|
| 507 |
+
def forward(self, x, return_all=False):
|
| 508 |
+
f = self.encode(x)
|
| 509 |
+
fn = torch.nn.functional.normalize(f, dim=1)
|
| 510 |
+
pn = torch.nn.functional.normalize(self.proto, dim=1)
|
| 511 |
+
w = torch.softmax((fn @ pn.t()) / self.tau, dim=1) # (B,K) router weights
|
| 512 |
+
Fm = sum(w[:, k:k + 1] * self.experts[k](f) for k in range(self.K))
|
| 513 |
+
v = self.vel(Fm)
|
| 514 |
+
if return_all:
|
| 515 |
+
return v, self.logvar(Fm).clamp(-6, 4), w
|
| 516 |
+
return v
|
submission.csv
CHANGED
|
The diff for this file is too large to render.
See raw diff
|
|
|
weights/best_ctx2tr.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2302748192bb8a027cd90e0e7582a8516d120716afc970e9918b09de2e99ac7c
|
| 3 |
+
size 8593623
|
weights/best_ctxtr.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7f48ab46b54aed65201f9e725a22b16ce62cd70f63af40fabd99a8de596e13d6
|
| 3 |
+
size 8593489
|
weights/best_moe2s.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:a66520353813c682991eac60d8e57ab79e9954f00d380577f5f293568d63166e
|
| 3 |
+
size 9295420
|
weights/best_moe2sK6.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:b4e4577d4b34edcec3f5d754eb5f0f40c89abebd199cce1794b445518ceb7ef1
|
| 3 |
+
size 10632438
|
weights/best_moe2s_s23.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:31b7d99e1300beaeb3ab85843e904cb69410d3c7937225758cfd2f4dd8b5ed89
|
| 3 |
+
size 9295872
|
weights/best_moe2s_s7.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:5eafeb20175be5fc5b4183c3d1c4e136b93cdd7e850197a04d0bcbb861667875
|
| 3 |
+
size 9295759
|
weights/best_moe3s.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:39498b8e129f258b7ca26aaa2beed5b86b7a69b040b29473d874c3b3ddf1028a
|
| 3 |
+
size 9295420
|
weights/best_ms.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:6f5aad881fe8d1ffb7a2e1cda507a6532784d8bf0a82fcb03a63be95a15d3ed6
|
| 3 |
+
size 3717403
|
weights/best_v18.pt
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:0d829aaf44660d70e1337f02a74842e41823084183e70644bcf8e11e6bc6efd6
|
| 3 |
+
size 5275025
|
weights/speedscale.npz
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:2c9eb048def7139a37e0fd9fbb0fb7cfb3da8b41eb76b37702a36ebc813f7ca9
|
| 3 |
+
size 506
|