File size: 984 Bytes
9313a90
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
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,)