Download data_processor.py from LeoChen085/SLIP: direct link, hf CLI and curl.
- Browser
- Download file 3.32 kB
-
https://huggingface.co/LeoChen085/SLIP/resolve/main/data_processor.py
- Command line
-
hf download hf://LeoChen085/SLIP/data_processor.py
-
curl -L -o data_processor.py https://huggingface.co/LeoChen085/SLIP/resolve/main/data_processor.py
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") | |