SLIP / data_processor.py
LeoChen085's picture
Add process_sensor; fix remote-code imports
96a058b verified
Raw History Blame Contribute Delete
3.32 kB
"""
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")