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