FallKLTN / scripts /train_gru.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw History Blame Contribute Delete
8.94 kB
#!/usr/bin/env python3
from __future__ import annotations
import argparse
import copy
import time
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import torch
from torch import nn
from torch.utils.data import DataLoader, TensorDataset
from fall_detection.experiment import prepare_experiment_data, split_summary
from fall_detection.features import CORE_FEATURE_INDICES
from fall_detection.metrics import (
choose_f1_threshold,
classification_metrics,
save_evaluation_plots,
save_predictions,
)
from fall_detection.models import PoseGRU, PoseTCN
from fall_detection.utils import set_seed, write_json
def make_loader(
features: np.ndarray,
labels: np.ndarray,
indices: np.ndarray,
batch_size: int,
shuffle: bool,
) -> DataLoader:
dataset = TensorDataset(
torch.from_numpy(features[indices]).float(),
torch.from_numpy(labels[indices]).float(),
)
return DataLoader(dataset, batch_size=batch_size, shuffle=shuffle)
@torch.no_grad()
def predict(model: nn.Module, loader: DataLoader, device: torch.device) -> np.ndarray:
model.eval()
probabilities = []
for features, _ in loader:
logits = model(features.to(device))
probabilities.append(torch.sigmoid(logits).cpu().numpy())
return np.concatenate(probabilities)
def main() -> None:
parser = argparse.ArgumentParser(description="Train the pose GRU")
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("--device", choices=["auto", "cpu", "cuda"], default="auto")
parser.add_argument("--seed", type=int)
parser.add_argument("--architecture", choices=["gru", "tcn"], default="gru")
args = parser.parse_args()
config, dataset, features, splits = prepare_experiment_data(
args.dataset, args.config, args.output, seed=args.seed
)
set_seed(config["seed"])
labels = dataset["labels"]
feature_indices = np.arange(features.shape[2], dtype=np.int64)
if args.architecture == "tcn":
feature_indices = CORE_FEATURE_INDICES
features = features[:, :, feature_indices]
train_mean = features[splits.train].mean(axis=(0, 1), keepdims=True)
train_std = features[splits.train].std(axis=(0, 1), keepdims=True)
train_std = np.maximum(train_std, 1e-5)
scaled = ((features - train_mean) / train_std).astype(np.float32)
if args.device == "auto":
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(args.device)
if device.type == "cuda" and not torch.cuda.is_available():
raise SystemExit("CUDA was requested but torch.cuda.is_available() is false")
training = config["training"]
batch_size = int(training["batch_size"])
train_loader = make_loader(scaled, labels, splits.train, batch_size, True)
val_loader = make_loader(scaled, labels, splits.val, batch_size, False)
test_loader = make_loader(scaled, labels, splits.test, batch_size, False)
model_config = config["model"]
if args.architecture == "gru":
model = PoseGRU(
input_size=features.shape[2],
hidden_size=int(model_config["hidden_size"]),
num_layers=int(model_config["num_layers"]),
dropout=float(model_config["dropout"]),
).to(device)
model_name = "pose_gru"
else:
model = PoseTCN(
input_size=features.shape[2],
channels=int(model_config["tcn_channels"]),
dropout=float(model_config["dropout"]),
).to(device)
model_name = "pose_tcn"
train_labels = labels[splits.train]
pos_weight = float(np.sum(train_labels == 0) / max(np.sum(train_labels == 1), 1))
criterion = nn.BCEWithLogitsLoss(pos_weight=torch.tensor(pos_weight, device=device))
optimizer = torch.optim.AdamW(
model.parameters(),
lr=float(training["learning_rate"]),
weight_decay=float(training["weight_decay"]),
)
best_loss = float("inf")
best_state = None
wait = 0
history = {"train_loss": [], "val_loss": []}
started = time.perf_counter()
for epoch in range(1, int(training["epochs"]) + 1):
model.train()
train_losses = []
for batch_features, batch_labels in train_loader:
batch_features = batch_features.to(device)
batch_labels = batch_labels.to(device)
optimizer.zero_grad(set_to_none=True)
loss = criterion(model(batch_features), batch_labels)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), max_norm=2.0)
optimizer.step()
train_losses.append(float(loss.item()))
model.eval()
val_losses = []
with torch.no_grad():
for batch_features, batch_labels in val_loader:
batch_features = batch_features.to(device)
batch_labels = batch_labels.to(device)
val_losses.append(float(criterion(model(batch_features), batch_labels).item()))
train_loss = float(np.mean(train_losses))
val_loss = float(np.mean(val_losses))
history["train_loss"].append(train_loss)
history["val_loss"].append(val_loss)
print(f"Epoch {epoch:02d}: train_loss={train_loss:.4f}, val_loss={val_loss:.4f}")
if val_loss < best_loss - 1e-4:
best_loss = val_loss
best_state = copy.deepcopy(model.state_dict())
wait = 0
else:
wait += 1
if wait >= int(training["patience"]):
print("Early stopping")
break
training_seconds = time.perf_counter() - started
if best_state is None:
raise RuntimeError("Training did not produce a checkpoint")
model.load_state_dict(best_state)
val_probabilities = predict(model, val_loader, device)
threshold = choose_f1_threshold(labels[splits.val], val_probabilities)
test_probabilities = predict(model, test_loader, device)
metrics = classification_metrics(labels[splits.test], test_probabilities, threshold)
metrics.update(
{
"model": model_name,
"architecture": args.architecture,
"seed": int(config["seed"]),
"device": str(device),
"device_name": torch.cuda.get_device_name(device) if device.type == "cuda" else "CPU",
"torch_version": torch.__version__,
"torch_cuda_version": torch.version.cuda,
"epochs_completed": len(history["train_loss"]),
"best_val_loss": best_loss,
"training_seconds": training_seconds,
"parameters": sum(parameter.numel() for parameter in model.parameters()),
"dataset": str(args.dataset),
"split": split_summary(labels, dataset["groups"], splits),
"result_scope": "test set only",
}
)
output_dir = args.output / model_name
output_dir.mkdir(parents=True, exist_ok=True)
torch.save(
{
"state_dict": {key: value.cpu() for key, value in best_state.items()},
"architecture": args.architecture,
"input_size": int(features.shape[2]),
"hidden_size": int(model_config["hidden_size"]),
"num_layers": int(model_config["num_layers"]),
"tcn_channels": int(model_config["tcn_channels"]),
"dropout": float(model_config["dropout"]),
"feature_indices": feature_indices,
"sequence_length": int(features.shape[1]),
"visibility_threshold": float(config["data"]["visibility_threshold"]),
"feature_mean": train_mean.squeeze().astype(np.float32),
"feature_std": train_std.squeeze().astype(np.float32),
"threshold": threshold,
},
output_dir / "model.pt",
)
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, model_name
)
plt.figure(figsize=(6, 4))
plt.plot(history["train_loss"], label="Train")
plt.plot(history["val_loss"], label="Validation")
plt.xlabel("Epoch")
plt.ylabel("BCE loss")
plt.title(f"Qua trinh huan luyen {model_name}")
plt.legend()
plt.tight_layout()
plt.savefig(output_dir / "training_curve.png", dpi=180)
plt.close()
print(
f"{model_name}: F1={metrics['f1']:.3f}, recall={metrics['recall']:.3f}, "
f"specificity={metrics['specificity']:.3f}, AUC={metrics['roc_auc']:.3f}"
)
if __name__ == "__main__":
main()