Skip to content

module multi

solverpy_learn.builder.plugins.multi

MultiTrains

Bases: SvmTrains

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/multi.py
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
class MultiTrains(SvmTrains):

   def __init__(self, dataname: str):
      self._trains: list["SvmTrains"] = []
      self._dataname = dataname
      self._pid = "trains"

   def represent(self) -> dict:
      return dict(
         cls=f"{self.__class__.__module__}.{self.__class__.__name__}",
         dataname=self._dataname,
         trains=[t.represent() for t in self._trains],
      )

   def dispatch(self, t: "SvmTrains"):
      self._trains.append(t)

   def apply(self, function: Callable[["SvmTrains"], None]) -> None:
      for t in self._trains:
         function(t)

   def connect(self, manager: "SyncManager") -> None:
      """Connect every underlying training-data collector."""
      self.apply(lambda x: x.connect(manager))

   def disconnect(self) -> None:
      """Disconnect every underlying training-data collector."""
      self.apply(lambda x: x.disconnect())

   def register(self, solver: "SolverPy") -> None:
      super().register(solver)
      self._solver = solver
      for t in self._trains:
         t._solver = solver

   def reset(
      self,
      dataname: (str | None) = None,
      filename: str = "train.in",
   ) -> None:
      if dataname:
         self._dataname = dataname
      self.apply(lambda x: x.reset(dataname=dataname, filename=filename))

   def path(
      self,
      dataname: (str | None) = None,
      filename: (str | None) = None,
   ) -> tuple[str, ...]:
      return tuple(t.path(dataname, filename) for t in self._trains)

   def exists(self) -> bool:
      return all(t.exists() for t in self._trains)

   def link(self, src: str | tuple[str]):
      assert isinstance(src, tuple)
      assert len(src) == len(self._trains)
      for (s, t) in zip(src, self._trains):
         t.link(s)

   def enable(self) -> None:
      self.apply(lambda x: x.enable())

   def disable(self) -> None:
      self.apply(lambda x: x.disable())

   def finished(self, *args: Any, **kwargs: Any) -> None:
      self.apply(lambda x: x.finished(*args, **kwargs))

   def extract(self, *args: Any, **kwargs: Any) -> None:
      self.apply(lambda x: x.extract(*args, **kwargs))

   def save(self, *args: Any, **kwargs: Any) -> None:
      self.apply(lambda x: x.save(*args, **kwargs))

   def stats(self, *args: Any, **kwargs: Any) -> None:
      self.apply(lambda x: x.stats(*args, **kwargs))

   def compress(self, *args: Any, **kwargs: Any) -> None:
      self.apply(lambda x: x.compress(*args, **kwargs))

   def train_data_snapshot(self) -> None:
      self.apply(lambda x: x.train_data_snapshot())

   def train_data_stats(
      self,
      dataset: str,
      paths: tuple[str, ...] | None = None,
   ) -> list[dict[str, Any]]:
      paths = paths or self.path()
      stats = [
         train.train_data_stats(dataset, path)
         for (train, path) in zip(self._trains, paths)
      ]
      return [stat for stat in stats if stat is not None]

   def merge(
      self,
      previous: str | tuple[str, ...],
      outfilename: str,
   ) -> None:
      assert len(previous) == len(self._trains)
      assert type(previous) is tuple
      for (t0, p0) in zip(self._trains, previous):
         t0.merge(p0, outfilename)

connect(manager: SyncManager) -> None

Connect every underlying training-data collector.

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/multi.py
34
35
36
def connect(self, manager: "SyncManager") -> None:
   """Connect every underlying training-data collector."""
   self.apply(lambda x: x.connect(manager))

disconnect() -> None

Disconnect every underlying training-data collector.

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/multi.py
38
39
40
def disconnect(self) -> None:
   """Disconnect every underlying training-data collector."""
   self.apply(lambda x: x.disconnect())