127 lines
3.7 KiB
Python
127 lines
3.7 KiB
Python
"""Definition of VisualCommunicationModel class."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from torch import nn
|
|
|
|
from shared.dto import (
|
|
AngleData,
|
|
ContactData,
|
|
DistanceData,
|
|
FramingData,
|
|
InformationValueData,
|
|
ModalityColorData,
|
|
ModalityDepthData,
|
|
ModalityLightingData,
|
|
PointOfViewData,
|
|
SalienceData,
|
|
VisualSyntaxData,
|
|
)
|
|
|
|
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):
|
|
super().__init__()
|
|
# store other models
|
|
self.resnet_head = ResNet18Head()
|
|
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):
|
|
"""Calculate model output on data."""
|
|
# generate visual representation
|
|
vis_rep = self.resnet_head(x)
|
|
# prepare result map
|
|
results = {}
|
|
# predict visual syntax
|
|
results['visual_syntax'] = VisualSyntaxData.from_list(
|
|
self.visual_syntax_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict contact
|
|
results['contact'] = ContactData.from_list(
|
|
self.contact_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict angle
|
|
results['angle'] = AngleData.from_list(
|
|
self.angle_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict point of view
|
|
results['point_of_view'] = PointOfViewData.from_list(
|
|
self.point_of_view_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict distance
|
|
results['distance'] = DistanceData.from_list(
|
|
self.distance_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict modality lighting
|
|
results['modality_lighting'] = ModalityLightingData.from_list(
|
|
self.modality_lighting_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict modality color
|
|
results['modality_color'] = ModalityColorData.from_list(
|
|
self.modality_color_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict modality depth
|
|
results['modality_depth'] = ModalityDepthData.from_list(
|
|
self.modality_depth_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict information value
|
|
results['information_value'] = InformationValueData.from_list(
|
|
self.information_value_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict framing
|
|
results['framing'] = FramingData.from_list(
|
|
self.framing_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
# predict salience
|
|
results['salience'] = SalienceData.from_list(
|
|
self.salience_tail(
|
|
vis_rep,
|
|
).cpu(),
|
|
)
|
|
return results
|