← points de vue
24 janvier 2023

SageMaker et EMR depuis Airflow : un orchestrateur qui ne calcule rien

Un DAG géré n’a ni la mémoire ni les droits pour porter un entraînement. Il déclenche, il attend, il réveille quelqu’un, et cette dernière partie est celle qu’on écrit toujours en dernier, souvent après l’incident.

La tentation, sur un Airflow géré, est de faire tourner le calcul dans le DAG : un `PythonOperator` qui charge un jeu de données et entraîne un modèle. Cela marche sur un échantillon et ne marche jamais ensuite. Les travailleurs d’un service géré sont dimensionnés pour planifier, pas pour calculer ; ils partagent leur mémoire avec le planificateur, et un entraînement qui gonfle emporte l’ordonnancement de toutes les autres chaînes avec lui.

La règle que j’applique depuis tient en une phrase : l’orchestrateur ne calcule rien. Il décrit un travail, le confie à un service qui sait le porter, et surveille. Tout le reste (la taille des machines, les droits, la durée de vie du cluster) devient de la configuration plutôt que du code.

01Décrire, déclencher, attendre

Deux formes reviennent. L’entraînement part chez le service de modèles, avec sa propre image et son propre type d’instance. La préparation lourde part sur un cluster éphémère qu’on crée, qu’on charge d’étapes, et qu’on éteint : le dernier opérateur étant celui qu’on oublie, et qui coûte le plus cher à oublier.

python
from airflow.providers.amazon.aws.operators.emr import (
    EmrAddStepsOperator, EmrCreateJobFlowOperator, EmrTerminateJobFlowOperator,
)
from airflow.providers.amazon.aws.operators.sagemaker import SageMakerTrainingOperator
from airflow.providers.amazon.aws.sensors.emr import EmrStepSensor

cluster = EmrCreateJobFlowOperator(task_id="cluster", job_flow_overrides=EMR_CONFIG)

prepare = EmrAddStepsOperator(
    task_id="prepare",
    job_flow_id=cluster.output,
    steps=SPARK_STEPS,
)

# The sensor is what makes the dependency true: without it the next task starts
# as soon as the step is *submitted*, not when it has finished.
wait = EmrStepSensor(
    task_id="wait",
    job_flow_id=cluster.output,
    step_id=prepare.output[0],
    mode="reschedule",  # gives the slot back instead of holding it
    poke_interval=60,
)

training = SageMakerTrainingOperator(
    task_id="training",
    config=TRAINING_CONFIG,
    wait_for_completion=True,
)

shutdown = EmrTerminateJobFlowOperator(
    task_id="shutdown",
    job_flow_id=cluster.output,
    trigger_rule="all_done",  # shut down even if the preparation failed
)

cluster >> prepare >> wait >> training
wait >> shutdown

Deux détails valent le paragraphe qu’ils prennent. Le capteur en mode `reschedule` rend son créneau entre deux vérifications au lieu de le garder occupé pendant deux heures ; sur un service géré où les créneaux sont comptés et facturés, la différence est visible sur la facture comme sur le débit. Et la règle de déclenchement `all_done` sur l’extinction est la seule qui garantisse qu’un cluster à plusieurs euros de l’heure ne survive pas à l’échec de l’étape qu’il servait.

02Ce qu’un service géré retire

Sur un Airflow géré, on ne choisit ni la version des dépendances ni le moment où l’environnement redémarre. Les paquets se déclarent dans un fichier déposé sur un stockage objet, et toute modification provoque un redémarrage de l’environnement : plusieurs dizaines de minutes pendant lesquelles rien ne s’ordonnance. La conséquence pratique est qu’on n’ajoute pas une dépendance à la légère, et qu’on préfère un opérateur du fournisseur à une bibliothèque qu’il faudrait installer.

03Réveiller quelqu’un, et pas pour rien

Un courriel d’échec n’est pas une alerte : personne ne le lit à trois heures du matin. La chaîne qui alimente un modèle en production mérite une astreinte, et les autres n’en méritent pas. C’est le tri qui compte, pas l’outil : une équipe réveillée pour une chaîne sans conséquence cesse de répondre aux vraies.

python
from airflow.providers.pagerduty.hooks.pagerduty_events import PagerdutyEventsHook

def page_on_call(context) -> None:
    instance = context["task_instance"]
    PagerdutyEventsHook(pagerduty_events_conn_id="pagerduty").create_event(
        summary=f"{instance.dag_id}.{instance.task_id} failed",
        severity="critical",
        source="airflow",
        # The dedup key avoids opening one incident per attempt: the three tries
        # of a single run wake exactly one person.
        dedup_key=f"{instance.dag_id}:{instance.task_id}:{context['run_id']}",
        custom_details={"log": instance.log_url},
    )

# Set on the task that matters, not on the whole DAG: the sorting is what keeps
# an on-call rota credible.
training.on_failure_callback = page_on_call
04Ce que je laisse de côté

Je ne compare pas les orchestrateurs. Le sujet est intéressant et il ne se tranche pas depuis une seule mission : ce qui a compté ici, c’est la contrainte du service géré, et elle est la même quelle que soit l’opinion qu’on a sur l’outil. Je ne dis rien non plus de la reprise d’un entraînement interrompu : nous relancions depuis le début, ce qui était supportable à cette échelle et ne l’aurait pas été au-delà.