Spaces:
Paused
Paused
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))
|