Codeseys/composer-replication-framework
0
1"""SageMakerExecutor — production boto3-backed serverless executor.2 3This is a fully-working cloud adapter (the sibling of `ModalSpawnExecutor`,4not the loud-failing `modal.py` / `hf_jobs.py` skeletons). It implements the5`ServerlessExecutor` Protocol against Amazon SageMaker Training Jobs via the6boto3 low-level `sagemaker` client.7 8Design choices9--------------10 111. **N independent single-instance jobs, NOT one multi-instance job.**12 SageMaker's *native* distributed training (``ResourceConfig.InstanceCount > 1``)13 groups instances into ONE job with an in-cluster NCCL/MPI fabric wired via14 ``/opt/ml/input/config/resourceconfig.json``. That is the WRONG model for15 DiLoCo replicas — it would couple replicas through SageMaker's intra-job16 network and break the "each replica is an independent DiLoCo worker that17 syncs only through S3" design. So ``launch_replicas`` submits N **separate**18 training jobs, each with ``ResourceConfig.InstanceCount == 1``, tagged with19 ``REPLICA_RANK=i`` / ``WORLD_SIZE=N`` via the ``Environment`` map. This20 mirrors ``ModalSpawnExecutor`` spawning N independent Modal calls.21 222. **Same S3 ``ObjectStoreAllReduce`` rendezvous — DiLoCo math untouched.**23 Cross-replica communication is EXCLUSIVELY the object-store rendezvous; the24 executor passes ``rendezvous_uri`` (an ``s3://...`` URI) through to25 ``replica_entrypoint.py`` unchanged. ``allreduce.py`` / ``MockManager`` /26 ``make_diloco_outer_loop`` / the trainer all stay byte-for-byte identical.27 283. **Stateless after launch; rank via ``Environment``.** Handle metadata is the29 ``training_job_name`` (plus submit timestamp). ``replica_entrypoint.py``30 already reads ``REPLICA_RANK`` from ``os.environ``, so the cleanest channel31 is the ``Environment`` map (string->string, max 100 entries, value <= 51232 chars). The container command is baked into the image entrypoint and the33 rendezvous args are passed via ``AlgorithmSpecification.ContainerArguments``.34 354. **``supports_inter_replica_network = False``.** Separate single-instance36 training jobs have no mutual network path by design — they rendezvous only37 through S3. (SageMaker's algo-N container fabric and38 ``EnableInterContainerTrafficEncryption`` only exist WITHIN a single39 multi-instance job, which this design deliberately does not use.)40 41Load-bearing gotcha — ``EnableNetworkIsolation`` MUST stay ``False``42--------------------------------------------------------------------43When ``EnableNetworkIsolation=True`` the training *container* has no outbound44network access. SageMaker's host-side processes still stage input channels and45ship CloudWatch logs, but the container itself cannot make S3 GET/PUT calls.46``ObjectStoreAllReduce`` needs live S3 PUT+GET every outer round, so network47isolation would silently dead-lock the allreduce poll loop until its timeout.48This executor pins ``EnableNetworkIsolation=False`` (the API default) and never49exposes it as a knob. The rendezvous bucket access must instead be granted on50the execution ``RoleArn`` — the SageMaker analog of EKS IRSA.51 52HyperPod <-> EKS 1:1 control-plane mapping (recommended hybrid)53---------------------------------------------------------------54Per the SageMaker docs: *"The high-level architecture of Amazon EKS support in55HyperPod involves a 1-to-1 mapping between an EKS cluster (control plane) and a56HyperPod cluster (worker nodes) within a VPC."*57(https://docs.aws.amazon.com/sagemaker/latest/dg/sagemaker-hyperpod-eks.html)58 59Consequence for this repo's hybrid: "use HyperPod for the inner GRPO trainer"60does NOT mean leaving EKS — it means attaching a HyperPod-managed61(auto-recovering, deep-health-checked, PyTorch-job auto-resume) node-group to62the SAME EKS control plane that runs the outer loop. A future ``EKSExecutor``63(kubernetes client, Indexed Jobs) therefore targets both plain Karpenter GPU64nodes AND HyperPod nodes transparently. ``SageMakerExecutor`` (ephemeral65Training Jobs via boto3) is the SEPARATE bursty-fallback inner-loop path for66when you don't want a persistent cluster: Training Jobs suit periodic /67smaller-model / pay-per-use runs; HyperPod suits continuous / large-model /68persistent runs. Both share the IDENTICAL S3 rendezvous, so a run can move69between them with zero trainer / loss / DiLoCo changes.70 71References72----------73- create_training_job: https://docs.aws.amazon.com/boto3/latest/reference/services/sagemaker/client/create_training_job.html74- describe_training_job: https://docs.aws.amazon.com/boto3/latest/reference/services/sagemaker/client/describe_training_job.html75- stop_training_job: https://docs.aws.amazon.com/boto3/latest/reference/services/sagemaker/client/stop_training_job.html76- network isolation: https://repost.aws/knowledge-center/sagemaker-access-network-isolation77- HyperPod-EKS: https://docs.aws.amazon.com/sagemaker/latest/dg/sagemaker-hyperpod-eks.html78- ADR-005 (executor protocol design)79"""80from __future__ import annotations81 82import json83import time84import uuid85from collections.abc import Callable, Mapping86from typing import Any87 88from composer_replication.diloco.serverless.executor import (89 ReplicaHandle,90)91 92# SageMaker TrainingJobStatus -> Protocol status vocabulary.93# describe_training_job's TrainingJobStatus is EXACTLY one of:94# 'InProgress' | 'Completed' | 'Failed' | 'Stopping' | 'Stopped'.95# We map Stopping -> 'running' (transient; still terminating, so collect()96# keeps waiting) and Stopped -> 'cancelled'.97_STATUS_MAP = {98 "InProgress": "running",99 "Completed": "succeeded",100 "Failed": "failed",101 "Stopping": "running",102 "Stopped": "cancelled",103}104 105# SecondaryStatus values that mean "queued / not yet executing user code" —106# used to refine an InProgress job into the Protocol's 'pending'.107_PENDING_SECONDARY = frozenset(108 {"Starting", "Pending", "LaunchingMLInstances", "PreparingTrainingStack"}109)110 111# Abstract Protocol GPU strings -> SageMaker instance types.112_GPU_INSTANCE_MAP = {113 "A100": "ml.p4d.24xlarge",114 "H100": "ml.p5.48xlarge",115 "H200": "ml.p5e.48xlarge",116 "B200": "ml.p6-b200.48xlarge",117 "L40S": "ml.g6e.12xlarge",118 "A10G": "ml.g5.2xlarge",119 "L4": "ml.g6.2xlarge",120}121 122_CLOUDWATCH_LOG_GROUP = "/aws/sagemaker/TrainingJobs"123 124 125class SageMakerExecutor:126 """Run replicas as N independent SageMaker Training Jobs.127 128 Implements the `ServerlessExecutor` Protocol against the boto3129 ``sagemaker`` client. Each replica is one single-instance training job;130 cross-replica communication happens only through the shared S3131 ``ObjectStoreAllReduce`` rendezvous.132 133 Args:134 role_arn: IAM execution role SageMaker assumes for the job. Must grant135 S3 access to the rendezvous + output buckets (the boto3 analog of136 EKS IRSA). The caller's credentials need ``iam:PassRole`` on it.137 image_uri: ECR image URI for the training container. The image must138 bake an entrypoint that runs139 ``python -m composer_replication.diloco.serverless.replica_entrypoint``140 (this executor also passes ``ContainerEntrypoint`` explicitly so a141 generic image works too).142 output_s3_path: ``s3://...`` prefix for ``OutputDataConfig.S3OutputPath``143 (model artifacts / failure output).144 instance_type: default SageMaker instance type when ``gpu`` is not145 mapped (e.g. ``"ml.g5.2xlarge"``). ``gpu=None`` at launch falls146 back to ``cpu_instance_type``.147 cpu_instance_type: instance type used when ``gpu`` is ``None`` (CPU148 smoke tests). Default ``"ml.m5.xlarge"``.149 volume_size_gb: ``ResourceConfig.VolumeSizeInGB`` per job.150 run_id: prefix for generated training-job names. Defaults to a short151 random token so names are unique per region+account.152 region: AWS region for the lazily-constructed boto3 clients. ``None``153 uses the ambient boto3 default-region resolution.154 sagemaker_client: inject a pre-built ``boto3.client('sagemaker')`` (or a155 mock) instead of constructing one. Used by tests.156 logs_client: inject a pre-built ``boto3.client('logs')`` (or a mock).157 158 Raises:159 RuntimeError: if boto3 is not installed and no client was injected.160 """161 162 backend_name = "sagemaker"163 # Separate single-instance jobs have no mutual network path — S3 only.164 supports_inter_replica_network = False165 166 def __init__(167 self,168 *,169 role_arn: str,170 image_uri: str,171 output_s3_path: str,172 instance_type: str = "ml.g5.2xlarge",173 cpu_instance_type: str = "ml.m5.xlarge",174 volume_size_gb: int = 100,175 run_id: str | None = None,176 region: str | None = None,177 sagemaker_client: Any = None,178 logs_client: Any = None,179 ) -> None:180 self.role_arn = role_arn181 self.image_uri = image_uri182 self.output_s3_path = output_s3_path183 self.instance_type = instance_type184 self.cpu_instance_type = cpu_instance_type185 self.volume_size_gb = volume_size_gb186 self.run_id = run_id or f"diloco-{uuid.uuid4().hex[:8]}"187 self._region = region188 189 # Lazy boto3 — only constructed if the caller didn't inject a client.190 # This keeps `import composer_replication.diloco.serverless` free of a191 # hard boto3 dependency (boto3 lives in the optional [aws] extra), and192 # lets tests inject a _MockSMClient with zero AWS calls.193 if sagemaker_client is None:194 sagemaker_client = self._make_boto3_client("sagemaker")195 self._client = sagemaker_client196 self._logs_client = logs_client # built lazily on first stream_logs()197 198 # rank -> {"job_name": str, "result": dict | None}199 self._handles: dict[int, dict[str, Any]] = {}200 201 # -----------------------------------------------------------------202 # boto3 plumbing (lazy)203 # -----------------------------------------------------------------204 205 def _make_boto3_client(self, service: str) -> Any:206 try:207 import boto3208 except ImportError as e:209 raise RuntimeError(210 "SageMakerExecutor requires boto3. Install with "211 "`pip install -e .[aws]` (or `pip install boto3`). "212 f"Got: {e!r}"213 ) from e214 if self._region is not None:215 return boto3.client(service, region_name=self._region)216 return boto3.client(service)217 218 def _map_gpu(self, gpu: str | None) -> str:219 """Translate the Protocol's abstract gpu string to an instance type.220 221 ``gpu=None`` -> ``cpu_instance_type`` (smoke tests). Unrecognized gpu222 strings fall back to ``instance_type`` (so a caller can pass a literal223 SageMaker instance type and it's honoured if not in the map).224 """225 if gpu is None:226 return self.cpu_instance_type227 if gpu in _GPU_INSTANCE_MAP:228 return _GPU_INSTANCE_MAP[gpu]229 # Caller may have passed a literal "ml.*" instance type.230 if gpu.startswith("ml."):231 return gpu232 return self.instance_type233 234 def _job_name(self, rank: int) -> str:235 """Build a unique, regex-safe training-job name (<= 63 chars).236 237 Pattern required by the API: ``[a-zA-Z0-9](-*[a-zA-Z0-9]){0,62}``.238 """239 name = f"{self.run_id}-r{rank:04d}-{int(time.time())}"240 return name[:63]241 242 # -----------------------------------------------------------------243 # ServerlessExecutor Protocol244 # -----------------------------------------------------------------245 246 def launch_replicas(247 self,248 n_replicas: int,249 entrypoint: str | Callable[..., Any],250 entrypoint_args: Mapping[str, Any],251 *,252 gpu: str | None = None,253 timeout: int = 3600,254 ) -> list[ReplicaHandle]:255 """Submit N independent single-instance SageMaker Training Jobs.256 257 Args:258 n_replicas: number of replicas (= number of training jobs).259 entrypoint: ignored — the container command is baked into the260 image / passed as ``ContainerEntrypoint``. Kept for Protocol261 compatibility.262 entrypoint_args: must contain ``rendezvous_uri`` (``s3://...``) and263 ``trainer_module``. Optional: ``trainer_fn`` (default264 ``"train"``), ``trainer_kwargs`` (dict, JSON-encoded into the265 container args). The conventional ``rank_env`` key (from266 ``LocalProcessExecutor``) is ignored — rank goes through the267 ``Environment`` map instead.268 gpu: abstract GPU spec mapped to an instance type via ``_map_gpu``.269 ``None`` -> CPU instance.270 timeout: ``StoppingCondition.MaxRuntimeInSeconds`` per job.271 272 Returns:273 ``list[ReplicaHandle]`` of length ``n_replicas`` in rank order274 (``handles[i].rank == i``).275 """276 del entrypoint # container command is baked / passed explicitly277 278 if n_replicas < 1:279 raise ValueError(f"n_replicas must be >= 1, got {n_replicas}")280 281 rendezvous_uri = entrypoint_args.get("rendezvous_uri")282 if not rendezvous_uri:283 raise ValueError(284 "entrypoint_args must include 'rendezvous_uri' (the s3:// "285 "ObjectStoreAllReduce rendezvous prefix)."286 )287 trainer_module = entrypoint_args.get("trainer_module")288 if not trainer_module:289 raise ValueError(290 "entrypoint_args must include 'trainer_module' (importable "291 "module path of the user's train function)."292 )293 trainer_fn = entrypoint_args.get("trainer_fn", "train")294 trainer_kwargs = entrypoint_args.get("trainer_kwargs", {})295 296 instance_type = self._map_gpu(gpu)297 298 # Container args: each element is a SINGLE token (StackOverflow299 # 77994925 — `['--world-size', '4']` NOT `['--world-size 4']`).300 container_args = [301 "--rendezvous", str(rendezvous_uri),302 "--world-size", str(n_replicas),303 "--trainer-module", str(trainer_module),304 "--trainer-fn", str(trainer_fn),305 "--trainer-kwargs-json", json.dumps(trainer_kwargs),306 ]307 308 handles: list[ReplicaHandle] = []309 for rank in range(n_replicas):310 job_name = self._job_name(rank)311 request = {312 "TrainingJobName": job_name,313 "AlgorithmSpecification": {314 "TrainingImage": self.image_uri,315 "TrainingInputMode": "File",316 "ContainerEntrypoint": [317 "python", "-m",318 "composer_replication.diloco.serverless.replica_entrypoint",319 ],320 "ContainerArguments": container_args,321 },322 "RoleArn": self.role_arn,323 # InputDataConfig intentionally omitted — the replica pulls324 # data via its own code / the S3 rendezvous, not SM channels.325 "OutputDataConfig": {"S3OutputPath": self.output_s3_path},326 "ResourceConfig": {327 "InstanceType": instance_type,328 "InstanceCount": 1,329 "VolumeSizeInGB": self.volume_size_gb,330 },331 "StoppingCondition": {"MaxRuntimeInSeconds": int(timeout)},332 # REPLICA_RANK / WORLD_SIZE injected as container env vars;333 # replica_entrypoint.py reads os.environ['REPLICA_RANK'].334 "Environment": {335 "REPLICA_RANK": str(rank),336 "WORLD_SIZE": str(n_replicas),337 "RENDEZVOUS_URI": str(rendezvous_uri),338 },339 # MUST stay False — True severs the container's S3 access and340 # dead-locks the allreduce poll loop. See module docstring.341 "EnableNetworkIsolation": False,342 }343 try:344 self._client.create_training_job(**request)345 except Exception as e:346 # Best-effort stop of already-launched siblings, then raise.347 for prior in handles:348 try:349 self.cancel(prior)350 except Exception:351 pass352 raise RuntimeError(353 f"SageMakerExecutor.launch_replicas failed at rank={rank} "354 f"of {n_replicas} (already-launched siblings stopped). "355 f"Underlying error: {e!r}"356 ) from e357 358 handle = ReplicaHandle(359 rank=rank,360 backend_name=self.backend_name,361 metadata={362 "training_job_name": job_name,363 "submit_ts": time.time(),364 },365 )366 self._handles[rank] = {"job_name": job_name, "result": None}367 handles.append(handle)368 369 return handles370 371 def poll(self, handle: ReplicaHandle) -> str:372 """Poll a training job's status.373 374 Returns one of: ``"pending"`` | ``"running"`` | ``"succeeded"`` |375 ``"failed"`` | ``"cancelled"``.376 377 Maps ``describe_training_job``'s ``TrainingJobStatus`` via378 ``_STATUS_MAP``; refines ``InProgress`` to ``"pending"`` while the job379 is still queued (``SecondaryStatus`` in ``_PENDING_SECONDARY``). A380 vanished job (``ResourceNotFound``) is treated as ``"cancelled"``.381 """382 meta = self._handles.get(handle.rank)383 if meta is None:384 return "cancelled"385 if meta["result"] is not None:386 return meta["result"]["status"]387 388 job_name = meta["job_name"]389 try:390 resp = self._client.describe_training_job(TrainingJobName=job_name)391 except Exception as e:392 if self._is_resource_not_found(e):393 return "cancelled"394 raise395 396 sm_status = resp.get("TrainingJobStatus", "InProgress")397 mapped = _STATUS_MAP.get(sm_status, "running")398 399 if sm_status == "InProgress":400 if resp.get("SecondaryStatus") in _PENDING_SECONDARY:401 return "pending"402 return "running"403 404 # Terminal — cache a result dict so collect()/repeat-poll are cheap.405 meta["result"] = self._terminal_result(handle.rank, sm_status, resp)406 return mapped407 408 def stream_logs(self, handle: ReplicaHandle, *, n_lines: int = 200) -> str:409 """Read recent CloudWatch logs for this replica's training job.410 411 SageMaker writes container stdout/stderr to the412 ``/aws/sagemaker/TrainingJobs`` log group, stream413 ``<job-name>/algo-<n>-<epoch>``. We discover the exact stream name by414 prefix then read the tail. Falls back to a CloudWatch console pointer415 on any error (mirrors ModalSpawnExecutor's dashboard-URL fallback).416 """417 meta = self._handles.get(handle.rank)418 if meta is None:419 return f"<replica {handle.rank}: no metadata>"420 job_name = meta["job_name"]421 422 try:423 logs = self._logs()424 prefix = f"{job_name}/"425 streams = logs.describe_log_streams(426 logGroupName=_CLOUDWATCH_LOG_GROUP,427 logStreamNamePrefix=prefix,428 orderBy="LastEventTime",429 descending=True,430 limit=1,431 )432 stream_list = streams.get("logStreams", [])433 if not stream_list:434 return (435 f"[rank {handle.rank}] job={job_name}: no CloudWatch log "436 f"stream yet (job pending / not started)."437 )438 stream_name = stream_list[0]["logStreamName"]439 events = logs.get_log_events(440 logGroupName=_CLOUDWATCH_LOG_GROUP,441 logStreamName=stream_name,442 limit=n_lines,443 startFromHead=False,444 )445 lines = [e.get("message", "") for e in events.get("events", [])]446 body = "\n".join(lines) if lines else "<no log events>"447 return f"[rank {handle.rank}] job={job_name} stream={stream_name}\n{body}"448 except Exception as e:449 region = self._region or "<region>"450 url = (451 f"https://{region}.console.aws.amazon.com/cloudwatch/home"452 f"?region={region}#logsV2:log-groups/log-group/"453 f"$252Faws$252Fsagemaker$252FTrainingJobs"454 )455 return (456 f"[rank {handle.rank}] job={job_name}: log fetch failed "457 f"({type(e).__name__}: {e!r}).\n CloudWatch console: {url}"458 )459 460 def cancel(self, handle: ReplicaHandle) -> None:461 """Best-effort stop of a training job.462 463 Calls ``stop_training_job`` (SIGTERM + 120s grace), swallowing464 ``ResourceNotFound`` and "already terminal" ``ValidationException`` so465 the contract — "no exception if already terminated" — holds.466 """467 meta = self._handles.get(handle.rank)468 if meta is None:469 return470 try:471 self._client.stop_training_job(TrainingJobName=meta["job_name"])472 except Exception as e:473 # R5: swallow ONLY already-terminated signals — a vanished job474 # (ResourceNotFound) or an already-Completed/Stopped job (boto3475 # raises ValidationException for "cannot stop a job in status X").476 # A genuinely unexpected error (AccessDenied, throttling that477 # outlived retries, malformed request) must propagate rather than478 # masquerade as a successful cancel.479 if self._is_resource_not_found(e) or self._is_already_terminal(e):480 return481 raise482 483 def collect(484 self,485 handles: list[ReplicaHandle],486 *,487 timeout: int | None = None,488 ) -> list[dict[str, Any]]:489 """Block until all replicas finish; return per-replica result dicts.490 491 Polls ``describe_training_job`` per handle until the job reaches a492 terminal status (``Completed`` / ``Failed`` / ``Stopped``) or the493 shared deadline elapses. Returns results aligned to the input handle494 order (Protocol contract; mirrors ``LocalProcessExecutor.collect``).495 496 Each result dict has at least497 ``{"rank", "status", "exit_code", "error"}``.498 """499 deadline = time.time() + (timeout if timeout is not None else 86400)500 poll_interval = 30.0501 results: list[dict[str, Any]] = []502 503 for h in handles:504 meta = self._handles.get(h.rank)505 if meta is None:506 results.append({507 "rank": h.rank,508 "status": "cancelled",509 "exit_code": None,510 "error": "handle has no metadata (cancelled or unknown)",511 "result": None,512 "training_job_name": h.metadata.get("training_job_name"),513 })514 continue515 516 # Already cached by an earlier poll()/collect().517 if meta["result"] is not None:518 results.append(meta["result"])519 continue520 521 job_name = meta["job_name"]522 result_dict: dict[str, Any] | None = None523 while True:524 try:525 resp = self._client.describe_training_job(526 TrainingJobName=job_name527 )528 except Exception as e:529 if self._is_resource_not_found(e):530 result_dict = {531 "rank": h.rank,532 "status": "cancelled",533 "exit_code": None,534 "error": "training job not found (deleted?)",535 "result": None,536 "training_job_name": job_name,537 }538 break539 raise540 541 sm_status = resp.get("TrainingJobStatus", "InProgress")542 if sm_status in ("Completed", "Failed", "Stopped"):543 result_dict = self._terminal_result(h.rank, sm_status, resp)544 break545 546 if time.time() >= deadline:547 result_dict = {548 "rank": h.rank,549 "status": "running",550 "exit_code": None,551 "error": "timeout before terminal",552 "result": None,553 "training_job_name": job_name,554 }555 break556 557 # Sleep, but never overrun the deadline.558 time.sleep(min(poll_interval, max(0.0, deadline - time.time())))559 560 # Cache only terminal results (not the timeout 'running' sentinel,561 # so a later collect() can re-check the job).562 if result_dict["status"] in ("succeeded", "failed", "cancelled"):563 meta["result"] = result_dict564 results.append(result_dict)565 566 return results567 568 # -----------------------------------------------------------------569 # Helpers570 # -----------------------------------------------------------------571 572 def _logs(self) -> Any:573 """Lazily build the CloudWatch Logs client (separate from sagemaker)."""574 if self._logs_client is None:575 self._logs_client = self._make_boto3_client("logs")576 return self._logs_client577 578 @staticmethod579 def _terminal_result(580 rank: int, sm_status: str, resp: Mapping[str, Any]581 ) -> dict[str, Any]:582 """Build a result dict from a terminal describe_training_job response."""583 mapped = _STATUS_MAP.get(sm_status, "failed")584 if sm_status == "Completed":585 exit_code: int | None = 0586 error = None587 elif sm_status == "Stopped":588 exit_code = None589 error = resp.get("FailureReason")590 else: # Failed591 exit_code = 1592 error = resp.get("FailureReason") or "training job failed"593 artifacts = resp.get("ModelArtifacts", {}) or {}594 return {595 "rank": rank,596 "status": mapped,597 "exit_code": exit_code,598 "error": error,599 "result": artifacts.get("S3ModelArtifacts"),600 "training_job_name": resp.get("TrainingJobName"),601 }602 603 def _is_resource_not_found(self, exc: Exception) -> bool:604 """True if ``exc`` is the boto3 ResourceNotFound for the sagemaker client.605 606 Handles both the typed client exception607 (``client.exceptions.ResourceNotFound``) and a generic botocore608 ``ClientError`` whose error code is ``ResourceNotFound`` /609 ``ValidationException`` naming a missing job — robust across whether a610 real boto3 client or a mock is in use.611 """612 rnf = getattr(getattr(self._client, "exceptions", None),613 "ResourceNotFound", None)614 if rnf is not None and isinstance(exc, rnf):615 return True616 # Generic botocore ClientError fallback.617 resp = getattr(exc, "response", None)618 if isinstance(resp, Mapping):619 code = resp.get("Error", {}).get("Code", "")620 if code in ("ResourceNotFound", "ValidationException"):621 return True622 return False623 624 def _is_already_terminal(self, exc: Exception) -> bool:625 """True if ``exc`` is the boto3 "cannot stop a job in status X" error.626 627 ``stop_training_job`` raises a ``ValidationException`` when the job is628 already Completed/Failed/Stopped — that is an idempotent no-op for629 cancel(), distinct from a genuinely unexpected error. Matched on the630 ClientError code + message text (robust to a mock raising a plain631 Exception whose message carries the phrase).632 """633 resp = getattr(exc, "response", None)634 if isinstance(resp, Mapping):635 err = resp.get("Error", {})636 if err.get("Code") == "ValidationException":637 return True638 msg = str(exc).lower()639 return (640 "cannot be stopped" in msg641 or "already" in msg and ("stopped" in msg or "complete" in msg or "terminal" in msg)642 )643 644 645__all__ = ["SageMakerExecutor"]646 