Container abstracting a list of callbacks. Arguments: callbacks: List of `Callback` instances. queue_length: Queue length for keeping running statistics over callback execution time.
| 184 | |
| 185 | |
| 186 | class CallbackList(object): |
| 187 | """Container abstracting a list of callbacks. |
| 188 | |
| 189 | Arguments: |
| 190 | callbacks: List of `Callback` instances. |
| 191 | queue_length: Queue length for keeping |
| 192 | running statistics over callback execution time. |
| 193 | """ |
| 194 | |
| 195 | def __init__(self, callbacks=None, queue_length=10): |
| 196 | callbacks = callbacks or [] |
| 197 | self.callbacks = [c for c in callbacks] |
| 198 | self.queue_length = queue_length |
| 199 | self.params = {} |
| 200 | self.model = None |
| 201 | self._reset_batch_timing() |
| 202 | |
| 203 | def _reset_batch_timing(self): |
| 204 | self._delta_t_batch = 0. |
| 205 | self._delta_ts = collections.defaultdict( |
| 206 | lambda: collections.deque([], maxlen=self.queue_length)) |
| 207 | |
| 208 | def append(self, callback): |
| 209 | self.callbacks.append(callback) |
| 210 | |
| 211 | def set_params(self, params): |
| 212 | self.params = params |
| 213 | for callback in self.callbacks: |
| 214 | callback.set_params(params) |
| 215 | |
| 216 | def set_model(self, model): |
| 217 | self.model = model |
| 218 | for callback in self.callbacks: |
| 219 | callback.set_model(model) |
| 220 | |
| 221 | def _call_batch_hook(self, mode, hook, batch, logs=None): |
| 222 | """Helper function for all batch_{begin | end} methods.""" |
| 223 | if not self.callbacks: |
| 224 | return |
| 225 | hook_name = 'on_{mode}_batch_{hook}'.format(mode=mode, hook=hook) |
| 226 | if hook == 'begin': |
| 227 | self._t_enter_batch = time.time() |
| 228 | if hook == 'end': |
| 229 | # Batch is ending, calculate batch time. |
| 230 | self._delta_t_batch = time.time() - self._t_enter_batch |
| 231 | |
| 232 | logs = logs or {} |
| 233 | t_before_callbacks = time.time() |
| 234 | for callback in self.callbacks: |
| 235 | batch_hook = getattr(callback, hook_name) |
| 236 | batch_hook(batch, logs) |
| 237 | self._delta_ts[hook_name].append(time.time() - t_before_callbacks) |
| 238 | |
| 239 | delta_t_median = np.median(self._delta_ts[hook_name]) |
| 240 | if (self._delta_t_batch > 0. and |
| 241 | delta_t_median > 0.95 * self._delta_t_batch and delta_t_median > 0.1): |
| 242 | logging.warning( |
| 243 | 'Method (%s) is slow compared ' |
no outgoing calls
no test coverage detected