added from_tensor method

This commit is contained in:
Brian Bjarke Jensen
2024-07-28 23:13:15 +02:00
parent 0bfbb1920c
commit cead3d24c8
+13 -13
View File
@@ -1,10 +1,9 @@
"""Definition of DataModel base class.""" """Definition of DataModel base class."""
from __future__ import annotations
import random import random
from pydantic import BaseModel from pydantic import BaseModel, ValidationError
from pydantic import ValidationError from torch import Tensor
class DataModel(BaseModel): class DataModel(BaseModel):
@@ -33,26 +32,27 @@ class DataModel(BaseModel):
raise ValidationError() raise ValidationError()
assert isinstance(option, str), 'option is not a string' assert isinstance(option, str), 'option is not a string'
allowed_options_list = cls.list_fields() allowed_options_list = cls.list_fields()
assert option in allowed_options_list, \ assert (
f"{option} is not among allowed fields {allowed_options_list}" option in allowed_options_list
), f"{option} is not among allowed fields {allowed_options_list}"
kwargs = {field: 0 for field in cls.list_fields()} kwargs = {field: 0 for field in cls.list_fields()}
kwargs[option] = 1 kwargs[option] = 1
return cls(**kwargs) return cls(**kwargs)
@classmethod @classmethod
def from_list(cls, data_list: list[float]): def from_tensor(cls, tensor: Tensor):
"""Instantiate from list of values.""" """Instantiate from list of values."""
kwargs = {key: val for key, val in zip(cls.list_fields(), data_list)} assert tensor.size(dim=0) == 1, f'tensor batch larger than 1: {tensor}'
data_list = [float(t.item()) for t in tensor[0]]
kwargs = dict(zip(cls.list_fields(), data_list))
return cls(**kwargs) return cls(**kwargs)
def __repr__(self) -> str: def __repr__(self) -> str:
model_dict = self.model_dump() model_dict = self.model_dump()
model_repr_str = f"{self.classname()}(" model_repr_str = f'{self.classname()}('
model_repr_str += ', '.join([ model_repr_str += ', '.join(
f"{field}={value:.3f}" [f'{field}={value:.3f}' for field, value in model_dict.items()],
for field, value )
in model_dict.items()
])
model_repr_str += ')' model_repr_str += ')'
return model_repr_str return model_repr_str