defined final model
This commit is contained in:
@@ -2,9 +2,29 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import torch.nn as nn
|
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 .resnet18_head import ResNet18Head
|
||||||
|
from .salience import SalienceTail
|
||||||
from .visual_syntax import VisualSyntaxTail
|
from .visual_syntax import VisualSyntaxTail
|
||||||
from core.data_models import VisualSyntaxData
|
from core.dto import AngleData
|
||||||
|
from core.dto import ContactData
|
||||||
|
from core.dto import DistanceData
|
||||||
|
from core.dto import FramingData
|
||||||
|
from core.dto import InformationValueData
|
||||||
|
from core.dto import ModalityColorData
|
||||||
|
from core.dto import ModalityDepthData
|
||||||
|
from core.dto import ModalityLightingData
|
||||||
|
from core.dto import PointOfViewData
|
||||||
|
from core.dto import SalienceData
|
||||||
|
from core.dto import VisualSyntaxData
|
||||||
|
|
||||||
|
|
||||||
class VisualCommunicationModel(nn.Module):
|
class VisualCommunicationModel(nn.Module):
|
||||||
@@ -13,6 +33,16 @@ class VisualCommunicationModel(nn.Module):
|
|||||||
# store other models
|
# store other models
|
||||||
self.resnet_head = ResNet18Head()
|
self.resnet_head = ResNet18Head()
|
||||||
self.visual_syntax_tail = VisualSyntaxTail()
|
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):
|
def forward(self, x):
|
||||||
# generate visual representation
|
# generate visual representation
|
||||||
@@ -20,10 +50,69 @@ class VisualCommunicationModel(nn.Module):
|
|||||||
# prepare result map
|
# prepare result map
|
||||||
results = {}
|
results = {}
|
||||||
# predict visual syntax
|
# predict visual syntax
|
||||||
visual_syntax_pred = self.visual_syntax_tail(vis_rep).cpu()
|
|
||||||
# visual_syntax_pred = float()
|
|
||||||
results['visual_syntax'] = VisualSyntaxData.from_list(
|
results['visual_syntax'] = VisualSyntaxData.from_list(
|
||||||
visual_syntax_pred,
|
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
|
return results
|
||||||
|
|||||||
Reference in New Issue
Block a user