File size: 4,291 Bytes
b744871
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
# waveform_agent/tools/user_data_loader.py
"""
User-data loader
================
Loads caller-supplied signals so the agent can run the same causal pipeline on
data that is not from VitalDB. Supported inputs:

  * .csv / .parquet  — a long-format table; one column identifies the case
    (via DatasetSpec.case_column), remaining numeric columns are treated as
    tracks. If no case column is given, the whole file is treated as one case.
  * .npz             — arrays keyed by track name, each a 1-D signal for a
    single case (optionally a 'caseid' array to segment multiple cases).

The output matches the VitalDB loader's contract (a persisted long-format
parquet frame described by LoadedData) so downstream preprocessing is agnostic
to the source.
"""

from __future__ import annotations

import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parent.parent))

from config import POLICY  # noqa: E402
from models.pipeline import DatasetSpec, LoadedData  # noqa: E402


def _load_table(path: Path):
    import pandas as pd

    suffix = path.suffix.lower()
    if suffix == ".csv":
        return pd.read_csv(path)
    if suffix in (".parquet", ".pq"):
        return pd.read_parquet(path)
    raise ValueError(f"Unsupported table format: {suffix}")


def _load_npz(path: Path):
    import numpy as np
    import pandas as pd

    npz = np.load(path, allow_pickle=False)
    keys = [k for k in npz.files if k != "caseid"]
    if not keys:
        raise ValueError("NPZ contains no signal arrays.")
    n = len(npz[keys[0]])
    data = {k: np.asarray(npz[k]).ravel()[:n] for k in keys}
    df = pd.DataFrame(data)
    df["caseid"] = np.asarray(npz["caseid"]).ravel()[:n] if "caseid" in npz.files else 0
    return df


def load_user_data(spec: DatasetSpec) -> LoadedData:
    """Load user-provided signals into a persisted long-format parquet frame.

    Args:
        spec: DatasetSpec with source='user' and user_path set. If
            spec.case_column is provided it segments cases; otherwise a single
            synthetic case id (0) is used.

    Returns:
        LoadedData pointing at the parquet frame written under the sandbox.
    """
    import pandas as pd

    if not spec.user_path:
        raise ValueError("source='user' requires spec.user_path.")
    path = Path(spec.user_path).expanduser().resolve()
    if not path.exists():
        raise FileNotFoundError(f"User data not found: {path}")

    POLICY.ensure_workdir()

    df = _load_npz(path) if path.suffix.lower() == ".npz" else _load_table(path)

    # Normalise the case identifier to a single 'caseid' column.
    if spec.case_column and spec.case_column in df.columns:
        df = df.rename(columns={spec.case_column: "caseid"})
    elif "caseid" not in df.columns:
        df["caseid"] = 0

    # Keep numeric tracks (+ caseid). Restrict to requested tracks if given.
    numeric_cols = [
        c for c in df.columns
        if c != "caseid" and pd.api.types.is_numeric_dtype(df[c])
    ]
    if spec.tracks:
        missing = [t for t in spec.tracks if t not in numeric_cols]
        if missing:
            raise ValueError(f"Requested tracks absent from user data: {missing}")
        numeric_cols = [t for t in spec.tracks if t in numeric_cols]

    keep = ["caseid", *numeric_cols]
    df = df[keep].copy()

    out = POLICY.artifact("data", "user_frame.parquet")
    df.to_parquet(out, index=False)

    return LoadedData(
        source="user",
        frame_path=str(out),
        n_cases=int(df["caseid"].nunique()),
        n_rows=int(len(df)),
        columns=list(df.columns),
        tracks_loaded=numeric_cols,
        notes=f"Loaded user file {path.name} ({path.suffix}).",
    )


if __name__ == "__main__":
    import argparse
    import json

    p = argparse.ArgumentParser(description="Load user-provided signals to parquet.")
    p.add_argument("path")
    p.add_argument("--case-column", default=None)
    p.add_argument("--tracks", nargs="*", default=None)
    args = p.parse_args()

    result = load_user_data(
        DatasetSpec(
            source="user",
            user_path=args.path,
            case_column=args.case_column,
            tracks=args.tracks or [],
        )
    )
    print(json.dumps(result.model_dump(), indent=2))