Skip to content

AbstractTask

Base class for all Python-based federated learning tasks.

starfish.controller.tasks.abstract_task.AbstractTask

Bases: ABC

Abstract Class of a Task. All FL tasks should inherit this class and override the abstract methods defined.

Agent hooks

If the task config contains an agent block with "enabled": true, an LLM-powered agent is consulted at key lifecycle points (post-training, pre/post-aggregation, on-failure). When the agent is disabled or unavailable the lifecycle proceeds identically to the non-agent path.

Source code in controller/starfish/controller/tasks/abstract_task.py
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
552
553
554
555
556
557
558
559
560
561
562
563
564
565
566
567
568
569
570
571
572
573
574
575
576
577
578
579
580
581
582
583
584
585
586
587
588
589
590
591
592
593
594
595
596
597
598
599
600
601
602
603
604
605
606
607
608
609
610
611
612
613
614
615
616
617
618
619
620
621
622
623
624
625
626
627
628
629
630
631
class AbstractTask(ABC):
    """
    Abstract Class of a Task.
    All FL tasks should inherit this class and override the abstract methods defined.

    Agent hooks
    -----------
    If the task config contains an ``agent`` block with ``"enabled": true``,
    an LLM-powered agent is consulted at key lifecycle points (post-training,
    pre/post-aggregation, on-failure).  When the agent is disabled or
    unavailable the lifecycle proceeds identically to the non-agent path.
    """
    project_id = None
    run_id = None
    batch_id = None
    cur_seq = None
    tasks = None
    role = None
    artifact = None
    status = None
    logger = None
    _agent_hooks = None

    def __init__(self, run):
        self.project_id = run['project']
        self.run_id = run['id']
        self.batch_id = run['batch']
        self.role = run['role']
        self.status = run['status']
        self.post_init(run)

    def method_call(self, name: str, *args, **kwargs):
        if hasattr(self, name) and callable(getattr(self, name)):
            func = getattr(self, name)
            func(*args, **kwargs)
        else:
            self.logger.warning("method {} not exists".format(name))

    def standby(self, *args, **kwargs):
        """
        Status: Standby,2
        Next Status: Preparing,3;Pending Failed,1
        Called when a task starts. Current participant will start to prepare the run.
        In this event, notify router.
        """
        try:
            run = args[0]
            self.post_init(run)
            s = inspect.currentframe().f_code.co_name
            if self.status == s:
                self.logger.warning(
                    "Already in status {}. Ignore message".format(s))
                return
            else:
                self.status = s
            if not self.is_first_round():
                valid = self.validate()
                if valid:
                    self.notify(3)
                else:
                    self.notify(1)
        except Exception as e:
            self.logger.warning("Exception in standby status: {}".format(e))
            self.notify(1)

    def preparing(self, *args, **kwargs):
        """
        Status: Preparing,3
        Next Status: Running,4; Pending Failed,1
        Called when a task starts. Current participant will start to prepare the run.
        In this event, input data files are validated.
        """
        try:
            s = inspect.currentframe().f_code.co_name
            if self.role == 'coordinator':
                self.status = s
                if self.runs_in_fails() or not self.prepare_data():
                    self.notify(1, param={'update_all': True})
                    return
                if self.runs_in_same_state('preparing'):
                    self.notify(4, param={'update_all': True})
            else:
                if not self.prepare_data():
                    self.notify(1, param={'update_all': False})
                    return
                if self.status == s:
                    self.logger.warning(
                        "Already in status {}. Ignore message".format(s))
                    return
                else:
                    self.status = s
        except Exception as e:
            self.logger.warning("Exception in preparing status: {}".format(e))
            self.notify(1, param={'update_all': True})

    def running(self, *args, **kwargs):
        """
        Status: Running,4
        Next Status: Pending Success,5;Pending Failed,1
        Called when all participants have prepared to run.
        In this event, input data files are used for training.
        """
        try:
            s = inspect.currentframe().f_code.co_name
            if self.status == s:
                self.logger.warning(
                    "Already in status {}. Ignore message".format(s))
                return
            else:
                self.status = s

            valid = self.training()
            if valid:
                # Agent hook: post-training summary
                if self._agent_hooks and self._agent_hooks.enabled:
                    try:
                        mid_url = gen_mid_artifacts_url(
                            self.run_id, self.cur_seq, self.get_round())
                        mid_data = {}
                        if mid_url and os.path.exists(mid_url):
                            with open(mid_url, 'r') as f:
                                mid_data = json.loads(f.readline())
                        self._agent_hooks.post_training(
                            self._get_task_type(), self.get_round(),
                            self.tasks[self.cur_seq - 1]['config'].get('total_round', 1),
                            mid_data, self.logger,
                        )
                    except Exception as hook_err:
                        self.logger.debug("Agent post_training hook error: %s", hook_err)
                self.notify(5)
            else:
                self.notify(1)
        except Exception as e:
            self.logger.warning("Exception in running status: {}".format(e))
            self.logger.debug(traceback.format_exc())
            self.notify(1)

    def pending_success(self, *args, **kwargs):
        """
        Status: Pending Success,5
        Next Status: Pending Aggregating,6; Standby,2
        Called when current participant successfully completes the task.
        In this event, output file and file will be uploaded to RS for forwarding to Coordinator.
        """
        try:
            if self.upload(False):
                self.notify(6)
        except Exception as e:
            self.logger.warning(
                "Exception in pending_success status: {}".format(e))

    def pending_aggregating(self, *args, **kwargs):
        """
        Status: Pending Aggregating, 6
        Next Status: Aggregating, 7; Failed, 0
        Participant: Do nothing
        Coordinator: Waiting for all participants been changed to this status and download the artifacts
        :return:
        """
        try:
            s = inspect.currentframe().f_code.co_name
            if self.role == 'coordinator':
                self.status = s
                if self.runs_in_fails():
                    self.notify(0, param={'update_all': True})
                if self.runs_in_same_state('pending_aggregating') and self.download_mid_artifacts():
                    self.notify(7, param={'update_all': True})
            else:
                if self.status == s:
                    self.logger.warning(
                        "Already in status {}. Ignore message".format(s))
                    return
                else:
                    self.status = s
        except Exception as e:
            self.logger.warning(
                "Exception in pending aggregating status: {}".format(e))
            self.notify(0, param={'update_all': True})

    def aggregating(self, *args, **kwargs):
        """
        Status: Aggregating, 7
        Next Status: Standby,2; Success,8; Failed,0
        Participant: Do nothing
        Coordinator: Aggregate artifacts from all participants and upload the final artifact.
        :return:
        """
        try:
            s = inspect.currentframe().f_code.co_name
            if self.role == 'coordinator':
                self.status = s
                if self.runs_in_fails():
                    self.notify(0, param={'update_all': True})

                # Agent hook: pre-aggregation outlier detection
                if self._agent_hooks and self._agent_hooks.enabled:
                    try:
                        self._agent_hooks.pre_aggregation(
                            self._get_task_type(), self.get_round(),
                            self.tasks[self.cur_seq - 1]['config'].get('total_round', 1),
                            [],  # mid-artifacts already on disk; pass empty for now
                            self.logger,
                        )
                    except Exception as hook_err:
                        self.logger.debug("Agent pre_aggregation hook error: %s", hook_err)

                if self.do_aggregate():
                    # Agent hook: post-aggregation convergence check
                    early_stop = False
                    if self._agent_hooks and self._agent_hooks.enabled:
                        try:
                            decision = self._agent_hooks.post_aggregation(
                                self._get_task_type(), self.get_round(),
                                self.tasks[self.cur_seq - 1]['config'].get('total_round', 1),
                                {},  # aggregated result
                                None,  # round history
                                self.logger,
                            )
                            if (decision and decision.get("converged")
                                    and not self.is_last_round()):
                                self.logger.info(
                                    "[Agent] Early stopping at round %d: %s",
                                    self.get_round(),
                                    decision.get("reason", "model converged"))
                                early_stop = True
                        except Exception as hook_err:
                            self.logger.debug("Agent post_aggregation hook error: %s", hook_err)

                    if early_stop:
                        self.notify(8, param={'update_all': True})
                    else:
                        is_last_round = self.is_last_round()
                        self.logger.debug(
                            "Is the last round? {}".format(is_last_round))
                        if is_last_round:
                            self.notify(8, param={'update_all': True})
                        else:
                            self.notify(
                                2, param={'increase_round': True, 'update_all': True})
                else:
                    self.notify(0, param={'update_all': True})
            else:
                if self.status == s:
                    self.logger.warning(
                        "Already in status {}. Ignore message".format(s))
                    return
                else:
                    self.status = s
        except Exception as e:
            self.logger.warning(
                "Exception in aggregating status: {}".format(e))
            self.logger.debug(traceback.print_exc())
            self.notify(0, param={'update_all': True})

    def pending_failed(self, *args, **kwargs):
        """
        Status: Pending Failed,1
        Next Status: Failed,0
        Called when current participant fails to complete the task.
        In this event, file will be uploaded to RS for forwarding to Coordinator.
        """
        try:
            # Agent hook: failure triage
            if self._agent_hooks and self._agent_hooks.enabled:
                try:
                    task_config = self.tasks[self.cur_seq - 1].get("config", {}) if self.tasks else {}
                    self._agent_hooks.on_failure(
                        self._get_task_type(), task_config,
                        self.get_round() or 0, self.role,
                        "Task entered pending_failed state", [],
                        self.logger,
                    )
                except Exception as hook_err:
                    self.logger.debug("Agent on_failure hook error: %s", hook_err)

            if self.upload(False):
                self.notify(0)
        except Exception as e:
            self.logger.warning(
                "Exception in pending failed status: {}".format(e))

    @abstractmethod
    def validate(self) -> bool:
        return True

    @abstractmethod
    def prepare_data(self) -> bool:
        return True

    @abstractmethod
    def training(self) -> bool:
        return True

    def download_mid_artifacts(self) -> bool:
        """
        :return:
        """
        task_round = self.get_round()
        if task_round:
            self.logger.debug("Downloading mid_artifacts")
            response = requests.get(
                '{0}/runs-action/download/?run={1}&task_seq={2}&round_seq={3}&all_runs={4}&type=mid_artifacts'.format(
                    router_url,
                    self.run_id,
                    self.cur_seq,
                    task_round,
                    1),
                auth=(router_username, router_password))

            if response.status_code == 404:
                self.logger.warning('No mid-artifacts found in router for project {} at batch {}'.format(
                    self.project_id, self.batch_id))
                return False

            if response.status_code == 200:
                content = response.content
                self.logger.debug(
                    'Saving all mid-artifacts to local for project {} at batch {}'.format(self.project_id,
                                                                                          self.batch_id))
                saved_url = download_all_mid_artifacts(
                    self.project_id, self.batch_id, content)
                if saved_url:
                    self.logger.debug(
                        'Successfully download and save all mid-artifacts to local for project {} at batch {} in {} dir'.format(
                            self.project_id, self.batch_id, saved_url))
                    return True
                self.logger.warning('Failed to save mid-artifacts to local for project {} at batch {}'.format(
                    self.project_id, self.batch_id))
        return False

    def download_artifact(self) -> bool:
        """ Download artifact of the last run
        :return:
        """
        seq_no, round_no = self.get_previous_seq_and_round()
        if seq_no and round_no:
            self.logger.debug("Downloading artifact")
            response = requests.get(
                '{0}/runs-action/download/?run={1}&task_seq={2}&round_seq={3}&all_runs={4}&type=artifacts'.format(
                    router_url,
                    self.run_id,
                    seq_no,
                    round_no,
                    0),
                auth=(router_username, router_password))

            if response.status_code == 404:
                self.logger.warning('No artifacts found in router for run {} at batch {}'.format(
                    self.project_id, self.batch_id))
                return False

            if response.status_code == 200:
                content = response.content
                self.logger.debug(
                    'Saving artifacts to local for run {} at seq {} and round {}'.format(self.run_id,
                                                                                         seq_no, round_no))
                saved_url = download_artifacts(
                    self.run_id, seq_no, round_no, content)
                if saved_url:
                    self.logger.debug(
                        'Successfully download and save artifacts to local for run {} at seq {} and round {} in {} dir'.format(
                            self.run_id, seq_no, round_no, saved_url))
                    return True
                self.logger.warning('Failed to save mid-artifacts to local for run {} at seq {} and round {}'.format(
                    self.run_id, seq_no, round_no))
        return False

    @abstractmethod
    def do_aggregate(self) -> bool:
        """
        :return:
        """
        return True

    def upload(self, is_artifact: bool) -> bool:
        # Here assume when round of task success, it should upload both mid-artifacts and logs to router
        task_round = self.get_round()
        if task_round:
            self.logger.debug('will upload logs and mid-artifacts')
            files_data = dict()
            data = dict()
            if is_artifact:
                artifact_url = gen_artifacts_url(
                    self.run_id, self.cur_seq, task_round)
                if artifact_url:
                    files_data['artifacts'] = read_file_from_url(artifact_url)
            else:
                mid_artifacts_url = gen_mid_artifacts_url(
                    self.run_id, self.cur_seq, task_round)
                logs_url = gen_logs_url(self.run_id, self.cur_seq, task_round)

                if mid_artifacts_url:
                    files_data['mid_artifacts'] = read_file_from_url(
                        mid_artifacts_url)
                if logs_url:
                    files_data['logs'] = read_file_from_url(logs_url)

            data['run'] = self.run_id
            data['task_seq'] = self.cur_seq
            data['round_seq'] = task_round
            if any(files_data.values()):

                response = requests.post('{0}/runs-action/upload/'.format(router_url, self.run_id),
                                         auth=(router_username,
                                               router_password),
                                         data=data,
                                         files=files_data)
                if response.status_code == 200:
                    self.logger.debug(
                        'Successfully upload logs and artifacts of run {} - task {} - round {}'.format(self.run_id,
                                                                                                       self.cur_seq,
                                                                                                       task_round))
                    return True
            else:
                self.logger.debug("Files data is empty. Ignore upload")
                return True

        return False

    def notify(self, next_state, param: dict = None):
        if param is None:
            param = dict()
        headers = {'Content-type': 'application/json'}
        param['status'] = next_state
        requests.put('{0}/runs/{1}/status/'.format(router_url, self.run_id),
                     headers=headers,
                     auth=(router_username, router_password),
                     data=json.dumps(param))

    def fetch_runs(self):
        runs_response = requests.get(
            '{0}/runs/detail/?batch={1}&project={2}&site_uid={3}'.format(
                router_url, self.batch_id, self.project_id, site_uid),
            auth=(router_username, router_password))
        if runs_response.ok:
            dic = runs_response.json()
            runs = dic['runs']
            return runs
        return None

    def is_last_round(self) -> bool:
        """
        Determine whether it is the last round.
        Logic:
        1. check tasks size vs the cur_seq
        2. check total_round and current_round inside tasks
        :return:
        """
        self.logger.debug(
            "Checking whether it is the last round. cur_seq: {}, total tasks: {}".format(self.cur_seq, len(self.tasks)))
        if self.cur_seq >= len(self.tasks):
            c = self.tasks[self.cur_seq - 1]['config']
            if 'total_round' in c and 'current_round' in c:
                total_round = c['total_round']
                current_round = c['current_round']
                self.logger.debug("Total Round: {}, Current Round: {}".format(
                    total_round, current_round))
                return current_round >= total_round
            else:
                return True
        else:
            return False

    def is_first_round(self) -> bool:
        self.logger.debug(
            "Checking whether it is the first round. cur_seq: {}, total tasks: {}".format(self.cur_seq,
                                                                                          len(self.tasks)))
        if self.cur_seq == 1:
            c = self.tasks[self.cur_seq - 1]['config']
            if 'current_round' in c:
                current_round = c['current_round']
                return current_round == 1
            else:
                return True
        else:
            return False

    def get_previous_seq_and_round(self):
        """
        Return the seq and round number of the previous round
        :return: [seq_no, round_no]
        """
        c = self.tasks[self.cur_seq - 1]['config']
        total_round = c['total_round']
        current_round = c['current_round']

        if self.cur_seq == 1:
            if current_round > 1:
                return self.cur_seq, current_round - 1
            else:
                return None, None
        else:
            if current_round == 1:
                return self.cur_seq - 1, total_round
            else:
                return self.cur_seq, current_round - 1

    def runs_in_same_state(self, expected_state) -> bool:
        """
        Coordinator: Check all participants are in expected status
        :param expected_state:
        :return:
        """
        self.logger.debug("Checking whether all runs are in the same status")
        runs = self.fetch_runs()
        self.logger.debug("expected state: {}. All runs: {}".format(
            expected_state, runs))
        if runs:
            for r in runs:
                if format_status(r['status']) != expected_state:
                    return False
            return True
        else:
            return False

    def runs_in_fails(self) -> bool:
        runs = self.fetch_runs()
        if runs:
            for r in runs:
                if format_status(r['status']) in ['pending_failed', 'failed']:
                    return True
        return False

    # This method can be used to get current round of current task

    def get_round(self):
        if self.tasks and len(self.tasks) > 0:
            cur_task = self.tasks[self.cur_seq - 1]
            if cur_task and len(cur_task) > 0 and 'config' in cur_task:
                return cur_task['config']['current_round']
        return None

    def save_artifacts(self, url, content):
        if content:
            create_if_not_exist(url)
            try:
                with open(url, 'w') as f:
                    f.write(content)
                    f.close()
                return True
            except Exception as e:
                self.logger.error(
                    'Error while saving mid_artifacts. due to {}'.format(e))
                return False

    def _init_agent_hooks(self):
        """Initialise agent hooks from the current task config (no-op if absent)."""
        from starfish.controller.agent.hooks import TaskAgentHooks
        try:
            task_config = self.tasks[self.cur_seq - 1].get("config", {}) if self.tasks else {}
            self._agent_hooks = TaskAgentHooks(task_config)
        except Exception:
            self._agent_hooks = TaskAgentHooks()

    def _get_task_type(self):
        """Return the model name for the current task sequence."""
        if self.tasks and self.cur_seq <= len(self.tasks):
            return self.tasks[self.cur_seq - 1].get("model", "Unknown")
        return "Unknown"

    def post_init(self, run):
        self.cur_seq = run['cur_seq']
        self.tasks = run['tasks']
        self._init_agent_hooks()
        cur_round = self.get_round()

        logger_name = 'logger-{}-{}-{}'.format(
            self.run_id, self.cur_seq, cur_round)

        if self.logger and self.logger.name == logger_name:
            self.logger.debug(
                'logger with name {} exists, will not init again'.format(logger_name))
            return

        url = gen_logs_url(self.run_id, self.cur_seq, cur_round)
        create_if_not_exist(url)

        logger = logging.getLogger(logger_name)
        logger.setLevel(logging.DEBUG)  # Set the logging level as needed

        # Create a log formatter
        log_formatter = logging.Formatter(
            '%(asctime)s [%(levelname)s] %(message)s')

        # Create a file handler to save logs to a file
        file_handler = logging.FileHandler(url)
        # Set the desired log level for the file handler
        file_handler.setLevel(logging.DEBUG)
        file_handler.setFormatter(log_formatter)

        # Create a console handler to print logs to the console
        console_handler = logging.StreamHandler()
        # Set the desired log level for the console handler
        console_handler.setLevel(logging.DEBUG)
        console_handler.setFormatter(log_formatter)

        # Add the handlers to the logger
        logger.addHandler(file_handler)
        logger.addHandler(console_handler)

        logger.debug(
            "Init logger for run {} - seq {} - round {} ".format(self.run_id, self.cur_seq, cur_round))
        self.logger = logger

    def read_dataset(self, run_id):
        return file_utils.load_dataset_by_run(run_id)

standby(*args, **kwargs)

Status: Standby,2 Next Status: Preparing,3;Pending Failed,1 Called when a task starts. Current participant will start to prepare the run. In this event, notify router.

Source code in controller/starfish/controller/tasks/abstract_task.py
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
def standby(self, *args, **kwargs):
    """
    Status: Standby,2
    Next Status: Preparing,3;Pending Failed,1
    Called when a task starts. Current participant will start to prepare the run.
    In this event, notify router.
    """
    try:
        run = args[0]
        self.post_init(run)
        s = inspect.currentframe().f_code.co_name
        if self.status == s:
            self.logger.warning(
                "Already in status {}. Ignore message".format(s))
            return
        else:
            self.status = s
        if not self.is_first_round():
            valid = self.validate()
            if valid:
                self.notify(3)
            else:
                self.notify(1)
    except Exception as e:
        self.logger.warning("Exception in standby status: {}".format(e))
        self.notify(1)

preparing(*args, **kwargs)

Status: Preparing,3 Next Status: Running,4; Pending Failed,1 Called when a task starts. Current participant will start to prepare the run. In this event, input data files are validated.

Source code in controller/starfish/controller/tasks/abstract_task.py
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
def preparing(self, *args, **kwargs):
    """
    Status: Preparing,3
    Next Status: Running,4; Pending Failed,1
    Called when a task starts. Current participant will start to prepare the run.
    In this event, input data files are validated.
    """
    try:
        s = inspect.currentframe().f_code.co_name
        if self.role == 'coordinator':
            self.status = s
            if self.runs_in_fails() or not self.prepare_data():
                self.notify(1, param={'update_all': True})
                return
            if self.runs_in_same_state('preparing'):
                self.notify(4, param={'update_all': True})
        else:
            if not self.prepare_data():
                self.notify(1, param={'update_all': False})
                return
            if self.status == s:
                self.logger.warning(
                    "Already in status {}. Ignore message".format(s))
                return
            else:
                self.status = s
    except Exception as e:
        self.logger.warning("Exception in preparing status: {}".format(e))
        self.notify(1, param={'update_all': True})

running(*args, **kwargs)

Status: Running,4 Next Status: Pending Success,5;Pending Failed,1 Called when all participants have prepared to run. In this event, input data files are used for training.

Source code in controller/starfish/controller/tasks/abstract_task.py
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
def running(self, *args, **kwargs):
    """
    Status: Running,4
    Next Status: Pending Success,5;Pending Failed,1
    Called when all participants have prepared to run.
    In this event, input data files are used for training.
    """
    try:
        s = inspect.currentframe().f_code.co_name
        if self.status == s:
            self.logger.warning(
                "Already in status {}. Ignore message".format(s))
            return
        else:
            self.status = s

        valid = self.training()
        if valid:
            # Agent hook: post-training summary
            if self._agent_hooks and self._agent_hooks.enabled:
                try:
                    mid_url = gen_mid_artifacts_url(
                        self.run_id, self.cur_seq, self.get_round())
                    mid_data = {}
                    if mid_url and os.path.exists(mid_url):
                        with open(mid_url, 'r') as f:
                            mid_data = json.loads(f.readline())
                    self._agent_hooks.post_training(
                        self._get_task_type(), self.get_round(),
                        self.tasks[self.cur_seq - 1]['config'].get('total_round', 1),
                        mid_data, self.logger,
                    )
                except Exception as hook_err:
                    self.logger.debug("Agent post_training hook error: %s", hook_err)
            self.notify(5)
        else:
            self.notify(1)
    except Exception as e:
        self.logger.warning("Exception in running status: {}".format(e))
        self.logger.debug(traceback.format_exc())
        self.notify(1)

pending_success(*args, **kwargs)

Status: Pending Success,5 Next Status: Pending Aggregating,6; Standby,2 Called when current participant successfully completes the task. In this event, output file and file will be uploaded to RS for forwarding to Coordinator.

Source code in controller/starfish/controller/tasks/abstract_task.py
163
164
165
166
167
168
169
170
171
172
173
174
175
def pending_success(self, *args, **kwargs):
    """
    Status: Pending Success,5
    Next Status: Pending Aggregating,6; Standby,2
    Called when current participant successfully completes the task.
    In this event, output file and file will be uploaded to RS for forwarding to Coordinator.
    """
    try:
        if self.upload(False):
            self.notify(6)
    except Exception as e:
        self.logger.warning(
            "Exception in pending_success status: {}".format(e))

pending_aggregating(*args, **kwargs)

Status: Pending Aggregating, 6 Next Status: Aggregating, 7; Failed, 0 Participant: Do nothing Coordinator: Waiting for all participants been changed to this status and download the artifacts :return:

Source code in controller/starfish/controller/tasks/abstract_task.py
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
def pending_aggregating(self, *args, **kwargs):
    """
    Status: Pending Aggregating, 6
    Next Status: Aggregating, 7; Failed, 0
    Participant: Do nothing
    Coordinator: Waiting for all participants been changed to this status and download the artifacts
    :return:
    """
    try:
        s = inspect.currentframe().f_code.co_name
        if self.role == 'coordinator':
            self.status = s
            if self.runs_in_fails():
                self.notify(0, param={'update_all': True})
            if self.runs_in_same_state('pending_aggregating') and self.download_mid_artifacts():
                self.notify(7, param={'update_all': True})
        else:
            if self.status == s:
                self.logger.warning(
                    "Already in status {}. Ignore message".format(s))
                return
            else:
                self.status = s
    except Exception as e:
        self.logger.warning(
            "Exception in pending aggregating status: {}".format(e))
        self.notify(0, param={'update_all': True})

aggregating(*args, **kwargs)

Status: Aggregating, 7 Next Status: Standby,2; Success,8; Failed,0 Participant: Do nothing Coordinator: Aggregate artifacts from all participants and upload the final artifact. :return:

Source code in controller/starfish/controller/tasks/abstract_task.py
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
def aggregating(self, *args, **kwargs):
    """
    Status: Aggregating, 7
    Next Status: Standby,2; Success,8; Failed,0
    Participant: Do nothing
    Coordinator: Aggregate artifacts from all participants and upload the final artifact.
    :return:
    """
    try:
        s = inspect.currentframe().f_code.co_name
        if self.role == 'coordinator':
            self.status = s
            if self.runs_in_fails():
                self.notify(0, param={'update_all': True})

            # Agent hook: pre-aggregation outlier detection
            if self._agent_hooks and self._agent_hooks.enabled:
                try:
                    self._agent_hooks.pre_aggregation(
                        self._get_task_type(), self.get_round(),
                        self.tasks[self.cur_seq - 1]['config'].get('total_round', 1),
                        [],  # mid-artifacts already on disk; pass empty for now
                        self.logger,
                    )
                except Exception as hook_err:
                    self.logger.debug("Agent pre_aggregation hook error: %s", hook_err)

            if self.do_aggregate():
                # Agent hook: post-aggregation convergence check
                early_stop = False
                if self._agent_hooks and self._agent_hooks.enabled:
                    try:
                        decision = self._agent_hooks.post_aggregation(
                            self._get_task_type(), self.get_round(),
                            self.tasks[self.cur_seq - 1]['config'].get('total_round', 1),
                            {},  # aggregated result
                            None,  # round history
                            self.logger,
                        )
                        if (decision and decision.get("converged")
                                and not self.is_last_round()):
                            self.logger.info(
                                "[Agent] Early stopping at round %d: %s",
                                self.get_round(),
                                decision.get("reason", "model converged"))
                            early_stop = True
                    except Exception as hook_err:
                        self.logger.debug("Agent post_aggregation hook error: %s", hook_err)

                if early_stop:
                    self.notify(8, param={'update_all': True})
                else:
                    is_last_round = self.is_last_round()
                    self.logger.debug(
                        "Is the last round? {}".format(is_last_round))
                    if is_last_round:
                        self.notify(8, param={'update_all': True})
                    else:
                        self.notify(
                            2, param={'increase_round': True, 'update_all': True})
            else:
                self.notify(0, param={'update_all': True})
        else:
            if self.status == s:
                self.logger.warning(
                    "Already in status {}. Ignore message".format(s))
                return
            else:
                self.status = s
    except Exception as e:
        self.logger.warning(
            "Exception in aggregating status: {}".format(e))
        self.logger.debug(traceback.print_exc())
        self.notify(0, param={'update_all': True})

pending_failed(*args, **kwargs)

Status: Pending Failed,1 Next Status: Failed,0 Called when current participant fails to complete the task. In this event, file will be uploaded to RS for forwarding to Coordinator.

Source code in controller/starfish/controller/tasks/abstract_task.py
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
def pending_failed(self, *args, **kwargs):
    """
    Status: Pending Failed,1
    Next Status: Failed,0
    Called when current participant fails to complete the task.
    In this event, file will be uploaded to RS for forwarding to Coordinator.
    """
    try:
        # Agent hook: failure triage
        if self._agent_hooks and self._agent_hooks.enabled:
            try:
                task_config = self.tasks[self.cur_seq - 1].get("config", {}) if self.tasks else {}
                self._agent_hooks.on_failure(
                    self._get_task_type(), task_config,
                    self.get_round() or 0, self.role,
                    "Task entered pending_failed state", [],
                    self.logger,
                )
            except Exception as hook_err:
                self.logger.debug("Agent on_failure hook error: %s", hook_err)

        if self.upload(False):
            self.notify(0)
    except Exception as e:
        self.logger.warning(
            "Exception in pending failed status: {}".format(e))

download_mid_artifacts()

:return:

Source code in controller/starfish/controller/tasks/abstract_task.py
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
def download_mid_artifacts(self) -> bool:
    """
    :return:
    """
    task_round = self.get_round()
    if task_round:
        self.logger.debug("Downloading mid_artifacts")
        response = requests.get(
            '{0}/runs-action/download/?run={1}&task_seq={2}&round_seq={3}&all_runs={4}&type=mid_artifacts'.format(
                router_url,
                self.run_id,
                self.cur_seq,
                task_round,
                1),
            auth=(router_username, router_password))

        if response.status_code == 404:
            self.logger.warning('No mid-artifacts found in router for project {} at batch {}'.format(
                self.project_id, self.batch_id))
            return False

        if response.status_code == 200:
            content = response.content
            self.logger.debug(
                'Saving all mid-artifacts to local for project {} at batch {}'.format(self.project_id,
                                                                                      self.batch_id))
            saved_url = download_all_mid_artifacts(
                self.project_id, self.batch_id, content)
            if saved_url:
                self.logger.debug(
                    'Successfully download and save all mid-artifacts to local for project {} at batch {} in {} dir'.format(
                        self.project_id, self.batch_id, saved_url))
                return True
            self.logger.warning('Failed to save mid-artifacts to local for project {} at batch {}'.format(
                self.project_id, self.batch_id))
    return False

download_artifact()

Download artifact of the last run :return:

Source code in controller/starfish/controller/tasks/abstract_task.py
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
def download_artifact(self) -> bool:
    """ Download artifact of the last run
    :return:
    """
    seq_no, round_no = self.get_previous_seq_and_round()
    if seq_no and round_no:
        self.logger.debug("Downloading artifact")
        response = requests.get(
            '{0}/runs-action/download/?run={1}&task_seq={2}&round_seq={3}&all_runs={4}&type=artifacts'.format(
                router_url,
                self.run_id,
                seq_no,
                round_no,
                0),
            auth=(router_username, router_password))

        if response.status_code == 404:
            self.logger.warning('No artifacts found in router for run {} at batch {}'.format(
                self.project_id, self.batch_id))
            return False

        if response.status_code == 200:
            content = response.content
            self.logger.debug(
                'Saving artifacts to local for run {} at seq {} and round {}'.format(self.run_id,
                                                                                     seq_no, round_no))
            saved_url = download_artifacts(
                self.run_id, seq_no, round_no, content)
            if saved_url:
                self.logger.debug(
                    'Successfully download and save artifacts to local for run {} at seq {} and round {} in {} dir'.format(
                        self.run_id, seq_no, round_no, saved_url))
                return True
            self.logger.warning('Failed to save mid-artifacts to local for run {} at seq {} and round {}'.format(
                self.run_id, seq_no, round_no))
    return False

do_aggregate() abstractmethod

:return:

Source code in controller/starfish/controller/tasks/abstract_task.py
393
394
395
396
397
398
@abstractmethod
def do_aggregate(self) -> bool:
    """
    :return:
    """
    return True

is_last_round()

Determine whether it is the last round. Logic: 1. check tasks size vs the cur_seq 2. check total_round and current_round inside tasks :return:

Source code in controller/starfish/controller/tasks/abstract_task.py
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
def is_last_round(self) -> bool:
    """
    Determine whether it is the last round.
    Logic:
    1. check tasks size vs the cur_seq
    2. check total_round and current_round inside tasks
    :return:
    """
    self.logger.debug(
        "Checking whether it is the last round. cur_seq: {}, total tasks: {}".format(self.cur_seq, len(self.tasks)))
    if self.cur_seq >= len(self.tasks):
        c = self.tasks[self.cur_seq - 1]['config']
        if 'total_round' in c and 'current_round' in c:
            total_round = c['total_round']
            current_round = c['current_round']
            self.logger.debug("Total Round: {}, Current Round: {}".format(
                total_round, current_round))
            return current_round >= total_round
        else:
            return True
    else:
        return False

get_previous_seq_and_round()

Return the seq and round number of the previous round :return: [seq_no, round_no]

Source code in controller/starfish/controller/tasks/abstract_task.py
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
def get_previous_seq_and_round(self):
    """
    Return the seq and round number of the previous round
    :return: [seq_no, round_no]
    """
    c = self.tasks[self.cur_seq - 1]['config']
    total_round = c['total_round']
    current_round = c['current_round']

    if self.cur_seq == 1:
        if current_round > 1:
            return self.cur_seq, current_round - 1
        else:
            return None, None
    else:
        if current_round == 1:
            return self.cur_seq - 1, total_round
        else:
            return self.cur_seq, current_round - 1

runs_in_same_state(expected_state)

Coordinator: Check all participants are in expected status :param expected_state: :return:

Source code in controller/starfish/controller/tasks/abstract_task.py
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
def runs_in_same_state(self, expected_state) -> bool:
    """
    Coordinator: Check all participants are in expected status
    :param expected_state:
    :return:
    """
    self.logger.debug("Checking whether all runs are in the same status")
    runs = self.fetch_runs()
    self.logger.debug("expected state: {}. All runs: {}".format(
        expected_state, runs))
    if runs:
        for r in runs:
            if format_status(r['status']) != expected_state:
                return False
        return True
    else:
        return False