File size: 2,360 Bytes
11fab85
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
"""tune/test 切分:UR Fall Detection Dataset 70 支影片的分層抽樣。



切分比例在規劃階段已與使用者確認並固定:tune = 10 falls + 13 adls(約 1/3),

test = 20 falls + 27 adls(約 2/3、佔多數以求指標穩定)。fall/adl 各自獨立

洗牌(分層),避免其中一類意外集中在同一個 split。



seed 固定為 42 且結果直接存進版控(``eval/splits.yaml``,repo 根目錄),

不在 notebook 裡即時重新產生——切分只決定一次,之後每次執行都讀同一份

名單,避免任何環境差異導致 test split 悄悄漂移(那樣所有 test 數字都不可信)。

"""

from __future__ import annotations

import random
from pathlib import Path

import yaml

from ..io.urfd import adl_sequences, fall_sequences

DEFAULT_SEED = 42
N_TUNE_FALLS = 10
N_TUNE_ADLS = 13


def generate_splits(

    seed: int = DEFAULT_SEED,

    n_tune_falls: int = N_TUNE_FALLS,

    n_tune_adls: int = N_TUNE_ADLS,

) -> dict:
    """分層洗牌切分:falls 與 adls 各自獨立洗牌,避免其中一類集中在同一 split。"""
    rng = random.Random(seed)
    falls = fall_sequences()
    adls = adl_sequences()
    rng.shuffle(falls)
    rng.shuffle(adls)
    return {
        "seed": seed,
        "tune": {
            "falls": sorted(falls[:n_tune_falls]),
            "adls": sorted(adls[:n_tune_adls]),
        },
        "test": {
            "falls": sorted(falls[n_tune_falls:]),
            "adls": sorted(adls[n_tune_adls:]),
        },
    }


def save_splits(splits: dict, path: str | Path) -> None:
    path = Path(path)
    path.parent.mkdir(parents=True, exist_ok=True)
    path.write_text(
        yaml.safe_dump(splits, sort_keys=False, allow_unicode=True), encoding="utf-8"
    )


def load_splits(path: str | Path) -> dict:
    with open(path, "r", encoding="utf-8") as f:
        return yaml.safe_load(f)


def split_of(sequence: str, splits: dict) -> str:
    """回傳某序列屬於 ``"tune"`` 或 ``"test"``;不在名單中就拋例外,

    及早抓出序列名拼字錯誤或名單過期。"""
    for split_name in ("tune", "test"):
        group = splits[split_name]
        if sequence in group["falls"] or sequence in group["adls"]:
            return split_name
    raise KeyError(f"{sequence} 不在任何 split 名單中")