FallKLTN / scripts /train_baselines.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw
History Blame Contribute Delete
3.98 kB
#!/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()