wip
This commit is contained in:
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user