Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
12 changes: 11 additions & 1 deletion leads/data_persistence/analyzer/inference.py
Original file line number Diff line number Diff line change
@@ -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, \
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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))]
Expand Down
3 changes: 3 additions & 0 deletions leads_vec_dp/run.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]())
Expand Down