removed unused class
This commit is contained in:
@@ -1,6 +1 @@
|
|||||||
"""Database classes module content."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from .dataset import Dataset
|
|
||||||
from .visual_communication import VisualCommunication
|
from .visual_communication import VisualCommunication
|
||||||
|
|||||||
@@ -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)
|
|
||||||
Reference in New Issue
Block a user