Skip to content
Open
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
153 changes: 153 additions & 0 deletions dashboard/index.html
Original file line number Diff line number Diff line change
Expand Up @@ -1243,6 +1243,62 @@ <h3>Revive Training</h3>
</div>
<label class="modal-label">LR</label>
<input type="text" class="resume-input" id="resume-lr" value="" placeholder="default (keep current)" oninput="validateResume()">
<details id="resume-advanced" style="margin:8px 0;background:rgba(255,255,255,0.04);border-radius:6px;padding:8px 10px">
<summary style="cursor:pointer;font-size:0.8em;color:#9e9e9e;user-select:none">Advanced</summary>
<div class="new-ft-grid4" style="display:grid;grid-template-columns:1fr 1fr 1fr 1fr;gap:12px;margin-top:8px">
<div>
<label class="modal-label">Warmup steps</label>
<input type="number" class="resume-input" id="resume-warmup-steps" value="" min="0" placeholder="default (keep current)" oninput="validateResume()">
</div>
<div>
<label class="modal-label">Weight decay</label>
<input type="text" class="resume-input" id="resume-weight-decay" value="" placeholder="default (keep current)" oninput="validateResume()">
</div>
<div>
<label class="modal-label">Timestep sampler</label>
<select class="resume-input" id="resume-timestep-sampler">
<option value="">default (keep current)</option>
<option value="uniform">uniform</option>
<option value="logit_normal">logit_normal</option>
<option value="trunc_logit_normal">trunc_logit_normal</option>
<option value="log_snr">log_snr</option>
<option value="log_snr_uniform">log_snr_uniform</option>
</select>
</div>
<div>
<label class="modal-label">Loss normalization</label>
<select class="resume-input" id="resume-loss-normalization">
<option value="">default (keep current)</option>
<option value="none">none</option>
<option value="timestep">timestep</option>
<option value="sample">sample</option>
<option value="sample_channel">sample_channel</option>
</select>
</div>
</div>
<div style="margin-top:10px;padding-top:10px;border-top:1px solid rgba(255,255,255,0.08)">
<label class="modal-label">Best checkpoint tracking</label>
<select class="resume-input" id="resume-best-ckpt-enabled" onchange="document.getElementById('resume-best-ckpt-row').style.display = this.value === 'true' ? 'grid' : 'none'">
<option value="">default (keep current)</option>
<option value="true">enabled</option>
<option value="false">disabled</option>
</select>
<div id="resume-best-ckpt-row" style="display:none;grid-template-columns:1fr 1fr 1fr;gap:12px;margin-top:8px">
<div>
<label class="modal-label">Best-ckpt warmup steps</label>
<input type="number" class="resume-input" id="resume-best-ckpt-warmup" value="" min="0" placeholder="default (keep current)" oninput="validateResume()">
</div>
<div>
<label class="modal-label">Keep last N best</label>
<input type="number" class="resume-input" id="resume-best-ckpt-keep-n" value="" min="1" placeholder="default (keep current)" oninput="validateResume()">
</div>
<div>
<label class="modal-label">Check every N epochs</label>
<input type="number" class="resume-input" id="resume-best-ckpt-check-every-n" value="" min="1" placeholder="default (keep current)" oninput="validateResume()">
</div>
</div>
</div>
</details>
<div class="resume-error" id="resume-error"></div>
<div class="modal-actions" style="margin-top:16px">
<button class="modal-cancel" onclick="closeResumeModal()">Cancel</button>
Expand Down Expand Up @@ -1380,6 +1436,60 @@ <h3>New Finetune</h3>
<input type="number" class="resume-input" id="new-ft-demo-every" value="1000" oninput="validateNewFt()">
</div>
</div>
<details id="new-ft-advanced" style="margin:8px 0;background:rgba(255,255,255,0.04);border-radius:6px;padding:8px 10px">
<summary style="cursor:pointer;font-size:0.8em;color:#9e9e9e;user-select:none">Advanced</summary>
<div class="new-ft-grid4" style="display:grid;grid-template-columns:1fr 1fr 1fr 1fr;gap:12px;margin-top:8px">
<div>
<label class="modal-label">Warmup steps</label>
<input type="number" class="resume-input" id="new-ft-warmup-steps" value="0" min="0" placeholder="0" oninput="validateNewFt()">
</div>
<div>
<label class="modal-label">Weight decay</label>
<input type="text" class="resume-input" id="new-ft-weight-decay" value="" placeholder="0.01 (default)" oninput="validateNewFt()">
</div>
<div>
<label class="modal-label">Timestep sampler</label>
<select class="resume-input" id="new-ft-timestep-sampler">
<option value="">default (uniform)</option>
<option value="uniform">uniform</option>
<option value="logit_normal">logit_normal</option>
<option value="trunc_logit_normal">trunc_logit_normal</option>
<option value="log_snr">log_snr</option>
<option value="log_snr_uniform">log_snr_uniform</option>
</select>
</div>
<div>
<label class="modal-label">Loss normalization</label>
<select class="resume-input" id="new-ft-loss-normalization">
<option value="">default (none)</option>
<option value="none">none</option>
<option value="timestep">timestep</option>
<option value="sample">sample</option>
<option value="sample_channel">sample_channel</option>
</select>
</div>
</div>
<div style="margin-top:10px;padding-top:10px;border-top:1px solid rgba(255,255,255,0.08)">
<label style="display:flex;align-items:center;gap:6px;font-size:0.85em;cursor:pointer">
<input type="checkbox" id="new-ft-best-ckpt-enabled" onchange="document.getElementById('new-ft-best-ckpt-row').style.display = this.checked ? 'grid' : 'none'">
Save best checkpoint (by EMA-smoothed training loss)
</label>
<div id="new-ft-best-ckpt-row" style="display:none;grid-template-columns:1fr 1fr;gap:12px;margin-top:8px">
<div>
<label class="modal-label">Best-ckpt warmup steps</label>
<input type="number" class="resume-input" id="new-ft-best-ckpt-warmup" value="1000" min="0" oninput="validateNewFt()">
</div>
<div>
<label class="modal-label">Keep last N best</label>
<input type="number" class="resume-input" id="new-ft-best-ckpt-keep-n" value="5" min="1" oninput="validateNewFt()">
</div>
<div>
<label class="modal-label">Check every N epochs</label>
<input type="number" class="resume-input" id="new-ft-best-ckpt-check-every-n" value="10" min="1" title="10 is good for small datasets (~13 samples). Larger datasets can use 1." oninput="validateNewFt()">
</div>
</div>
</div>
</details>
<div class="modal-actions" style="margin-top:16px;display:flex;align-items:center">
<div id="new-ft-ckpt-estimate" style="color:var(--yellow);flex:1;font-size:0.85em"></div>
<button class="modal-cancel" onclick="closeNewFinetuneModal()">Cancel</button>
Expand Down Expand Up @@ -4626,6 +4736,12 @@ <h3>Download Audio Selection</h3>
} else if (lrStr && isNaN(Number(lrStr))) {
err.textContent = 'LR must be a number (e.g. 1e-4)';
btn.disabled = true;
} else if ((() => { const s = document.getElementById('resume-weight-decay').value.trim(); return s && (isNaN(Number(s)) || Number(s) < 0); })()) {
err.textContent = 'Weight decay must be a non-negative number';
btn.disabled = true;
} else if ((() => { const s = document.getElementById('resume-warmup-steps').value.trim(); return s && (isNaN(parseInt(s, 10)) || parseInt(s, 10) < 0); })()) {
err.textContent = 'Warmup steps must be a non-negative integer';
btn.disabled = true;
} else {
err.textContent = '';
btn.disabled = false;
Expand Down Expand Up @@ -4653,6 +4769,22 @@ <h3>Download Audio Selection</h3>
const payload = {max_steps: val, gpu: _resumeSelectedGpu, batch_size: batch, checkpoint_every: ckptEvery, demo_every: demoEvery};
if (ckptPath) payload.checkpoint_path = ckptPath;
if (lrStr) payload.lr = lrStr;
const weightDecayStr = document.getElementById('resume-weight-decay').value.trim();
if (weightDecayStr) payload.weight_decay = Number(weightDecayStr);
const warmupStepsStr = document.getElementById('resume-warmup-steps').value.trim();
if (warmupStepsStr) payload.warmup_steps = parseInt(warmupStepsStr, 10);
const timestepSampler = document.getElementById('resume-timestep-sampler').value;
if (timestepSampler) payload.timestep_sampler = timestepSampler;
const lossNormalization = document.getElementById('resume-loss-normalization').value;
if (lossNormalization) payload.loss_normalization = lossNormalization;
const bcEnabledStr = document.getElementById('resume-best-ckpt-enabled').value;
if (bcEnabledStr) payload.best_checkpoint_enabled = bcEnabledStr === 'true';
const bcWarmup = document.getElementById('resume-best-ckpt-warmup').value.trim();
if (bcWarmup) payload.best_checkpoint_warmup_steps = parseInt(bcWarmup, 10);
const bcKeepN = document.getElementById('resume-best-ckpt-keep-n').value.trim();
if (bcKeepN) payload.best_checkpoint_keep_n = parseInt(bcKeepN, 10);
const bcCheckEvery = document.getElementById('resume-best-ckpt-check-every-n').value.trim();
if (bcCheckEvery) payload.best_checkpoint_check_every_n_epochs = parseInt(bcCheckEvery, 10);
if (!isNaN(cropLen) && cropLen > 0) payload.latent_crop_length = cropLen;
payload.random_crop = cropMode === 'random';
try {
Expand Down Expand Up @@ -6132,6 +6264,8 @@ <h3>Download Audio Selection</h3>
const ckptEvery = parseInt(document.getElementById('new-ft-ckpt-every').value, 10);
const demoEvery = parseInt(document.getElementById('new-ft-demo-every').value, 10);
const lrStr = document.getElementById('new-ft-lr').value.trim();
const weightDecayStr = document.getElementById('new-ft-weight-decay').value.trim();
const warmupStepsStr = document.getElementById('new-ft-warmup-steps').value.trim();
const el = document.getElementById('new-ft-ckpt-estimate');
const btn = document.getElementById('new-ft-next-btn');
const baseModelKey = document.getElementById('new-ft-base-model').value;
Expand All @@ -6149,6 +6283,8 @@ <h3>Download Audio Selection</h3>
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';
Expand Down Expand Up @@ -6189,6 +6325,23 @@ <h3>Download Audio Selection</h3>
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;
Expand Down
98 changes: 98 additions & 0 deletions dashboard/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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", {})
Expand Down Expand Up @@ -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", []):
Expand Down
Loading