Files
visual_critical_discourse_a…/model/src/models/visual_communication.py
T
2024-10-20 20:15:19 +00:00

63 lines
2.4 KiB
Python

"""Definition of VisualCommunicationModel class."""
from __future__ import annotations
from torch import nn
from shared.mongodb.classes import ModelData
from .angle import AngleTail
from .contact import ContactTail
from .distance import DistanceTail
from .framing import FramingTail
from .information_value import InformationValueTail
from .modality_color import ModalityColorTail
from .modality_depth import ModalityDepthTail
from .modality_lighting import ModalityLightingTail
from .point_of_view import PointOfViewTail
from .resnet18_head import ResNet18Head
from .salience import SalienceTail
from .visual_syntax import VisualSyntaxTail
class VisualCommunicationModel(nn.Module):
"""Visual communication model."""
def __init__(self, download_resnet_weights: bool = False):
super().__init__()
# store other models
self.resnet_head = ResNet18Head(download_resnet_weights)
self.visual_syntax_tail = VisualSyntaxTail()
self.contact_tail = ContactTail()
self.angle_tail = AngleTail()
self.point_of_view_tail = PointOfViewTail()
self.distance_tail = DistanceTail()
self.modality_lighting_tail = ModalityLightingTail()
self.modality_color_tail = ModalityColorTail()
self.modality_depth_tail = ModalityDepthTail()
self.information_value_tail = InformationValueTail()
self.framing_tail = FramingTail()
self.salience_tail = SalienceTail()
def forward(self, x) -> ModelData:
"""Calculate model output on data."""
# generate visual representation
features = self.resnet_head(x)
# make predictions
prediction_dict = {
'visual_syntax': self.visual_syntax_tail(features),
'contact': self.contact_tail(features),
'angle': self.angle_tail(features),
'point_of_view': self.point_of_view_tail(features),
'distance': self.distance_tail(features),
'modality_lighting': self.modality_lighting_tail(features),
'modality_color': self.modality_color_tail(features),
'modality_depth': self.modality_depth_tail(features),
'information_value': self.information_value_tail(features),
'framing': self.framing_tail(features),
'salience': self.salience_tail(features),
}
# convert to respective classes
data = ModelData.from_prediction_dict(prediction_dict)
return data