From da3e436da08f530f39c0791d0a9d82d964bc00a7 Mon Sep 17 00:00:00 2001 From: Asger Gitz-Johansen Date: Sun, 2 Aug 2026 18:05:26 +0200 Subject: [PATCH] wip --- src/sleep_detection/main.py | 28 ++++++++++++++++++++-------- 1 file changed, 20 insertions(+), 8 deletions(-) diff --git a/src/sleep_detection/main.py b/src/sleep_detection/main.py index b08204d..82ee880 100644 --- a/src/sleep_detection/main.py +++ b/src/sleep_detection/main.py @@ -23,15 +23,15 @@ from sleep_detection.sleep_epoch_classifer import SleepEpochClassifier def main() -> None: """The Main entrypoint function.""" cli = parse_args() - session = load_session(cli.a, cli.H, cli.l) + sessions = load_sessions(cli.a, cli.H, cli.l) - inputs = torch.tensor(session[["timestamp", "heartrate", "motion-avg"]].to_numpy(), dtype=torch.float32) + inputs = torch.tensor(sessions[["timestamp", "heartrate", "motion-avg"]].to_numpy(), dtype=torch.float32) mean = inputs.mean(dim=0, keepdim=True) # shape (1, n_features) std = inputs.std(dim=0, keepdim=True) + 1e-8 normalized_inputs = (inputs - mean) / std logger.info(f"inputs={normalized_inputs}") - targets = torch.tensor(session["stage"].to_numpy(), dtype=torch.long) + targets = torch.tensor(sessions["stage"].to_numpy(), dtype=torch.long) logger.info(f"targets={targets}") net = SleepEpochClassifier(n_features=3, hidden_size=16, n_classes=3) @@ -44,13 +44,25 @@ def parse_args() -> Namespace: parser = ArgumentParser(description="Training program.") parser.add_argument("-v", "--verbose", help="Increase verbosity.", action="store_true") parser.add_argument("-V", "--version", action="store_true", help="Show version") - parser.add_argument("-a", type=Path, required=True, help="Acceleration data file(s).") - parser.add_argument("-H", type=Path, required=True, help="Heartrate data file(s).") - parser.add_argument("-l", type=Path, required=True, help="Labels data file(s).") + parser.add_argument("-a", type=Path, nargs="+", required=True, help="Acceleration data file(s).") + parser.add_argument("-H", type=Path, nargs="+", required=True, help="Heartrate data file(s).") + parser.add_argument("-l", type=Path, nargs="+", required=True, help="Labels data file(s).") parser.add_argument("-o", "--output", type=Path, required=True, help="Output .pt file.") return parser.parse_args() +def load_sessions( + acceleration_files: list[Path], heartrate_files: list[Path], label_files: list[Path] +) -> pandas.DataFrame: + if len(acceleration_files) != len(heartrate_files) != len(label_files): + raise RuntimeError("unmatched session file sets.") + result = pandas.DataFrame() + for i in range(len(acceleration_files)): + result = pandas.concat([result, load_session(acceleration_files[i], heartrate_files[i], label_files[i])]) + logger.info(f"sessions={result}") + return result + + def load_session(acceleration_file: Path, heartrate_file: Path, label_file: Path) -> pandas.DataFrame: """Load a session triplet.""" adf = load_acceleration_data(acceleration_file) @@ -108,13 +120,13 @@ def normalize_sleep_stage(stage) -> int: if pandas.isna(stage): return int(numpy.nan) value = int(stage) - if value == 0: + if value <= 0: return 0 if 1 <= value <= 4: return 1 if 5 <= value <= 6: return 2 - return int(numpy.nan) + raise ValueError(stage) def train(input: Tensor, target: Tensor, net: Module) -> None: