diff --git a/sagemaker-train/src/sagemaker/train/agent_rft_job.py b/sagemaker-train/src/sagemaker/train/agent_rft_job.py index edc5c1976c..b4f14c293a 100644 --- a/sagemaker-train/src/sagemaker/train/agent_rft_job.py +++ b/sagemaker-train/src/sagemaker/train/agent_rft_job.py @@ -22,6 +22,15 @@ from sagemaker.core.telemetry.telemetry_logging import _telemetry_emitter from sagemaker.core.telemetry.constants import Feature +from sagemaker.train.common_utils.log_streamer import ( + AGENT_RFT_LOG_GROUP, + LogStreamer, + _resolve_start_time_ms, + _validate_poll, + stream_log_loop, +) +from sagemaker.train.defaults import TrainDefaults + logger = logging.getLogger(__name__) JOB_CATEGORY = "AgentRFT" @@ -108,6 +117,37 @@ def wait(self, poll: int = 5, timeout: Optional[int] = 3000, max_log_lines: int _job_wait(self._job, poll=poll, timeout=timeout, description=self.description, max_log_lines=max_log_lines) + def stream_logs(self, poll: int = 5, start_time=None) -> None: + """Stream CloudWatch logs for this job in real-time. + + Polls ``/aws/sagemaker/Job/AgentRFT`` and exits when the job + reaches a terminal status or the user interrupts with Ctrl+C. + + :param poll: Seconds between CloudWatch polling cycles (1-300). + :param start_time: Stream from this timestamp. Accepts datetime or + epoch milliseconds (int). If None, streams from the beginning. + :raises ValueError: If poll is out of range. + """ + _validate_poll(poll) + start_ms = _resolve_start_time_ms(start_time) + sagemaker_session = self.sagemaker_session or TrainDefaults.get_sagemaker_session() + + streamer = LogStreamer( + log_group=AGENT_RFT_LOG_GROUP, + job_name=self.job_name, + sagemaker_session=sagemaker_session, + start_time_ms=start_ms, + ) + + logger.info("Streaming logs for job: %s", self.job_name) + logger.info("Log group: %s", AGENT_RFT_LOG_GROUP) + + def _get_status() -> str: + self._job.refresh() + return self._job.job_status + + stream_log_loop(streamer, poll, _get_status) + def stop(self): """Stop the job via StopJob API.""" self._job.stop() diff --git a/sagemaker-train/src/sagemaker/train/base_trainer.py b/sagemaker-train/src/sagemaker/train/base_trainer.py index c0da67a880..e22cb10723 100644 --- a/sagemaker-train/src/sagemaker/train/base_trainer.py +++ b/sagemaker-train/src/sagemaker/train/base_trainer.py @@ -15,6 +15,7 @@ import yaml import boto3 +from botocore.exceptions import ClientError from sagemaker.core.helper.session_helper import Session from sagemaker.core.training.configs import Tag, Networking, InputData, Channel, OutputDataConfig, HyperPodCompute @@ -651,55 +652,41 @@ def stream_logs(self, poll: int = 5, start_time: Optional[Any] = None) -> None: if isinstance(compute, HyperPodCompute): self._stream_logs_smhp(training_job, compute, poll, start_time_ms) else: - self._stream_logs_smtj(training_job, poll) + self._stream_logs_smtj(training_job, poll, start_time_ms) - def _stream_logs_smtj(self, training_job, poll: int) -> None: - """Stream logs for an SMTJ training job using MultiLogStreamHandler.""" + def _stream_logs_smtj(self, training_job, poll: int, start_time_ms=None) -> None: + """Stream logs for an SMTJ training job.""" + from sagemaker.train.common_utils.log_streamer import ( + LogStreamer, + stream_log_loop, + ) - # Resolve job name if hasattr(training_job, 'training_job_name'): job_name = training_job.training_job_name else: job_name = str(training_job) log_group = "/aws/sagemaker/TrainingJobs" - instance_count = 1 - if hasattr(self, 'compute') and self.compute and hasattr(self.compute, 'instance_count'): - instance_count = self.compute.instance_count or 1 - - handler = MultiLogStreamHandler( - log_group_name=log_group, - log_stream_name_prefix=job_name, - expected_stream_count=instance_count, + + sagemaker_session = TrainDefaults.get_sagemaker_session( + sagemaker_session=self.sagemaker_session ) - logger.info(f"Streaming logs for job: {job_name}") - logger.info(f"Log group: {log_group}") + streamer = LogStreamer( + log_group=log_group, + job_name=job_name, + sagemaker_session=sagemaker_session, + start_time_ms=start_time_ms, + ) - terminal_statuses = {"Completed", "Failed", "Stopped"} + logger.info("Streaming logs for job: %s", job_name) + logger.info("Log group: %s", log_group) - while True: - for stream_name, event in handler.get_latest_log_events(): - message = event.get("message", "").rstrip() - if message: - logger.info(message) + def _get_status() -> str: + job = TrainingJob.get(training_job_name=job_name) + return job.training_job_status - # Check job status - try: - job = TrainingJob.get(training_job_name=job_name) - status = job.training_job_status - if status in terminal_statuses: - # Final flush - for stream_name, event in handler.get_latest_log_events(): - message = event.get("message", "").rstrip() - if message: - logger.info(message) - logger.info(f"Job {job_name} finished with status: {status}") - return - except Exception: - pass - - time.sleep(poll) + stream_log_loop(streamer, poll, _get_status) def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None) -> None: """Stream logs for a HyperPod job using filter_log_events polling.""" @@ -735,6 +722,7 @@ def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None last_timestamp = int(time.time() * 1000) seen_event_ids = set() + empty_cycles = 0 while True: try: params = { @@ -746,6 +734,8 @@ def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None response = logs_client.filter_log_events(**params) events = response.get("events", []) + if events: + empty_cycles = 0 for event in events: event_id = event.get("eventId", "") if event_id not in seen_event_ids: @@ -756,11 +746,32 @@ def _stream_logs_smhp(self, training_job, compute, poll: int, start_time_ms=None ts = event.get("timestamp", 0) if ts > last_timestamp: last_timestamp = ts + if not events: + empty_cycles += 1 + if empty_cycles == 3: + logger.info("No log events yet, still waiting...") + except ClientError as e: + error_code = e.response.get("Error", {}).get("Code", "") + if error_code == "AccessDeniedException": + raise + if error_code == "ResourceNotFoundException": + empty_cycles += 1 + if empty_cycles == 1: + logger.info("Waiting for log group to become available...") + elif empty_cycles >= 60: + logger.warning( + "Log group %s still not found after %d attempts. " + "Check IAM permissions for logs:FilterLogEvents.", + log_group, + empty_cycles, + ) + else: + logger.debug(f"Error fetching HP logs: {e}") except Exception as e: logger.debug(f"Error fetching HP logs: {e}") # Note: HyperPod jobs don't have a simple status API to poll for completion. - # This polls till the user interrupts with Ctrl+C. + # This polls till the user interrupts with Ctrl+C. try: time.sleep(poll) except KeyboardInterrupt: diff --git a/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py b/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py new file mode 100644 index 0000000000..c120bd9131 --- /dev/null +++ b/sagemaker-train/src/sagemaker/train/common_utils/log_streamer.py @@ -0,0 +1,308 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""CloudWatch log streaming utility for SageMaker training and evaluation jobs.""" +from __future__ import annotations + +import logging +import time +from datetime import datetime, timezone +from typing import Callable, Optional + +from botocore.exceptions import ClientError + +from sagemaker.train.defaults import TrainDefaults + +logger = logging.getLogger(__name__) + +JOB_LOG_GROUP_PREFIX = "/aws/sagemaker/Job" +AGENT_RFT_LOG_GROUP = f"{JOB_LOG_GROUP_PREFIX}/AgentRFT" +AGENT_RFT_EVAL_LOG_GROUP = f"{JOB_LOG_GROUP_PREFIX}/AgentRFTEvaluation" +TERMINAL_STATUSES = frozenset({"Completed", "Succeeded", "Failed", "Stopped"}) + +_MIN_POLL = 1 +_MAX_POLL = 300 +_SMHP_STREAM_PREFIX = "SagemakerHyperPodTrainingJob" + + +def _resolve_start_time_ms(start_time: datetime | int | None) -> int | None: + """Convert start_time to epoch milliseconds. + + :param start_time: datetime, epoch milliseconds int, or None. + :returns: Epoch milliseconds or None. + :raises TypeError: If start_time is not a supported type. + """ + if start_time is None: + return None + if isinstance(start_time, datetime): + return int(start_time.timestamp() * 1000) + if isinstance(start_time, int): + return start_time + raise TypeError( + f"start_time must be datetime or int (epoch ms), got: {type(start_time).__name__}" + ) + + +def _validate_poll(poll: int) -> None: + """Validate poll interval. + + :raises ValueError: If poll is out of range. + """ + if not isinstance(poll, int) or poll < _MIN_POLL or poll > _MAX_POLL: + raise ValueError(f"poll must be between {_MIN_POLL} and {_MAX_POLL} seconds, got: {poll!r}") + + +def _format_timestamp(ts_ms: int) -> str: + """Format epoch milliseconds to ISO timestamp string.""" + return datetime.fromtimestamp(ts_ms / 1000, tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S") + + +class LogStreamer: + """Fetches new CloudWatch log events since last call. + + This is a low-level utility that does NOT own the polling loop, + does NOT check job status, and does NOT handle timeouts. The caller + is responsible for all of that. + + :param log_group: CloudWatch log group name. + :param job_name: Job name used as log stream prefix (stream mode) + or in the filter pattern (filter mode). + :param sagemaker_session: SageMaker session for obtaining CW client. + :param filter_pattern: If provided, uses filter_log_events API (HyperPod). + If None, uses describe_log_streams + get_log_events (stream mode). + :param start_time_ms: Initial start time as epoch milliseconds. + """ + + def __init__( + self, + log_group: str, + job_name: str, + sagemaker_session=None, + filter_pattern: str | None = None, + start_time_ms: int | None = None, + ): + self._log_group = log_group + self._job_name = job_name + self._filter_pattern = filter_pattern + self._start_time_ms = start_time_ms + + session = sagemaker_session or TrainDefaults.get_sagemaker_session() + region = session.boto_session.region_name + self._logs_client = session.boto_session.client("logs", region_name=region) + + # Stream mode state + self._stream_handlers: list[dict] | None = None + + # Filter mode state + self._last_timestamp_ms: int | None = start_time_ms + self._last_event_ids: set[str] = set() + + def poll_once(self) -> list[tuple[int, str]]: + """Fetch and return new log event messages since last poll. + + :returns: List of (timestamp_ms, message) tuples. May be empty. + :raises: ClientError with ResourceNotFoundException or + AccessDeniedException propagates to caller. Other ClientErrors + are caught internally and result in an empty list. + """ + try: + if self._filter_pattern is not None: + return self._poll_filter_mode() + return self._poll_stream_mode() + except ClientError as e: + code = e.response["Error"]["Code"] + if code in ("ResourceNotFoundException", "AccessDeniedException"): + raise + logger.debug("Transient CloudWatch error: %s", e) + return [] + + def _poll_filter_mode(self) -> list[tuple[int, str]]: + """Poll using filter_log_events (HyperPod style).""" + params = { + "logGroupName": self._log_group, + "logStreamNamePrefix": _SMHP_STREAM_PREFIX, + "filterPattern": self._filter_pattern, + } + if self._last_timestamp_ms is not None: + params["startTime"] = self._last_timestamp_ms + + response = self._logs_client.filter_log_events(**params) + events = response.get("events", []) + + results = [] + for event in events: + event_id = event.get("eventId", "") + ts = event.get("timestamp", 0) + + # Dedup: skip events from same millisecond already seen + if ts == self._last_timestamp_ms and event_id in self._last_event_ids: + continue + + message = event.get("message", "").rstrip() + if message: + results.append((ts, message)) + + # Advance cursor + if ts > (self._last_timestamp_ms or 0): + self._last_timestamp_ms = ts + self._last_event_ids = {event_id} + elif ts == self._last_timestamp_ms: + self._last_event_ids.add(event_id) + + return results + + def _poll_stream_mode(self) -> list[tuple[int, str]]: + """Poll using describe_log_streams + get_log_events per stream.""" + if self._stream_handlers is None: + self._stream_handlers = self._discover_streams() + if not self._stream_handlers: + return [] + + results = [] + for handler in self._stream_handlers: + events = self._get_events_for_stream(handler) + results.extend(events) + + return results + + def _discover_streams(self) -> list[dict]: + """Discover log streams matching the job name prefix.""" + handlers = [] + kwargs = { + "logGroupName": self._log_group, + "logStreamNamePrefix": self._job_name, + } + paginator = self._logs_client.get_paginator("describe_log_streams") + for page in paginator.paginate(**kwargs): + for stream in page.get("logStreams", []): + handlers.append({ + "stream_name": stream["logStreamName"], + "next_token": None, + "started": False, + }) + return handlers + + def _get_events_for_stream(self, handler: dict) -> list[tuple[int, str]]: + """Get new events from a single log stream.""" + kwargs = { + "logGroupName": self._log_group, + "logStreamName": handler["stream_name"], + "startFromHead": True, + } + if handler["next_token"]: + kwargs["nextToken"] = handler["next_token"] + elif not handler["started"] and self._start_time_ms is not None: + kwargs["startTime"] = self._start_time_ms + + response = self._logs_client.get_log_events(**kwargs) + new_token = response.get("nextForwardToken") + + # Same token means caught up — no new events + if new_token == handler["next_token"]: + return [] + + handler["next_token"] = new_token + handler["started"] = True + + results = [] + for event in response.get("events", []): + message = event.get("message", "").rstrip() + ts = event.get("timestamp", 0) + if message: + results.append((ts, message)) + + return results + + +def stream_log_loop( + streamer: LogStreamer, + poll: int, + status_fn: Callable[[], str], +) -> None: + """Run the standard log streaming loop. + + Polls the streamer, prints events, checks job status, and exits + when terminal. Handles ResourceNotFoundException with escalating + feedback and propagates AccessDeniedException. + + :param streamer: A configured LogStreamer instance. + :param poll: Seconds between polls. + :param status_fn: Callable that returns the current job status string. + """ + status = status_fn() + if status in TERMINAL_STATUSES: + logger.info("Job already in terminal state: %s", status) + try: + while True: + events = streamer.poll_once() + if not events: + break + for ts_ms, message in events: + logger.info("[%s] %s", _format_timestamp(ts_ms), message) + except ClientError: + pass + logger.info("Job finished with status: %s", status) + return + + empty_cycles = 0 + max_empty_cycles = max(300 // poll, 1) # ~5 minutes + warn_cycle = max(30 // poll, 1) # ~30 seconds + + while True: + try: + events = streamer.poll_once() + except ClientError as e: + error_code = e.response["Error"]["Code"] + if error_code == "AccessDeniedException": + raise + if error_code == "ResourceNotFoundException": + status = status_fn() + if status in TERMINAL_STATUSES: + logger.info("Job finished before producing logs.") + return + empty_cycles += 1 + if empty_cycles == 1: + logger.info("Waiting for log group to become available...") + elif empty_cycles >= max_empty_cycles: + raise RuntimeError( + f"Log group not found after {empty_cycles * poll}s. " + "Check IAM permissions for logs:GetLogEvents and " + "logs:DescribeLogStreams." + ) + time.sleep(poll) + continue + raise + + if events: + empty_cycles = 0 + for ts_ms, message in events: + logger.info("[%s] %s", _format_timestamp(ts_ms), message) + else: + empty_cycles += 1 + if empty_cycles == warn_cycle: + logger.info("No log events yet, still waiting...") + elif empty_cycles >= max_empty_cycles: + logger.warning("No log events found after 5 minutes.") + return + + status = status_fn() + if status in TERMINAL_STATUSES: + for ts_ms, message in streamer.poll_once(): + logger.info("[%s] %s", _format_timestamp(ts_ms), message) + logger.info("Job finished with status: %s", status) + return + + try: + time.sleep(poll) + except KeyboardInterrupt: + logger.info("Streaming stopped by user.") + return diff --git a/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py b/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py index c966f4a8ff..1d05b80a08 100644 --- a/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py +++ b/sagemaker-train/src/sagemaker/train/evaluate/base_evaluator.py @@ -9,9 +9,11 @@ import logging import re +import time from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union -from pydantic import BaseModel, validator +from botocore.exceptions import ClientError +from pydantic import BaseModel, PrivateAttr, validator from sagemaker.core.common_utils import TagsDict from sagemaker.core.helper.iam_role_resolver import ( @@ -32,6 +34,14 @@ _resolve_mlflow_resource_arn, _is_nova_model, ) +from sagemaker.train.common_utils.cloudwatch_metrics import _get_smhp_log_group +from sagemaker.train.common_utils.log_streamer import ( + LogStreamer, + _format_timestamp, + _resolve_start_time_ms, + _validate_poll, + stream_log_loop, +) from sagemaker.train.common_utils.recipe_utils import resolve_recipe, get_resolved_recipe_from_context from sagemaker.train.common_utils.validator import validate_hyperpod_compute from sagemaker.train.defaults import TrainDefaults @@ -145,7 +155,9 @@ class BaseEvaluator(BaseModel): training_image: Optional[str] = None recipe: Optional[str] = None overrides: Optional[Dict[str, Any]] = None - + + _latest_execution: Any = PrivateAttr(default=None) + class Config: arbitrary_types_allowed = True @@ -984,7 +996,8 @@ def _start_execution( region=region, tags=tags ) - + + self._latest_execution = execution return execution def _get_effective_hyperparameters(self) -> Dict[str, Any]: @@ -1083,6 +1096,155 @@ def evaluate(self, dry_run: bool = False) -> Any: """ raise NotImplementedError("Subclasses must implement evaluate method") + def stream_logs(self, poll: int = 5, start_time=None) -> None: + """Stream CloudWatch logs for the latest evaluation execution. + + Dispatches to the underlying execution object's ``stream_logs()`` + for pipeline-based evaluations, or streams directly from the + HyperPod cluster log group for HyperPod evaluations. + + :param poll: Seconds between CloudWatch polling cycles (1-300). + :param start_time: Stream from this timestamp. Accepts datetime or + epoch milliseconds (int). If None, defaults depend on the backend. + :raises ValueError: If no evaluation has been executed yet or poll + is out of range. + """ + if self._latest_execution is None: + raise ValueError( + "No evaluation executed yet. Call .evaluate() first, " + "then call .stream_logs() to stream logs." + ) + _validate_poll(poll) + + if isinstance(self.compute, HyperPodCompute): + self._stream_logs_hyperpod(self._latest_execution, poll, start_time) + else: + self._stream_logs_pipeline(self._latest_execution, poll, start_time) + + def _stream_logs_hyperpod(self, job_name: str, poll: int, start_time) -> None: + """Stream logs for a HyperPod evaluation job.""" + sagemaker_session = TrainDefaults.get_sagemaker_session( + sagemaker_session=self.sagemaker_session + ) + log_group = _get_smhp_log_group( + self.compute.cluster_name, sagemaker_session.sagemaker_client + ) + + start_ms = _resolve_start_time_ms(start_time) + if start_ms is None: + start_ms = int(time.time() * 1000) + + streamer = LogStreamer( + log_group=log_group, + job_name=job_name, + sagemaker_session=sagemaker_session, + filter_pattern=f'"{job_name}"', + start_time_ms=start_ms, + ) + + _logger.info("Streaming logs for HyperPod eval job: %s", job_name) + _logger.info("Log group: %s", log_group) + _logger.info("Press Ctrl+C to stop streaming.") + + empty_cycles = 0 + while True: + try: + events = streamer.poll_once() + except ClientError as e: + error_code = e.response["Error"]["Code"] + if error_code == "AccessDeniedException": + raise + if error_code == "ResourceNotFoundException": + empty_cycles += 1 + if empty_cycles == 1: + _logger.info("Waiting for log group to become available...") + elif empty_cycles >= 60: + _logger.warning( + "Log group still not found after %d attempts. " + "Check IAM permissions.", + empty_cycles, + ) + time.sleep(poll) + continue + raise + + if events: + empty_cycles = 0 + for ts_ms, message in events: + _logger.info("[%s] %s", _format_timestamp(ts_ms), message) + else: + empty_cycles += 1 + if empty_cycles == 3: + _logger.info("No log events yet, still waiting...") + + # HyperPod jobs don't have a simple status API to poll for completion. + # This polls till the user interrupts with Ctrl+C. + try: + time.sleep(poll) + except KeyboardInterrupt: + _logger.info("Streaming stopped by user.") + return + + def _stream_logs_pipeline(self, execution, poll: int, start_time) -> None: + """Stream logs for a pipeline-based evaluation.""" + + start_ms = _resolve_start_time_ms(start_time) + + execution.refresh() + job_arn = self._find_job_arn(execution) + if not job_arn: + _logger.warning("No pipeline step with a job ARN found.") + return + + log_group = self._log_group_for_step_arn(job_arn) + job_name = self._job_name_from_arn(job_arn) + + sagemaker_session = TrainDefaults.get_sagemaker_session( + sagemaker_session=self.sagemaker_session + ) + + streamer = LogStreamer( + log_group=log_group, + job_name=job_name, + sagemaker_session=sagemaker_session, + start_time_ms=start_ms, + ) + + _logger.info("Streaming logs for eval step: %s", job_name) + _logger.info("Log group: %s", log_group) + + def _get_status() -> str: + execution.refresh() + return execution.status.overall_status + + stream_log_loop(streamer, poll, _get_status) + + @staticmethod + def _find_job_arn(execution) -> str | None: + """Find the first pipeline step with a job ARN.""" + for step in execution.status.step_details: + if step.job_arn: + return step.job_arn + return None + + @staticmethod + def _log_group_for_step_arn(arn: str) -> str: + """Resolve CloudWatch log group from a pipeline step's job ARN.""" + if ":training-job/" in arn: + return "/aws/sagemaker/TrainingJobs" + elif ":job/" in arn: + # Only MTRL evals use Job-API steps in pipelines today + return "/aws/sagemaker/Job/AgentRFTEvaluation" + return "/aws/sagemaker/TrainingJobs" + + @staticmethod + def _job_name_from_arn(arn: str) -> str | None: + """Extract job name from a SageMaker resource ARN.""" + for prefix in (":training-job/", ":job/", ":processing-job/", ":transform-job/"): + if prefix in arn: + return arn.split(prefix, 1)[1].split("/")[0] + return None + # ─── Shared SMTJ evaluation helpers ───────────────────────────────────────── def _get_smtj_session_and_role(self): @@ -1866,4 +2028,6 @@ def _submit_hyperpod_eval_job(self, override_parameters=None, base_job_name=None if not matched: raise ValueError(f"Could not find job name in output: {start_result.stdout}") - return matched.group(1) + job_name = matched.group(1) + self._latest_execution = job_name + return job_name diff --git a/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py b/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py index 8794a2a625..c18ef7320c 100644 --- a/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py +++ b/sagemaker-train/src/sagemaker/train/multi_turn_rl_trainer.py @@ -41,6 +41,13 @@ _validate_s3_path_exists, ) from sagemaker.train.common_utils.constants import MIN_MLFLOW_VERSION +from sagemaker.train.common_utils.log_streamer import ( + AGENT_RFT_LOG_GROUP, + LogStreamer, + _resolve_start_time_ms, + _validate_poll, + stream_log_loop, +) from sagemaker.train.common_utils.recipe_utils import _list_hub_models_by_recipe, _is_nova_model from sagemaker.train.constants import get_sagemaker_hub_name from sagemaker.train.defaults import TrainDefaults @@ -350,6 +357,46 @@ def output_model_package_arn(self) -> str | None: return self._latest_job.output_model_package_arn return None + def stream_logs(self, poll: int = 5, start_time=None) -> None: + """Stream CloudWatch logs for the latest MTRL training job. + + Polls the correct log group (``/aws/sagemaker/Job/AgentRFT``) and + checks job status via the Job API. + + :param poll: Seconds between CloudWatch polling cycles (1-300). + :param start_time: Stream from this timestamp. Accepts datetime or + epoch milliseconds (int). If None, streams from the beginning. + :raises ValueError: If no training job has been launched yet or poll + is out of range. + """ + if self._latest_job is None: + raise ValueError( + "No training job found. Call .train(wait=False) first, " + "then call .stream_logs() to stream logs in real-time." + ) + _validate_poll(poll) + + job_name = self._latest_job.job_name + start_ms = _resolve_start_time_ms(start_time) + sagemaker_session = TrainDefaults.get_sagemaker_session( + sagemaker_session=self.sagemaker_session + ) + + streamer = LogStreamer( + log_group=AGENT_RFT_LOG_GROUP, + job_name=job_name, + sagemaker_session=sagemaker_session, + start_time_ms=start_ms, + ) + + logger.info("Streaming logs for job: %s", job_name) + logger.info("Log group: %s", AGENT_RFT_LOG_GROUP) + + def _get_status() -> str: + return Job.get(job_name=job_name, job_category=JOB_CATEGORY).job_status + + stream_log_loop(streamer, poll, _get_status) + @classmethod @_telemetry_emitter( feature=Feature.MODEL_CUSTOMIZATION, func_name="MultiTurnRLTrainer.attach" diff --git a/sagemaker-train/tests/integ/train/test_stream_logs_evaluator.py b/sagemaker-train/tests/integ/train/test_stream_logs_evaluator.py new file mode 100644 index 0000000000..fc84437325 --- /dev/null +++ b/sagemaker-train/tests/integ/train/test_stream_logs_evaluator.py @@ -0,0 +1,150 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Integration tests for evaluator stream_logs()""" +from __future__ import annotations + +import logging +import time + +import boto3 +import pytest + +from sagemaker.core.helper.session_helper import Session +from sagemaker.train.evaluate.benchmark_evaluator import BenchMarkEvaluator, get_benchmarks +from sagemaker.train.evaluate.custom_scorer_evaluator import CustomScorerEvaluator, get_builtin_metrics +from sagemaker.train.evaluate.llm_as_judge_evaluator import LLMAsJudgeEvaluator +from sagemaker.train.evaluate.execution import ( + EvaluationPipelineExecution, + PipelineExecutionStatus, + StepDetail, +) + +logger = logging.getLogger(__name__) + +REGION = "us-west-2" + +S3_OUTPUT = "s3://sagemaker-us-west-2-729646638167/model-customization/eval/" +MODEL_PACKAGE_ARN = "arn:aws:sagemaker:us-west-2:729646638167:model-package/sdk-test-finetuned-models/1" +DATASET_S3 = "s3://sagemaker-us-west-2-729646638167/model-customization/eval/zc_test.jsonl" + +BENCHMARK_EXECUTION_ARN = "arn:aws:sagemaker:us-west-2:729646638167:pipeline/SagemakerEvaluation-BenchmarkEvaluation-499b3c7e-e456-4297-9dc0-cc5737137c9c/execution/p1gtwhjm9dzt" +BENCHMARK_STEP_ARN = "arn:aws:sagemaker:us-west-2:729646638167:training-job/pipelines-p1gtwhjm9dzt-EvaluateCustomModel-XEdt5h2gQC" + +CUSTOM_SCORER_EXECUTION_ARN = "arn:aws:sagemaker:us-west-2:729646638167:pipeline/SagemakerEvaluation-CustomScorerEvaluation-2d0fde36-af0f-49d7-8b8e-a5e11352dc1f/execution/yca2ij65mlhr" +CUSTOM_SCORER_STEP_ARN = "arn:aws:sagemaker:us-west-2:729646638167:training-job/pipelines-yca2ij65mlhr-EvaluateCustomModel-MlMUskwbNB" + +LLMAJ_EXECUTION_ARN = "arn:aws:sagemaker:us-west-2:729646638167:pipeline/SagemakerEvaluation-LLMAJEvaluation-ac7a1fe7-fe8a-445c-8aa5-702b3d6b7771/execution/hmk0lcu6ufzc" +LLMAJ_STEP_ARN = "arn:aws:sagemaker:us-west-2:729646638167:training-job/pipelines-hmk0lcu6ufzc-EvaluateCustomModelM-6UaY2bgNL5" + + + +@pytest.fixture(scope="module") +def sagemaker_session(): + boto_session = boto3.Session(region_name=REGION) + return Session(boto_session=boto_session) + + +def _make_execution(execution_arn: str, step_name: str, step_arn: str): + """Construct a completed execution from known pipeline step details.""" + return EvaluationPipelineExecution( + arn=execution_arn, + name=execution_arn.split("/execution/")[1], + status=PipelineExecutionStatus( + overall_status="Succeeded", + step_details=[ + StepDetail( + name=step_name, + status="Succeeded", + display_name=step_name, + job_arn=step_arn, + ) + ], + ), + ) + + +class TestEvaluatorStreamLogsFromCompletedJobs: + """Verify evaluator.stream_logs() works for each evaluator type. + + Each test uses the actual evaluator class that produced the pipeline + execution, with real step ARNs that have CloudWatch logs available. + """ + + def test_benchmark_evaluator_stream_logs(self, sagemaker_session): + """BenchMarkEvaluator.stream_logs() on a completed benchmark pipeline.""" + Benchmark = get_benchmarks() + evaluator = BenchMarkEvaluator( + benchmark=Benchmark.MMLU, + model=MODEL_PACKAGE_ARN, + s3_output_path=S3_OUTPUT, + sagemaker_session=sagemaker_session, + ) + evaluator._latest_execution = _make_execution( + BENCHMARK_EXECUTION_ARN, "EvaluateCustomModel", BENCHMARK_STEP_ARN + ) + + start = time.time() + evaluator.stream_logs(poll=2) + elapsed = time.time() - start + + assert elapsed < 30, ( + f"stream_logs() took {elapsed:.1f}s — should exit quickly for completed job" + ) + print(f"✓ BenchMarkEvaluator.stream_logs() completed in {elapsed:.1f}s") + + def test_custom_scorer_evaluator_stream_logs(self, sagemaker_session): + """CustomScorerEvaluator.stream_logs() on a completed custom scorer pipeline.""" + BuiltInMetric = get_builtin_metrics() + evaluator = CustomScorerEvaluator( + evaluator=BuiltInMetric.PRIME_MATH, + dataset=DATASET_S3, + model=MODEL_PACKAGE_ARN, + s3_output_path=S3_OUTPUT, + sagemaker_session=sagemaker_session, + ) + evaluator._latest_execution = _make_execution( + CUSTOM_SCORER_EXECUTION_ARN, "EvaluateCustomModel", CUSTOM_SCORER_STEP_ARN + ) + + start = time.time() + evaluator.stream_logs(poll=2) + elapsed = time.time() - start + + assert elapsed < 30, ( + f"stream_logs() took {elapsed:.1f}s — should exit quickly for completed job" + ) + print(f"✓ CustomScorerEvaluator.stream_logs() completed in {elapsed:.1f}s") + + def test_llm_as_judge_evaluator_stream_logs(self, sagemaker_session): + """LLMAsJudgeEvaluator.stream_logs() on a completed LLMAJ pipeline.""" + evaluator = LLMAsJudgeEvaluator( + model=MODEL_PACKAGE_ARN, + evaluator_model="amazon.nova-pro-v1:0", + dataset=DATASET_S3, + builtin_metrics=["Completeness"], + s3_output_path=S3_OUTPUT, + sagemaker_session=sagemaker_session, + evaluate_base_model=False, + ) + evaluator._latest_execution = _make_execution( + LLMAJ_EXECUTION_ARN, "EvaluateCustomModelMetrics", LLMAJ_STEP_ARN + ) + + start = time.time() + evaluator.stream_logs(poll=2) + elapsed = time.time() - start + + assert elapsed < 30, ( + f"stream_logs() took {elapsed:.1f}s — should exit quickly for completed job" + ) + print(f"✓ LLMAsJudgeEvaluator.stream_logs() completed in {elapsed:.1f}s") diff --git a/sagemaker-train/tests/integ/train/test_stream_logs_trainer.py b/sagemaker-train/tests/integ/train/test_stream_logs_trainer.py new file mode 100644 index 0000000000..b6a952bf49 --- /dev/null +++ b/sagemaker-train/tests/integ/train/test_stream_logs_trainer.py @@ -0,0 +1,118 @@ +# Copyright Amazon.com, Inc. or its affiliates. All Rights Reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"). You +# may not use this file except in compliance with the License. A copy of +# the License is located at +# +# http://aws.amazon.com/apache2.0/ +# +# or in the "license" file accompanying this file. This file is +# distributed on an "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF +# ANY KIND, either express or implied. See the License for the specific +# language governing permissions and limitations under the License. +"""Integration tests for trainer stream_logs()""" +from __future__ import annotations + +import logging +import time + +import boto3 +import pytest + +from sagemaker.core.helper.session_helper import Session +from sagemaker.core.resources import TrainingJob +from sagemaker.train.agent_rft_job import AgentRFTJob +from sagemaker.train.sft_trainer import SFTTrainer + +logger = logging.getLogger(__name__) + +REGION = "us-west-2" +MTRL_JOB_NAME = "mock-oss-test-mtrl-20260729120959" +SERVERFUL_JOB_NAME = "pytorch-training-260729-1927-002-95b83cb6" + + + +@pytest.fixture(scope="module") +def sagemaker_session(): + boto_session = boto3.Session(region_name=REGION) + return Session(boto_session=boto_session) + + + +class TestMTRLStreamLogs: + """Verify AgentRFTJob.stream_logs() on a completed MTRL job.""" + + def test_stream_logs_exits_on_completed(self, sagemaker_session): + """AgentRFTJob.stream_logs() exits quickly for a completed job.""" + job = AgentRFTJob.get(MTRL_JOB_NAME, session=sagemaker_session.boto_session) + assert job.job_status == "Completed" + + start = time.time() + job.stream_logs(poll=2) + elapsed = time.time() - start + + assert elapsed < 30, ( + f"stream_logs() took {elapsed:.1f}s — should exit quickly for completed job" + ) + print(f"✓ AgentRFTJob.stream_logs() completed in {elapsed:.1f}s") + + def test_stream_logs_with_start_time(self, sagemaker_session): + """AgentRFTJob.stream_logs() respects start_time parameter.""" + job = AgentRFTJob.get(MTRL_JOB_NAME, session=sagemaker_session.boto_session) + + future_ms = int((time.time() + 86400) * 1000) + start = time.time() + job.stream_logs(poll=2, start_time=future_ms) + elapsed = time.time() - start + + assert elapsed < 30 + print(f"✓ stream_logs(start_time=future) completed in {elapsed:.1f}s") + + + +class TestServerfulSMTJStreamLogs: + """Verify BaseTrainer.stream_logs() on a completed serverful training job.""" + + def test_stream_logs_exits_on_completed(self, sagemaker_session): + """stream_logs() exits quickly for a completed serverful job.""" + tj = TrainingJob.get( + training_job_name=SERVERFUL_JOB_NAME, + session=sagemaker_session.boto_session, + ) + assert tj.training_job_status == "Completed" + + trainer = SFTTrainer.__new__(SFTTrainer) + trainer._latest_training_job = tj + trainer.sagemaker_session = sagemaker_session + trainer.compute = None + + start = time.time() + trainer.stream_logs(poll=2) + elapsed = time.time() - start + + assert elapsed < 30, ( + f"stream_logs() took {elapsed:.1f}s — should exit quickly for completed job" + ) + print(f"✓ Serverful SMTJ stream_logs() completed in {elapsed:.1f}s") + + + +class TestStreamLogsValidation: + """Verify input validation for stream_logs() parameters.""" + + def test_poll_validation(self, sagemaker_session): + """stream_logs() rejects invalid poll values.""" + job = AgentRFTJob.get(MTRL_JOB_NAME, session=sagemaker_session.boto_session) + + with pytest.raises(ValueError, match="poll must be between"): + job.stream_logs(poll=0) + + with pytest.raises(ValueError, match="poll must be between"): + job.stream_logs(poll=400) + + def test_start_time_type_validation(self, sagemaker_session): + """stream_logs() rejects invalid start_time types.""" + job = AgentRFTJob.get(MTRL_JOB_NAME, session=sagemaker_session.boto_session) + + with pytest.raises(TypeError, match="start_time must be datetime or int"): + job.stream_logs(start_time="2025-01-01") diff --git a/sagemaker-train/tests/unit/train/test_log_streamer.py b/sagemaker-train/tests/unit/train/test_log_streamer.py new file mode 100644 index 0000000000..f09040e2f0 --- /dev/null +++ b/sagemaker-train/tests/unit/train/test_log_streamer.py @@ -0,0 +1,445 @@ +"""Unit tests for LogStreamer utility.""" +from __future__ import annotations + +from datetime import datetime, timezone +from unittest.mock import MagicMock, patch + +import pytest +from botocore.exceptions import ClientError + +from sagemaker.train.common_utils.log_streamer import ( + LogStreamer, + _resolve_start_time_ms, + _validate_poll, + stream_log_loop, +) + + +class TestResolveStartTimeMs: + def test_datetime_converts_to_ms(self): + dt = datetime(2023, 11, 14, 22, 13, 20, tzinfo=timezone.utc) + result = _resolve_start_time_ms(dt) + assert result == int(dt.timestamp() * 1000) + + def test_invalid_type_raises_typeerror(self): + with pytest.raises(TypeError, match="start_time must be datetime or int"): + _resolve_start_time_ms("2023-11-14") + + +class TestValidatePoll: + def test_poll_too_low(self): + with pytest.raises(ValueError, match="poll must be between 1 and 300"): + _validate_poll(0) + + def test_poll_too_high(self): + with pytest.raises(ValueError, match="poll must be between 1 and 300"): + _validate_poll(301) + + def test_poll_not_int(self): + with pytest.raises(ValueError, match="poll must be between 1 and 300"): + _validate_poll(5.0) + + +def _make_mock_session(): + session = MagicMock() + session.boto_session.region_name = "us-west-2" + return session + + +def _make_client_error(code, message="error"): + return ClientError( + {"Error": {"Code": code, "Message": message}}, + "operation_name", + ) + + +class TestLogStreamerStreamMode: + def test_poll_once_returns_events(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + # describe_log_streams returns one stream + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "my-job/algo-1"}]} + ] + # get_log_events returns events + logs_client.get_log_events.return_value = { + "events": [ + {"timestamp": 1700000000000, "message": "Training started\n"}, + {"timestamp": 1700000001000, "message": "Epoch 1\n"}, + ], + "nextForwardToken": "token-2", + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + ) + events = streamer.poll_once() + + assert len(events) == 2 + assert events[0] == (1700000000000, "Training started") + assert events[1] == (1700000001000, "Epoch 1") + + def test_poll_once_returns_empty_when_caught_up(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "my-job/algo-1"}]} + ] + # First call returns events + logs_client.get_log_events.return_value = { + "events": [{"timestamp": 1700000000000, "message": "line1\n"}], + "nextForwardToken": "token-A", + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + ) + streamer.poll_once() + + # Second call: same token means caught up + logs_client.get_log_events.return_value = { + "events": [], + "nextForwardToken": "token-A", + } + events = streamer.poll_once() + assert events == [] + + def test_poll_once_returns_empty_when_no_streams(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": []} + ] + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + ) + events = streamer.poll_once() + assert events == [] + + def test_start_time_ms_passed_to_first_call(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "my-job/algo-1"}]} + ] + logs_client.get_log_events.return_value = { + "events": [], + "nextForwardToken": "token-1", + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + start_time_ms=1700000000000, + ) + streamer.poll_once() + + call_kwargs = logs_client.get_log_events.call_args[1] + assert call_kwargs["startTime"] == 1700000000000 + + +class TestLogStreamerFilterMode: + def test_poll_once_with_filter_pattern(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.filter_log_events.return_value = { + "events": [ + {"eventId": "e1", "timestamp": 1700000000000, "message": "log line 1\n"}, + {"eventId": "e2", "timestamp": 1700000001000, "message": "log line 2\n"}, + ] + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Clusters/my-cluster/abc123", + job_name="my-hp-job", + sagemaker_session=session, + filter_pattern='"my-hp-job"', + start_time_ms=1700000000000, + ) + events = streamer.poll_once() + + assert len(events) == 2 + assert events[0] == (1700000000000, "log line 1") + assert events[1] == (1700000001000, "log line 2") + + # Verify filter_log_events was called with correct params + call_kwargs = logs_client.filter_log_events.call_args[1] + assert call_kwargs["logGroupName"] == "/aws/sagemaker/Clusters/my-cluster/abc123" + assert call_kwargs["filterPattern"] == '"my-hp-job"' + assert call_kwargs["startTime"] == 1700000000000 + assert call_kwargs["logStreamNamePrefix"] == "SagemakerHyperPodTrainingJob" + + def test_dedup_within_same_millisecond(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + # First poll + logs_client.filter_log_events.return_value = { + "events": [ + {"eventId": "e1", "timestamp": 1700000000000, "message": "line 1\n"}, + ] + } + streamer = LogStreamer( + log_group="/aws/sagemaker/Clusters/c/id", + job_name="job", + sagemaker_session=session, + filter_pattern='"job"', + ) + events = streamer.poll_once() + assert len(events) == 1 + + # Second poll returns same event again + logs_client.filter_log_events.return_value = { + "events": [ + {"eventId": "e1", "timestamp": 1700000000000, "message": "line 1\n"}, + {"eventId": "e2", "timestamp": 1700000000000, "message": "line 2\n"}, + ] + } + events = streamer.poll_once() + # e1 should be deduped, only e2 is new + assert len(events) == 1 + assert events[0][1] == "line 2" + + def test_dedup_clears_on_timestamp_advance(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + streamer = LogStreamer( + log_group="/aws/sagemaker/Clusters/c/id", + job_name="job", + sagemaker_session=session, + filter_pattern='"job"', + ) + + # First poll + logs_client.filter_log_events.return_value = { + "events": [ + {"eventId": "e1", "timestamp": 1000, "message": "a\n"}, + ] + } + streamer.poll_once() + + # Second poll: new timestamp + logs_client.filter_log_events.return_value = { + "events": [ + {"eventId": "e2", "timestamp": 2000, "message": "b\n"}, + ] + } + events = streamer.poll_once() + assert len(events) == 1 + assert events[0] == (2000, "b") + + +class TestStreamLogLoop: + """Tests for the shared stream_log_loop function.""" + + def test_early_exit_when_already_terminal(self): + """stream_log_loop returns immediately if status_fn reports terminal.""" + + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "job/algo-1"}]} + ] + logs_client.get_log_events.return_value = { + "events": [{"timestamp": 1000, "message": "done\n"}], + "nextForwardToken": "t1", + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="job", + sagemaker_session=session, + ) + + call_count = [0] + + def _status(): + call_count[0] += 1 + return "Completed" + + stream_log_loop(streamer, poll=1, status_fn=_status) + # status_fn called exactly once (the upfront check) + assert call_count[0] == 1 + + def test_exits_when_status_becomes_terminal(self): + """stream_log_loop exits after status transitions to terminal.""" + + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "job/algo-1"}]} + ] + # Return events then empty (caught up) + logs_client.get_log_events.side_effect = [ + {"events": [{"timestamp": 1000, "message": "training\n"}], "nextForwardToken": "t1"}, + {"events": [], "nextForwardToken": "t1"}, + ] + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="job", + sagemaker_session=session, + ) + + statuses = iter(["Training", "Completed"]) + + with patch("time.sleep"): + stream_log_loop(streamer, poll=1, status_fn=lambda: next(statuses)) + + def test_empty_cycles_feedback(self): + """stream_log_loop logs 'No log events yet' after ~30s of empty polls.""" + + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "job/algo-1"}]} + ] + logs_client.get_log_events.return_value = { + "events": [], "nextForwardToken": "t1", + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="job", + sagemaker_session=session, + ) + + call_count = [0] + + def _status(): + call_count[0] += 1 + # With poll=5, warn_cycle=6. Become terminal after 8 calls. + return "Completed" if call_count[0] >= 8 else "Training" + + with patch("time.sleep"): + with patch("sagemaker.train.common_utils.log_streamer.logger") as mock_logger: + stream_log_loop(streamer, poll=5, status_fn=_status) + info_calls = [str(c) for c in mock_logger.info.call_args_list] + assert any("No log events yet" in c for c in info_calls) + + def test_resource_not_found_with_running_job_retries(self): + """stream_log_loop retries on ResourceNotFoundException while job runs.""" + + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + # First call raises ResourceNotFound, second returns events + logs_client.get_paginator.return_value.paginate.side_effect = [ + _make_client_error("ResourceNotFoundException"), + [{"logStreams": [{"logStreamName": "job/algo-1"}]}], + ] + logs_client.get_log_events.return_value = { + "events": [{"timestamp": 1000, "message": "hi\n"}], + "nextForwardToken": "t1", + } + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="job", + sagemaker_session=session, + ) + + statuses = iter(["Training", "Training", "Completed"]) + + with patch("time.sleep"): + stream_log_loop(streamer, poll=1, status_fn=lambda: next(statuses)) + + def test_access_denied_propagates(self): + """stream_log_loop raises AccessDeniedException.""" + + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.side_effect = _make_client_error( + "AccessDeniedException" + ) + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="job", + sagemaker_session=session, + ) + + with pytest.raises(ClientError): + stream_log_loop(streamer, poll=1, status_fn=lambda: "Training") + + +class TestLogStreamerErrorHandling: + def test_resource_not_found_propagates(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.get_paginator.return_value.paginate.side_effect = _make_client_error( + "ResourceNotFoundException" + ) + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + ) + with pytest.raises(ClientError) as exc_info: + streamer.poll_once() + assert exc_info.value.response["Error"]["Code"] == "ResourceNotFoundException" + + def test_access_denied_propagates(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.get_paginator.return_value.paginate.side_effect = _make_client_error( + "AccessDeniedException" + ) + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + ) + with pytest.raises(ClientError) as exc_info: + streamer.poll_once() + assert exc_info.value.response["Error"]["Code"] == "AccessDeniedException" + + def test_throttling_returns_empty(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.get_paginator.return_value.paginate.side_effect = _make_client_error( + "ThrottlingException" + ) + + streamer = LogStreamer( + log_group="/aws/sagemaker/Job/AgentRFT", + job_name="my-job", + sagemaker_session=session, + ) + events = streamer.poll_once() + assert events == [] + + def test_filter_mode_access_denied_propagates(self): + session = _make_mock_session() + logs_client = session.boto_session.client.return_value + + logs_client.filter_log_events.side_effect = _make_client_error( + "AccessDeniedException" + ) + + streamer = LogStreamer( + log_group="/aws/sagemaker/Clusters/c/id", + job_name="job", + sagemaker_session=session, + filter_pattern='"job"', + ) + with pytest.raises(ClientError): + streamer.poll_once() diff --git a/sagemaker-train/tests/unit/train/test_stream_logs.py b/sagemaker-train/tests/unit/train/test_stream_logs.py new file mode 100644 index 0000000000..8015d0c073 --- /dev/null +++ b/sagemaker-train/tests/unit/train/test_stream_logs.py @@ -0,0 +1,204 @@ +"""Unit tests for stream_logs() on AgentRFTJob, MultiTurnRLTrainer, and evaluators.""" +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest +from botocore.exceptions import ClientError + +from sagemaker.train.agent_rft_job import AgentRFTJob +from sagemaker.train.multi_turn_rl_trainer import MultiTurnRLTrainer +from sagemaker.train.evaluate.base_evaluator import BaseEvaluator + + +def _make_client_error(code, message="error"): + return ClientError( + {"Error": {"Code": code, "Message": message}}, + "operation_name", + ) + + +def _make_mock_job(**overrides): + job = MagicMock() + job.job_name = "test-mtrl-job" + job.job_arn = "arn:aws:sagemaker:us-west-2:123456789012:job/test-mtrl-job" + job.job_status = "Training" + job.job_category = "AgentRFT" + job.job_config_document = "{}" + for k, v in overrides.items(): + setattr(job, k, v) + return job + + +class TestAgentRFTJobStreamLogs: + @patch("sagemaker.train.defaults.TrainDefaults.get_sagemaker_session") + def test_stream_logs_exits_on_completed(self, mock_get_session): + """stream_logs exits when job status reaches Completed.""" + mock_session = MagicMock() + mock_session.boto_session.region_name = "us-west-2" + mock_get_session.return_value = mock_session + + logs_client = mock_session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "test-mtrl-job/algo-1"}]} + ] + logs_client.get_log_events.return_value = { + "events": [{"timestamp": 1700000000000, "message": "done\n"}], + "nextForwardToken": "token-1", + } + + mock_job = _make_mock_job() + # Job becomes Completed after refresh + mock_job.refresh.side_effect = lambda: setattr(mock_job, "job_status", "Completed") + + rft_job = AgentRFTJob(mock_job) + + with patch("time.sleep"): + rft_job.stream_logs(poll=1) + + # Verify it exited (didn't hang) + assert mock_job.refresh.called + + @patch("sagemaker.train.defaults.TrainDefaults.get_sagemaker_session") + def test_stream_logs_raises_on_access_denied(self, mock_get_session): + """stream_logs propagates AccessDeniedException.""" + mock_session = MagicMock() + mock_session.boto_session.region_name = "us-west-2" + mock_get_session.return_value = mock_session + + logs_client = mock_session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.side_effect = _make_client_error( + "AccessDeniedException" + ) + + mock_job = _make_mock_job() + rft_job = AgentRFTJob(mock_job) + + with pytest.raises(ClientError) as exc_info: + rft_job.stream_logs(poll=1) + assert exc_info.value.response["Error"]["Code"] == "AccessDeniedException" + + @patch("sagemaker.train.defaults.TrainDefaults.get_sagemaker_session") + def test_stream_logs_handles_resource_not_found_terminal(self, mock_get_session): + """stream_logs returns cleanly when log group not found and job is terminal.""" + mock_session = MagicMock() + mock_session.boto_session.region_name = "us-west-2" + mock_get_session.return_value = mock_session + + logs_client = mock_session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.side_effect = _make_client_error( + "ResourceNotFoundException" + ) + + mock_job = _make_mock_job(job_status="Failed") + rft_job = AgentRFTJob(mock_job) + + with patch("time.sleep"): + rft_job.stream_logs(poll=1) + + def test_stream_logs_validates_poll(self): + """stream_logs raises ValueError for invalid poll.""" + mock_job = _make_mock_job() + rft_job = AgentRFTJob(mock_job) + + with pytest.raises(ValueError, match="poll must be between 1 and 300"): + rft_job.stream_logs(poll=0) + + with pytest.raises(ValueError, match="poll must be between 1 and 300"): + rft_job.stream_logs(poll=500) + + def test_stream_logs_validates_start_time_type(self): + """stream_logs raises TypeError for invalid start_time.""" + mock_job = _make_mock_job() + rft_job = AgentRFTJob(mock_job) + + with pytest.raises(TypeError, match="start_time must be datetime or int"): + rft_job.stream_logs(start_time="2023-01-01") + + +class TestMultiTurnRLTrainerStreamLogs: + @patch("sagemaker.train.defaults.TrainDefaults.get_sagemaker_session") + @patch("sagemaker.core.resources.Job.get") + def test_stream_logs_uses_correct_log_group(self, mock_job_get, mock_get_session): + """MultiTurnRLTrainer.stream_logs uses /aws/sagemaker/Job/AgentRFT.""" + mock_session = MagicMock() + mock_session.boto_session.region_name = "us-west-2" + mock_get_session.return_value = mock_session + + logs_client = mock_session.boto_session.client.return_value + logs_client.get_paginator.return_value.paginate.return_value = [ + {"logStreams": [{"logStreamName": "my-job/algo-1"}]} + ] + logs_client.get_log_events.return_value = { + "events": [], + "nextForwardToken": "token-1", + } + + # Job.get returns terminal status + mock_job_obj = MagicMock() + mock_job_obj.job_status = "Completed" + mock_job_get.return_value = mock_job_obj + + # Create trainer with _latest_job set + trainer = MagicMock(spec=MultiTurnRLTrainer) + trainer._latest_job = MagicMock() + trainer._latest_job.job_name = "my-job" + trainer.sagemaker_session = None + + # Call the actual method + with patch("time.sleep"): + MultiTurnRLTrainer.stream_logs(trainer, poll=1) + + # Verify Job.get was called with correct category + mock_job_get.assert_called_with(job_name="my-job", job_category="AgentRFT") + + def test_stream_logs_raises_when_no_job(self): + """MultiTurnRLTrainer.stream_logs raises ValueError if no job exists.""" + trainer = MagicMock(spec=MultiTurnRLTrainer) + trainer._latest_job = None + + with pytest.raises(ValueError, match="No training job found"): + MultiTurnRLTrainer.stream_logs(trainer) + + +class TestBaseEvaluatorStreamLogs: + def test_stream_logs_raises_when_no_execution(self): + """BaseEvaluator.stream_logs raises ValueError if no evaluation executed.""" + evaluator = MagicMock(spec=BaseEvaluator) + evaluator._latest_execution = None + + with pytest.raises(ValueError, match="No evaluation executed yet"): + BaseEvaluator.stream_logs(evaluator) + + def test_stream_logs_validates_poll(self): + """BaseEvaluator.stream_logs validates poll parameter.""" + evaluator = MagicMock(spec=BaseEvaluator) + evaluator._latest_execution = MagicMock() + + with pytest.raises(ValueError, match="poll must be between 1 and 300"): + BaseEvaluator.stream_logs(evaluator, poll=0) + + def test_log_group_for_training_job_arn(self): + """_log_group_for_step_arn resolves training-job ARNs correctly.""" + arn = "arn:aws:sagemaker:us-west-2:123456789012:training-job/my-eval-job" + assert BaseEvaluator._log_group_for_step_arn(arn) == "/aws/sagemaker/TrainingJobs" + + def test_log_group_for_job_arn(self): + """_log_group_for_step_arn resolves Job API ARNs correctly.""" + arn = "arn:aws:sagemaker:us-west-2:123456789012:job/my-eval-job" + assert BaseEvaluator._log_group_for_step_arn(arn) == "/aws/sagemaker/Job/AgentRFTEvaluation" + + def test_job_name_from_training_job_arn(self): + """_job_name_from_arn extracts job name from training-job ARN.""" + arn = "arn:aws:sagemaker:us-west-2:123456789012:training-job/my-eval-job-abc123" + assert BaseEvaluator._job_name_from_arn(arn) == "my-eval-job-abc123" + + def test_job_name_from_job_arn(self): + """_job_name_from_arn extracts job name from Job API ARN.""" + arn = "arn:aws:sagemaker:us-west-2:123456789012:job/my-mtrl-eval-job" + assert BaseEvaluator._job_name_from_arn(arn) == "my-mtrl-eval-job" + + def test_job_name_from_unknown_arn_returns_none(self): + """_job_name_from_arn returns None for unrecognized ARN format.""" + arn = "arn:aws:sagemaker:us-west-2:123456789012:pipeline/my-pipeline" + assert BaseEvaluator._job_name_from_arn(arn) is None diff --git a/v3-examples/model-customization-examples/benchmark_demo.ipynb b/v3-examples/model-customization-examples/benchmark_demo.ipynb index 1a48a3ff07..2e6f14e359 100644 --- a/v3-examples/model-customization-examples/benchmark_demo.ipynb +++ b/v3-examples/model-customization-examples/benchmark_demo.ipynb @@ -222,6 +222,9 @@ "source": [ "# Run evaluation with configured parameters\n", "execution = evaluator.evaluate()\n", + "\n", + "# Stream evaluation logs\n", + "# evaluator.stream_logs()\n", "pprint(execution)\n", "\n", "print(f\"\\nPipeline Execution ARN: {execution.arn}\")\n", diff --git a/v3-examples/model-customization-examples/custom_scorer_demo.ipynb b/v3-examples/model-customization-examples/custom_scorer_demo.ipynb index a930d17961..23ca69f296 100644 --- a/v3-examples/model-customization-examples/custom_scorer_demo.ipynb +++ b/v3-examples/model-customization-examples/custom_scorer_demo.ipynb @@ -212,6 +212,9 @@ "# Start evaluation\n", "execution = evaluator.evaluate()\n", "\n", + "# Stream evaluation logs\n", + "# evaluator.stream_logs()\n", + "\n", "print(\"\\n✓ Evaluation execution started successfully!\")\n", "print(f\" Execution Name: {execution.name}\")\n", "print(f\" Pipeline Execution ARN: {execution.arn}\")\n", diff --git a/v3-examples/model-customization-examples/inspect_ai_evaluation_demo.ipynb b/v3-examples/model-customization-examples/inspect_ai_evaluation_demo.ipynb index af94fe0bc8..0ce5ffbbfe 100644 --- a/v3-examples/model-customization-examples/inspect_ai_evaluation_demo.ipynb +++ b/v3-examples/model-customization-examples/inspect_ai_evaluation_demo.ipynb @@ -275,6 +275,9 @@ "# Start evaluation\n", "execution = evaluator.evaluate()\n", "\n", + "# Stream evaluation logs\n", + "# evaluator.stream_logs()\n", + "\n", "print(f\"\\n\u2713 Evaluation started!\")\n", "print(f\" Execution ARN: {execution.arn}\")\n", "print(f\" Status: {execution.status.overall_status}\")" diff --git a/v3-examples/model-customization-examples/llm_as_judge_demo.ipynb b/v3-examples/model-customization-examples/llm_as_judge_demo.ipynb index 3e4a28c4ba..ad21827f8e 100644 --- a/v3-examples/model-customization-examples/llm_as_judge_demo.ipynb +++ b/v3-examples/model-customization-examples/llm_as_judge_demo.ipynb @@ -320,6 +320,9 @@ "# Run evaluation\n", "execution = evaluator.evaluate()\n", "\n", + "# Stream evaluation logs\n", + "# evaluator.stream_logs()\n", + "\n", "print(f\"\u2705 Evaluation job started!\")\n", "print(f\"Job ARN: {execution.arn}\")\n", "print(f\"Job Name: {execution.name}\")\n", diff --git a/v3-examples/model-customization-examples/mtrl_finetuning_example_notebook_v3_prod.ipynb b/v3-examples/model-customization-examples/mtrl_finetuning_example_notebook_v3_prod.ipynb index 564fade811..67debc9b5a 100644 --- a/v3-examples/model-customization-examples/mtrl_finetuning_example_notebook_v3_prod.ipynb +++ b/v3-examples/model-customization-examples/mtrl_finetuning_example_notebook_v3_prod.ipynb @@ -476,7 +476,8 @@ "job = trainer.train(wait=True)\n", "\n", "# Or launch without waiting\n", - "# job = trainer.train(wait=False)" + "# job = trainer.train(wait=False)", + "# job.stream_logs()\n" ] }, { @@ -545,7 +546,10 @@ "\n", "# Wait for completion if still running\n", "if existing_job.job_status == \"InProgress\":\n", - " existing_job.wait()" + " existing_job.wait()", + "\n", + "# Stream logs from an in-progress or completed job\n", + "# existing_job.stream_logs()\n" ] }, { @@ -641,7 +645,10 @@ "execution_count": null, "source": [ "evaluation.wait()\n", - "print(f\"Evaluation status: {evaluation.status.overall_status}\")" + "print(f\"Evaluation status: {evaluation.status.overall_status}\")", + "\n", + "# Stream evaluation logs\n", + "# evaluator.stream_logs()\n" ] }, {