68 lines
2.1 KiB
Python
68 lines
2.1 KiB
Python
"""Definition of DataModel base class."""
|
|
|
|
import random
|
|
|
|
from pydantic import BaseModel, ValidationError
|
|
from torch import Tensor
|
|
|
|
|
|
class DataModel(BaseModel):
|
|
"""DataModel base class."""
|
|
|
|
@classmethod
|
|
def classname(cls) -> str:
|
|
"""Return classname."""
|
|
return cls.__name__
|
|
|
|
@classmethod
|
|
def list_fields(cls) -> list[str]:
|
|
"""List options that are stored as attributes."""
|
|
return list(cls.model_fields.keys())
|
|
|
|
@classmethod
|
|
def from_random(cls):
|
|
"""Instantiate with random numbers."""
|
|
kwargs = {field: random.random() for field in cls.list_fields()}
|
|
return cls(**kwargs)
|
|
|
|
@classmethod
|
|
def from_choice(cls, option: str):
|
|
"""Instantiate from choice."""
|
|
if option is None:
|
|
raise ValidationError()
|
|
assert isinstance(option, str), 'option is not a string'
|
|
allowed_options_list = cls.list_fields()
|
|
assert (
|
|
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[option] = 1
|
|
return cls(**kwargs)
|
|
|
|
@classmethod
|
|
def from_tensor(cls, tensor: Tensor):
|
|
"""Instantiate from list of values."""
|
|
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)
|
|
|
|
def __repr__(self) -> str:
|
|
model_dict = self.model_dump()
|
|
model_repr_str = f'{self.classname()}('
|
|
model_repr_str += ', '.join(
|
|
[f'{field}={value:.3f}' for field, value in model_dict.items()],
|
|
)
|
|
model_repr_str += ')'
|
|
return model_repr_str
|
|
|
|
def highest_score_field(self) -> str:
|
|
"""Return name of field with highest score."""
|
|
model_dict = self.model_dump()
|
|
return max(model_dict, key=lambda k: model_dict[k])
|
|
|
|
def highest_score_value(self) -> float:
|
|
"""Return value of field with highest score."""
|
|
model_dict = self.model_dump()
|
|
return max(model_dict.values())
|