"""Definition of database Dataset class.""" from __future__ import annotations import logging import random from datetime import UTC, datetime from pydantic import BaseModel, Field from pymongo.collection import Collection class Dataset(BaseModel): """Database Dataset model.""" create_time: datetime = Field(default_factory=lambda: datetime.now(UTC)) 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, ) logging.debug('finished') return dataset def save( self, collection: Collection, ) -> None: """Save dataset to database.""" res = collection.insert_one( document=self.model_dump(), ) logging.debug('inserted document: %s', res)