Skip to content
Draft
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
43 changes: 24 additions & 19 deletions src/detectmatelibrary/common/core.py
Original file line number Diff line number Diff line change
@@ -1,27 +1,19 @@
from detectmatelibrary.common._core_op._fit_logic import FitLogicState, StatesL
from detectmatelibrary.common._core_op._schema_pipeline import SchemaPipeline
from detectmatelibrary.common._core_op._fed_component import FedOperations
from detectmatelibrary.common._core_op._basic_component import (
Component,
TInput,
TOutput,
)
from detectmatelibrary.common._core_op._basic_component import (Component, TInput, TOutput)
from detectmatelibrary.common._core_op._fit_logic import FitLogic

from detectmatelibrary.utils.data_buffer import DataBuffer, ArgsBuffer, BufferMode
from detectmatelibrary.utils.disk_backed_queue import DiskBackedQueue
from detectmatelibrary.utils.id_generator import SimpleIDGenerator

from detectmatelibrary.common._config import BasicConfig

from detectmatelibrary.schemas import BaseSchema

from detectmatelibrary.tools.logging import logger, setup_logging


from typing import Any
from pydantic import Field

from detectmatelibrary.common._other_op._persistency_components import PersistencyOp
from pathlib import Path
from tempfile import TemporaryDirectory


setup_logging()
Expand All @@ -31,23 +23,35 @@

class TrainBuffer:
def __init__(self) -> None:
self.buffer: list[BaseSchema | list[BaseSchema]] = []

def __len__(self) -> int:
return len(self.buffer)
queue_directory = Path("/tmp/detectmate-queues")
queue_directory.mkdir(parents=True, exist_ok=True)
self._temp_directory = TemporaryDirectory(dir=queue_directory)
queue_path = Path(self._temp_directory.name) / "queue"
try:
self.buffer: DiskBackedQueue = DiskBackedQueue(queue_path)
except BaseException:
self._temp_directory.cleanup()
raise

def __add__(self, elem: BaseSchema | list[BaseSchema]) -> "TrainBuffer":
self.buffer.append(elem)
self.buffer.add_line(elem)
return self

def __next__(self) -> BaseSchema | list[BaseSchema]:
if len(self.buffer) == 0:
res = self.buffer.get_line()
if res is None:
raise StopIteration
return self.buffer.pop(0)
return res

def __iter__(self) -> "TrainBuffer":
return self

def close(self) -> None:
try:
self.buffer.close()
finally:
self._temp_directory.cleanup()


# Core component ################################################

Expand Down Expand Up @@ -149,6 +153,7 @@ def process(self, data: BaseSchema | bytes) -> BaseSchema | bytes | None:
self.train(input_=data_buffered)
elif self.fitlogic.finish_training():
logger.debug(f"<<{self.name}>> finalizing training")
self.buffer_train.close()
self.post_train()

output_ = self.output_schema()
Expand Down
Loading
Loading