add_model_resnet18 #31
@@ -2,9 +2,29 @@ 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 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):
|
||||
@@ -13,6 +33,16 @@ class VisualCommunicationModel(nn.Module):
|
||||
# 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
|
||||
@@ -20,10 +50,69 @@ class VisualCommunicationModel(nn.Module):
|
||||
# prepare result map
|
||||
results = {}
|
||||
# predict visual syntax
|
||||
visual_syntax_pred = self.visual_syntax_tail(vis_rep).cpu()
|
||||
# visual_syntax_pred = float()
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user