From 0fbdd743a07c3cd287339fac28cd24f2b98b6435 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Fri, 3 Jul 2026 11:10:40 +0000 Subject: [PATCH] Report grad norm from the trainers and chart it instead of LR; celebrate completion in the run header --- .../core/training/diffusion_dit_trainer.py | 6 ++- .../core/training/diffusion_lora_trainer.py | 7 ++- .../training/diffusion_training_service.py | 47 ++++++++++++++----- studio/backend/models/training.py | 9 +++- studio/backend/routes/training.py | 1 + studio/frontend/src/features/images/api.ts | 3 ++ .../images/train/diffusion-charts.tsx | 47 ++++++++++--------- .../images/train/diffusion-train-panel.tsx | 12 +++-- 8 files changed, 89 insertions(+), 43 deletions(-) diff --git a/studio/backend/core/training/diffusion_dit_trainer.py b/studio/backend/core/training/diffusion_dit_trainer.py index 5221425e4b..5776f87652 100644 --- a/studio/backend/core/training/diffusion_dit_trainer.py +++ b/studio/backend/core/training/diffusion_dit_trainer.py @@ -1157,8 +1157,11 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto (loss / cfg.gradient_accumulation_steps).backward() step_loss += float(loss.detach()) / cfg.gradient_accumulation_steps + grad_norm = None if cfg.max_grad_norm and cfg.max_grad_norm > 0: - torch.nn.utils.clip_grad_norm_(lora_params, cfg.max_grad_norm) + # clip_grad_norm_ returns the total PRE-clip norm: the health signal the UI + # charts (an exploding norm shows up here even while the clip caps the update). + grad_norm = float(torch.nn.utils.clip_grad_norm_(lora_params, cfg.max_grad_norm)) optimizer.step() lr_sched.step() @@ -1185,6 +1188,7 @@ def _train_dit(cfg, spec, pairs, rng, device, weight_dtype, on_event, _check_sto loss = round(step_loss, 5), avg_loss = round(running_loss / done, 5), learning_rate = lr_sched.get_last_lr()[0], + grad_norm = round(grad_norm, 5) if grad_norm is not None else None, samples_per_second = sps, peak_memory_gb = peak_gb or None, ) diff --git a/studio/backend/core/training/diffusion_lora_trainer.py b/studio/backend/core/training/diffusion_lora_trainer.py index 7c7ca79d51..10b6cd89cd 100644 --- a/studio/backend/core/training/diffusion_lora_trainer.py +++ b/studio/backend/core/training/diffusion_lora_trainer.py @@ -509,8 +509,12 @@ def run_diffusion_lora_training( # max_grad_norm <= 0 means "disable clipping" (the Studio payload sends 0.0 for that); # passing 0.0 to clip_grad_norm_ would scale every gradient to zero (no learning). + grad_norm = None if cfg.max_grad_norm and cfg.max_grad_norm > 0: - torch.nn.utils.clip_grad_norm_(lora_params, cfg.max_grad_norm) + # The returned value is the total PRE-clip norm, reported to the UI chart. + grad_norm = float( + torch.nn.utils.clip_grad_norm_(lora_params, cfg.max_grad_norm) + ) optimizer.step() lr_sched.step() @@ -535,6 +539,7 @@ def run_diffusion_lora_training( loss = round(step_loss, 5), avg_loss = round(running_loss / done, 5), learning_rate = lr_sched.get_last_lr()[0], + grad_norm = round(grad_norm, 5) if grad_norm is not None else None, samples_per_second = samples_per_second, peak_memory_gb = peak_gb or None, ) diff --git a/studio/backend/core/training/diffusion_training_service.py b/studio/backend/core/training/diffusion_training_service.py index fa44cceabe..2251038b78 100644 --- a/studio/backend/core/training/diffusion_training_service.py +++ b/studio/backend/core/training/diffusion_training_service.py @@ -67,6 +67,7 @@ def _idle_state() -> dict[str, Any]: "loss": None, "avg_loss": None, "learning_rate": None, + "grad_norm": None, "num_images": None, "in_model_load": False, "output_dir": None, @@ -82,16 +83,21 @@ def _idle_state() -> dict[str, Any]: "metric_steps": [], "metric_loss": [], "metric_lr": [], + "metric_grad_norm": [], } -def _append_metric(state: dict[str, Any], step: Any, loss: Any, lr: Any) -> None: - """Append one (step, loss, lr) point to the bounded history arrays on ``state``. +def _append_metric( + state: dict[str, Any], step: Any, loss: Any, lr: Any, grad_norm: Any = None +) -> None: + """Append one (step, loss, lr, grad_norm) point to the bounded history arrays on + ``state``. Only records finite, positive-step points (mirrors the LLM trainer, which logs history only for step > 0 with a real loss). When the arrays hit ``_METRIC_CAP`` they are decimated in place (keep every other point) so appends stay bounded without losing the - curve's shape. lr may be None (kept as None so the LR series can be sparse).""" + curve's shape. lr / grad_norm may be None (kept as None so those series can be sparse + while staying index-aligned with ``steps``).""" try: istep = int(step) except (TypeError, ValueError): @@ -104,22 +110,34 @@ def _append_metric(state: dict[str, Any], step: Any, loss: Any, lr: Any) -> None return if floss != floss: # NaN guard return - flr: Optional[float] - try: - flr = float(lr) if lr is not None else None - except (TypeError, ValueError): - flr = None + + def _opt_float(v: Any) -> Optional[float]: + try: + return float(v) if v is not None else None + except (TypeError, ValueError): + return None + + flr = _opt_float(lr) + fgn = _opt_float(grad_norm) steps = state["metric_steps"] losses = state["metric_loss"] lrs = state["metric_lr"] + gns = state["metric_grad_norm"] if len(steps) >= _METRIC_CAP: state["metric_steps"] = steps[::2] state["metric_loss"] = losses[::2] state["metric_lr"] = lrs[::2] - steps, losses, lrs = state["metric_steps"], state["metric_loss"], state["metric_lr"] + state["metric_grad_norm"] = gns[::2] + steps, losses, lrs, gns = ( + state["metric_steps"], + state["metric_loss"], + state["metric_lr"], + state["metric_grad_norm"], + ) steps.append(istep) losses.append(floss) lrs.append(flr) + gns.append(fgn) class DiffusionTrainingService: @@ -315,6 +333,7 @@ class DiffusionTrainingService: loss = ev.get("loss", s["loss"]), avg_loss = ev.get("avg_loss", s["avg_loss"]), learning_rate = ev.get("learning_rate", s["learning_rate"]), + grad_norm = ev.get("grad_norm", s["grad_norm"]), message = "Training...", ) # Fold optional perf fields (emitted by the trainers) so the UI can show @@ -323,8 +342,14 @@ class DiffusionTrainingService: s["samples_per_second"] = ev.get("samples_per_second") if ev.get("peak_memory_gb") is not None: s["peak_memory_gb"] = ev.get("peak_memory_gb") - # Retain a bounded (step, loss, lr) history for the live loss chart. - _append_metric(s, ev.get("step"), ev.get("loss"), ev.get("learning_rate")) + # Retain a bounded (step, loss, lr, grad_norm) history for the live charts. + _append_metric( + s, + ev.get("step"), + ev.get("loss"), + ev.get("learning_rate"), + ev.get("grad_norm"), + ) elif etype == "complete": # Reset in_model_load: a stop during model load emits complete without a # preceding model_load_completed, which would otherwise leave a stale diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index 58a9779609..41452ddc89 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -753,12 +753,14 @@ class DiffusionTrainingStartResponse(BaseModel): class DiffusionMetricHistory(BaseModel): - """Paired step-indexed history arrays for the live training charts. ``lr`` entries may - be null so a sparse learning-rate series still aligns with ``steps`` by index.""" + """Paired step-indexed history arrays for the live training charts. ``lr`` and + ``grad_norm`` entries may be null so those sparse series still align with ``steps`` + by index.""" steps: List[int] = Field(default_factory = list) loss: List[float] = Field(default_factory = list) lr: List[Optional[float]] = Field(default_factory = list) + grad_norm: List[Optional[float]] = Field(default_factory = list) class DiffusionTrainingStatusResponse(BaseModel): @@ -773,6 +775,9 @@ class DiffusionTrainingStatusResponse(BaseModel): loss: Optional[float] = None avg_loss: Optional[float] = None learning_rate: Optional[float] = None + # Total pre-clip gradient norm from the last optimizer step (the training health + # signal the UI charts alongside the loss). + grad_norm: Optional[float] = None num_images: Optional[int] = None in_model_load: bool = False output_dir: Optional[str] = None diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 271abebd78..b16d95add0 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -1270,6 +1270,7 @@ async def diffusion_training_status(current_subject: str = Depends(get_current_s steps = snap.pop("metric_steps", []), loss = snap.pop("metric_loss", []), lr = snap.pop("metric_lr", []), + grad_norm = snap.pop("metric_grad_norm", []), ) return DiffusionTrainingStatusResponse(**snap, metric_history = metric_history) diff --git a/studio/frontend/src/features/images/api.ts b/studio/frontend/src/features/images/api.ts index bd6c09dd33..25e4f103e4 100644 --- a/studio/frontend/src/features/images/api.ts +++ b/studio/frontend/src/features/images/api.ts @@ -304,6 +304,8 @@ export interface DiffusionMetricHistory { steps: number[]; loss: number[]; lr: Array; + // Total pre-clip gradient norm per step (the training health signal the charts show). + grad_norm?: Array; } // A snapshot of the current diffusion training job (GET /api/train/diffusion/status). @@ -317,6 +319,7 @@ export interface DiffusionTrainingStatus { loss: number | null; avg_loss: number | null; learning_rate: number | null; + grad_norm?: number | null; num_images: number | null; in_model_load: boolean; output_dir: string | null; diff --git a/studio/frontend/src/features/images/train/diffusion-charts.tsx b/studio/frontend/src/features/images/train/diffusion-charts.tsx index daf762f056..7deb70f122 100644 --- a/studio/frontend/src/features/images/train/diffusion-charts.tsx +++ b/studio/frontend/src/features/images/train/diffusion-charts.tsx @@ -4,12 +4,13 @@ import { type ReactElement, useMemo } from "react"; import type { TrainingSeriesPoint } from "@/features/training"; -// The loss + LR cards are pure presentational (props only), so reuse them directly. We do -// NOT reuse ChartsSection/ChartsContent: those also render Grad Norm and an Eval Loss card, -// which are meaningless for diffusion LoRA training and showed as an empty card and an -// "Evaluation not configured" placeholder. This is a diffusion-only two-card layout. +// The loss + grad-norm cards are pure presentational (props only), so reuse them directly. +// We do NOT reuse ChartsSection/ChartsContent: those also render an LR and an Eval Loss +// card, which add little for diffusion LoRA training (the LR curve is the deterministic +// schedule the user just picked; eval is not configured). This is a diffusion-only +// two-card layout: Training Loss + Grad Norm (the actual training health signal). // eslint-disable-next-line no-restricted-imports -import { LearningRateChartCard } from "@/features/studio/sections/charts/learning-rate-chart-card"; +import { GradNormChartCard } from "@/features/studio/sections/charts/grad-norm-chart-card"; // eslint-disable-next-line no-restricted-imports import { TrainingLossChartCard } from "@/features/studio/sections/charts/training-loss-chart-card"; // eslint-disable-next-line no-restricted-imports @@ -43,16 +44,16 @@ function fullStepDomain(steps: number[]): [number, number] { return [min, max]; } -// A diffusion-only metrics view: just Training Loss and Learning Rate, side by side, with a -// note under the loss card explaining why per-step loss looks noisy. Always renders both -// cards (even with no data) so the Train tab can show them grayed before a run starts; the -// parent applies the grayed treatment via a wrapper, so we never early-return null here. +// A diffusion-only metrics view: Training Loss and Grad Norm, side by side, with a note +// under the loss card explaining why per-step loss looks noisy. Always renders both cards +// (even with no data) so the parent can decide when to mount them; we never early-return +// null here. export function DiffusionCharts({ lossHistory, - lrHistory, + gradNormHistory, }: { lossHistory: TrainingSeriesPoint[]; - lrHistory: TrainingSeriesPoint[]; + gradNormHistory: TrainingSeriesPoint[]; }): ReactElement { const lossItems = useMemo(() => toLossItems(lossHistory), [lossHistory]); const smoothed = useMemo( @@ -73,23 +74,23 @@ export function DiffusionCharts({ [reducedLoss], ); - const lrData = useMemo( + const gradData = useMemo( () => compressSeries( - lrHistory + gradNormHistory .filter((p) => Number.isFinite(p.value)) - .map((p) => ({ step: p.step, lr: p.value, displayLr: p.value })), + .map((p) => ({ step: p.step, gradNorm: p.value, displayGradNorm: p.value })), MAX_RENDER_POINTS, ), - [lrHistory], + [gradNormHistory], ); const steps = useMemo(() => { const set = new Set(); for (const p of lossData) set.add(p.step); - for (const p of lrData) set.add(p.step); + for (const p of gradData) set.add(p.step); return Array.from(set).sort((a, b) => a - b); - }, [lossData, lrData]); + }, [lossData, gradData]); const stepDomain = useMemo(() => fullStepDomain(steps), [steps]); const xAxisTicks = useMemo( @@ -101,9 +102,9 @@ export function DiffusionCharts({ () => buildYDomain(lossData.flatMap((p) => [p.displayLoss, p.displaySmoothed])), [lossData], ); - const lrDomain = useMemo( - () => buildYDomain(lrData.map((p) => p.displayLr)), - [lrData], + const gradDomain = useMemo( + () => buildYDomain(gradData.map((p) => p.displayGradNorm)), + [gradData], ); const avgRaw = @@ -131,9 +132,9 @@ export function DiffusionCharts({ the smoothed line for the trend, not the raw jitter.

- ({ step, value: h.loss[i] })).filter((p) => p.value != null); }, [status?.metric_history]); - const lrHistory: TrainingSeriesPoint[] = useMemo(() => { + const gradNormHistory: TrainingSeriesPoint[] = useMemo(() => { const h = status?.metric_history; - if (!h) return []; + if (!h?.grad_norm) return []; return h.steps - .map((step, i) => ({ step, value: h.lr[i] })) + .map((step, i) => ({ step, value: h.grad_norm?.[i] ?? null })) .filter((p): p is TrainingSeriesPoint => p.value != null); }, [status?.metric_history]); @@ -1029,7 +1029,9 @@ export function DiffusionTrainPanel({ <>
- {status?.status} + + {status?.status === "completed" ? "Training complete \u{1F389}" : status?.status} + {(status?.total_steps ?? 0) > 0 ? `${status?.step}/${status?.total_steps} steps` @@ -1073,7 +1075,7 @@ export function DiffusionTrainPanel({ )}
- + )}