else if (isNaN(ckptEvery) || ckptEvery < 1) error = 'Ckpt every must be >= 1';
else if (isNaN(demoEvery) || demoEvery < 1) error = 'Demo every must be >= 1';
else if (lrStr && isNaN(Number(lrStr))) error = 'LR must be a number (e.g. 1e-4)';
+ else if (weightDecayStr && (isNaN(Number(weightDecayStr)) || Number(weightDecayStr) < 0)) error = 'Weight decay must be a non-negative number';
+ else if (warmupStepsStr && (isNaN(parseInt(warmupStepsStr, 10)) || parseInt(warmupStepsStr, 10) < 0)) error = 'Warmup steps must be a non-negative integer';
else {
const cropLen = parseInt(document.getElementById('new-ft-crop-length').value, 10);
if (isNaN(cropLen) || cropLen < 1) error = 'Latent seq length must be >= 1';
@@ -6189,6 +6325,23 @@
Download Audio Selection
if (includeStr) payload.lora_include = includeStr;
if (excludeStr) payload.lora_exclude = excludeStr;
if (lrStr) payload.lr = lrStr;
+ const weightDecayStr = document.getElementById('new-ft-weight-decay').value.trim();
+ if (weightDecayStr) payload.weight_decay = Number(weightDecayStr);
+ const warmupStepsStr = document.getElementById('new-ft-warmup-steps').value.trim();
+ if (warmupStepsStr && parseInt(warmupStepsStr, 10) > 0) payload.warmup_steps = parseInt(warmupStepsStr, 10);
+ const timestepSampler = document.getElementById('new-ft-timestep-sampler').value;
+ if (timestepSampler) payload.timestep_sampler = timestepSampler;
+ const lossNormalization = document.getElementById('new-ft-loss-normalization').value;
+ if (lossNormalization) payload.loss_normalization = lossNormalization;
+ payload.best_checkpoint_enabled = document.getElementById('new-ft-best-ckpt-enabled').checked;
+ if (payload.best_checkpoint_enabled) {
+ const bcWarmup = document.getElementById('new-ft-best-ckpt-warmup').value.trim();
+ if (bcWarmup) payload.best_checkpoint_warmup_steps = parseInt(bcWarmup, 10);
+ const bcKeepN = document.getElementById('new-ft-best-ckpt-keep-n').value.trim();
+ if (bcKeepN) payload.best_checkpoint_keep_n = parseInt(bcKeepN, 10);
+ const bcCheckEvery = document.getElementById('new-ft-best-ckpt-check-every-n').value.trim();
+ if (bcCheckEvery) payload.best_checkpoint_check_every_n_epochs = parseInt(bcCheckEvery, 10);
+ }
const basePrecision = document.getElementById('new-ft-base-precision').value;
if (basePrecision) payload.base_precision = basePrecision;
if (datasetId) payload.dataset_id = datasetId;
diff --git a/dashboard/server.py b/dashboard/server.py
index bcc3682..404e4f7 100644
--- a/dashboard/server.py
+++ b/dashboard/server.py
@@ -86,6 +86,102 @@ def _atomic_write_json(path, data):
f.flush()
os.fsync(f.fileno())
os.replace(str(tmp), str(path))
+
+
+# Valid values for the two enum-like advanced overrides — kept here so both
+# the New Finetune and Resume handlers validate identically.
+_VALID_TIMESTEP_SAMPLERS = {
+ "uniform", "logit_normal", "trunc_logit_normal", "log_snr", "log_snr_uniform",
+}
+_VALID_LOSS_NORMALIZATIONS = {"none", "timestep", "sample", "sample_channel"}
+
+
+def _warmup_base_for_steps(warmup_steps, target=0.99):
+ """Convert an intuitive "warmup steps" count into the `warmup` base
+ InverseLR actually takes (a (0,1) exponential-decay base, NOT a step
+ count — see underfit/training/optim.py). Returns the base such that the
+ warmup ramp reaches `target` fraction of full LR by `warmup_steps`.
+
+ warmup factor at step n is (1 - base**(n+1)); solving for base when
+ n = warmup_steps and factor = target gives the formula below.
+ """
+ if not warmup_steps or warmup_steps <= 0:
+ return 0.0
+ return (1 - target) ** (1.0 / (warmup_steps + 1))
+
+
+def _apply_advanced_training_overrides(cfg, body):
+ """Apply the four 'advanced' training overrides (warmup steps, timestep
+ sampler, loss normalization, weight decay) from a New Finetune / Resume
+ request body onto a run's training config dict, in place.
+
+ Mirrors the existing `lr` override pattern used by both handlers: only
+ touches a field if the request actually supplied it, so omitting a field
+ leaves whatever was already in `cfg` (e.g. restored from a previous
+ run's config on resume) untouched.
+ """
+ training_cfg = cfg.setdefault("training", {})
+
+ warmup_steps = body.get("warmup_steps")
+ if warmup_steps not in (None, ""):
+ warmup_steps = int(warmup_steps)
+ opt_cfg = (training_cfg.setdefault("optimizer_configs", {})
+ .setdefault("diffusion", {}))
+ if warmup_steps > 0:
+ # inv_gamma huge + power=1 => negligible post-warmup decay, so
+ # this only adds the warmup ramp without changing the flat-LR
+ # behavior runs already have today (see _warmup_base_for_steps).
+ opt_cfg["scheduler"] = {
+ "type": "InverseLR",
+ "config": {
+ "warmup": _warmup_base_for_steps(warmup_steps),
+ "inv_gamma": 1e9,
+ "power": 1.0,
+ },
+ }
+ else:
+ opt_cfg.pop("scheduler", None)
+
+ weight_decay = body.get("weight_decay")
+ if weight_decay not in (None, ""):
+ (training_cfg.setdefault("optimizer_configs", {})
+ .setdefault("diffusion", {}).setdefault("optimizer", {})
+ .setdefault("config", {}))["weight_decay"] = float(weight_decay)
+
+ # Best-checkpoint tracking: saves a copy of the regular checkpoint as
+ # "best" whenever an EMA of the training loss hits a new low (after an
+ # optional warmup), keeping only the most recent N best files. See
+ # _BestCheckpointTracker in underfit/training/loop.py. Off by default —
+ # only touched if the request explicitly includes the enabled flag
+ # (checkbox semantics: always present as true/false, never blank).
+ if "best_checkpoint_enabled" in body:
+ training_cfg["best_checkpoint_enabled"] = bool(body["best_checkpoint_enabled"])
+
+ best_ckpt_warmup = body.get("best_checkpoint_warmup_steps")
+ if best_ckpt_warmup not in (None, ""):
+ training_cfg["best_checkpoint_warmup_steps"] = int(best_ckpt_warmup)
+
+ best_ckpt_keep_n = body.get("best_checkpoint_keep_n")
+ if best_ckpt_keep_n not in (None, ""):
+ training_cfg["best_checkpoint_keep_n"] = int(best_ckpt_keep_n)
+
+ best_ckpt_check_every_n = body.get("best_checkpoint_check_every_n_epochs")
+ if best_ckpt_check_every_n not in (None, ""):
+ training_cfg["best_checkpoint_check_every_n_epochs"] = int(best_ckpt_check_every_n)
+
+ timestep_sampler = body.get("timestep_sampler")
+ if timestep_sampler:
+ if timestep_sampler not in _VALID_TIMESTEP_SAMPLERS:
+ raise ValueError(f"invalid timestep_sampler: {timestep_sampler!r}")
+ training_cfg["timestep_sampler"] = timestep_sampler
+
+ loss_normalization = body.get("loss_normalization")
+ if loss_normalization:
+ if loss_normalization not in _VALID_LOSS_NORMALIZATIONS:
+ raise ValueError(f"invalid loss_normalization: {loss_normalization!r}")
+ training_cfg["loss_normalization"] = loss_normalization
+
+
AUDIO_DIR = STATE_DIR / "audio" # generated demo MP3s + spectrogram JPGs
# Base-model files (SA3 RF + ARC, T5Gemma) — defaults to STATE_DIR/models,
@@ -3974,6 +4070,7 @@ def _handle_new_finetune(self, body):
if lr_raw:
lr_val = float(lr_raw)
cfg["training"].setdefault("optimizer_configs", {}).setdefault("diffusion", {}).setdefault("optimizer", {}).setdefault("config", {})["lr"] = lr_val
+ _apply_advanced_training_overrides(cfg, body)
# Inject ARC path for demos during training.
if mi.get("arc_ckpt"):
demo_config = cfg["training"].setdefault("demo", {})
@@ -4311,6 +4408,7 @@ def _handle_resume(self, run_id, body):
if lr_raw:
lr_val = float(lr_raw)
cfg.setdefault("training", {}).setdefault("optimizer_configs", {}).setdefault("diffusion", {}).setdefault("optimizer", {}).setdefault("config", {})["lr"] = lr_val
+ _apply_advanced_training_overrides(cfg, body)
# Assign fresh random seeds 10-100 to each demo on resume
import random as _rng
for entry in cfg.get("training", {}).get("demo", {}).get("demo_cond", []):
diff --git a/underfit/training/loop.py b/underfit/training/loop.py
index 51325b5..d9b99f9 100644
--- a/underfit/training/loop.py
+++ b/underfit/training/loop.py
@@ -12,6 +12,7 @@
import math
import os
import re
+import shutil
import signal
import struct
import sys
@@ -174,6 +175,77 @@ def close(self):
self._f = None
+class _BestCheckpointTracker:
+ """Tracks an EMA of the training loss and saves a best checkpoint
+ whenever the EMA hits a new low at the end of each epoch (after an
+ optional warmup period). Keeps only the last `keep_n` best files.
+
+ Design notes:
+ - Checked at each epoch boundary (not per-step and not only at
+ interval checkpoints) — gives finer granularity than save_every
+ without the cost of writing on every single step.
+ - Compares an EMA of per-step loss, not raw loss — raw loss is too
+ noisy; a single lucky batch would trigger a spurious "best."
+ - `best_so_far` survives a resume: every regular checkpoint stamps
+ the current best into its safetensors metadata (best_ema_loss
+ kwarg on save_lora_step), so resuming from any checkpoint recovers
+ the right value without a separate sidecar file.
+ - Does its own save_lora_step call (not a copy of a regular
+ checkpoint) since epoch boundaries don't align with save_every.
+ - Keeps only the last `keep_n` best files (oldest deleted first).
+ - underfit has no validation loss, so this is a train-loss proxy —
+ it can reflect overfitting, not genuine held-out generalization.
+ """
+ def __init__(self, *, alpha=0.05, warmup_steps=0, keep_n=5, best_so_far=None,
+ check_every_n_epochs=1):
+ self.alpha = alpha
+ self.warmup_steps = warmup_steps
+ self.keep_n = keep_n
+ self.check_every_n_epochs = max(1, int(check_every_n_epochs))
+ self.ema_loss = None
+ self.best_so_far = best_so_far
+ self._saved_paths = [] # oldest-first, for rotation
+
+ def update(self, loss_value):
+ """Call every step with the raw per-step loss; updates the EMA."""
+ if self.ema_loss is None:
+ self.ema_loss = loss_value
+ else:
+ self.ema_loss = self.alpha * loss_value + (1 - self.alpha) * self.ema_loss
+
+ def maybe_save_best(self, *, backend, model, saved_lora_cfg, base_model_name,
+ global_step, epoch, checkpoint_dir, run_label):
+ """Call at each epoch boundary. If the current EMA loss is a new
+ low (and warmup has passed), saves a best checkpoint via its own
+ save_lora_step call. Returns the saved path, or None if skipped.
+ """
+ if self.ema_loss is None or global_step < self.warmup_steps:
+ return None
+ is_new_best = self.best_so_far is None or self.ema_loss < self.best_so_far
+ if not is_new_best:
+ return None
+ self.best_so_far = self.ema_loss
+
+ best_dir = os.path.join(checkpoint_dir, "best")
+ os.makedirs(best_dir, exist_ok=True)
+ best_name = (f"{run_label}-best-step={global_step}-epoch={epoch}.safetensors"
+ if run_label else f"best-step={global_step}-epoch={epoch}.safetensors")
+ best_path = os.path.join(best_dir, best_name)
+
+ save_lora_step(backend, model, saved_lora_cfg, best_path,
+ step=global_step, epoch=epoch, base_model=base_model_name,
+ best_ema_loss=self.best_so_far)
+ self._saved_paths.append(best_path)
+
+ while len(self._saved_paths) > self.keep_n:
+ old_path = self._saved_paths.pop(0)
+ try:
+ os.remove(old_path)
+ except OSError:
+ pass
+ return best_path
+
+
def _explain_model_load_error(exc, model_config):
"""Print friendly help for known model-load failure modes before the
traceback bubbles up. Catches gated/missing HuggingFace repos (typically
@@ -421,8 +493,36 @@ def run_training(args, backend):
os.makedirs(checkpoint_dir, exist_ok=True)
run_label = re.sub(r"-\d{14}$", "", args.name) if args.name else None
- # --- Loss-by-timestep log ---
- lbt_log = _LossByTimestepLog(os.path.join(os.getcwd(), "loss_by_timestep.bin"))
+ # --- demos dir: where the dashboard reads loss_by_timestep.bin from.
+ # server.py hardcodes: RUNS_DIR / run_id / "demos" / "loss_by_timestep.bin"
+ # We must write there, not in a custom metrics/ dir, or the dashboard
+ # Loss chart stays blank. Also added to HF sync WATCH list in the notebook
+ # so it persists across sessions.
+ demos_dir = None
+ if args.save_dir and args.name:
+ demos_dir = os.path.join(args.save_dir, args.name, "demos")
+ os.makedirs(demos_dir, exist_ok=True)
+
+ lbt_path = (os.path.join(demos_dir, "loss_by_timestep.bin")
+ if demos_dir else os.path.join(os.getcwd(), "loss_by_timestep.bin"))
+ lbt_log = _LossByTimestepLog(lbt_path)
+
+ # --- TensorBoard (optional — graceful no-op if tensorboard not installed)
+ # Writes tfevents files into a tb/ dir alongside the checkpoints, which
+ # falls under runs/** so the HF sync picks them up automatically.
+ # HF Hub renders a TensorBoard tab on the model page once tfevents files
+ # are present. Logs: loss, grad_norm, lora_magnitude, lr, ema_loss.
+ tb_writer = None
+ if checkpoint_dir:
+ tb_dir = os.path.join(os.path.dirname(checkpoint_dir), "tb")
+ os.makedirs(tb_dir, exist_ok=True)
+ try:
+ from torch.utils.tensorboard import SummaryWriter
+ tb_writer = SummaryWriter(log_dir=tb_dir, purge_step=step_offset or 0)
+ print(f"[startup] TensorBoard logging -> {tb_dir}", flush=True)
+ except ImportError:
+ print("[startup] tensorboard not installed — skipping TB logging "
+ "(pip install tensorboard to enable)", flush=True)
# --- SIGUSR1 manual save ---
manual_save_requested = [False]
@@ -468,6 +568,36 @@ def _request_save(*_):
demo_config = training_config.get("demo", {}) or {}
demo_every = int(demo_config.get("demo_every", 0))
last_demo_step = -1
+
+ # --- Best-checkpoint tracking ---
+ # best_so_far is recovered from the resumed checkpoint's metadata (every
+ # checkpoint stamps the current best, not just literal-best ones — see
+ # save_lora_step's best_ema_loss kwarg) so a resume doesn't reset
+ # progress. training_config.best_ema_loss is an explicit override,
+ # same precedence convention as step_offset/epoch_offset above.
+ has_config_best = "best_ema_loss" in training_config
+ config_best = float(training_config["best_ema_loss"]) if has_config_best else None
+ meta_best = None
+ if resume_metadata and "best_ema_loss" in resume_metadata:
+ try:
+ meta_best = float(resume_metadata["best_ema_loss"])
+ except (TypeError, ValueError):
+ pass
+ best_so_far = config_best if has_config_best else meta_best
+
+ best_ckpt_warmup_steps = int(training_config.get("best_checkpoint_warmup_steps", 0) or 0)
+ best_ckpt_keep_n = int(training_config.get("best_checkpoint_keep_n", 5) or 5)
+ best_ckpt_check_every_n = int(training_config.get("best_checkpoint_check_every_n_epochs", 10) or 10)
+ best_tracker = _BestCheckpointTracker(
+ warmup_steps=best_ckpt_warmup_steps,
+ keep_n=best_ckpt_keep_n,
+ check_every_n_epochs=best_ckpt_check_every_n,
+ best_so_far=best_so_far,
+ ) if best_ckpt_warmup_steps >= 0 and training_config.get("best_checkpoint_enabled", False) else None
+ if best_tracker is not None:
+ print(f"[startup] Best-checkpoint tracking enabled (warmup={best_ckpt_warmup_steps} steps, "
+ f"keep_n={best_ckpt_keep_n}, check_every={best_ckpt_check_every_n} epochs, "
+ f"resumed best={best_so_far})", flush=True)
raw_step = 0
epoch = epoch_offset
@@ -644,6 +774,19 @@ def _request_save(*_):
# --- Loss-by-timestep ---
lbt_log.write(global_step, t.detach().float().mean().item(), loss.item())
+ if best_tracker is not None:
+ best_tracker.update(loss.item())
+
+ # --- TensorBoard ---
+ if tb_writer is not None:
+ tb_writer.add_scalar("train/loss", loss.item(), global_step)
+ tb_writer.add_scalar("train/lr", lr, global_step)
+ if grad_norm is not None:
+ tb_writer.add_scalar("train/grad_norm", grad_norm, global_step)
+ if lora_mag is not None:
+ tb_writer.add_scalar("train/lora_magnitude", lora_mag, global_step)
+ if best_tracker is not None and best_tracker.ema_loss is not None:
+ tb_writer.add_scalar("train/ema_loss", best_tracker.ema_loss, global_step)
raw_step += 1
global_step = raw_step + step_offset
@@ -676,7 +819,9 @@ def _request_save(*_):
if manual_save_requested[0]:
manual_save_requested[0] = False
out = os.path.join(checkpoint_dir, _ckpt_filename(run_label, global_step, epoch))
- save_lora_step(backend, model, saved_lora_cfg, out, step=global_step, epoch=epoch, base_model=base_model_name)
+ current_best = best_tracker.best_so_far if best_tracker is not None else None
+ save_lora_step(backend, model, saved_lora_cfg, out, step=global_step, epoch=epoch,
+ base_model=base_model_name, best_ema_loss=current_best)
print(f"✓ Saved checkpoint -- {os.path.basename(out)}", flush=True)
if demo_will_fire:
@@ -700,6 +845,19 @@ def _request_save(*_):
if raw_step >= max_steps:
break
+
+ # --- End of epoch: check for new best checkpoint ---
+ if (best_tracker is not None and checkpoint_dir
+ and epoch % best_tracker.check_every_n_epochs == 0):
+ best_path = best_tracker.maybe_save_best(
+ backend=backend, model=model,
+ saved_lora_cfg=saved_lora_cfg, base_model_name=base_model_name,
+ global_step=global_step, epoch=epoch,
+ checkpoint_dir=checkpoint_dir, run_label=run_label,
+ )
+ if best_path:
+ print(f" ★ New best (EMA loss={best_tracker.best_so_far:.6f}, "
+ f"epoch={epoch}) -- {os.path.basename(best_path)}", flush=True)
epoch += 1
# Final save (skip if the last regular save already covered this
@@ -711,8 +869,23 @@ def _request_save(*_):
and global_step > 0
and global_step % save_every != 0):
out = os.path.join(checkpoint_dir, _ckpt_filename(run_label, global_step, epoch))
- save_lora_step(backend, model, saved_lora_cfg, out, step=global_step, epoch=epoch, base_model=base_model_name)
+ current_best = best_tracker.best_so_far if best_tracker is not None else None
+ save_lora_step(backend, model, saved_lora_cfg, out, step=global_step, epoch=epoch,
+ base_model=base_model_name, best_ema_loss=current_best)
print(f"✓ Saved checkpoint -- {os.path.basename(out)} (final)", flush=True)
+
+ if best_tracker is not None:
+ best_path = best_tracker.maybe_save_best(
+ backend=backend, model=model,
+ saved_lora_cfg=saved_lora_cfg, base_model_name=base_model_name,
+ global_step=global_step, epoch=epoch,
+ checkpoint_dir=checkpoint_dir, run_label=run_label,
+ )
+ if best_path:
+ print(f" ★ New best (EMA loss={best_tracker.best_so_far:.6f}) -- "
+ f"{os.path.basename(best_path)}", flush=True)
finally:
lbt_log.close()
+ if tb_writer is not None:
+ tb_writer.close()
print("Training done", flush=True)
diff --git a/underfit/training/lora.py b/underfit/training/lora.py
index c8b9da5..cafef9a 100644
--- a/underfit/training/lora.py
+++ b/underfit/training/lora.py
@@ -88,7 +88,7 @@ def apply_lora_from_config(backend, model, lora_config, lora_state_dict=None,
def save_lora_step(backend, model, lora_save_config, out_path,
- *, step=None, epoch=None, base_model=None):
+ *, step=None, epoch=None, base_model=None, best_ema_loss=None):
"""Save LoRA weights to out_path as a .safetensors file with config metadata.
`step` and `epoch` are folded into the saved metadata (under the "step"
@@ -98,6 +98,12 @@ def save_lora_step(backend, model, lora_save_config, out_path,
`base_model` (e.g. "sa3-medium") goes into metadata too — used by the
dashboard's "Start from a previous LoRA" upload flow to verify the seed
is shape-compatible with the user's selected base model.
+
+ `best_ema_loss`, if given, is the best EMA-smoothed training loss seen
+ so far in this run (across any prior resumes) — stamped into every
+ checkpoint (not just the literal best one) so a future resume from
+ *any* checkpoint can recover "best so far as of that point" without a
+ separate sidecar file. See _BestCheckpointTracker in loop.py.
"""
lora_mod = backend.lora_module()
state_dict = {
@@ -111,6 +117,8 @@ def save_lora_step(backend, model, lora_save_config, out_path,
enriched_cfg["epoch"] = int(epoch)
if base_model:
enriched_cfg["base_model"] = str(base_model)
+ if best_ema_loss is not None:
+ enriched_cfg["best_ema_loss"] = float(best_ema_loss)
Path(out_path).parent.mkdir(parents=True, exist_ok=True)
lora_mod.save_lora_safetensors(state_dict, enriched_cfg, out_path)