Skip to content

module trains

solverpy_learn.builder.plugins.trains

Trains

Bases: Managed

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/trains.py
 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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
class Trains(Managed):

   def __init__(self, dataname: str, filename: str = "train.in", **kwargs: Any):
      Managed.__init__(
         self,
         pid="trains",
         dataname=dataname,
         filename=filename,
         **kwargs,
      )
      self._lock = None
      self._enabled = True
      self.reset(dataname, filename)

   def connect(self, manager: "SyncManager") -> None:
      """Create process-shared state from the session-owned Manager."""
      if self._lock is None:
         self._lock = manager.Lock()

   def disconnect(self) -> None:
      """Discard process-shared proxies before their Manager is shut down."""
      self._lock = None

   def represent(self) -> dict[str, Any]:
      return dict(
         cls=f"{self.__class__.__module__}.{self.__class__.__name__}",
         dataname=self._dataname,
         filename=self._filename,
      )

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

   def path(
      self,
      dataname: (str | None) = None,
      filename: (str | None) = None,
   ) -> Any:
      dataname = dataname or self._dataname
      filename = filename or self._filename
      return os.path.join(bids.dbpath(NAME), dataname, filename)

   def exists(self) -> bool:
      return os.path.isfile(self.path())

   def link(self, src: str):
      if not os.path.isfile(src):
         logger.warning(f"Link source not found: {src}.")
         return
      rellink(src, self.path())

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

   def finished(
      self,
      instance: tuple[str, str],
      strategy: str,
      output: str,
      result: dict[str, Any],
   ):
      if not (output and self._solver.solved(result)):
         return
      samples = self.extract(instance, strategy, output, result)
      self.save(instance, strategy, samples)

   def extract(
      self,
      instance: tuple[str, str],
      strategy: str,
      output: str,
      result: dict[str, Any],
   ) -> Any:
      del instance, strategy, output, result  # unused arguments
      "Extract training samples from `output`."
      raise NotImplementedError()

   def save(
      self,
      instance: tuple[str, str],
      strategy: str,
      samples: str,
   ) -> None:
      if (not samples) or (not self._enabled):
         return
      if self._lock is None:
         raise RuntimeError("Trains must be connected before evaluation")
      self._lock.acquire()
      try:
         os.makedirs(os.path.dirname(self.path()), exist_ok=True)
         with open(self.path(), "a") as fa:
            fa.write(samples)
            self.stats(instance, strategy, samples)
      finally:
         self._lock.release()

   def stats(
      self,
      instance: tuple[str, str],
      strategy: str,
      samples: str,
   ):
      "Save optional statistics."
      del instance, strategy, samples  # unused arguments
      pass

connect(manager: SyncManager) -> None

Create process-shared state from the session-owned Manager.

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/trains.py
39
40
41
42
def connect(self, manager: "SyncManager") -> None:
   """Create process-shared state from the session-owned Manager."""
   if self._lock is None:
      self._lock = manager.Lock()

disconnect() -> None

Discard process-shared proxies before their Manager is shut down.

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/trains.py
44
45
46
def disconnect(self) -> None:
   """Discard process-shared proxies before their Manager is shut down."""
   self._lock = None

stats(instance: tuple[str, str], strategy: str, samples: str)

Save optional statistics.

Source code in packages/solverpy-learn/src/solverpy_learn/builder/plugins/trains.py
128
129
130
131
132
133
134
135
136
def stats(
   self,
   instance: tuple[str, str],
   strategy: str,
   samples: str,
):
   "Save optional statistics."
   del instance, strategy, samples  # unused arguments
   pass