diff --git a/shared/mongodb/classes/__init__.py b/shared/mongodb/classes/__init__.py index 129df79..af91d95 100755 --- a/shared/mongodb/classes/__init__.py +++ b/shared/mongodb/classes/__init__.py @@ -1,6 +1 @@ -"""Database classes module content.""" - -from __future__ import annotations - -from .dataset import Dataset from .visual_communication import VisualCommunication diff --git a/shared/mongodb/classes/dataset.py b/shared/mongodb/classes/dataset.py deleted file mode 100755 index 633c533..0000000 --- a/shared/mongodb/classes/dataset.py +++ /dev/null @@ -1,67 +0,0 @@ -"""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)