diff --git a/src/fev/task.py b/src/fev/task.py index 6eb995c..a1c73d4 100644 --- a/src/fev/task.py +++ b/src/fev/task.py @@ -340,6 +340,11 @@ class Task: name of 2 parent directories for local or S3-based datasets. This field is only here for convenience and is not used for any validation when computing the results. + dataset_description : str | None, default None + Text description of the dataset. + column_descriptions : dict[str, str] | None, default None + Text description of each column used by the task. If provided, the keys must exactly match the target, + dynamic and static columns of the task. Examples -------- @@ -379,6 +384,8 @@ class Task: past_dynamic_columns: list[str] = dataclasses.field(default_factory=list) static_columns: list[str] = dataclasses.field(default_factory=list) task_name: str | None = None + dataset_description: str | None = None + column_descriptions: dict[str, str] | None = None def __post_init__(self): if self.task_name is None: @@ -458,6 +465,16 @@ def __post_init__(self): "`generate_univariate_targets_from` cannot be used for multivariate tasks (when `target` is a list)" ) + if self.column_descriptions is not None: + expected = set(self.target_columns + self.dynamic_columns + self.static_columns) + missing = sorted(expected - set(self.column_descriptions)) + unexpected = sorted(set(self.column_descriptions) - expected) + if missing or unexpected: + raise ValueError( + "`column_descriptions` must have exactly one entry per target, dynamic and static column of the " + f"task. Missing: {missing}. Unexpected: {unexpected}." + ) + # Attributes computed after the dataset is loaded self._full_dataset: datasets.Dataset | None = None self._freq: str | None = None diff --git a/test/test_task.py b/test/test_task.py index eeef7b8..d446371 100644 --- a/test/test_task.py +++ b/test/test_task.py @@ -510,3 +510,48 @@ def test_when_datasets_prefix_is_set_then_s3_dataset_is_loaded_from_mirror(tmp_p assert loaded_dataset[0]["id"] == "series_0" assert task.dataset_path == dataset_path assert task.to_dict()["dataset_path"] == dataset_path + + +def test_when_descriptions_not_provided_then_they_default_to_none(): + task = fev.Task(dataset_path="my_dataset", horizon=12) + assert task.dataset_description is None + assert task.column_descriptions is None + assert task.to_dict()["dataset_description"] is None + assert task.to_dict()["column_descriptions"] is None + + +def test_when_column_descriptions_match_task_columns_then_task_is_created(): + column_descriptions = {"OT": "oil temperature", "HULL": "load", "LULL": "load", "store": "store id"} + task = fev.Task( + dataset_path="my_dataset", + horizon=12, + target="OT", + known_dynamic_columns=["HULL"], + past_dynamic_columns=["LULL"], + static_columns=["store"], + dataset_description="Transformer data.", + column_descriptions=column_descriptions, + ) + assert task.dataset_description == "Transformer data." + assert task.column_descriptions == column_descriptions + assert fev.Task(**task.to_dict()) == task + + +@pytest.mark.parametrize( + "column_descriptions", + [ + {"OT": "oil temperature"}, + {"OT": "oil temperature", "HULL": "load", "LULL": "load"}, + {"target": "oil temperature", "HULL": "load"}, + {}, + ], +) +def test_when_column_descriptions_do_not_match_task_columns_then_validation_error_is_raised(column_descriptions): + with pytest.raises(pydantic.ValidationError, match="column_descriptions"): + fev.Task( + dataset_path="my_dataset", + horizon=12, + target="OT", + known_dynamic_columns=["HULL"], + column_descriptions=column_descriptions, + )