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)
|