File size: 3,983 Bytes
9313a90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
from __future__ import annotations

import argparse
import time
from pathlib import Path

import joblib
import numpy as np
from sklearn.ensemble import RandomForestClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler

from fall_detection.experiment import prepare_experiment_data, split_summary
from fall_detection.features import summarize_sequence_features
from fall_detection.metrics import (
    choose_f1_threshold,
    classification_metrics,
    save_evaluation_plots,
    save_predictions,
)
from fall_detection.utils import set_seed, write_json


def main() -> None:
    parser = argparse.ArgumentParser(description="Train classical pose baselines")
    parser.add_argument("--dataset", default="data/processed/urfd_pose.npz")
    parser.add_argument("--config", default="configs/default.yaml")
    parser.add_argument("--output", type=Path, default=Path("artifacts/experiments/urfd"))
    parser.add_argument("--seed", type=int)
    args = parser.parse_args()

    config, dataset, sequence_features, splits = prepare_experiment_data(
        args.dataset, args.config, args.output, seed=args.seed
    )
    set_seed(config["seed"])
    features = summarize_sequence_features(sequence_features)
    labels = dataset["labels"]
    models = {
        "logistic_regression": Pipeline(
            [
                ("scaler", StandardScaler()),
                (
                    "classifier",
                    LogisticRegression(
                        C=1.0,
                        max_iter=2000,
                        class_weight="balanced",
                        random_state=config["seed"],
                    ),
                ),
            ]
        ),
        "random_forest": RandomForestClassifier(
            n_estimators=300,
            max_depth=10,
            min_samples_leaf=2,
            class_weight="balanced",
            n_jobs=-1,
            random_state=config["seed"],
        ),
    }

    for name, model in models.items():
        output_dir = args.output / name
        output_dir.mkdir(parents=True, exist_ok=True)
        started = time.perf_counter()
        model.fit(features[splits.train], labels[splits.train])
        training_seconds = time.perf_counter() - started
        val_probabilities = model.predict_proba(features[splits.val])[:, 1]
        threshold = choose_f1_threshold(labels[splits.val], val_probabilities)
        test_probabilities = model.predict_proba(features[splits.test])[:, 1]
        metrics = classification_metrics(labels[splits.test], test_probabilities, threshold)
        metrics.update(
            {
                "model": name,
                "seed": int(config["seed"]),
                "training_seconds": training_seconds,
                "dataset": str(args.dataset),
                "split": split_summary(labels, dataset["groups"], splits),
                "result_scope": "test set only",
            }
        )
        joblib.dump(
            {
                "model": model,
                "threshold": threshold,
                "sequence_length": int(sequence_features.shape[1]),
                "visibility_threshold": config["data"]["visibility_threshold"],
            },
            output_dir / "model.joblib",
        )
        write_json(output_dir / "metrics.json", metrics)
        save_predictions(
            output_dir / "predictions.csv",
            labels[splits.test],
            test_probabilities,
            dataset["sources"][splits.test],
            threshold,
        )
        save_evaluation_plots(
            labels[splits.test], test_probabilities, threshold, output_dir, name
        )
        print(
            f"{name}: F1={metrics['f1']:.3f}, recall={metrics['recall']:.3f}, "
            f"specificity={metrics['specificity']:.3f}, AUC={metrics['roc_auc']:.3f}"
        )


if __name__ == "__main__":
    main()