renamed module

This commit is contained in:
brian
2024-10-20 20:00:13 +00:00
parent d1ee43e135
commit e2338ad710
26 changed files with 63 additions and 82 deletions
+67
View File
@@ -0,0 +1,67 @@
"""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)