added function to randomly generate prediction data

This commit is contained in:
Brian Bjarke Jensen
2024-02-24 11:58:00 +01:00
parent 14be5cfffc
commit 8196628981
+23
View File
@@ -4,6 +4,9 @@ from PIL import Image
from io import BytesIO from io import BytesIO
from pathlib import Path from pathlib import Path
from base64 import b64encode from base64 import b64encode
import random
import logging
from typing import List
from src.model_experiential import ExperientialModelOutput from src.model_experiential import ExperientialModelOutput
from src.model_interpersonal import ( from src.model_interpersonal import (
@@ -34,6 +37,21 @@ class ModelOutputs(BaseModel):
framing: FramingModelOutput framing: FramingModelOutput
salience: SalienceModelOutput 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): class VisualCommunication(BaseModel):
name: str name: str
@@ -83,6 +101,11 @@ class VisualCommunication(BaseModel):
img_enc = b64encode(buffer.getvalue()).decode("utf-8") img_enc = b64encode(buffer.getvalue()).decode("utf-8")
return img_enc 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): class NoDocumentFoundException(Exception):
pass pass