From dc976540bfb5d018b2c87d4fd4f6d877eed2a8d2 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Wed, 20 Mar 2024 19:29:05 +0100 Subject: [PATCH] added definition of dataset --- core/database/classes/dataset.py | 52 ++++++++++++++++++++++++++++++++ 1 file changed, 52 insertions(+) create mode 100644 core/database/classes/dataset.py diff --git a/core/database/classes/dataset.py b/core/database/classes/dataset.py new file mode 100644 index 0000000..cea1d16 --- /dev/null +++ b/core/database/classes/dataset.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import random + +from pydantic import BaseModel + + +class Dataset(BaseModel): + train_names: list[str] + test_names: list[str] + validation_names: list[str] + + @classmethod + def fraction_map(cls) -> dict[str, float]: + """Dict with train, test and validation fractions.""" + # define map + split_map = { + 'train': 0.7, + 'test': 0.2, + 'validation': 0.1, + } + # sanity check + assert sum(split_map.values()) == 1.0 + return split_map + + @classmethod + def new_from_name_list( + cls, + name_list: list[str], + ) -> Dataset: + """Generate new dataset from list of filenames.""" + # calculate split fractions + fraction_map = cls.fraction_map() + num_total = len(name_list) + num_validation = round(num_total * fraction_map['validation']) + num_test = round(num_total * fraction_map['test']) + # split data + validation_name_list = random.choices(name_list, k=num_validation) + name_list = [ + name for name in name_list if name not in validation_name_list + ] + test_name_list = random.choices(name_list, k=num_test) + train_name_list = [ + name for name in name_list if name not in test_name_list + ] + # instantiate object + dataset = Dataset( + train_names=train_name_list, + test_names=test_name_list, + validation_names=validation_name_list, + ) + return dataset