File size: 3,322 Bytes
96a058b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
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")