added method to instantiate from prediction dict
This commit is contained in:
+57
-24
@@ -1,4 +1,5 @@
|
|||||||
"""Definition of ModelData data model."""
|
"""Definition of ModelData data model."""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from .angle import AngleData
|
from .angle import AngleData
|
||||||
@@ -17,6 +18,7 @@ from .visual_syntax import VisualSyntaxData
|
|||||||
|
|
||||||
class ModelData(DataModel):
|
class ModelData(DataModel):
|
||||||
"""ModelData model for data IO with combined ML model."""
|
"""ModelData model for data IO with combined ML model."""
|
||||||
|
|
||||||
visual_syntax: VisualSyntaxData
|
visual_syntax: VisualSyntaxData
|
||||||
contact: ContactData
|
contact: ContactData
|
||||||
angle: AngleData
|
angle: AngleData
|
||||||
@@ -34,8 +36,50 @@ class ModelData(DataModel):
|
|||||||
"""Instantiate with random numbers."""
|
"""Instantiate with random numbers."""
|
||||||
kwargs = {
|
kwargs = {
|
||||||
field: field_info.annotation.from_random() # type: ignore
|
field: field_info.annotation.from_random() # type: ignore
|
||||||
for field, field_info
|
for field, field_info in cls.model_fields.items()
|
||||||
in cls.model_fields.items()
|
}
|
||||||
|
return cls(**kwargs)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def from_prediction_dict(
|
||||||
|
cls,
|
||||||
|
prediction_dict: dict,
|
||||||
|
) -> ModelData:
|
||||||
|
"""Instantiate from prediction dictionary."""
|
||||||
|
kwargs = {
|
||||||
|
'visual_syntax': VisualSyntaxData.from_tensor(
|
||||||
|
prediction_dict['visual_syntax'],
|
||||||
|
),
|
||||||
|
'contact': ContactData.from_tensor(
|
||||||
|
prediction_dict['contact'],
|
||||||
|
),
|
||||||
|
'angle': AngleData.from_tensor(
|
||||||
|
prediction_dict['angle'],
|
||||||
|
),
|
||||||
|
'point_of_view': PointOfViewData.from_tensor(
|
||||||
|
prediction_dict['point_of_view'],
|
||||||
|
),
|
||||||
|
'distance': DistanceData.from_tensor(
|
||||||
|
prediction_dict['distance'],
|
||||||
|
),
|
||||||
|
'modality_lighting': ModalityLightingData.from_tensor(
|
||||||
|
prediction_dict['modality_lighting'],
|
||||||
|
),
|
||||||
|
'modality_color': ModalityColorData.from_tensor(
|
||||||
|
prediction_dict['modality_color'],
|
||||||
|
),
|
||||||
|
'modality_depth': ModalityDepthData.from_tensor(
|
||||||
|
prediction_dict['modality_depth'],
|
||||||
|
),
|
||||||
|
'information_value': InformationValueData.from_tensor(
|
||||||
|
prediction_dict['information_value'],
|
||||||
|
),
|
||||||
|
'framing': FramingData.from_tensor(
|
||||||
|
prediction_dict['framing'],
|
||||||
|
),
|
||||||
|
'salience': SalienceData.from_tensor(
|
||||||
|
prediction_dict['salience'],
|
||||||
|
),
|
||||||
}
|
}
|
||||||
return cls(**kwargs)
|
return cls(**kwargs)
|
||||||
|
|
||||||
@@ -56,27 +100,16 @@ class ModelData(DataModel):
|
|||||||
) -> ModelData:
|
) -> ModelData:
|
||||||
"""Instantiate from annotation."""
|
"""Instantiate from annotation."""
|
||||||
kwargs = {
|
kwargs = {
|
||||||
'visual_syntax': VisualSyntaxData
|
'visual_syntax': VisualSyntaxData.from_choice(visual_syntax),
|
||||||
.from_choice(visual_syntax),
|
'contact': ContactData.from_choice(contact),
|
||||||
'contact': ContactData
|
'angle': AngleData.from_choice(angle),
|
||||||
.from_choice(contact),
|
'point_of_view': PointOfViewData.from_choice(point_of_view),
|
||||||
'angle': AngleData
|
'distance': DistanceData.from_choice(distance),
|
||||||
.from_choice(angle),
|
'modality_lighting': ModalityLightingData.from_choice(modality_lighting),
|
||||||
'point_of_view': PointOfViewData
|
'modality_color': ModalityColorData.from_choice(modality_color),
|
||||||
.from_choice(point_of_view),
|
'modality_depth': ModalityDepthData.from_choice(modality_depth),
|
||||||
'distance': DistanceData
|
'information_value': InformationValueData.from_choice(information_value),
|
||||||
.from_choice(distance),
|
'framing': FramingData.from_choice(framing),
|
||||||
'modality_lighting': ModalityLightingData
|
'salience': SalienceData.from_choice(salience),
|
||||||
.from_choice(modality_lighting),
|
|
||||||
'modality_color': ModalityColorData
|
|
||||||
.from_choice(modality_color),
|
|
||||||
'modality_depth': ModalityDepthData
|
|
||||||
.from_choice(modality_depth),
|
|
||||||
'information_value': InformationValueData
|
|
||||||
.from_choice(information_value),
|
|
||||||
'framing': FramingData
|
|
||||||
.from_choice(framing),
|
|
||||||
'salience': SalienceData
|
|
||||||
.from_choice(salience),
|
|
||||||
}
|
}
|
||||||
return cls(**kwargs)
|
return cls(**kwargs)
|
||||||
|
|||||||
Reference in New Issue
Block a user