tensorless 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
tensorless/runtime.py ADDED
@@ -0,0 +1,158 @@
1
+ """Runtime inference wrapper.
2
+
3
+ `tl.load("model.tl")` returns a `LoadedModel`, which knows how to rebuild
4
+ the exact architecture used at training time, load the weights, and
5
+ expose a task-appropriate prediction API:
6
+
7
+ - text-generation -> `.generate(prompt)` and `.chat()`
8
+ - text-classification -> `.predict(text)`
9
+ - classification/regression (tabular) -> `.predict(record_or_records)`
10
+ """
11
+
12
+ from __future__ import annotations
13
+
14
+ from typing import Any, Dict, List, Union
15
+
16
+ import torch
17
+
18
+ from .models.registry import build_model
19
+ from .tokenization.char_tokenizer import CharTokenizer
20
+ from .data.tabular import TabularPreprocessor
21
+ from .devices.device import get_torch_device
22
+ from .errors import ModelError
23
+ from .serialization.tl_format import load_tl
24
+
25
+
26
+ class LoadedModel:
27
+ def __init__(self, payload: Dict[str, Any], device: str = None):
28
+ self.payload = payload
29
+ self.task: str = payload["task"]
30
+ self.model_type: str = payload["model_type"]
31
+ self.config: Dict[str, Any] = payload["config"]
32
+ self.meta: Dict[str, Any] = payload["meta"]
33
+ self.metrics: Dict[str, Any] = payload.get("metrics", {})
34
+ self.dataset_fingerprint: str = payload.get("dataset_fingerprint")
35
+
36
+ self.tokenizer = None
37
+ if payload.get("tokenizer_state") is not None:
38
+ self.tokenizer = CharTokenizer.from_state_dict(payload["tokenizer_state"])
39
+
40
+ self.preprocessor = None
41
+ if payload.get("preprocessor_state") is not None:
42
+ self.preprocessor = TabularPreprocessor.from_state_dict(payload["preprocessor_state"])
43
+
44
+ device_name = device or self.config.get("device", "cpu")
45
+ self.device = get_torch_device(device_name)
46
+
47
+ self.model = build_model(self.task, self.model_type, self.config, self.meta)
48
+ self.model.load_state_dict(payload["model_state_dict"])
49
+ self.model.to(self.device)
50
+ self.model.eval()
51
+
52
+ # ------------------------------------------------------------------
53
+ # Text generation
54
+ # ------------------------------------------------------------------
55
+ def generate(self, prompt: str = "", max_new_tokens: int = 200, temperature: float = 0.8, top_k: int = 40) -> str:
56
+ if self.task != "text-generation":
57
+ raise ModelError(f"generate() is only available for text-generation models, not '{self.task}'.")
58
+ ids = self.tokenizer.encode(prompt, add_special_tokens=True)[:-1] # drop trailing eos
59
+ if not ids:
60
+ ids = [self.tokenizer.bos_id]
61
+ input_ids = torch.tensor([ids], dtype=torch.long, device=self.device)
62
+ out = self.model.generate(
63
+ input_ids,
64
+ max_new_tokens=max_new_tokens,
65
+ temperature=temperature,
66
+ top_k=top_k,
67
+ eos_id=self.tokenizer.eos_id,
68
+ )
69
+ return self.tokenizer.decode(out[0].tolist())
70
+
71
+ def chat(self) -> None:
72
+ """Interactive terminal chat loop for text-generation models."""
73
+ if self.task != "text-generation":
74
+ raise ModelError(f"chat() is only available for text-generation models, not '{self.task}'.")
75
+ print("Tensorless interactive chat. Type 'exit' or Ctrl+C to quit.")
76
+ while True:
77
+ try:
78
+ prompt = input("> ")
79
+ except (EOFError, KeyboardInterrupt):
80
+ print()
81
+ break
82
+ if prompt.strip().lower() in ("exit", "quit"):
83
+ break
84
+ reply = self.generate(prompt, max_new_tokens=200)
85
+ print(reply)
86
+
87
+ # ------------------------------------------------------------------
88
+ # Text classification
89
+ # ------------------------------------------------------------------
90
+ def _predict_text_classification(self, texts: List[str]) -> List[str]:
91
+ block_size = self.config["max_seq_len"]
92
+ batch_ids, batch_mask = [], []
93
+ for t in texts:
94
+ ids = self.tokenizer.encode(t, add_special_tokens=True)[:block_size]
95
+ mask = [1] * len(ids)
96
+ if len(ids) < block_size:
97
+ pad_len = block_size - len(ids)
98
+ ids = ids + [self.tokenizer.pad_id] * pad_len
99
+ mask = mask + [0] * pad_len
100
+ batch_ids.append(ids)
101
+ batch_mask.append(mask)
102
+ input_ids = torch.tensor(batch_ids, dtype=torch.long, device=self.device)
103
+ attn_mask = torch.tensor(batch_mask, dtype=torch.long, device=self.device)
104
+ with torch.no_grad():
105
+ logits = self.model(input_ids, attention_mask=attn_mask)
106
+ pred_idx = logits.argmax(dim=-1).tolist()
107
+ classes = self.meta["classes"]
108
+ return [classes[i] for i in pred_idx]
109
+
110
+ # ------------------------------------------------------------------
111
+ # Tabular classification / regression
112
+ # ------------------------------------------------------------------
113
+ def _predict_tabular(self, records: List[Dict[str, Any]]) -> List[Any]:
114
+ transformed = self.preprocessor.transform(records, with_target=False)
115
+ numeric = transformed["numeric"].to(self.device)
116
+ categorical = transformed["categorical"].to(self.device)
117
+ with torch.no_grad():
118
+ out = self.model(numeric, categorical)
119
+ if self.task == "classification":
120
+ pred_idx = out.argmax(dim=-1)
121
+ return self.preprocessor.inverse_target(pred_idx)
122
+ else:
123
+ return self.preprocessor.inverse_target(out)
124
+
125
+ # ------------------------------------------------------------------
126
+ # Unified predict()
127
+ # ------------------------------------------------------------------
128
+ def predict(self, x: Union[str, Dict[str, Any], List[Any]]) -> Any:
129
+ single = not isinstance(x, list)
130
+ items = [x] if single else x
131
+
132
+ if self.task == "text-classification":
133
+ preds = self._predict_text_classification(items)
134
+ elif self.task in ("classification", "regression"):
135
+ preds = self._predict_tabular(items)
136
+ elif self.task == "text-generation":
137
+ preds = [self.generate(prompt=str(i)) for i in items]
138
+ else:
139
+ raise ModelError(f"predict() not supported for task '{self.task}'.")
140
+
141
+ return preds[0] if single else preds
142
+
143
+ def info(self) -> Dict[str, Any]:
144
+ return {
145
+ "task": self.task,
146
+ "model_type": self.model_type,
147
+ "tensorless_version": self.payload.get("tensorless_version"),
148
+ "tl_format_version": self.payload.get("tl_format_version"),
149
+ "config": self.config,
150
+ "metrics": self.metrics,
151
+ "training_complete": self.payload.get("training_complete"),
152
+ "n_parameters": sum(p.numel() for p in self.model.parameters()),
153
+ }
154
+
155
+
156
+ def load_model(path: str, device: str = None) -> LoadedModel:
157
+ payload = load_tl(path)
158
+ return LoadedModel(payload, device=device)
@@ -0,0 +1,3 @@
1
+ from .tl_format import save_tl, load_tl
2
+
3
+ __all__ = ["save_tl", "load_tl"]
@@ -0,0 +1,89 @@
1
+ """The `.tl` file format.
2
+
3
+ A `.tl` file is a single portable file (a torch pickle archive under the
4
+ hood) containing everything needed for inference on a *different*
5
+ machine, with no access to the original dataset or training code:
6
+
7
+ {
8
+ "tl_format_version": int,
9
+ "tensorless_version": str,
10
+ "task": str, # e.g. "text-generation"
11
+ "model_type": str, # e.g. "transformer"
12
+ "config": {...}, # resolved training config
13
+ "meta": {...}, # vocab_size / n_classes / column info
14
+ "model_state_dict": {...},
15
+ "tokenizer_state": {...} | None,
16
+ "preprocessor_state": {...} | None,
17
+ "dataset_fingerprint": str,
18
+ "training_complete": bool,
19
+ "metrics": {...},
20
+ }
21
+
22
+ We deliberately use a single file (rather than a directory/zip of many
23
+ files) so users can `scp`/email/upload one `model.tl` and have it just
24
+ work elsewhere, per the framework's "portable single file" requirement.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import os
30
+ from typing import Any, Dict
31
+
32
+ import torch
33
+
34
+ from .._version import __version__, TL_FORMAT_VERSION
35
+ from ..errors import SerializationError
36
+
37
+ REQUIRED_KEYS = (
38
+ "tl_format_version",
39
+ "task",
40
+ "model_type",
41
+ "config",
42
+ "meta",
43
+ "model_state_dict",
44
+ )
45
+
46
+
47
+ def save_tl(path: str, payload: Dict[str, Any]) -> None:
48
+ payload = dict(payload)
49
+ payload.setdefault("tl_format_version", TL_FORMAT_VERSION)
50
+ payload.setdefault("tensorless_version", __version__)
51
+
52
+ parent = os.path.dirname(os.path.abspath(path))
53
+ if parent:
54
+ os.makedirs(parent, exist_ok=True)
55
+
56
+ tmp_path = path + ".tmp"
57
+ try:
58
+ torch.save(payload, tmp_path)
59
+ os.replace(tmp_path, path)
60
+ except Exception as e:
61
+ if os.path.exists(tmp_path):
62
+ os.remove(tmp_path)
63
+ raise SerializationError(f"Failed to write .tl file to '{path}': {e}") from e
64
+
65
+
66
+ def load_tl(path: str, map_location: str = "cpu") -> Dict[str, Any]:
67
+ if not os.path.isfile(path):
68
+ raise SerializationError(f"'{path}' does not exist.")
69
+ try:
70
+ payload = torch.load(path, map_location=map_location, weights_only=False)
71
+ except Exception as e:
72
+ raise SerializationError(f"Failed to read .tl file '{path}': {e}") from e
73
+
74
+ missing = [k for k in REQUIRED_KEYS if k not in payload]
75
+ if missing:
76
+ raise SerializationError(
77
+ f"'{path}' is missing required field(s) {missing}. It may be "
78
+ f"corrupt or not a valid Tensorless .tl file."
79
+ )
80
+
81
+ file_version = payload.get("tl_format_version")
82
+ if file_version > TL_FORMAT_VERSION:
83
+ raise SerializationError(
84
+ f"'{path}' was created with a newer .tl format (v{file_version}) "
85
+ f"than this installed version of Tensorless supports "
86
+ f"(v{TL_FORMAT_VERSION}). Please upgrade Tensorless."
87
+ )
88
+
89
+ return payload
@@ -0,0 +1,3 @@
1
+ from .char_tokenizer import CharTokenizer
2
+
3
+ __all__ = ["CharTokenizer"]
@@ -0,0 +1,84 @@
1
+ """Character-level tokenizer.
2
+
3
+ Tensorless defaults to a character-level tokenizer because it:
4
+ - requires no external vocabulary files or training corpus assumptions
5
+ - works on any UTF-8 text out of the box (any language, code, etc.)
6
+ - has a small, easily-portable vocabulary that fits directly inside a
7
+ `.tl` file
8
+
9
+ This is deliberately simple rather than a full BPE tokenizer -- Tensorless
10
+ optimizes for "it just works with zero setup" over maximal efficiency.
11
+ Advanced users can plug in their own tokenizer via `model_type=` /
12
+ extension points documented in the developer docs.
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ from typing import Dict, List, Sequence
18
+
19
+ PAD = "<pad>"
20
+ UNK = "<unk>"
21
+ BOS = "<bos>"
22
+ EOS = "<eos>"
23
+ SPECIAL_TOKENS = [PAD, UNK, BOS, EOS]
24
+
25
+
26
+ class CharTokenizer:
27
+ def __init__(self, vocab: List[str] = None):
28
+ self.vocab: List[str] = vocab if vocab is not None else list(SPECIAL_TOKENS)
29
+ self._stoi: Dict[str, int] = {c: i for i, c in enumerate(self.vocab)}
30
+
31
+ @property
32
+ def vocab_size(self) -> int:
33
+ return len(self.vocab)
34
+
35
+ @property
36
+ def pad_id(self) -> int:
37
+ return self._stoi[PAD]
38
+
39
+ @property
40
+ def unk_id(self) -> int:
41
+ return self._stoi[UNK]
42
+
43
+ @property
44
+ def bos_id(self) -> int:
45
+ return self._stoi[BOS]
46
+
47
+ @property
48
+ def eos_id(self) -> int:
49
+ return self._stoi[EOS]
50
+
51
+ @classmethod
52
+ def build(cls, texts: Sequence[str]) -> "CharTokenizer":
53
+ chars = set()
54
+ for t in texts:
55
+ chars.update(t)
56
+ vocab = list(SPECIAL_TOKENS) + sorted(chars)
57
+ return cls(vocab)
58
+
59
+ def encode(self, text: str, add_special_tokens: bool = True) -> List[int]:
60
+ ids = [self._stoi.get(c, self.unk_id) for c in text]
61
+ if add_special_tokens:
62
+ ids = [self.bos_id] + ids + [self.eos_id]
63
+ return ids
64
+
65
+ def decode(self, ids: Sequence[int], skip_special_tokens: bool = True) -> str:
66
+ chars = []
67
+ special_ids = {self.pad_id, self.bos_id, self.eos_id} if skip_special_tokens else set()
68
+ for i in ids:
69
+ if i in special_ids:
70
+ continue
71
+ if 0 <= i < len(self.vocab):
72
+ tok = self.vocab[i]
73
+ if tok == UNK:
74
+ chars.append("\ufffd")
75
+ elif tok not in SPECIAL_TOKENS:
76
+ chars.append(tok)
77
+ return "".join(chars)
78
+
79
+ def state_dict(self) -> Dict:
80
+ return {"vocab": self.vocab}
81
+
82
+ @classmethod
83
+ def from_state_dict(cls, state: Dict) -> "CharTokenizer":
84
+ return cls(vocab=list(state["vocab"]))
@@ -0,0 +1,4 @@
1
+ from .trainer import run_training
2
+ from .early_stopping import EarlyStopping
3
+
4
+ __all__ = ["run_training", "EarlyStopping"]
@@ -0,0 +1,215 @@
1
+ """Turns a loaded `Dataset` + resolved config into PyTorch-ready tensors,
2
+ train/val splits, and the tokenizer/preprocessor that produced them.
3
+
4
+ Kept separate from `trainer.py` so the "how do I turn data into tensors"
5
+ logic can be tested and extended independently of the training loop.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import random
11
+ from dataclasses import dataclass
12
+ from typing import Any, Dict, List, Optional, Tuple
13
+
14
+ import torch
15
+ from torch.utils.data import DataLoader, Dataset as TorchDataset
16
+
17
+ from ..data.loader import Dataset
18
+ from ..data.tabular import TabularPreprocessor
19
+ from ..tokenization.char_tokenizer import CharTokenizer
20
+ from ..auto.detector import target_column
21
+ from ..errors import DataError
22
+
23
+
24
+ @dataclass
25
+ class PreparedData:
26
+ train_loader: DataLoader
27
+ val_loader: Optional[DataLoader]
28
+ meta: Dict[str, Any]
29
+ tokenizer: Optional[CharTokenizer] = None
30
+ preprocessor: Optional[TabularPreprocessor] = None
31
+
32
+
33
+ class _LMChunkDataset(TorchDataset):
34
+ def __init__(self, token_ids: List[int], block_size: int):
35
+ self.ids = token_ids
36
+ self.block_size = block_size
37
+
38
+ def __len__(self) -> int:
39
+ return max(0, len(self.ids) - self.block_size)
40
+
41
+ def __getitem__(self, idx: int):
42
+ chunk = self.ids[idx: idx + self.block_size + 1]
43
+ x = torch.tensor(chunk[:-1], dtype=torch.long)
44
+ y = torch.tensor(chunk[1:], dtype=torch.long)
45
+ return x, y
46
+
47
+
48
+ class _ClsTextDataset(TorchDataset):
49
+ def __init__(self, input_ids: List[List[int]], attn_masks: List[List[int]], labels: List[int]):
50
+ self.input_ids = input_ids
51
+ self.attn_masks = attn_masks
52
+ self.labels = labels
53
+
54
+ def __len__(self) -> int:
55
+ return len(self.labels)
56
+
57
+ def __getitem__(self, idx: int):
58
+ return (
59
+ torch.tensor(self.input_ids[idx], dtype=torch.long),
60
+ torch.tensor(self.attn_masks[idx], dtype=torch.long),
61
+ torch.tensor(self.labels[idx], dtype=torch.long),
62
+ )
63
+
64
+
65
+ class _TabularDataset(TorchDataset):
66
+ def __init__(self, numeric: torch.Tensor, categorical: torch.Tensor, target: torch.Tensor):
67
+ self.numeric = numeric
68
+ self.categorical = categorical
69
+ self.target = target
70
+
71
+ def __len__(self) -> int:
72
+ return self.numeric.shape[0]
73
+
74
+ def __getitem__(self, idx: int):
75
+ return self.numeric[idx], self.categorical[idx], self.target[idx]
76
+
77
+
78
+ def _split_indices(n: int, val_split: float, seed: int) -> Tuple[List[int], List[int]]:
79
+ idx = list(range(n))
80
+ random.Random(seed).shuffle(idx)
81
+ n_val = int(n * val_split)
82
+ val_idx = idx[:n_val]
83
+ train_idx = idx[n_val:]
84
+ if not train_idx:
85
+ train_idx = idx
86
+ val_idx = []
87
+ return train_idx, val_idx
88
+
89
+
90
+ def prepare_text_generation(
91
+ ds: Dataset, cfg: Dict[str, Any], tokenizer: Optional[CharTokenizer] = None
92
+ ) -> PreparedData:
93
+ tokenizer = tokenizer or CharTokenizer.build(ds.texts)
94
+ all_ids: List[int] = []
95
+ for t in ds.texts:
96
+ all_ids.extend(tokenizer.encode(t, add_special_tokens=True))
97
+
98
+ block_size = cfg["max_seq_len"]
99
+ if len(all_ids) <= block_size:
100
+ # Pad tiny corpora so we have at least one training example.
101
+ all_ids = all_ids + [tokenizer.pad_id] * (block_size + 1 - len(all_ids))
102
+
103
+ n_val_tokens = int(len(all_ids) * cfg["val_split"])
104
+ if n_val_tokens > block_size + 1:
105
+ split_point = len(all_ids) - n_val_tokens
106
+ train_ids, val_ids = all_ids[:split_point], all_ids[split_point:]
107
+ else:
108
+ train_ids, val_ids = all_ids, None
109
+
110
+ train_ds = _LMChunkDataset(train_ids, block_size)
111
+ train_loader = DataLoader(
112
+ train_ds, batch_size=cfg["batch_size"], shuffle=True, drop_last=False
113
+ )
114
+
115
+ val_loader = None
116
+ if val_ids is not None:
117
+ val_ds = _LMChunkDataset(val_ids, block_size)
118
+ if len(val_ds) > 0:
119
+ val_loader = DataLoader(val_ds, batch_size=cfg["batch_size"], shuffle=False)
120
+
121
+ meta = {
122
+ "vocab_size": tokenizer.vocab_size,
123
+ "pad_id": tokenizer.pad_id,
124
+ "n_classes": 0,
125
+ }
126
+ return PreparedData(train_loader=train_loader, val_loader=val_loader, meta=meta, tokenizer=tokenizer)
127
+
128
+
129
+ def prepare_text_classification(
130
+ ds: Dataset,
131
+ cfg: Dict[str, Any],
132
+ tokenizer: Optional[CharTokenizer] = None,
133
+ classes: Optional[List[str]] = None,
134
+ ) -> PreparedData:
135
+ tokenizer = tokenizer or CharTokenizer.build(ds.texts)
136
+ classes = classes or sorted(set(ds.labels))
137
+ label2id = {c: i for i, c in enumerate(classes)}
138
+
139
+ block_size = cfg["max_seq_len"]
140
+ input_ids, attn_masks, labels = [], [], []
141
+ for text, label in zip(ds.texts, ds.labels):
142
+ ids = tokenizer.encode(text, add_special_tokens=True)[:block_size]
143
+ mask = [1] * len(ids)
144
+ if len(ids) < block_size:
145
+ pad_len = block_size - len(ids)
146
+ ids = ids + [tokenizer.pad_id] * pad_len
147
+ mask = mask + [0] * pad_len
148
+ input_ids.append(ids)
149
+ attn_masks.append(mask)
150
+ labels.append(label2id.get(label, 0))
151
+
152
+ n = len(labels)
153
+ train_idx, val_idx = _split_indices(n, cfg["val_split"], cfg["seed"])
154
+
155
+ def subset(indices):
156
+ return _ClsTextDataset(
157
+ [input_ids[i] for i in indices],
158
+ [attn_masks[i] for i in indices],
159
+ [labels[i] for i in indices],
160
+ )
161
+
162
+ train_loader = DataLoader(subset(train_idx), batch_size=cfg["batch_size"], shuffle=True)
163
+ val_loader = (
164
+ DataLoader(subset(val_idx), batch_size=cfg["batch_size"], shuffle=False) if val_idx else None
165
+ )
166
+
167
+ meta = {
168
+ "vocab_size": tokenizer.vocab_size,
169
+ "pad_id": tokenizer.pad_id,
170
+ "n_classes": len(classes),
171
+ "classes": classes,
172
+ }
173
+ return PreparedData(train_loader=train_loader, val_loader=val_loader, meta=meta, tokenizer=tokenizer)
174
+
175
+
176
+ def prepare_tabular(
177
+ ds: Dataset,
178
+ cfg: Dict[str, Any],
179
+ task: str,
180
+ preprocessor: Optional[TabularPreprocessor] = None,
181
+ ) -> PreparedData:
182
+ target_col = target_column(ds)
183
+ if target_col is None:
184
+ raise DataError("Could not determine a target column for tabular training.")
185
+
186
+ prep = preprocessor or TabularPreprocessor().fit(ds.records, ds.columns, target_col, task)
187
+ transformed = prep.transform(ds.records, with_target=True)
188
+
189
+ n = transformed["numeric"].shape[0]
190
+ train_idx, val_idx = _split_indices(n, cfg["val_split"], cfg["seed"])
191
+ train_idx_t = torch.tensor(train_idx, dtype=torch.long)
192
+
193
+ train_ds = _TabularDataset(
194
+ transformed["numeric"][train_idx_t],
195
+ transformed["categorical"][train_idx_t],
196
+ transformed["target"][train_idx_t],
197
+ )
198
+ train_loader = DataLoader(train_ds, batch_size=cfg["batch_size"], shuffle=True)
199
+
200
+ val_loader = None
201
+ if val_idx:
202
+ val_idx_t = torch.tensor(val_idx, dtype=torch.long)
203
+ val_ds = _TabularDataset(
204
+ transformed["numeric"][val_idx_t],
205
+ transformed["categorical"][val_idx_t],
206
+ transformed["target"][val_idx_t],
207
+ )
208
+ val_loader = DataLoader(val_ds, batch_size=cfg["batch_size"], shuffle=False)
209
+
210
+ meta = {
211
+ "n_numeric": len(prep.numeric_columns),
212
+ "categorical_vocab_sizes": prep.categorical_vocab_sizes(),
213
+ "n_classes": len(prep.classes) if task == "classification" else 0,
214
+ }
215
+ return PreparedData(train_loader=train_loader, val_loader=val_loader, meta=meta, preprocessor=prep)
@@ -0,0 +1,31 @@
1
+ """Early stopping.
2
+
3
+ Tracks a validation metric (lower is better, e.g. loss) and reports when
4
+ training should stop because it hasn't improved by at least `min_delta`
5
+ for `patience` consecutive checks.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+
11
+ class EarlyStopping:
12
+ def __init__(self, patience: int = 5, min_delta: float = 1e-4):
13
+ self.patience = patience
14
+ self.min_delta = min_delta
15
+ self.best: float = float("inf")
16
+ self.best_state: dict = None
17
+ self.num_bad_checks = 0
18
+ self.should_stop = False
19
+
20
+ def step(self, value: float, state: dict = None) -> bool:
21
+ """Returns True if `value` is a new best."""
22
+ if value < self.best - self.min_delta:
23
+ self.best = value
24
+ self.best_state = state
25
+ self.num_bad_checks = 0
26
+ return True
27
+ else:
28
+ self.num_bad_checks += 1
29
+ if self.num_bad_checks >= self.patience:
30
+ self.should_stop = True
31
+ return False