How Marin Handles Training Loss Spikes and Health Alerts

Marin detects training loss spikes by comparing 5-minute recent windows against 60-minute baselines in the Levanter telemetry table, triggering health alerts when recent loss floors exceed statistically derived thresholds defined in infra/grafana/src/loss_spikes.py.

The marin-community/marin repository implements a statistically-grounded monitoring system that surfaces training loss spikes and health alerts for large-scale machine learning jobs. By analyzing telemetry data from active hero runs, Marin computes dynamic thresholds to distinguish between normal variance and genuine loss divergence while preventing false positives during warmup.

Time-Window Strategy for Statistical Comparison

Marin's detection algorithm relies on two distinct time windows to establish context and detect deviations:

  • Baseline Window: 60 minutes (_BASELINE_LOOKUP) capturing average loss, standard deviation, and minimum loss (the "floor") over extended history
  • Recent Window: 5 minutes (_RECENT_WINDOW) measuring the same statistics over the most recent period

This dual-window approach provides the statistical grounding necessary to differentiate between transient noise and persistent spikes. The system requires a minimum of 20 baseline samples (_MIN_BASELINE_SAMPLES) and 5 recent samples (_MIN_RECENT_SAMPLES) before rendering a definitive health verdict.

Core Detection Functions in loss_spikes.py

The implementation resides in infra/grafana/src/loss_spikes.py, which provides three primary functions for the detection pipeline.

Building Queries with loss_window_query

The loss_window_query function (lines 55-89) constructs a single SQL query that retrieves aggregated metrics for both time windows simultaneously. This query operates against the Levanter telemetry table, fetching the required statistics for specified run IDs and optional executions in one database round-trip.

Processing Results with windows_by_run

Once the query executes, windows_by_run (lines 97-113) converts the PyArrow table result into a structured mapping from (cluster, run_id) tuples to LossWindows dataclass instances. This transformation normalizes the raw telemetry into a format suitable for comparative analysis.

Determining Status with loss_spike_reason

The loss_spike_reason function (lines 15-36) implements the classification logic. It evaluates four distinct conditions in sequence to determine the health status of a training run.

The Four Health Classification States

Marin categorizes every monitored run into one of four distinct states based on statistical analysis of the two time windows:

warming_up

The run has insufficient data—not enough baseline samples (fewer than 20) or recent samples (fewer than 5), or the baseline floor is missing. This state indicates the job is still stabilizing and suppresses false alerts during initialization.

not_finite

The recent loss or recent peak contains NaN or infinite values (indicated by _diverged flag). This state signals numerical divergence in the training process and immediately flags the run as abnormal.

spiking

The recent floor exceeds the baseline floor plus the maximum of either a 0.05 minimum rise (_MIN_RISE_) or six standard deviations (_SIGMA_FACTOR_ = 6.0). Mathematically: recent_floor > baseline + max(0.05, 6.0 × baseline_stddev). When this condition triggers, the system returns ("spiking", 1) and fires a health alert.

healthy

None of the above conditions apply. The training loss remains within statistically normal bounds, returning ("healthy", 0).

The 6.0 sigma factor combined with the 0.05 minimum rise threshold (line 35) creates a robust band that accommodates noisy training curves while catching genuine divergence.

Generating Grafana-Compatible Alerts

Once classification completes, loss_spike_alert_rows (lines 43-52) transforms the results into structured alert rows consumable by monitoring systems. Each row contains:

{
  "cluster": "<cluster>",
  "job": "<root_job>",
  "run": "<run_id>",
  "reason": "<warming_up|not_finite|spiking|healthy>",
  "value": <0|1>
}

These JSON objects feed directly into Grafana alerts, enabling operators to receive notifications when training jobs exhibit abnormal loss behavior. The binary value field (0 for healthy/warming, 1 for spiking/not_finite) simplifies alert threshold configuration in downstream monitoring tools.

Code Example: Querying and Alerting on Loss Spikes

The following example demonstrates the complete workflow from query construction to alert generation:

from datetime import datetime
from infra.grafana.src.loss_spikes import (
    loss_window_query,
    loss_spike_alert_rows,
)
import pyarrow as pa

# 1️⃣ Build the query for the current moment and a set of hero runs.

now = datetime.utcnow()
run_ids = ["run-123", "run-456"]
query = loss_window_query(now, runs=run_ids)

# 2️⃣ Execute the query against the telemetry backend (example only).

#    The execution returns a pyarrow.Table with the aggregated columns.

#    loss_windows = execute_sql(query)   # <-- your DB client here

# 3️⃣ Convert the results into alert rows.

#    Assume `loss_windows` is a pyarrow.Table from step 2.

alert_rows = loss_spike_alert_rows(runs=tuple_of_HeroRun_objects, loss_windows=loss_windows)

# 4️⃣ Send the rows to Grafana (or any monitoring system).

for row in alert_rows:
    print(row)   # {'cluster': 'fleet', 'job': 'train', 'run': 'run-123', 'reason': 'spiking', 'value': 1}

Summary

  • Marin employs a dual-window analysis (60-minute baseline vs. 5-minute recent) to detect training loss spikes in the Levanter telemetry table.
  • The detection logic requires minimum sample thresholds (20 baseline, 5 recent) before classification to prevent false positives during warmup.
  • Four distinct health states—warming_up, not_finite, spiking, and healthy—provide granular visibility into training stability.
  • The loss_spike_reason function applies a 6.0 sigma factor with a 0.05 minimum rise to statistically isolate genuine spikes from noise.
  • loss_spike_alert_rows generates standardized JSON output that integrates directly with Grafana for operational alerting.

Frequently Asked Questions

How does Marin avoid false positives from normal training noise?

Marin uses a conservative statistical threshold combining a 6.0 sigma factor (_SIGMA_FACTOR_) and a 0.05 minimum rise (_MIN_RISE_). This means a spike is only flagged if the recent loss floor exceeds the baseline by at least 0.05 or six standard deviations, whichever is larger. This dual-threshold approach accommodates naturally volatile training curves while catching genuine divergence, as implemented in infra/grafana/src/loss_spikes.py at line 35.

What happens when a training run first starts and lacks sufficient data?

The system returns a warming_up status when baseline samples fall below 20 or recent samples fall below 5, or when the baseline floor is missing. This suppresses alerts during the initialization phase, ensuring that transient startup behavior does not trigger false health alerts while the model stabilizes.

How does Marin handle numerical divergence or NaN values in loss calculations?

If the recent loss or recent peak contains NaN or infinite values (detected via the _diverged flag), the loss_spike_reason function immediately categorizes the run as not_finite. This state indicates critical numerical instability in the training process and triggers an alert regardless of the statistical window comparisons.

What data source does Marin query to monitor training loss?

Marin queries the Levanter telemetry table (referenced via LEVANTER_METRICS_TABLE in the hero runs module) for active hero runs. The loss_window_query function in infra/grafana/src/loss_spikes.py constructs SQL queries that aggregate this telemetry data into the baseline and recent windows used for statistical comparison.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →