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

1import logging 

2import os 

3import sys 

4from typing import Any, Dict, List, Optional, Type 

5import numpy as np 

6 

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 

13 

14 

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 

26 

27 for k, v in params.items(): 

28 setattr(self, k, v) 

29 

30 self._freeze() 

31 

32 

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

37 

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 

43 

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 

51 

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] 

58 

59 for hook in self.hooks: 

60 hook.pre_setup(step=None, level_number=None) 

61 

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

65 

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

68 

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 ) 

79 

80 self.base_convergence_controllers: List[Type[Any]] = [CheckConvergence] 

81 self.setup_convergence_controllers(description) 

82 

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 

89 

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

95 

96 assert type(level) is int 

97 

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 

110 

111 std_formatter = logging.Formatter(fmt='%(name)s - %(levelname)s: %(message)s') 

112 

113 if level <= logging.DEBUG: 

114 import warnings 

115 

116 warnings.warn('Running with debug output will degrade performance as all output is immediately flushed.') 

117 

118 class StreamFlushingHandler(logging.StreamHandler): 

119 """ 

120 This will immediately flush any messages to the output. 

121 """ 

122 

123 def emit(self, record: logging.LogRecord) -> None: 

124 super().emit(record) 

125 self.flush() 

126 

127 std_handler = StreamFlushingHandler(sys.stdout) 

128 else: 

129 std_handler = logging.StreamHandler(sys.stdout) 

130 

131 std_handler.setFormatter(std_formatter) 

132 

133 # instantiate logger 

134 logger = logging.getLogger('') 

135 

136 # remove handlers from previous calls to controller 

137 for handler in logger.handlers[:]: 

138 logger.removeHandler(handler) 

139 

140 logger.setLevel(level) 

141 logger.addHandler(std_handler) 

142 if log_to_file: 

143 logger.addHandler(file_handler) 

144 else: 

145 pass 

146 

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. 

151 

152 Args: 

153 hook (pySDC.Hook): A hook class that is derived from the core hook class 

154 

155 Returns: 

156 None 

157 """ 

158 if hook not in [type(me) for me in self.hooks]: 

159 self.__hooks += [hook()] 

160 

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) 

185 

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 

189 

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

195 

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) 

207 

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' 

216 

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__ 

242 

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) 

261 

262 out += '\n' 

263 out += self.get_convergence_controllers_as_table(description) 

264 out += '\n' 

265 self.logger.info(out) 

266 

267 out = '----------------------------------------------------------------------------------------------------' 

268 self.logger.info(out) 

269 out = 'Setup overview (--> user-defined, -> dependency) -- END\n' 

270 self.logger.info(out) 

271 

272 def run(self, u0: Any, t0: float, Tend: float) -> Any: 

273 """ 

274 Abstract interface to the run() method 

275 

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

282 

283 @property 

284 def hooks(self) -> List[Any]: 

285 """ 

286 Getter for the hooks 

287 

288 Returns: 

289 pySDC.Hooks.hooks: hooks 

290 """ 

291 return self.__hooks 

292 

293 @property 

294 def steps(self) -> List[Any]: 

295 """ 

296 Getter for the steps this controller owns. 

297 

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. 

301 

302 Returns: 

303 list: the steps owned by this controller 

304 """ 

305 return self.MS if hasattr(self, 'MS') else [self.S] 

306 

307 def check_variable_coefficients(self, num_procs: int) -> None: 

308 """ 

309 Reject k-dependent QDelta coefficients outside plain SDC. 

310 

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. 

315 

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. 

319 

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. 

322 

323 Args: 

324 num_procs (int): number of parallel time steps 

325 

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 

332 

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 ) 

342 

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. 

346 

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. 

350 

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 

355 

356 Returns: 

357 bool: whether this step takes part 

358 """ 

359 return time < Tend - 10 * np.finfo(float).eps 

360 

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. 

366 

367 Args: 

368 description (dict): The description object used to instantiate the controller 

369 

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', {}) 

377 

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) 

381 

382 return None 

383 

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. 

394 

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 

400 

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} 

406 

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

410 

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

414 

415 return None 

416 

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. 

420 

421 Args: 

422 description (dict): Description of the problem 

423 

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

432 

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 = ' -> ' 

440 

441 out += f'\n{user_added}|{i:3} | {C.params.control_order:5} | {type(C).__name__}' 

442 

443 return out 

444 

445 def return_stats(self) -> Dict[Any, Any]: 

446 """ 

447 Return the merged stats from all hooks 

448 

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