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,)