""" Convert raw time series into the sensor input SLIP expects. Usage: from data_processor import process_sensor # or model.process_sensor(x) when loaded from HF sensors = process_sensor(torch.randn(2, 3, 320), patch_size=32) emb = model.get_sensor_embedding(sensors["input_ids"], sensors["attention_mask"], sensors["time_index"]) """ import torch import torch.nn.functional as F def process_sensor(x, patch_size=32, normalize=True): """ x: (bs, nvar, L) tensor or array of raw values. NaNs are treated as missing. Returns {"input_ids", "attention_mask", "time_index"}: nested lists [bs][nvar] of (num_patches, patch_size) tensors, the flexi-patch format SLIP takes. """ x = torch.as_tensor(x, dtype=torch.float32) if x.dim() != 3: raise ValueError(f"expected (bs, nvar, L), got shape {tuple(x.shape)}") mask = ~torch.isnan(x) x = torch.nan_to_num(x, nan=0.0) if normalize: # same per-channel scaling as pretraining: x / mean(|x|) mean_abs = x.abs().sum(-1, keepdim=True) / mask.sum(-1, keepdim=True).clamp(min=1) x = x / (mean_abs + 1e-6) L = x.shape[-1] time_index = (torch.arange(1, L + 1, device=x.device, dtype=torch.float32) / (L + 1)).expand_as(x) # left-pad to a multiple of patch_size, like util.dataset.prepare_patches pad = (patch_size - L % patch_size) % patch_size x = F.pad(x, (pad, 0), value=0.0) mask = F.pad(mask, (pad, 0), value=False) time_index = F.pad(time_index, (pad, 0), value=0.0) shape = (*x.shape[:2], -1, patch_size) # bs, nvar, num_patches, patch_size return { "input_ids": [list(s) for s in x.reshape(shape)], "attention_mask": [list(s) for s in mask.reshape(shape)], "time_index": [list(s) for s in time_index.reshape(shape)], } if __name__ == "__main__": x = torch.randn(2, 3, 100) * 5 out = process_sensor(x, patch_size=32) # shape: 100 -> padded to 128 -> 4 patches of 32 assert len(out["input_ids"]) == 2 and len(out["input_ids"][0]) == 3 assert out["input_ids"][0][0].shape == (4, 32) # left padding: first 28 values are zero, masked out, time 0 assert not out["attention_mask"][0][0].flatten()[:28].any() assert out["attention_mask"][0][0].flatten()[28:].all() assert (out["time_index"][1][2].flatten()[:28] == 0).all() # normalization: mean |x| per channel ~ 1 assert abs(out["input_ids"][0][1].abs().sum().item() / 100 - 1) < 1e-4 # NaN -> value 0, mask False x_nan = x.clone() x_nan[0, 0, 50] = float("nan") m = process_sensor(x_nan)["attention_mask"][0][0].flatten() assert not m[28 + 50] and m[28 + 49] # parity with the training pipeline (needs the repo's util/ deps) try: from util.dataset import prepare_patches except ImportError as e: print(f"skipped parity check ({e})") else: for b in range(2): ref = x[b] / (x[b].abs().mean(-1, keepdim=True) + 1e-6) s, msk, t = prepare_patches(ref, 32) for v in range(3): assert torch.allclose(out["input_ids"][b][v], s[v], atol=1e-5) assert torch.equal(out["attention_mask"][b][v], msk[v]) assert torch.allclose(out["time_index"][b][v], t[v]) print("data_processor self-check passed")