Files
visual_critical_discourse_a…/model/src/models/visual_communication.py
T
Brian Bjarke Jensen fd9140093d
Code Quality Pipeline / Test (push) Successful in 4m5s
CI Pipeline / Test (pull_request) Failing after 3m42s
CI Pipeline / Build and Publish (pull_request) Has been skipped
moved webui and updated poetry packages
2024-05-20 19:19:29 +02:00

119 lines
3.8 KiB
Python

from __future__ import annotations
import torch.nn as nn
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
from shared.dto import AngleData
from shared.dto import ContactData
from shared.dto import DistanceData
from shared.dto import FramingData
from shared.dto import InformationValueData
from shared.dto import ModalityColorData
from shared.dto import ModalityDepthData
from shared.dto import ModalityLightingData
from shared.dto import PointOfViewData
from shared.dto import SalienceData
from shared.dto import VisualSyntaxData
class VisualCommunicationModel(nn.Module):
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):
# 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