Spaces:
Paused
Paused
| # waveform_agent/tools/preprocess.py | |
| """ | |
| Signal preprocessing | |
| ==================== | |
| Turns a long-format frame of tracks (from either loader) into a tidy | |
| point-treatment panel with columns: A (treatment), Y_next (outcome at t+1), | |
| the confounder tracks, and any static covariates. | |
| The point-treatment construction follows the reference DoWhy cross-check: at a | |
| single step the time-varying confounders (map/hr/spo2 …) are *pre-treatment* | |
| w.r.t. A_t, so backdoor adjustment on them is valid. See | |
| skills/signal_preprocess/references/REFERENCE.md for why this differs from the | |
| sustained-regime (g-methods) case. | |
| Steps, in order, per case: | |
| 1. optional band-pass / notch filtering of raw high-frequency waveform tracks | |
| 2. forward-fill gaps in numeric tracks | |
| 3. physiologic range filtering (drop implausible rows) | |
| 4. binarise treatment: A = 1[treatment_track > threshold] | |
| 5. build outcome: Y_next = outcome_track.shift(-1) (drop last row) | |
| 6. subsample by stride (cut autocorrelation for point-treatment analysis) | |
| Heavy imports (pandas, numpy, scipy) are deferred into functions. | |
| """ | |
| 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 LoadedData, PreprocessConfig, PreprocessResult # noqa: E402 | |
| def _filter_waveform(sig, low_high, notch, fs: float): | |
| """Apply an optional Butterworth band-pass and mains-notch to a 1-D signal.""" | |
| from scipy.signal import butter, filtfilt, iirnotch | |
| out = sig | |
| if low_high is not None: | |
| lo, hi = low_high | |
| nyq = 0.5 * fs | |
| b, a = butter(4, [lo / nyq, hi / nyq], btype="band") | |
| out = filtfilt(b, a, out) | |
| if notch is not None: | |
| b, a = iirnotch(w0=notch / (0.5 * fs), Q=30.0) | |
| out = filtfilt(b, a, out) | |
| return out | |
| def _build_case_panel(case_df, cfg: PreprocessConfig, fs: float, steps: list[str]): | |
| import numpy as np | |
| import pandas as pd | |
| df = case_df.copy() | |
| # 1. optional waveform filtering (applied to confounder/outcome signal tracks) | |
| if cfg.bandpass_hz is not None or cfg.notch_hz is not None: | |
| for col in {cfg.outcome_track, *cfg.confounder_tracks}: | |
| if col in df.columns and df[col].notna().sum() > 30: | |
| try: | |
| df[col] = _filter_waveform( | |
| df[col].to_numpy(dtype=float), | |
| cfg.bandpass_hz, cfg.notch_hz, fs, | |
| ) | |
| except Exception: # noqa: BLE001 - filtering is best-effort | |
| pass | |
| if "filter_waveform" not in steps: | |
| steps.append("filter_waveform") | |
| # 2. forward-fill numeric tracks | |
| signal_cols = [c for c in {cfg.outcome_track, *cfg.confounder_tracks} if c in df.columns] | |
| df[signal_cols] = df[signal_cols].ffill() | |
| if cfg.treatment_track in df.columns: | |
| df[cfg.treatment_track] = df[cfg.treatment_track].fillna(0.0).clip(lower=0) | |
| if "ffill" not in steps: | |
| steps.append("ffill") | |
| # 3. physiologic range filtering | |
| if cfg.physiologic_ranges: | |
| mask = pd.Series(True, index=df.index) | |
| for col, (lo, hi) in cfg.physiologic_ranges.items(): | |
| if col in df.columns: | |
| mask &= df[col].between(lo, hi) | |
| df = df[mask] | |
| if "range_filter" not in steps: | |
| steps.append("range_filter") | |
| df = df.dropna(subset=signal_cols).reset_index(drop=True) | |
| if len(df) < 3: | |
| return None | |
| # 4. binarise treatment | |
| df["A"] = (df[cfg.treatment_track] > cfg.treatment_threshold).astype(int) | |
| if "binarise_treatment" not in steps: | |
| steps.append("binarise_treatment") | |
| # 5. next-step outcome | |
| df["Y_next"] = df[cfg.outcome_track].shift(-1) | |
| df = df.iloc[:-1].copy() | |
| if "outcome_shift" not in steps: | |
| steps.append("outcome_shift") | |
| # 6. subsample to cut autocorrelation | |
| if cfg.stride > 1: | |
| df = df.iloc[:: cfg.stride].copy() | |
| if "stride_subsample" not in steps: | |
| steps.append("stride_subsample") | |
| return df | |
| def preprocess( | |
| loaded: LoadedData, cfg: PreprocessConfig, resample_sec: float = 10.0 | |
| ) -> PreprocessResult: | |
| """Build the analysis-ready point-treatment panel and persist it. | |
| Args: | |
| loaded: LoadedData from a loader (its frame_path is read). | |
| cfg: PreprocessConfig describing treatment/outcome/confounders + filters. | |
| resample_sec: sampling interval of the loaded numeric tracks (for filter fs). | |
| Returns: | |
| PreprocessResult pointing at the persisted panel parquet. | |
| """ | |
| import pandas as pd | |
| POLICY.ensure_workdir() | |
| frame = pd.read_parquet(loaded.frame_path) | |
| fs = 1.0 / resample_sec if resample_sec else 1.0 | |
| steps: list[str] = [] | |
| warnings: list[str] = [] | |
| confounders = [c for c in cfg.confounder_tracks if c in frame.columns] | |
| confounders += [c for c in cfg.static_covariates if c in frame.columns] | |
| panels = [] | |
| for _cid, case_df in frame.groupby("caseid"): | |
| p = _build_case_panel(case_df, cfg, fs, steps) | |
| if p is not None: | |
| panels.append(p) | |
| if not panels: | |
| raise RuntimeError("Preprocessing produced no usable rows.") | |
| panel = pd.concat(panels, ignore_index=True) | |
| need = ["A", "Y_next", *confounders] | |
| panel = panel.dropna(subset=need).reset_index(drop=True) | |
| treated_fraction = float(panel["A"].mean()) if len(panel) else 0.0 | |
| if treated_fraction < 0.02 or treated_fraction > 0.98: | |
| warnings.append( | |
| f"Positivity concern: treated fraction={treated_fraction:.3f}." | |
| ) | |
| out = POLICY.artifact("data", "panel.parquet") | |
| keep_cols = ["caseid", "A", "Y_next", *confounders] | |
| panel[[c for c in keep_cols if c in panel.columns]].to_parquet(out, index=False) | |
| return PreprocessResult( | |
| panel_path=str(out), | |
| n_rows=int(len(panel)), | |
| n_cases=int(panel["caseid"].nunique()) if "caseid" in panel.columns else 1, | |
| treatment="A", | |
| outcome="Y_next", | |
| confounders=confounders, | |
| treated_fraction=treated_fraction, | |
| steps_applied=steps, | |
| warnings=warnings, | |
| ) | |