72 lines
2.1 KiB
Python
Executable File
72 lines
2.1 KiB
Python
Executable File
"""Definition of database Dataset class."""
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import random
|
|
from datetime import datetime
|
|
from datetime import UTC
|
|
|
|
from pydantic import BaseModel
|
|
from pydantic import 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)
|