Coverage for pySDC/core/controller.py: 99%
187 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-03 11:35 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-03 11:35 +0000
1import logging
2import os
3import sys
4from typing import Any, Dict, List, Optional, Type
5import numpy as np
7from pySDC.core.base_transfer import BaseTransfer
8from pySDC.core.errors import ControllerError
9from pySDC.helpers.pysdc_helper import FrozenClass
10from pySDC.core.check_convergence import CheckConvergence
11from pySDC.core.default_hook import DefaultHooks
12from pySDC.core.timings import CPUTimings
15# short helper class to add params as attributes
16class _Pars(FrozenClass):
17 def __init__(self, params: Dict[str, Any]) -> None:
18 self.mssdc_jac: bool = True
19 self.predict_type: Optional[str] = None
20 self.all_to_done: bool = False
21 self.logger_level: int = 20
22 self.log_to_file: bool = False
23 self.dump_setup: bool = True
24 self.fname: str = 'run_pid' + str(os.getpid()) + '.log'
25 self.use_iteration_estimator: bool = False
27 for k, v in params.items():
28 setattr(self, k, v)
30 self._freeze()
33class Controller(object):
34 """
35 Abstract base class of the controllers, which set up hooks and convergence controllers and run the steps in time.
36 """
38 def __init__(
39 self, controller_params: Dict[str, Any], description: Dict[str, Any], useMPI: Optional[bool] = None
40 ) -> None:
41 """
42 Initialization routine for the base controller
44 Args:
45 controller_params (dict): parameter set for the controller and the steps
46 description (dict): description of the problem, sweeper, levels, ... passed to the steps
47 useMPI (bool): whether the controller communicates via MPI
48 """
49 self.useMPI: Optional[bool] = useMPI
50 self.description: Dict[str, Any] = description
52 # check if we have a hook on this list. If not, use default class.
53 self.__hooks: List[Any] = []
54 hook_classes: List[Type[Any]] = [DefaultHooks, CPUTimings]
55 user_hooks = controller_params.get('hook_class', [])
56 hook_classes += user_hooks if type(user_hooks) == list else [user_hooks]
57 [self.add_hook(hook) for hook in hook_classes]
59 for hook in self.hooks:
60 hook.pre_setup(step=None, level_number=None)
62 # the hooks go into the parameters, not into the caller's dictionary, where dump_setup would mark the
63 # defaults as user-defined and a controller reusing the dictionary would get them again
64 self.params: _Pars = _Pars({**controller_params, 'hook_class': hook_classes})
66 self.__setup_custom_logger(self.params.logger_level, self.params.log_to_file, self.params.fname)
67 self.logger: logging.Logger = logging.getLogger('controller')
69 if self.params.use_iteration_estimator:
70 raise ControllerError(
71 'The `use_iteration_estimator` controller parameter has been removed. Its only '
72 'implementation lived in `controller_MPI`, was never switched on anywhere, had no '
73 'tests, and deadlocked when used: every rank but the last posted a broadcast that '
74 'the last rank only matched if the estimate happened to fire. Use the '
75 '`CheckIterationEstimatorNonMPI` convergence controller instead, as '
76 '`pySDC/tutorial/step_8/C_iteration_estimator.py` does. Note it is, as its name '
77 'says, not yet available under MPI.'
78 )
80 self.base_convergence_controllers: List[Type[Any]] = [CheckConvergence]
81 self.setup_convergence_controllers(description)
83 @staticmethod
84 def __setup_custom_logger(
85 level: Optional[int] = None, log_to_file: Optional[bool] = None, fname: Optional[str] = None
86 ) -> None:
87 """
88 Helper function to set main parameters for the logging facility
90 Args:
91 level (int): level of logging
92 log_to_file (bool): flag to turn on/off logging to file
93 fname (str):
94 """
96 assert type(level) is int
98 # specify formats and handlers
99 if log_to_file:
100 file_formatter = logging.Formatter(
101 fmt='%(asctime)s - %(name)s - %(module)s - %(funcName)s - %(lineno)d - %(levelname)s: %(message)s'
102 )
103 if os.path.isfile(fname):
104 file_handler = logging.FileHandler(fname, mode='a')
105 else:
106 file_handler = logging.FileHandler(fname, mode='w')
107 file_handler.setFormatter(file_formatter)
108 else:
109 file_handler = None
111 std_formatter = logging.Formatter(fmt='%(name)s - %(levelname)s: %(message)s')
113 if level <= logging.DEBUG:
114 import warnings
116 warnings.warn('Running with debug output will degrade performance as all output is immediately flushed.')
118 class StreamFlushingHandler(logging.StreamHandler):
119 """
120 This will immediately flush any messages to the output.
121 """
123 def emit(self, record: logging.LogRecord) -> None:
124 super().emit(record)
125 self.flush()
127 std_handler = StreamFlushingHandler(sys.stdout)
128 else:
129 std_handler = logging.StreamHandler(sys.stdout)
131 std_handler.setFormatter(std_formatter)
133 # instantiate logger
134 logger = logging.getLogger('')
136 # remove handlers from previous calls to controller
137 for handler in logger.handlers[:]:
138 logger.removeHandler(handler)
140 logger.setLevel(level)
141 logger.addHandler(std_handler)
142 if log_to_file:
143 logger.addHandler(file_handler)
144 else:
145 pass
147 def add_hook(self, hook: Type[Any]) -> None:
148 """
149 Add a hook to the controller which will be called in addition to all other hooks whenever something happens.
150 The hook is only added if a hook of the same class is not already present.
152 Args:
153 hook (pySDC.Hook): A hook class that is derived from the core hook class
155 Returns:
156 None
157 """
158 if hook not in [type(me) for me in self.hooks]:
159 self.__hooks += [hook()]
161 def welcome_message(self) -> None:
162 """Log the pySDC welcome banner at info level."""
163 out = (
164 "Welcome to the one and only, really very astonishing and 87.3% bug free"
165 + "\n"
166 + r" _____ _____ _____ "
167 + "\n"
168 + r" / ____| __ \ / ____|"
169 + "\n"
170 + r" _ __ _ _| (___ | | | | | "
171 + "\n"
172 + r" | '_ \| | | |\___ \| | | | | "
173 + "\n"
174 + r" | |_) | |_| |____) | |__| | |____ "
175 + "\n"
176 + r" | .__/ \__, |_____/|_____/ \_____|"
177 + "\n"
178 + r" | | __/ | "
179 + "\n"
180 + r" |_| |___/ "
181 + "\n"
182 + r" "
183 )
184 self.logger.info(out)
186 def dump_setup(self, step: Any, controller_params: Dict[str, Any], description: Dict[str, Any]) -> None:
187 """
188 Helper function to dump the setup used for this controller
190 Args:
191 step (pySDC.Step.step): the step instance (will/should be the first one only)
192 controller_params (dict): controller parameters
193 description (dict): description of the problem
194 """
196 self.welcome_message()
197 out = 'Setup overview (--> user-defined, -> dependency) -- BEGIN'
198 self.logger.info(out)
199 out = '----------------------------------------------------------------------------------------------------\n\n'
200 out += 'Controller: %s\n' % self.__class__
201 for k, v in sorted(vars(self.params).items()):
202 if not k.startswith('_'):
203 if k in controller_params:
204 out += '--> %s = %s\n' % (k, v)
205 else:
206 out += ' %s = %s\n' % (k, v)
208 out += '\nStep: %s\n' % step.__class__
209 for k, v in sorted(vars(step.params).items()):
210 if not k.startswith('_'):
211 if k in description.get('step_params', {}):
212 out += '--> %s = %s\n' % (k, v)
213 else:
214 out += ' %s = %s\n' % (k, v)
215 out += f' Number of steps: {step.status.time_size}\n'
217 out += ' Level: %s\n' % step.levels[0].__class__
218 for L in step.levels:
219 out += ' Level %2i\n' % L.level_index
220 for k, v in sorted(vars(L.params).items()):
221 if not k.startswith('_'):
222 if k in description['level_params']:
223 out += '--> %s = %s\n' % (k, v)
224 else:
225 out += ' %s = %s\n' % (k, v)
226 out += '--> Problem: %s\n' % L.prob.__class__
227 for k, v in sorted(L.prob.params.items()):
228 if k in description['problem_params']:
229 out += '--> %s = %s\n' % (k, v)
230 else:
231 out += ' %s = %s\n' % (k, v)
232 out += ' -> Data type u: %s\n' % L.prob.dtype_u
233 out += ' -> Data type f: %s\n' % L.prob.dtype_f
234 out += '--> Sweeper: %s\n' % L.sweep.__class__
235 for k, v in sorted(vars(L.sweep.params).items()):
236 if not k.startswith('_'):
237 if k in description['sweeper_params']:
238 out += '--> %s = %s\n' % (k, v)
239 else:
240 out += ' %s = %s\n' % (k, v)
241 out += ' -> Collocation: %s\n' % L.sweep.coll.__class__
243 if len(step.levels) > 1:
244 if 'base_transfer_class' in description and description['base_transfer_class'] is not BaseTransfer:
245 out += '--> Base Transfer: %s\n' % step.base_transfer.__class__
246 else:
247 out += ' Base Transfer: %s\n' % step.base_transfer.__class__
248 for k, v in sorted(vars(step.base_transfer.params).items()):
249 if not k.startswith('_'):
250 if k in description['base_transfer_params']:
251 out += '--> %s = %s\n' % (k, v)
252 else:
253 out += ' %s = %s\n' % (k, v)
254 out += '--> Space Transfer: %s\n' % step.base_transfer.space_transfer.__class__
255 for k, v in sorted(vars(step.base_transfer.space_transfer.params).items()):
256 if not k.startswith('_'):
257 if k in description['space_transfer_params']:
258 out += '--> %s = %s\n' % (k, v)
259 else:
260 out += ' %s = %s\n' % (k, v)
262 out += '\n'
263 out += self.get_convergence_controllers_as_table(description)
264 out += '\n'
265 self.logger.info(out)
267 out = '----------------------------------------------------------------------------------------------------'
268 self.logger.info(out)
269 out = 'Setup overview (--> user-defined, -> dependency) -- END\n'
270 self.logger.info(out)
272 def run(self, u0: Any, t0: float, Tend: float) -> Any:
273 """
274 Abstract interface to the run() method
276 Args:
277 u0: initial values
278 t0 (float): starting time
279 Tend (float): ending time
280 """
281 raise NotImplementedError('ERROR: controller has to implement run(self, u0, t0, Tend)')
283 @property
284 def hooks(self) -> List[Any]:
285 """
286 Getter for the hooks
288 Returns:
289 pySDC.Hooks.hooks: hooks
290 """
291 return self.__hooks
293 @property
294 def steps(self) -> List[Any]:
295 """
296 Getter for the steps this controller owns.
298 Controllers that hold the whole block expose them as `MS`; MPI controllers hold a single step
299 as `S`. Dispatch on which of those exists rather than on the class name, so that subclasses
300 keep working.
302 Returns:
303 list: the steps owned by this controller
304 """
305 return self.MS if hasattr(self, 'MS') else [self.S]
307 def check_variable_coefficients(self, num_procs: int) -> None:
308 """
309 Reject k-dependent QDelta coefficients outside plain SDC.
311 MIN-SR-FLEX and the Jumper variants vary QDelta with the sweep index, and the nilpotency
312 argument behind them is derived for SDC, where that index *is* the SDC iteration count.
313 Anything needing more iterations (PFASST) or fewer (MLSDC) breaks that identity and would
314 need its own analysis first.
316 They are also only refreshed by `Sweeper.updateVariableCoeffs`, which runs on the finest
317 level of the Jacobi sweep alone, so on a coarse level or on the Gauss-Seidel path they
318 silently degrade to a fixed preconditioner. Failing loudly beats either of those.
320 Note this gates parallelism across *steps* and the number of *levels*. Parallelism across
321 collocation nodes (`generic_implicit_MPI` and friends) is still SDC and stays allowed.
323 Args:
324 num_procs (int): number of parallel time steps
326 Raises:
327 ControllerError: if a k-dependent QDelta is combined with multiple levels or steps
328 """
329 S = self.steps[0]
330 if len(S.levels) == 1 and num_procs == 1:
331 return
333 for level in S.levels:
334 for name in ['genQI', 'genQE']:
335 generator = getattr(level.sweep, name, None)
336 if generator is not None and generator.isKDependent():
337 raise ControllerError(
338 f'{type(generator).__name__} varies QDelta with the sweep index and is only '
339 f'verified for SDC, but you have {len(S.levels)} level(s) and {num_procs} '
340 f'step(s). Use a preconditioner with fixed coefficients, e.g. MIN-SR-S or LU.'
341 )
343 def step_is_active(self, time: float, block_start: float, Tend: float) -> bool:
344 """
345 Whether a step starting at `time`, in a block starting at `block_start`, still has work.
347 The default is that a step runs if it starts before the end of the interval, so a block
348 may be run partially. An algorithm that couples its steps too tightly to drop one of them
349 answers this from `block_start` instead and runs the block whole.
351 Args:
352 time (float): when this step starts
353 block_start (float): when the first step of this step's block starts
354 Tend (float): ending time
356 Returns:
357 bool: whether this step takes part
358 """
359 return time < Tend - 10 * np.finfo(float).eps
361 def setup_convergence_controllers(self, description: Dict[str, Any]) -> None:
362 '''
363 Setup variables needed for convergence controllers, notably a list containing all of them and a list containing
364 their order. Also, we add the `CheckConvergence` convergence controller, which takes care of maximum iteration
365 count or a residual based stopping criterion, as well as all convergence controllers added to the description.
367 Args:
368 description (dict): The description object used to instantiate the controller
370 Returns:
371 None
372 '''
373 self.convergence_controllers: List[Any] = []
374 # List of indices specifying the order of convergence controllers
375 self.convergence_controller_order: List[int] = []
376 conv_classes = description.get('convergence_controllers', {})
378 # instantiate the convergence controllers
379 for conv_class, params in conv_classes.items():
380 self.add_convergence_controller(conv_class, description=description, params=params)
382 return None
384 def add_convergence_controller(
385 self,
386 convergence_controller: Type[Any],
387 description: Dict[str, Any],
388 params: Optional[Dict[str, Any]] = None,
389 allow_double: bool = False,
390 ) -> None:
391 '''
392 Add an individual convergence controller to the list of convergence controllers and instantiate it.
393 Afterwards, the order of the convergence controllers is updated.
395 Args:
396 convergence_controller (pySDC.ConvergenceController): The convergence controller to be added
397 description (dict): The description object used to instantiate the controller
398 params (dict): Parameters for the convergence controller
399 allow_double (bool): Allow adding the same convergence controller multiple times
401 Returns:
402 None
403 '''
404 # check if we passed any sort of special params
405 params = {**({} if params is None else params), 'useMPI': self.useMPI}
407 # check if we already have the convergence controller or if we want to have it multiple times
408 if convergence_controller not in [type(me) for me in self.convergence_controllers] or allow_double:
409 self.convergence_controllers.append(convergence_controller(self, params, description))
411 # update ordering
412 orders = [C.params.control_order for C in self.convergence_controllers]
413 self.convergence_controller_order = np.arange(len(self.convergence_controllers))[np.argsort(orders)]
415 return None
417 def get_convergence_controllers_as_table(self, description: Dict[str, Any]) -> str:
418 '''
419 This function is for debugging purposes to keep track of the different convergence controllers and their order.
421 Args:
422 description (dict): Description of the problem
424 Returns:
425 str: Table of convergence controllers as a string
426 '''
427 out = 'Active convergence controllers:'
428 out += '\n | # | order | convergence controller'
429 out += '\n----+----+-------+---------------------------------------------------------------------------------------'
430 for i in range(len(self.convergence_controllers)):
431 C = self.convergence_controllers[self.convergence_controller_order[i]]
433 # figure out how the convergence controller was added
434 if type(C) in description.get('convergence_controllers', {}).keys(): # added by user
435 user_added = '--> '
436 elif type(C) in self.base_convergence_controllers: # added by default
437 user_added = ' '
438 else: # added as dependency
439 user_added = ' -> '
441 out += f'\n{user_added}|{i:3} | {C.params.control_order:5} | {type(C).__name__}'
443 return out
445 def return_stats(self) -> Dict[Any, Any]:
446 """
447 Return the merged stats from all hooks
449 Returns:
450 dict: Merged stats from all hooks
451 """
452 stats = {}
453 for hook in self.hooks:
454 stats = {**stats, **hook.return_stats()}
455 return stats