From 0be3e6f5252e797d1e7b4ff1a69ff10b7fda5715 Mon Sep 17 00:00:00 2001 From: Shine1i Date: Tue, 17 Feb 2026 18:26:59 +0100 Subject: [PATCH] feat: integrate gradient norm tracking in training runtime and metrics - Enhanced chart logic to filter and visualize finite gradient norm values. --- studio/backend/core/training/training.py | 13 +++++ studio/backend/models/responses.py | 2 + studio/backend/models/training.py | 3 +- studio/backend/routes/training.py | 35 ++++++++++-- .../studio/sections/charts-content.tsx | 26 ++++++--- .../training/stores/training-runtime-store.ts | 53 ++++++++++++++----- .../src/features/training/types/runtime.ts | 4 ++ 7 files changed, 111 insertions(+), 25 deletions(-) diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 51d411be78..9b589d215d 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -4,6 +4,7 @@ Training backend for FastAPI integration import matplotlib.pyplot as plt from typing import Any, Generator, Tuple import logging +import math from .trainer import get_trainer, TrainingProgress from utils.hardware import clear_gpu_cache @@ -28,6 +29,8 @@ class TrainingBackend: self.loss_history = [] self.lr_history = [] self.step_history = [] + self.grad_norm_history = [] + self.grad_norm_step_history = [] self.eval_loss_history = [] self.eval_step_history = [] self.eval_enabled = False @@ -43,6 +46,14 @@ class TrainingBackend: self.loss_history.append(progress.loss) self.lr_history.append(progress.learning_rate) self.step_history.append(progress.step) + if progress.step >= 0 and progress.grad_norm is not None: + try: + grad_norm = float(progress.grad_norm) + except (TypeError, ValueError): + grad_norm = None + if grad_norm is not None and math.isfinite(grad_norm): + self.grad_norm_history.append(grad_norm) + self.grad_norm_step_history.append(progress.step) if progress.eval_loss is not None: self.eval_loss_history.append(progress.eval_loss) self.eval_step_history.append(progress.step) @@ -144,6 +155,8 @@ class TrainingBackend: self.loss_history = [] self.lr_history = [] self.step_history = [] + self.grad_norm_history = [] + self.grad_norm_step_history = [] self.eval_loss_history = [] self.eval_step_history = [] self.eval_enabled = False diff --git a/studio/backend/models/responses.py b/studio/backend/models/responses.py index 2aa798c5c9..c72dfc54dd 100644 --- a/studio/backend/models/responses.py +++ b/studio/backend/models/responses.py @@ -19,6 +19,8 @@ class TrainingMetricsResponse(BaseModel): loss_history: List[float] = Field(default_factory=list, description="Loss values per step") lr_history: List[float] = Field(default_factory=list, description="Learning rate per step") step_history: List[int] = Field(default_factory=list, description="Step numbers") + grad_norm_history: List[float] = Field(default_factory=list, description="Gradient norm values") + grad_norm_step_history: List[int] = Field(default_factory=list, description="Step numbers for gradient norm values") current_loss: Optional[float] = Field(None, description="Most recent loss value") current_lr: Optional[float] = Field(None, description="Most recent learning rate") current_step: Optional[int] = Field(None, description="Most recent step number") diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index fd04baf3a0..2b989e6a82 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -104,7 +104,7 @@ class TrainingStatus(BaseModel): metric_history: Optional[dict] = Field( None, description="Full metric history arrays for chart recovery after SSE reconnection. " - "Keys: 'steps', 'loss', 'lr' — each a list of numeric values.", + "Keys: 'steps', 'loss', 'lr', 'grad_norm', 'grad_norm_steps' — each a list of numeric values.", ) @@ -122,4 +122,3 @@ class TrainingProgress(BaseModel): grad_norm: Optional[float] = Field(None, description="L2 norm of gradients, computed before gradient clipping") num_tokens: Optional[int] = Field(None, description="Total number of tokens processed so far") eval_loss: Optional[float] = Field(None, description="Eval loss from the most recent evaluation step") - diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index c09b814fa1..1770999284 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -325,6 +325,8 @@ async def reset_training( backend.loss_history = [] backend.lr_history = [] backend.step_history = [] + backend.grad_norm_history = [] + backend.grad_norm_step_history = [] return {"status": "ok"} except Exception as e: logger.error(f"Error resetting training: {e}", exc_info=True) @@ -408,6 +410,8 @@ async def get_training_status( "steps": list(backend.step_history), "loss": list(backend.loss_history), "lr": list(backend.lr_history), + "grad_norm": list(getattr(backend, "grad_norm_history", [])), + "grad_norm_steps": list(getattr(backend, "grad_norm_step_history", [])), "eval_loss": list(backend.eval_loss_history), "eval_steps": list(backend.eval_step_history), } @@ -445,6 +449,8 @@ async def get_training_metrics( loss_history = backend.loss_history lr_history = backend.lr_history step_history = backend.step_history + grad_norm_history = getattr(backend, "grad_norm_history", []) + grad_norm_step_history = getattr(backend, "grad_norm_step_history", []) # Get current values current_loss = loss_history[-1] if loss_history else None @@ -455,6 +461,8 @@ async def get_training_metrics( loss_history=loss_history, lr_history=lr_history, step_history=step_history, + grad_norm_history=grad_norm_history, + grad_norm_step_history=grad_norm_step_history, current_loss=current_loss, current_lr=current_lr, current_step=current_step, @@ -505,6 +513,8 @@ async def stream_training_progress( total_steps: int, epoch: Optional[float] = None, progress: Optional[Any] = None, + grad_norm_override: Optional[float] = None, + eval_loss_override: Optional[float] = None, ) -> TrainingProgress: total = max(total_steps, 0) if step < 0 or total == 0: @@ -517,9 +527,13 @@ async def stream_training_progress( # Get actual values from progress object if available elapsed_seconds = getattr(progress, 'elapsed_seconds', None) if progress else None eta_seconds = getattr(progress, 'eta_seconds', None) if progress else None - grad_norm = getattr(progress, 'grad_norm', None) if progress else None + grad_norm = grad_norm_override + if grad_norm is None and progress: + grad_norm = getattr(progress, 'grad_norm', None) num_tokens = getattr(progress, 'num_tokens', None) if progress else None - eval_loss = getattr(progress, 'eval_loss', None) if progress else None + eval_loss = eval_loss_override + if eval_loss is None and progress: + eval_loss = getattr(progress, 'eval_loss', None) return TrainingProgress( job_id=job_id, @@ -558,6 +572,13 @@ async def stream_training_progress( # ── Replay missed steps on reconnect ───────────────────── if resume_from_step is not None and backend.step_history: replayed = 0 + grad_norm_by_step = { + step_val: grad_val + for step_val, grad_val in zip( + getattr(backend, "grad_norm_step_history", []), + getattr(backend, "grad_norm_history", []), + ) + } for i, step_val in enumerate(backend.step_history): if step_val > resume_from_step: loss_val = backend.loss_history[i] if i < len(backend.loss_history) else 0.0 @@ -567,7 +588,15 @@ async def stream_training_progress( ) total_replay = getattr(tp_replay, "total_steps", step_val) if tp_replay else step_val epoch_replay = getattr(tp_replay, "epoch", None) if tp_replay else None - payload = build_progress(step_val, loss_val, lr_val, total_replay, epoch_replay, progress=tp_replay) + payload = build_progress( + step_val, + loss_val, + lr_val, + total_replay, + epoch_replay, + progress=tp_replay, + grad_norm_override=grad_norm_by_step.get(step_val), + ) yield format_sse(payload.model_dump_json(), event="progress", event_id=step_val) replayed += 1 if replayed: diff --git a/studio/frontend/src/features/studio/sections/charts-content.tsx b/studio/frontend/src/features/studio/sections/charts-content.tsx index 61cbb0a3ea..bb7a7a131c 100644 --- a/studio/frontend/src/features/studio/sections/charts-content.tsx +++ b/studio/frontend/src/features/studio/sections/charts-content.tsx @@ -117,12 +117,13 @@ function buildStepTicks(min: number, max: number, targetCount = 6): number[] { } function buildYDomain(values: number[]): [number, number] { - if (values.length === 0) { + const finiteValues = values.filter((value) => Number.isFinite(value)); + if (finiteValues.length === 0) { return [0, 1]; } - const min = Math.min(...values); - const max = Math.max(...values); + const min = Math.min(...finiteValues); + const max = Math.max(...finiteValues); if (min === max) { const base = Math.abs(min); @@ -240,7 +241,8 @@ export function ChartsContent({ (point) => point.step >= visibleStepDomain[0] && point.step <= visibleStepDomain[1], ) - .map((point) => point.gradNorm), + .map((point) => point.gradNorm) + .filter((value) => Number.isFinite(value)), [reducedGradNormData, visibleStepDomain], ); @@ -251,7 +253,8 @@ export function ChartsContent({ (point) => point.step >= visibleStepDomain[0] && point.step <= visibleStepDomain[1], ) - .map((point) => point.lr), + .map((point) => point.lr) + .filter((value) => Number.isFinite(value)), [reducedLrData, visibleStepDomain], ); @@ -345,6 +348,7 @@ export function ChartsContent({ @@ -444,6 +448,7 @@ export function ChartsContent({ @@ -509,6 +514,7 @@ export function ChartsContent({ @@ -536,7 +542,10 @@ export function ChartsContent({ tickMargin={4} fontSize={10} width={52} - tickFormatter={(value) => Number(value).toExponential(0)} + tickFormatter={(value) => { + const num = Number(value); + return Number.isFinite(num) ? num.toExponential(0) : "0e+0"; + }} /> `Step ${payload?.[0]?.payload?.step ?? ""}` } - formatter={(value) => [Number(value).toExponential(3), "LR"]} + formatter={(value) => { + const num = Number(value); + return [Number.isFinite(num) ? num.toExponential(3) : "0e+0", "LR"]; + }} /> } /> diff --git a/studio/frontend/src/features/training/stores/training-runtime-store.ts b/studio/frontend/src/features/training/stores/training-runtime-store.ts index fd0d3d5733..6a3756a8ca 100644 --- a/studio/frontend/src/features/training/stores/training-runtime-store.ts +++ b/studio/frontend/src/features/training/stores/training-runtime-store.ts @@ -56,6 +56,11 @@ function toSeries(steps: number[], values: number[]): TrainingSeriesPoint[] { return sortSeries(points); } +function toFiniteNumber(value: unknown): number | null { + if (typeof value !== "number") return null; + return Number.isFinite(value) ? value : null; +} + function upsertPoint( points: TrainingSeriesPoint[], step: number, @@ -74,22 +79,32 @@ function upsertPoint( function applyMetricHistoryFromStatus(payload: TrainingStatusResponse): { lossHistory: TrainingSeriesPoint[] | null; lrHistory: TrainingSeriesPoint[] | null; + gradNormHistory: TrainingSeriesPoint[] | null; evalLossHistory: TrainingSeriesPoint[] | null; } { const history = payload.metric_history; if (!history || !history.steps?.length) { - return { lossHistory: null, lrHistory: null, evalLossHistory: null }; + return { + lossHistory: null, + lrHistory: null, + gradNormHistory: null, + evalLossHistory: null, + }; } const steps = history.steps; const lossHistory = history.loss ? toSeries(steps, history.loss) : null; const lrHistory = history.lr ? toSeries(steps, history.lr) : null; + const gradNormHistory = + history.grad_norm && history.grad_norm_steps + ? toSeries(history.grad_norm_steps, history.grad_norm) + : null; const evalLossHistory = history.eval_loss && history.eval_steps ? toSeries(history.eval_steps, history.eval_loss) : null; - return { lossHistory, lrHistory, evalLossHistory }; + return { lossHistory, lrHistory, gradNormHistory, evalLossHistory }; } export const useTrainingRuntimeStore = create()((set) => ({ @@ -163,6 +178,7 @@ export const useTrainingRuntimeStore = create()((set) => ( typeof detailEpoch === "number" ? detailEpoch : state.currentEpoch, lossHistory: metricHistory.lossHistory ?? state.lossHistory, lrHistory: metricHistory.lrHistory ?? state.lrHistory, + gradNormHistory: metricHistory.gradNormHistory ?? state.gradNormHistory, evalLossHistory: metricHistory.evalLossHistory ?? state.evalLossHistory, }; }), @@ -171,6 +187,10 @@ export const useTrainingRuntimeStore = create()((set) => ( set((state) => { const lossHistory = toSeries(payload.step_history, payload.loss_history); const lrHistory = toSeries(payload.step_history, payload.lr_history); + const gradNormHistory = toSeries( + payload.grad_norm_step_history, + payload.grad_norm_history, + ); const latestStep = payload.current_step ?? (payload.step_history.length > 0 @@ -181,6 +201,8 @@ export const useTrainingRuntimeStore = create()((set) => ( ...state, lossHistory: lossHistory.length > 0 ? lossHistory : state.lossHistory, lrHistory: lrHistory.length > 0 ? lrHistory : state.lrHistory, + gradNormHistory: + gradNormHistory.length > 0 ? gradNormHistory : state.gradNormHistory, currentStep: typeof latestStep === "number" ? Math.max(latestStep, state.currentStep) @@ -199,36 +221,41 @@ export const useTrainingRuntimeStore = create()((set) => ( applyProgress: (payload: TrainingProgressPayload, eventId?: number) => set((state) => { const step = Math.max(payload.step, 0); + const currentLoss = toFiniteNumber(payload.loss); + const currentLearningRate = toFiniteNumber(payload.learning_rate); + const currentGradNorm = toFiniteNumber(payload.grad_norm); + const evalLoss = toFiniteNumber(payload.eval_loss); + return { ...state, jobId: payload.job_id || state.jobId, currentStep: step, totalSteps: Math.max(payload.total_steps, state.totalSteps), - currentLoss: payload.loss, - currentLearningRate: payload.learning_rate, + currentLoss: currentLoss ?? state.currentLoss, + currentLearningRate: currentLearningRate ?? state.currentLearningRate, progressPercent: payload.progress_percent, currentEpoch: payload.epoch ?? state.currentEpoch, elapsedSeconds: payload.elapsed_seconds, etaSeconds: payload.eta_seconds, - currentGradNorm: payload.grad_norm, + currentGradNorm, currentNumTokens: payload.num_tokens, firstStepReceived: state.firstStepReceived || step > 0, lastEventId: typeof eventId === "number" ? eventId : state.lastEventId, lossHistory: - step > 0 - ? upsertPoint(state.lossHistory, step, payload.loss) + step > 0 && currentLoss !== null + ? upsertPoint(state.lossHistory, step, currentLoss) : state.lossHistory, lrHistory: - step > 0 - ? upsertPoint(state.lrHistory, step, payload.learning_rate) + step > 0 && currentLearningRate !== null + ? upsertPoint(state.lrHistory, step, currentLearningRate) : state.lrHistory, gradNormHistory: - step > 0 && typeof payload.grad_norm === "number" - ? upsertPoint(state.gradNormHistory, step, payload.grad_norm) + step > 0 && currentGradNorm !== null + ? upsertPoint(state.gradNormHistory, step, currentGradNorm) : state.gradNormHistory, evalLossHistory: - step > 0 && typeof payload.eval_loss === "number" - ? upsertPoint(state.evalLossHistory, step, payload.eval_loss) + step > 0 && evalLoss !== null + ? upsertPoint(state.evalLossHistory, step, evalLoss) : state.evalLossHistory, }; }), diff --git a/studio/frontend/src/features/training/types/runtime.ts b/studio/frontend/src/features/training/types/runtime.ts index 8e0b43bd52..df24c020ca 100644 --- a/studio/frontend/src/features/training/types/runtime.ts +++ b/studio/frontend/src/features/training/types/runtime.ts @@ -26,6 +26,8 @@ export interface TrainingStatusResponse { steps?: number[]; loss?: number[]; lr?: number[]; + grad_norm?: number[]; + grad_norm_steps?: number[]; eval_loss?: number[]; eval_steps?: number[]; } | null; @@ -35,6 +37,8 @@ export interface TrainingMetricsResponse { loss_history: number[]; lr_history: number[]; step_history: number[]; + grad_norm_history: number[]; + grad_norm_step_history: number[]; current_loss: number | null; current_lr: number | null; current_step: number | null;