Download scripts/train_baselines.py from minhy112/FallKLTN: direct link, hf CLI and curl.
- Browser
- Download file 3.98 kB
-
https://huggingface.co/minhy112/FallKLTN/resolve/main/scripts/train_baselines.py
- Command line
-
hf download hf://minhy112/FallKLTN/scripts/train_baselines.py
-
curl -L -o train_baselines.py https://huggingface.co/minhy112/FallKLTN/resolve/main/scripts/train_baselines.py
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() | |