Files
visual_critical_discourse_a…/core/database/classes.py
T

166 lines
5.3 KiB
Python

from __future__ import annotations
import logging
from base64 import b64decode
from base64 import b64encode
from io import BytesIO
from pathlib import Path
from PIL import Image
from pydantic import BaseModel
from pydantic import field_serializer
from pydantic import field_validator
from src.model_experiential import (
VisualSyntaxModelOutput,
)
from src.model_interpersonal import AngleModelOutput
from src.model_interpersonal import ContactModelOutput
from src.model_interpersonal import DistanceModelOutput
from src.model_interpersonal import ModalityColorModelOutput
from src.model_interpersonal import ModalityDepthModelOutput
from src.model_interpersonal import ModalityLightingModelOutput
from src.model_interpersonal import PointOfViewModelOutput
from src.model_textual import FramingModelOutput
from src.model_textual import InformationValueModelOutput
from src.model_textual import SalienceModelOutput
class NoDocumentFoundException(Exception):
pass
class ModelOutputs(BaseModel):
visual_syntax: VisualSyntaxModelOutput
contact: ContactModelOutput
angle: AngleModelOutput
point_of_view: PointOfViewModelOutput
distance: DistanceModelOutput
modality_lighting: ModalityLightingModelOutput
modality_color: ModalityColorModelOutput
modality_depth: ModalityDepthModelOutput
information_value: InformationValueModelOutput
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() # type: ignore
for field, field_info
in cls.model_fields.items()
}
return cls(**kwargs)
@classmethod
def from_annotations(
cls,
visual_syntax: str,
contact: str,
angle: str,
point_of_view: str,
distance: str,
modality_lighting: str,
modality_color: str,
modality_depth: str,
information_value: str,
framing: str,
salience: str,
) -> ModelOutputs:
"""Instantiate from annotation."""
kwargs = {
'visual_syntax': VisualSyntaxModelOutput
.from_choice(visual_syntax),
'contact': ContactModelOutput
.from_choice(contact),
'angle': AngleModelOutput
.from_choice(angle),
'point_of_view': PointOfViewModelOutput
.from_choice(point_of_view),
'distance': DistanceModelOutput
.from_choice(distance),
'modality_lighting': ModalityLightingModelOutput
.from_choice(modality_lighting),
'modality_color': ModalityColorModelOutput
.from_choice(modality_color),
'modality_depth': ModalityDepthModelOutput
.from_choice(modality_depth),
'information_value': InformationValueModelOutput
.from_choice(information_value),
'framing': FramingModelOutput
.from_choice(framing),
'salience': SalienceModelOutput
.from_choice(salience),
}
return cls(**kwargs)
class VisualCommunication(BaseModel):
name: str
image: Image.Image
annotation: ModelOutputs | None = None
prediction: ModelOutputs | None = None
class Config:
arbitrary_types_allowed = True
@classmethod
def classname(cls) -> str:
"""Return classname."""
return cls.__name__
@classmethod
def from_file(cls, path: Path) -> VisualCommunication:
"""Instantiate from file."""
name = path.stem
image = Image.open(path)
image.load()
return VisualCommunication(name=name, image=image)
@classmethod
def decode_image(cls, content: str) -> Image.Image:
"""Decode image."""
_, content_data = content.split(',')
return Image.open(BytesIO(b64decode(content_data)))
@field_serializer('image')
def serialize_image(image: Image.Image) -> bytes: # type: ignore
buffer = BytesIO()
image.save(buffer, format='JPEG')
return buffer.getvalue()
@field_validator('image', mode='before')
@classmethod
def convert_to_image(
cls,
image: Image.Image | BytesIO | bytes,
) -> Image.Image:
if isinstance(image, bytes):
image = BytesIO(image)
if isinstance(image, BytesIO):
image = Image.open(image)
return image
def __repr__(self) -> str:
return f"{self.classname()}(name='{self.name}')"
def webencoded_image(self) -> str:
"""Convert image to be displayed on webpage."""
# convert images to bytes string
buffer = BytesIO()
self.image.save(buffer, format='png')
img_enc = b64encode(buffer.getvalue()).decode('utf-8')
return f"data:image/png;base64, {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()