Source code for omegaml.runtimes.runtime

  1import logging
  2from copy import deepcopy
  3from socket import gethostname
  4
  5from celery import Celery
  6from celery.events import EventReceiver
  7
  8from omegaml.mongoshim import mongo_url
  9from omegaml.util import dict_merge
 10
 11logger = logging.getLogger(__name__)
 12
 13
 14class CeleryTask:
 15    """
 16    A thin wrapper for a Celery.Task object
 17
 18    This is so that we can collect common delay arguments on the
 19    .task() call
 20    """
 21
 22    def __init__(self, task, kwargs):
 23        """
 24
 25        Args:
 26            task (Celery.Task): the celery task object
 27            kwargs (dict): optional, the kwargs to pass to apply_async
 28        """
 29        self.task = task
 30        self.kwargs = dict(kwargs)
 31
 32    def _apply_kwargs(self, task_kwargs, celery_kwargs):
 33        # update task_kwargs from runtime's passed on kwargs
 34        # update celery_kwargs to match celery routing semantics
 35        task_kwargs.update(self.kwargs.get('task', {}))
 36        celery_kwargs.update(self.kwargs.get('routing', {}))
 37        if 'label' in celery_kwargs:
 38            celery_kwargs['queue'] = celery_kwargs['label']
 39            del celery_kwargs['label']
 40
 41    def _apply_auth(self, args, kwargs, celery_kwargs):
 42        from omegaml.client.auth import AuthenticationEnv
 43
 44        AuthenticationEnv.active().taskauth(args, kwargs, celery_kwargs)
 45
 46    def apply_async(self, args=None, kwargs=None, **celery_kwargs):
 47        """
 48
 49        Args:
 50            args (tuple): the task args
 51            kwargs (dict): the task kwargs, passed as task.apply_async(kwargs=kwargs)
 52            celery_kwargs (dict): apply_async kwargs, passed as task.apply_async(..., **celery_kwargs)
 53
 54        Returns:
 55            AsyncResult
 56        """
 57        args = args or tuple()
 58        kwargs = kwargs or {}
 59        self._apply_kwargs(kwargs, celery_kwargs)
 60        self._apply_auth(args, kwargs, celery_kwargs)
 61        return self.task.apply_async(args=args, kwargs=kwargs, **celery_kwargs)
 62
 63    def delay(self, *args, **kwargs):
 64        """
 65        submit the task with args and kwargs to pass on
 66
 67        This calls task.apply_async and passes on the self.kwargs.
 68        """
 69        return self.apply_async(args=args, kwargs=kwargs)
 70
 71    def signature(self, args=None, kwargs=None, immutable=False, **celery_kwargs):
 72        """return the task signature with all kwargs and celery_kwargs applied"""
 73        self._apply_kwargs(kwargs, celery_kwargs)
 74        sig = self.task.signature(args=args, kwargs=kwargs, **celery_kwargs, immutable=immutable)
 75        return sig
 76
 77    def run(self, *args, **kwargs):
 78        return self.delay(*args, **kwargs)
 79
 80
[docs] 81class OmegaRuntime: 82 """ 83 omegaml compute cluster gateway 84 """ 85 86 def __init__(self, omega, bucket=None, defaults=None, celeryconf=None): 87 from omegaml.util import settings 88 89 self.omega = omega 90 defaults = defaults or settings() 91 self.bucket = bucket 92 self.pure_python = getattr(defaults, 'OMEGA_FORCE_PYTHON_CLIENT', False) 93 self.pure_python = self.pure_python or self._client_is_pure_python() 94 self._create_celery_app(defaults, celeryconf=celeryconf) 95 # temporary requirements, use .require() to set 96 self._require_kwargs = dict(task={}, routing={}) 97 # fixed default arguments, use .require(always=True) to set 98 self._task_default_kwargs = dict(task={}, routing={}) 99 # default routing label 100 self._default_label = self.celeryapp.conf.get('CELERY_DEFAULT_QUEUE') 101 102 def __repr__(self): 103 return f'OmegaRuntime({self.omega.__repr__()})' 104 105 @property 106 def auth(self): 107 return None 108 109 @property 110 def _common_kwargs(self): 111 common = deepcopy(self._task_default_kwargs) 112 common['task'].update(pure_python=self.pure_python, __bucket=self.bucket) 113 common['task'].update(self._require_kwargs['task']) 114 common['routing'].update(self._require_kwargs['routing']) 115 return common 116 117 @property 118 def _inspect(self): 119 return self.celeryapp.control.inspect() 120 121 @property 122 def is_local(self): 123 return self.celeryapp.conf['CELERY_ALWAYS_EAGER'] 124
[docs] 125 def mode(self, local=None, logging=None): 126 """specify runtime modes 127 128 Args: 129 local (bool): if True, all execution will run locally, else on 130 the configured remote cluster 131 logging (bool|str|tuple): if True, will set the root logger output 132 at INFO level; a single string is the name of the logger, 133 typically a module name; a tuple (logger, level) will select 134 logger and the level. Valid levels are INFO, WARNING, ERROR, 135 CRITICAL, DEBUG 136 137 Usage:: 138 139 # run all runtime tasks locally 140 om.runtime.mode(local=True) 141 142 # enable logging both in local and remote mode 143 om.runtime.mode(logging=True) 144 145 # select a specific module and level) 146 om.runtime.mode(logging=('sklearn', 'DEBUG')) 147 148 # disable logging 149 om.runtime.mode(logging=False) 150 """ 151 if isinstance(local, bool): 152 self.celeryapp.conf['CELERY_ALWAYS_EAGER'] = local 153 self._task_default_kwargs['task']['__logging'] = logging 154 return self
155 156 def _create_celery_app(self, defaults, celeryconf=None): 157 # initialize celery as a runtimes 158 taskpkgs = defaults.OMEGA_CELERY_IMPORTS 159 celeryconf = dict(celeryconf or defaults.OMEGA_CELERY_CONFIG) 160 # ensure we use current value 161 celeryconf['CELERY_ALWAYS_EAGER'] = bool(defaults.OMEGA_LOCAL_RUNTIME) 162 if celeryconf['CELERY_RESULT_BACKEND'].startswith('mongodb://'): 163 celeryconf['CELERY_RESULT_BACKEND'] = mongo_url(self.omega, drop_kwargs=['uuidRepresentation']) 164 # initialize ssl configuration 165 if celeryconf.get('BROKER_USE_SSL'): 166 # celery > 5 requires ssl options to be specific 167 # https://docs.celeryq.dev/en/stable/userguide/configuration.html#std-setting-broker_use_ssl 168 # https://github.com/celery/kombu/issues/1493 169 # https://docs.python.org/dev/library/ssl.html#ssl.wrap_socket 170 # https://www.openssl.org/docs/man3.0/man3/SSL_CTX_set_default_verify_paths.html 171 # env variables: 172 # SSL_CERT_FILE, CA_CERTS_PATH 173 self._apply_broker_ssl(celeryconf) 174 self.celeryapp = Celery('omegaml') 175 self.celeryapp.config_from_object(celeryconf) 176 # needed to get it to actually load the tasks 177 # https://stackoverflow.com/a/35735471 178 self.celeryapp.autodiscover_tasks(taskpkgs, force=True) 179 self.celeryapp.finalize() 180 181 def _apply_broker_ssl(self, celeryconf): 182 # hook to apply broker ssl options 183 pass 184 185 def _client_is_pure_python(self): 186 try: 187 pass 188 except Exception as e: 189 logging.getLogger().info(e) 190 return True 191 else: 192 return False 193 194 def _sanitize_require(self, value): 195 # convert value into dict(label=value) 196 if isinstance(value, str): 197 return dict(label=value) 198 if isinstance(value, (list, tuple)): 199 return dict(*value) 200 return value 201
[docs] 202 def require(self, label=None, always=False, routing=None, task=None, logging=None, override=True, **kwargs): 203 """ 204 specify requirements for the task execution 205 206 Use this to specify resource or routing requirements on the next task 207 call sent to the runtime. Any requirements will be reset after the 208 call has been submitted. 209 210 Args: 211 always (bool): if True requirements will persist across task calls. defaults to False 212 label (str): the label required by the worker to have a runtime task dispatched to it. 213 'local' is equivalent to calling self.mode(local=True). 214 task (dict): if specified applied to the task's kwargs, passed as task.apply_async(..., kwargs=task) 215 routing (dict): if specified applied to the task's routing, passed as task.apply_async(..., **routing) 216 logging (str|tuple): if specified, same as runtime.mode(logging=...) 217 override (bool): if True overrides previously set .require(), defaults to True 218 kwargs: requirements specification that the runtime understands 219 220 Usage: 221 om.runtime.require(label='gpu').model('foo').fit(...) 222 223 See Also: 224 - CeleryTask.apply_async 225 - celery.app.task.Task.apply_async, specifically the kwargs= and **options 226 227 Returns: 228 self 229 """ 230 if label: 231 # avoid overriding a previous local() call by an erronous label 232 assert isinstance(label, str), "label must be valid, run om.runtime.labels() to list of active workers" 233 if label == 'local': 234 self.mode(local=True) 235 elif override: 236 self.mode(local=False) 237 # update routing, don't replace (#416) 238 routing = routing or {} 239 routing.update({'label': label or self._default_label}) 240 task = task or {} 241 routing = routing or {} 242 if task or routing: 243 if not override: 244 # override not allowed, remove previously existing 245 ex_task = dict( 246 **self._task_default_kwargs['task'], 247 **self._require_kwargs['task'], 248 ) 249 ex_routing = dict( 250 **self._task_default_kwargs['routing'], 251 **self._require_kwargs['routing'], 252 ) 253 exists_or_none = lambda k, d: k not in d or d.get(k, False) is None 254 task = {k: v for k, v in task.items() if exists_or_none(k, ex_task)} 255 routing = {k: v for k, v in routing.items() if exists_or_none(k, ex_routing)} 256 if always: 257 self._task_default_kwargs['routing'].update(routing) 258 self._task_default_kwargs['task'].update(task) 259 else: 260 self._require_kwargs['routing'].update(routing) 261 self._require_kwargs['task'].update(task) 262 else: 263 # FIXME this does not work as expected (will only reset if both task and routing are False) 264 if not task: 265 self._require_kwargs['task'] = {} 266 if not routing: 267 self._require_kwargs['routing'] = {} 268 if logging is not None: 269 self.mode(logging=logging) 270 return self
271
[docs] 272 def model(self, modelname, require=None): 273 """ 274 return a model for remote execution 275 276 Args: 277 modelname (str): the name of the object in om.models 278 require (dict): routing requirements for this job 279 280 Returns: 281 OmegaModelProxy 282 """ 283 from omegaml.runtimes.proxies.modelproxy import OmegaModelProxy 284 285 self.require(**self._sanitize_require(require)) if require else None 286 return OmegaModelProxy(modelname, runtime=self)
287
[docs] 288 def job(self, jobname, require=None): 289 """ 290 return a job for remote execution 291 292 Args: 293 jobname (str): the name of the object in om.jobs 294 require (dict): routing requirements for this job 295 296 Returns: 297 OmegaJobProxy 298 """ 299 from omegaml.runtimes.proxies.jobproxy import OmegaJobProxy 300 301 self.require(**self._sanitize_require(require)) if require else None 302 return OmegaJobProxy(jobname, runtime=self)
303
[docs] 304 def script(self, scriptname, require=None): 305 """ 306 return a script for remote execution 307 308 Args: 309 scriptname (str): the name of object in om.scripts 310 require (dict): routing requirements for this job 311 312 Returns: 313 OmegaScriptProxy 314 """ 315 from omegaml.runtimes.proxies.scriptproxy import OmegaScriptProxy 316 317 self.require(**self._sanitize_require(require)) if require else None 318 return OmegaScriptProxy(scriptname, runtime=self)
319
[docs] 320 def experiment(self, experiment, provider=None, implied_run=True, recreate=False, **tracker_kwargs): 321 """set the tracking backend and experiment 322 323 Args: 324 experiment (str): the name of the experiment 325 provider (str): the name of the provider 326 tracker_kwargs (dict): additional kwargs for the tracker 327 recreate (bool): if True, recreate the experiment (i.e. drop and recreate, 328 this is useful to change the provider or other settings. All previous data will 329 be kept) 330 331 Returns: 332 OmegaTrackingProxy 333 """ 334 from omegaml.runtimes.proxies.trackingproxy import OmegaTrackingProxy 335 336 # tracker implied_run means we are using the currently active run, i.e. with block will call exp.start() 337 tracker = OmegaTrackingProxy( 338 experiment, provider=provider, runtime=self, implied_run=implied_run, recreate=recreate, **tracker_kwargs 339 ) 340 return tracker
341
[docs] 342 def task(self, name, **kwargs): 343 """ 344 retrieve the task function from the celery instance 345 346 Args: 347 name (str): a registered celery task as ``module.tasks.task_name`` 348 kwargs (dict): routing keywords to CeleryTask.apply_async 349 350 Returns: 351 CeleryTask 352 """ 353 taskfn = self.celeryapp.tasks.get(name) or self.celeryapp.tasks.get(f'omega_{name}') 354 assert taskfn is not None, "cannot find task {name} in Celery runtime".format(**locals()) 355 kwargs = dict_merge(self._common_kwargs, dict(routing=kwargs)) 356 task = CeleryTask(taskfn, kwargs) 357 self._require_kwargs = dict(routing={}, task={}) 358 return task
359 360 @property 361 def tasks(self): 362 """return registered task names 363 364 .. versionadded:: 0.18.2 365 """ 366 return list(self.celeryapp.tasks.keys()) 367 368 def result(self, task_id, wait=True): 369 from celery.result import AsyncResult 370 371 promise = AsyncResult(task_id, app=self.celeryapp) 372 return promise.get() if wait else promise 373
[docs] 374 def settings(self, require=None): 375 """return the runtimes's cluster settings""" 376 self.require(**require) if require else None 377 return self.task('omegaml.tasks.omega_settings').delay().get()
378
[docs] 379 def ping(self, *args, require=None, wait=True, timeout=10, **kwargs): 380 """ 381 ping the runtime 382 383 Args: 384 args (tuple): task args 385 require (dict): routing requirements for this job 386 wait (bool): if True, wait for the task to return, else return 387 AsyncResult 388 timeout (int): if wait is True, the timeout in seconds, defaults to 10 389 kwargs (dict): task kwargs, as accepted by CeleryTask.apply_async 390 391 Returns: 392 * response (dict) for wait=True 393 * AsyncResult for wait=False 394 """ 395 self.require(**require) if require else None 396 promise = self.task('omegaml.tasks.omega_ping').delay(*args, **kwargs) 397 return promise.get(timeout=timeout) if wait else promise
398
[docs] 399 def enable_hostqueues(self): 400 """enable a worker-specific queue on every worker host 401 402 Returns: 403 list of labels (one entry for each hostname) 404 """ 405 control = self.celeryapp.control 406 inspect = control.inspect() 407 active = inspect.active() 408 queues = [] 409 for worker in active.keys(): 410 hostname = worker.split('@')[-1] 411 control.cancel_consumer(hostname) 412 control.add_consumer(hostname, destination=[worker]) 413 queues.append(hostname) 414 return queues
415
[docs] 416 def workers(self): 417 """list of workers 418 419 Returns: 420 dict of workers => list of active tasks 421 422 See Also: 423 celery Inspect.active() 424 """ 425 local_worker = { 426 gethostname(): [ 427 { 428 'name': 'local', 429 'is_local': True, 430 } 431 ] 432 } 433 celery_workers = self._inspect.active() or {} 434 return dict_merge(local_worker, celery_workers)
435
[docs] 436 def queues(self): 437 """list queues 438 439 Returns: 440 dict of workers => list of queues 441 442 See Also: 443 celery Inspect.active_queues() 444 """ 445 local_q = {gethostname(): [{'name': 'local', 'is_local': True}]} 446 celery_qs = self._inspect.active_queues() or {} 447 return dict_merge(local_q, celery_qs)
448
[docs] 449 def labels(self): 450 """list available labels 451 452 Returns: 453 dict of workers => list of lables 454 """ 455 return {worker: [q.get('name') for q in queues] for worker, queues in self.queues().items()}
456
[docs] 457 def stats(self): 458 """worker statistics 459 460 Returns: 461 dict of workers => dict of stats 462 463 See Also: 464 celery Inspect.stats() 465 """ 466 return self._inspect.stats()
467
[docs] 468 def status(self): 469 """current cluster status 470 471 This collects key information from .labels(), .stats() and the latest 472 worker heartbeat. Note that loadavg is only available if the worker has 473 recently sent a heartbeat and may not be accurate across the cluster. 474 475 Returns: 476 snapshot (dict): a snapshot of the cluster status 477 '<worker>': { 478 'loadavg': [0.0, 0.0, 0.0], # load average in % seen by the worker (1, 5, 15 min) 479 'processes': 1, # number of active worker processes 480 'concurrency': 1, # max concurrency 481 'uptime': 0, # uptime in seconds 482 'processed': Counter(task=n), # number of tasks processed 483 'queues': ['default'], # list of queues (labels) the worker is listening on 484 } 485 """ 486 labels = self.labels() 487 stats = self.stats() 488 heartbeat = self.events.latest() 489 snapshot = { 490 worker: { 491 'loadavg': heartbeat.get('loadavg', []), 492 'processes': stats[worker]['pool']['processes'], 493 'concurrency': stats[worker]['pool']['max-concurrency'], 494 'uptime': stats[worker]['uptime'], 495 'processed': stats[worker]['total'], 496 'queues': labels[worker], 497 } 498 for worker in labels 499 if worker in stats 500 } 501 return snapshot
502 503 @property 504 def events(self): 505 return CeleryEventStream(self.celeryapp) 506
[docs] 507 def callback(self, script_name, always=False, **kwargs): 508 """Add a callback to a registered script 509 510 The callback will be triggered upon successful or failed 511 execution of the runtime tasks. The script syntax is:: 512 513 # script.py 514 def run(om, state=None, result=None, **kwargs): 515 # state (str): 'SUCCESS'|'ERROR' 516 # result (obj): the task's serialized result 517 518 Args: 519 script_name (str): the name of the script (in om.scripts) 520 always (bool): if True always apply this callback, defaults to False 521 **kwargs: and other kwargs to pass on to the script 522 523 Returns: 524 self 525 """ 526 success_sig = ( 527 self 528 .script(script_name) 529 .task(as_callback=True) 530 .signature( 531 args=['SUCCESS', script_name], 532 kwargs=kwargs, 533 immutable=False, 534 ) 535 ) 536 error_sig = ( 537 self 538 .script(script_name) 539 .task(as_callback=True) 540 .signature( 541 args=['ERROR', script_name], 542 kwargs=kwargs, 543 immutable=False, 544 ) 545 ) 546 547 if always: 548 self._task_default_kwargs['routing']['link'] = success_sig 549 self._task_default_kwargs['routing']['link_error'] = error_sig 550 else: 551 self._require_kwargs['routing']['link'] = success_sig 552 self._require_kwargs['routing']['link_error'] = error_sig 553 return self
554 555 556class CeleryEventStream: 557 def __init__(self, app, limit=None, timeout=5, wakeup=False): 558 self.app = app 559 self.limit = limit 560 self.timeout = timeout 561 self.wakeup = wakeup 562 self.max_size = 100 563 self.buffer = [] 564 565 def handle(self, event): 566 self.buffer.append(event) 567 if len(self.buffer) > self.max_size: 568 self.buffer = self.buffer[-1 * self.max_size :] 569 570 def listen(self, handlers=None, limit=None, timeout=None): 571 # Connect to the broker using Kombu (Celery's underlying messaging system) 572 handlers = handlers or {'worker-heartbeat': self.handle} 573 limit = limit or self.limit 574 timeout = timeout or self.timeout 575 with self.app.connection() as conn: 576 # Create the EventReceiver to listen to all events 577 recv = EventReceiver(conn, handlers=handlers, app=self.app) 578 recv.capture(limit=limit, timeout=timeout, wakeup=self.wakeup) 579 580 def latest(self, timeout=None): 581 while len(self.buffer) == 0: 582 self.listen(limit=1, timeout=timeout) 583 return self.buffer[-1] 584 585 586# apply mixins 587from omegaml.runtimes.mixins.swagger import SwaggerGenerator 588from omegaml.runtimes.mixins.taskcanvas import canvas_chain, canvas_chord, canvas_group 589 590OmegaRuntime.sequence = canvas_chain 591OmegaRuntime.parallel = canvas_group 592OmegaRuntime.mapreduce = canvas_chord 593OmegaRuntime.swagger = SwaggerGenerator.build_swagger