Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions sagemaker-train/src/sagemaker/train/agent_rft_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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()
Expand Down
85 changes: 48 additions & 37 deletions sagemaker-train/src/sagemaker/train/base_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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."""
Expand Down Expand Up @@ -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 = {
Expand All @@ -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:
Expand All @@ -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:
Expand Down
Loading