FallKLTN / tests /test_split_and_model.py
minhy112's picture
Upload fall detection code, trained models, and repeated experiments
9313a90 verified
Raw History Blame Contribute Delete
984 Bytes
import numpy as np
import torch
from fall_detection.data import group_train_val_test_split
from fall_detection.models import PoseGRU, PoseTCN
def test_group_split_has_no_leakage() -> None:
groups = np.repeat(np.asarray([f"g{i}" for i in range(20)]), 4)
labels = np.tile(np.asarray([0, 1, 0, 1]), 20)
splits = group_train_val_test_split(labels, groups, seed=7)
group_sets = [set(groups[index]) for index in (splits.train, splits.val, splits.test)]
assert group_sets[0].isdisjoint(group_sets[1])
assert group_sets[0].isdisjoint(group_sets[2])
assert group_sets[1].isdisjoint(group_sets[2])
def test_gru_returns_one_logit_per_sequence() -> None:
model = PoseGRU(input_size=140, hidden_size=16)
result = model(torch.randn(5, 40, 140))
assert result.shape == (5,)
def test_tcn_returns_one_logit_per_sequence() -> None:
model = PoseTCN(input_size=60, channels=16)
result = model(torch.randn(5, 40, 60))
assert result.shape == (5,)