diff --git a/README.md b/README.md index 683ebe4..9382817 100644 --- a/README.md +++ b/README.md @@ -622,6 +622,8 @@ inferences: repeat: 100 # default: 1 enhanced: true # default: false assume_initial_zeros: true # default: false + clear: # default: [] + - speed methods: - safe-speed - speed-by-acceleration diff --git a/leads/data_persistence/analyzer/inference.py b/leads/data_persistence/analyzer/inference.py index 36c2fb3..5617d0f 100644 --- a/leads/data_persistence/analyzer/inference.py +++ b/leads/data_persistence/analyzer/inference.py @@ -1,5 +1,6 @@ from abc import ABCMeta as _ABCMeta, abstractmethod as _abstractmethod -from typing import Any as _Any, override as _override, Generator as _Generator, Literal as _Literal +from typing import Any as _Any, override as _override, Generator as _Generator, Literal as _Literal, \ + Sequence as _Sequence from leads.data import distance_between from leads.data_persistence.analyzer.utils import time_invalid, speed_invalid, acceleration_invalid, \ @@ -266,6 +267,13 @@ def __init__(self, file: str, chunk_size: int = 100) -> None: super().__init__(file, chunk_size) self._raw_data: tuple[dict[str, _Any], ...] = () self._inferred_data: list[dict[str, _Any]] = [] + self._clear: set[str] = set() + + def clear(self, entry: str) -> None: + self._clear.add(entry) + + def clear_all(self, entries: _Sequence[str]) -> None: + self._clear = set(entries) @_override def __len__(self) -> int: @@ -308,6 +316,8 @@ def load(self) -> None: super().load() raw_data = [] for row in super().__iter__(): + for clear_entry in self._clear: + row[clear_entry] = None raw_data.append(row) self._raw_data = tuple(raw_data) self._inferred_data = [{} for _ in range(len(raw_data))] diff --git a/leads_vec_dp/run.py b/leads_vec_dp/run.py index a8c8a5a..87447c1 100644 --- a/leads_vec_dp/run.py +++ b/leads_vec_dp/run.py @@ -40,6 +40,9 @@ def run(target: str) -> int: if "inferences" in target: dataset = _InferredDataset(target["dataset"]) inferences = target["inferences"] + if "clear" in inferences: + dataset.clear_all(inferences["clear"]) + inferences.pop("clear") methods = [] for method in inferences["methods"]: methods.append(INFERENCE_METHODS[method]())