From 81966289816bc55002f43a4e10d58ac8886c6031 Mon Sep 17 00:00:00 2001 From: Brian Bjarke Jensen Date: Sat, 24 Feb 2024 11:58:00 +0100 Subject: [PATCH] added function to randomly generate prediction data --- src/database/classes.py | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/src/database/classes.py b/src/database/classes.py index 1a88dbb..568b415 100644 --- a/src/database/classes.py +++ b/src/database/classes.py @@ -4,6 +4,9 @@ from PIL import Image from io import BytesIO from pathlib import Path from base64 import b64encode +import random +import logging +from typing import List from src.model_experiential import ExperientialModelOutput from src.model_interpersonal import ( @@ -34,6 +37,21 @@ class ModelOutputs(BaseModel): framing: FramingModelOutput salience: SalienceModelOutput + @classmethod + def list_fields(cls) -> List[str]: + """List options that are stored as attributes.""" + return list(cls.model_fields.keys()) + + @classmethod + def from_random(cls) -> ModelOutputs: + """Instantiate with random numbers.""" + kwargs = { + field: field_info.annotation.from_random() + for field, field_info + in cls.model_fields.items() + } + return cls(**kwargs) + class VisualCommunication(BaseModel): name: str @@ -83,6 +101,11 @@ class VisualCommunication(BaseModel): img_enc = b64encode(buffer.getvalue()).decode("utf-8") return img_enc + def generate_random_prediction(self, force: bool = False) -> None: + """Generate random prediction values.""" + if not force and self.prediction is not None: + logging.warning("set force=True to overwrite existing values.") + self.prediction = ModelOutputs.from_random() class NoDocumentFoundException(Exception): pass