63 lines
2.4 KiB
Python
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
|