Merge branch 'nightly' into integrate/exports-page

This commit is contained in:
Roland Tannous 2026-02-17 22:09:36 +04:00 committed by GitHub
commit b8171a86ac
22 changed files with 733 additions and 196 deletions

3
requirements/base.txt Normal file
View file

@ -0,0 +1,3 @@
# Core unsloth packages
unsloth-zoo
unsloth

View file

@ -0,0 +1,13 @@
# Audio extras (installed with --no-deps --no-cache-dir)
descript-audio-codec
descript-audiotools
julius
torchcodec
snac
# TRL and related packages
trl==0.23.1
git+https://github.com/meta-pytorch/OpenEnv.git
executorch==1.0.1
torch-c-dlpack-ext
sentence_transformers==5.2.0

56
requirements/extras.txt Normal file
View file

@ -0,0 +1,56 @@
# OpenEnv dependencies
tomli
tomli-w
# ExecuTorch dependencies
ruamel.yaml
coremltools
expecttest
flatbuffers
hydra-core
hypothesis
kgb
parameterized
pytest<9.0
pytest-json-report
pytest-rerunfailures==15.1
pytest-xdist
# Also needed by sentence_transformers
scikit-learn==1.7.1
# Additional extras
pybind11
langid
jiwer
omegaconf
einx
pyloudnorm
openai-whisper
uroman
MeCab
loguru
flatten_dict
ffmpy
randomname
argbind
tiktoken
ftfy
importlib-resources
librosa
markdown2
matplotlib
pystoi
soundfile
tensorboard
torch-stoi
evaluate
timm
transformers-cfg
open_spiel
addict
easydict
einops
tabulate
fastmcp>=2.0.0
openai>=2.7.2
websockets>=13.0,<14

View file

@ -0,0 +1,7 @@
# Torch AO overrides (installed with --force-reinstall --no-cache-dir)
torchao==0.14.0
transformers==4.57.1
pytorch_tokenizers
# Kernel packages
kernels

13
requirements/studio.txt Normal file
View file

@ -0,0 +1,13 @@
# Studio UI backend dependencies
typer
fastapi
uvicorn
pydantic
matplotlib
pandas
nest_asyncio
datasets==4.3.0
pyjwt
easydict
addict
gradio>=4.0.0

View file

@ -0,0 +1,2 @@
# Triton kernels (installed with --no-deps, from source)
triton_kernels @ git+https://github.com/triton-lang/triton.git@release/3.6.x#subdirectory=python/triton_kernels

View file

@ -132,17 +132,27 @@ fi
BEST_VER=$("$BEST_PY" --version 2>&1 | awk '{print $2}')
echo "✅ Using $BEST_PY ($BEST_VER) — compatible (≤ 3.12.x)"
# Always start fresh to preserve correct install order
rm -rf .venv
"$BEST_PY" -m venv .venv
source .venv/bin/activate
run_quiet "pip upgrade" pip install --upgrade pip
echo " Installing unsloth-zoo + unsloth..."
run_quiet "pip install unsloth" pip install unsloth-zoo unsloth
echo " Installing llama-cpp deps..."
run_quiet "pip install llama-cpp deps" pip install gguf==0.17.1 protobuf==6.33.5 sentencepiece==0.2.1 mistral_common==1.9.0
run_quiet "pip install unsloth" pip install -r "$SCRIPT_DIR/requirements/base.txt"
echo " Installing additional unsloth dependencies..."
run_quiet "pip install extras" pip install --no-cache-dir -r "$SCRIPT_DIR/requirements/extras.txt"
run_quiet "pip install extras" pip install --no-deps --no-cache-dir -r "$SCRIPT_DIR/requirements/extras-no-deps.txt"
run_quiet "pip install torchao+transformers" pip install --force-reinstall --no-cache-dir -r "$SCRIPT_DIR/requirements/overrides.txt"
run_quiet "pip install triton_kernels" pip install --no-deps -r "$SCRIPT_DIR/requirements/triton-kernels.txt"
# Patch: override llama_cpp.py with fix from unsloth-zoo branch
LLAMA_CPP_DST="$(pip show unsloth-zoo | grep -i '^Location:' | awk '{print $2}')/unsloth_zoo/llama_cpp.py"
curl -sSL "https://raw.githubusercontent.com/unslothai/unsloth-zoo/refs/heads/main/unsloth_zoo/llama_cpp.py" \
-o "$LLAMA_CPP_DST"
echo " Installing studio dependencies..."
run_quiet "pip install extras" pip install typer fastapi uvicorn pydantic matplotlib pandas nest_asyncio "datasets==4.3.0" pyjwt easydict addict
run_quiet "pip install studio" pip install -r "$SCRIPT_DIR/requirements/studio.txt"
echo "✅ Python dependencies installed"
# ── 7. Add shell alias ──
# Note: venv activation does NOT persist across terminal sessions.
# This alias hardcodes the venv python path so users don't need to activate.

View file

@ -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

View file

@ -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")

View file

@ -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")

View file

@ -1,7 +0,0 @@
fastapi>=0.100.0
uvicorn>=0.27.0
pydantic>=2.0
torch
psutil
nest-asyncio>=1.5.8

View file

@ -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:

View file

@ -45,7 +45,7 @@ import {
Search01Icon,
} from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { useMemo, useRef, useState } from "react";
import { useEffect, useMemo, useRef, useState } from "react";
import { useShallow } from "zustand/react/shallow";
export function ModelSelectionStep() {
@ -53,6 +53,7 @@ export function ModelSelectionStep() {
modelType,
selectedModel,
setSelectedModel,
ensureModelDefaultsLoaded,
trainingMethod,
setTrainingMethod,
hfToken,
@ -62,6 +63,7 @@ export function ModelSelectionStep() {
modelType: s.modelType,
selectedModel: s.selectedModel,
setSelectedModel: s.setSelectedModel,
ensureModelDefaultsLoaded: s.ensureModelDefaultsLoaded,
trainingMethod: s.trainingMethod,
setTrainingMethod: s.setTrainingMethod,
hfToken: s.hfToken,
@ -91,6 +93,10 @@ export function ModelSelectionStep() {
hfResults.length,
);
useEffect(() => {
ensureModelDefaultsLoaded();
}, [selectedModel, ensureModelDefaultsLoaded]);
return (
<FieldGroup>
<Field>

View file

@ -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({
<LineChart
data={reducedLossData}
syncId={CHART_SYNC_ID}
syncMethod="value"
accessibilityLayer={true}
margin={{ left: 0, right: 8 }}
>
@ -444,6 +448,7 @@ export function ChartsContent({
<LineChart
data={reducedGradNormData}
syncId={CHART_SYNC_ID}
syncMethod="value"
accessibilityLayer={true}
margin={{ left: 0, right: 8 }}
>
@ -509,6 +514,7 @@ export function ChartsContent({
<LineChart
data={reducedLrData}
syncId={CHART_SYNC_ID}
syncMethod="value"
accessibilityLayer={true}
margin={{ left: 0, right: 8 }}
>
@ -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";
}}
/>
<ChartTooltip
content={
@ -544,7 +553,10 @@ export function ChartsContent({
labelFormatter={(_value, payload) =>
`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"];
}}
/>
}
/>

View file

@ -8,7 +8,7 @@ import {
useTrainingRuntimeStore,
} from "@/features/training";
import { GuidedTour, useGuidedTourController } from "@/features/tour";
import { studioTourSteps, studioTrainingTourSteps } from "@/features/studio/tour";
import { studioTourSteps, studioTrainingTourSteps } from "./tour";
import { ArrowLeft01Icon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { type ReactElement, useEffect } from "react";
@ -31,6 +31,10 @@ export function StudioPage(): ReactElement {
const { dismissTrainingRun } = useTrainingActions();
const config = useTrainingConfigStore();
const selectedModel = useTrainingConfigStore((s) => s.selectedModel);
const ensureModelDefaultsLoaded = useTrainingConfigStore(
(s) => s.ensureModelDefaultsLoaded,
);
const dialogOpen = useDatasetPreviewDialogStore((s) => s.open);
const dialogMode = useDatasetPreviewDialogStore((s) => s.mode);
const dialogInitial = useDatasetPreviewDialogStore((s) => s.initialData);
@ -48,9 +52,14 @@ export function StudioPage(): ReactElement {
autoWhen: isConfigTour,
});
const setTourOpen = tour.setOpen;
useEffect(() => {
tour.setOpen(false);
}, [showTrainingView, tour.setOpen]);
setTourOpen(false);
}, [showTrainingView, setTourOpen]);
useEffect(() => {
ensureModelDefaultsLoaded();
}, [selectedModel, ensureModelDefaultsLoaded]);
return (
<div className="min-h-screen bg-background">

View file

@ -21,7 +21,7 @@ export async function checkDatasetFormat({
body: JSON.stringify({
dataset_name: datasetName,
hf_token: hfToken || undefined,
config: subset || undefined, // backend currently ignores, safe to send
subset: subset || undefined,
split: split || "train",
is_vlm: !!isVlm,
}),

View file

@ -1,8 +1,61 @@
import { authFetch } from "@/features/auth";
interface VisionCheckResponse {
model_name: string;
is_vision: boolean;
model_name: string;
is_vision: boolean;
}
interface BackendTrainingDefaults {
max_seq_length?: number;
num_epochs?: number;
learning_rate?: number | string;
batch_size?: number;
gradient_accumulation_steps?: number;
warmup_steps?: number;
max_steps?: number;
save_steps?: number;
eval_steps?: number;
weight_decay?: number;
random_seed?: number;
packing?: boolean;
train_on_completions?: boolean;
gradient_checkpointing?: "none" | "true" | "unsloth";
}
interface BackendLoraDefaults {
lora_r?: number;
lora_alpha?: number;
lora_dropout?: number;
target_modules?: string[];
use_rslora?: boolean;
use_loftq?: boolean;
finetune_vision_layers?: boolean;
finetune_language_layers?: boolean;
finetune_attention_modules?: boolean;
finetune_mlp_modules?: boolean;
}
interface BackendLoggingDefaults {
enable_wandb?: boolean;
wandb_project?: string;
enable_tensorboard?: boolean;
tensorboard_dir?: string;
log_frequency?: number;
}
export interface BackendModelConfig {
training?: BackendTrainingDefaults;
lora?: BackendLoraDefaults;
logging?: BackendLoggingDefaults;
}
export interface ModelConfigResponse {
id: string;
model_name?: string | null;
config?: BackendModelConfig | null;
is_vision: boolean;
is_lora: boolean;
base_model?: string | null;
}
/**
@ -10,12 +63,24 @@ interface VisionCheckResponse {
* Calls GET /api/models/check-vision/{model_name}.
*/
export async function checkVisionModel(modelName: string): Promise<boolean> {
const encoded = encodeURIComponent(modelName);
const response = await authFetch(`/api/models/check-vision/${encoded}`);
if (!response.ok) {
// If the check fails (e.g. network error), default to non-vision
return false;
}
const data = (await response.json()) as VisionCheckResponse;
return data.is_vision;
const encoded = encodeURIComponent(modelName);
const response = await authFetch(`/api/models/check-vision/${encoded}`);
if (!response.ok) {
// If the check fails (e.g. network error), default to non-vision
return false;
}
const data = (await response.json()) as VisionCheckResponse;
return data.is_vision;
}
export async function getModelConfig(
modelName: string,
signal?: AbortSignal,
): Promise<ModelConfigResponse> {
const encoded = encodeURIComponent(modelName);
const response = await authFetch(`/api/models/config/${encoded}`, { signal });
if (!response.ok) {
throw new Error(`Failed to fetch model config (${response.status})`);
}
return (await response.json()) as ModelConfigResponse;
}

View file

@ -0,0 +1,180 @@
import type { BackendModelConfig } from "../api/models-api";
import type { TrainingConfigState } from "../types/config";
type ModelDefaultsPatch = Partial<
Pick<
TrainingConfigState,
| "epochs"
| "contextLength"
| "learningRate"
| "loraRank"
| "loraAlpha"
| "loraDropout"
| "loraVariant"
| "batchSize"
| "gradientAccumulation"
| "weightDecay"
| "warmupSteps"
| "maxSteps"
| "saveSteps"
| "evalSteps"
| "packing"
| "trainOnCompletions"
| "gradientCheckpointing"
| "randomSeed"
| "enableWandb"
| "wandbProject"
| "enableTensorboard"
| "tensorboardDir"
| "logFrequency"
| "finetuneVisionLayers"
| "finetuneLanguageLayers"
| "finetuneAttentionModules"
| "finetuneMLPModules"
| "targetModules"
>
>;
function toNumber(value: unknown): number | undefined {
if (typeof value === "number" && Number.isFinite(value)) return value;
if (typeof value === "string") {
const parsed = Number(value);
if (Number.isFinite(parsed)) return parsed;
}
return undefined;
}
function toBoolean(value: unknown): boolean | undefined {
if (typeof value === "boolean") return value;
return undefined;
}
function toStringValue(value: unknown): string | undefined {
if (typeof value === "string") return value;
return undefined;
}
function toStringArray(value: unknown): string[] | undefined {
if (!Array.isArray(value)) return undefined;
const result = value.filter((item): item is string => typeof item === "string");
return result.length > 0 ? result : undefined;
}
function toGradientCheckpointing(
value: unknown,
): TrainingConfigState["gradientCheckpointing"] | undefined {
if (value === "none" || value === "true" || value === "unsloth") return value;
return undefined;
}
export function mapBackendModelConfigToTrainingPatch(
config?: BackendModelConfig | null,
): ModelDefaultsPatch {
if (!config) return {};
const patch: ModelDefaultsPatch = {};
const training = config.training;
const lora = config.lora;
const logging = config.logging;
const maxSeqLength = toNumber(training?.max_seq_length);
if (maxSeqLength !== undefined) patch.contextLength = maxSeqLength;
const numEpochs = toNumber(training?.num_epochs);
if (numEpochs !== undefined) patch.epochs = numEpochs;
const learningRate = toNumber(training?.learning_rate);
if (learningRate !== undefined) patch.learningRate = learningRate;
const batchSize = toNumber(training?.batch_size);
if (batchSize !== undefined) patch.batchSize = batchSize;
const gradAccum = toNumber(training?.gradient_accumulation_steps);
if (gradAccum !== undefined) patch.gradientAccumulation = gradAccum;
const warmupSteps = toNumber(training?.warmup_steps);
if (warmupSteps !== undefined) patch.warmupSteps = warmupSteps;
const maxSteps = toNumber(training?.max_steps);
if (maxSteps !== undefined) patch.maxSteps = maxSteps;
const saveSteps = toNumber(training?.save_steps);
if (saveSteps !== undefined) patch.saveSteps = saveSteps;
const evalSteps = toNumber(training?.eval_steps);
if (evalSteps !== undefined) patch.evalSteps = evalSteps;
const weightDecay = toNumber(training?.weight_decay);
if (weightDecay !== undefined) patch.weightDecay = weightDecay;
const randomSeed = toNumber(training?.random_seed);
if (randomSeed !== undefined) patch.randomSeed = randomSeed;
const packing = toBoolean(training?.packing);
if (packing !== undefined) patch.packing = packing;
const trainOnCompletions = toBoolean(training?.train_on_completions);
if (trainOnCompletions !== undefined) {
patch.trainOnCompletions = trainOnCompletions;
}
const gradientCheckpointing = toGradientCheckpointing(
training?.gradient_checkpointing,
);
if (gradientCheckpointing !== undefined) {
patch.gradientCheckpointing = gradientCheckpointing;
}
const loraRank = toNumber(lora?.lora_r);
if (loraRank !== undefined) patch.loraRank = loraRank;
const loraAlpha = toNumber(lora?.lora_alpha);
if (loraAlpha !== undefined) patch.loraAlpha = loraAlpha;
const loraDropout = toNumber(lora?.lora_dropout);
if (loraDropout !== undefined) patch.loraDropout = loraDropout;
const targetModules = toStringArray(lora?.target_modules);
if (targetModules !== undefined) patch.targetModules = targetModules;
if (lora?.use_loftq === true) patch.loraVariant = "loftq";
else if (lora?.use_rslora === true) patch.loraVariant = "rslora";
else if (lora) patch.loraVariant = "lora";
const finetuneVisionLayers = toBoolean(lora?.finetune_vision_layers);
if (finetuneVisionLayers !== undefined) {
patch.finetuneVisionLayers = finetuneVisionLayers;
}
const finetuneLanguageLayers = toBoolean(lora?.finetune_language_layers);
if (finetuneLanguageLayers !== undefined) {
patch.finetuneLanguageLayers = finetuneLanguageLayers;
}
const finetuneAttentionModules = toBoolean(lora?.finetune_attention_modules);
if (finetuneAttentionModules !== undefined) {
patch.finetuneAttentionModules = finetuneAttentionModules;
}
const finetuneMLPModules = toBoolean(lora?.finetune_mlp_modules);
if (finetuneMLPModules !== undefined) {
patch.finetuneMLPModules = finetuneMLPModules;
}
const enableWandb = toBoolean(logging?.enable_wandb);
if (enableWandb !== undefined) patch.enableWandb = enableWandb;
const wandbProject = toStringValue(logging?.wandb_project);
if (wandbProject !== undefined) patch.wandbProject = wandbProject;
const enableTensorboard = toBoolean(logging?.enable_tensorboard);
if (enableTensorboard !== undefined) patch.enableTensorboard = enableTensorboard;
const tensorboardDir = toStringValue(logging?.tensorboard_dir);
if (tensorboardDir !== undefined) patch.tensorboardDir = tensorboardDir;
const logFrequency = toNumber(logging?.log_frequency);
if (logFrequency !== undefined) patch.logFrequency = logFrequency;
return patch;
}

View file

@ -2,9 +2,10 @@ import { DEFAULT_HYPERPARAMS, STEPS } from "@/config/training";
import type { StepNumber } from "@/types/training";
import { create } from "zustand";
import { persist } from "zustand/middleware";
import type { TrainingConfigState, TrainingConfigStore } from "../types/config";
import { checkVisionModel } from "../api/models-api";
import { checkDatasetFormat } from "../api/datasets-api";
import { checkVisionModel, getModelConfig } from "../api/models-api";
import { mapBackendModelConfigToTrainingPatch } from "../lib/model-defaults";
import type { TrainingConfigState, TrainingConfigStore } from "../types/config";
const MIN_STEP: StepNumber = 1;
const MAX_STEP: StepNumber = STEPS.length as StepNumber;
@ -28,18 +29,41 @@ const initialState: TrainingConfigState = {
uploadedFile: null,
isCheckingVision: false,
isVisionModel: false,
isLoadingModelDefaults: false,
modelDefaultsError: null,
modelDefaultsAppliedFor: null,
isCheckingDataset: false,
isDatasetMultimodal: null,
...DEFAULT_HYPERPARAMS,
};
// AbortController for in-flight vision checks so rapid model changes
// cancel stale requests.
let _visionCheckController: AbortController | null = null;
// AbortController for in-flight dataset multimodal checks.
let _datasetCheckController: AbortController | null = null;
// AbortController for in-flight model default loads.
let _modelConfigController: AbortController | null = null;
const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set([
"modelType",
"isCheckingVision",
"isVisionModel",
"isLoadingModelDefaults",
"modelDefaultsError",
"isCheckingDataset",
"isDatasetMultimodal",
]);
function partializePersistedState(
state: TrainingConfigStore,
): Partial<TrainingConfigStore> {
return Object.fromEntries(
Object.entries(state).filter(([key]) => {
const stateKey = key as keyof TrainingConfigState;
return !NON_PERSISTED_STATE_KEYS.has(stateKey);
}),
) as Partial<TrainingConfigStore>;
}
function clampStep(step: number): StepNumber {
return Math.min(MAX_STEP, Math.max(MIN_STEP, step)) as StepNumber;
}
@ -64,167 +88,233 @@ function canProceedForStep(state: TrainingConfigState): boolean {
export const useTrainingConfigStore = create<TrainingConfigStore>()(
persist(
(set, get) => ({
...initialState,
setStep: (step) => set({ currentStep: step }),
nextStep: () => set({ currentStep: clampStep(get().currentStep + 1) }),
prevStep: () => set({ currentStep: clampStep(get().currentStep - 1) }),
setModelType: (modelType) => set({ modelType, selectedModel: null }),
setSelectedModel: (selectedModel) => {
set({ selectedModel });
// Cancel any in-flight vision check
_visionCheckController?.abort();
_visionCheckController = null;
if (!selectedModel) {
set({ isCheckingVision: false });
return;
}
// Fire async backend check to determine if model is vision
(set, get) => {
const loadAndApplyModelDefaults = (modelName: string) => {
_modelConfigController?.abort();
const controller = new AbortController();
_visionCheckController = controller;
set({ isCheckingVision: true });
_modelConfigController = controller;
set({
isLoadingModelDefaults: true,
isCheckingVision: true,
modelDefaultsError: null,
});
checkVisionModel(selectedModel)
.then((isVision) => {
// Only apply if this is still the active check
void getModelConfig(modelName, controller.signal)
.then((modelDetails) => {
if (controller.signal.aborted) return;
if (get().selectedModel !== modelName) return;
set({
isVisionModel: isVision,
...mapBackendModelConfigToTrainingPatch(modelDetails.config),
isVisionModel: modelDetails.is_vision,
isLoadingModelDefaults: false,
isCheckingVision: false,
modelDefaultsError: null,
modelDefaultsAppliedFor: modelName,
});
})
.catch(() => {
.catch((error) => {
if (controller.signal.aborted) return;
// On error, default to text and stop loading
set({ isCheckingVision: false });
});
},
setTrainingMethod: (trainingMethod) => set({ trainingMethod }),
setHfToken: (hfToken) => set({ hfToken }),
setDatasetSource: (datasetSource) => set({ datasetSource }),
setDatasetFormat: (datasetFormat) => set({ datasetFormat }),
setDataset: (dataset) => {
// Cancel any in-flight dataset check
_datasetCheckController?.abort();
_datasetCheckController = null;
set({
dataset,
datasetSubset: null,
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isCheckingDataset: false,
});
},
setDatasetSubset: (datasetSubset) => {
_datasetCheckController?.abort();
_datasetCheckController = null;
set({
datasetSubset,
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isCheckingDataset: false,
});
},
setDatasetSplit: (datasetSplit) => {
_datasetCheckController?.abort();
_datasetCheckController = null;
set({
datasetSplit,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isCheckingDataset: false,
});
// Trigger async dataset multimodal check
const state = get();
const datasetName = state.datasetSource === "huggingface"
? state.dataset
: state.uploadedFile;
if (!datasetName) return;
if (get().selectedModel !== modelName) return;
const controller = new AbortController();
_datasetCheckController = controller;
set({ isCheckingDataset: true });
checkDatasetFormat({
datasetName,
hfToken: state.hfToken.trim() || null,
subset: state.datasetSubset,
split: datasetSplit || "train",
})
.then((res) => {
if (controller.signal.aborted) return;
set({
isDatasetMultimodal: !!res.is_multimodal,
isCheckingDataset: false,
isLoadingModelDefaults: false,
modelDefaultsError:
error instanceof Error
? error.message
: "Failed to load model defaults",
});
})
.catch(() => {
if (controller.signal.aborted) return;
set({ isDatasetMultimodal: null, isCheckingDataset: false });
// Fallback vision check if config endpoint fails.
void checkVisionModel(modelName)
.then((isVision) => {
if (get().selectedModel !== modelName) return;
set({
isVisionModel: isVision,
isCheckingVision: false,
});
})
.catch(() => {
if (get().selectedModel !== modelName) return;
set({ isCheckingVision: false });
});
});
},
setDatasetManualMapping: (datasetManualMapping) =>
set({ datasetManualMapping }),
setUploadedFile: (uploadedFile) => set({ uploadedFile }),
setEpochs: (epochs) => set({ epochs }),
setContextLength: (contextLength) => set({ contextLength }),
setLearningRate: (learningRate) => set({ learningRate }),
setLoraRank: (loraRank) => set({ loraRank }),
setLoraAlpha: (loraAlpha) => set({ loraAlpha }),
setLoraDropout: (loraDropout) => set({ loraDropout }),
setLoraVariant: (loraVariant) => set({ loraVariant }),
setBatchSize: (batchSize) => set({ batchSize }),
setGradientAccumulation: (gradientAccumulation) =>
set({ gradientAccumulation }),
setWeightDecay: (weightDecay) => set({ weightDecay }),
setWarmupSteps: (warmupSteps) => set({ warmupSteps }),
setMaxSteps: (maxSteps) => set({ maxSteps }),
setSaveSteps: (saveSteps) => set({ saveSteps }),
setEvalSteps: (evalSteps) => set({ evalSteps }),
setPacking: (packing) => set({ packing }),
setTrainOnCompletions: (trainOnCompletions) =>
set({ trainOnCompletions }),
setGradientCheckpointing: (gradientCheckpointing) =>
set({ gradientCheckpointing }),
setRandomSeed: (randomSeed) => set({ randomSeed }),
setEnableWandb: (enableWandb) => set({ enableWandb }),
setWandbToken: (wandbToken) => set({ wandbToken }),
setWandbProject: (wandbProject) => set({ wandbProject }),
setEnableTensorboard: (enableTensorboard) => set({ enableTensorboard }),
setTensorboardDir: (tensorboardDir) => set({ tensorboardDir }),
setLogFrequency: (logFrequency) => set({ logFrequency }),
setFinetuneVisionLayers: (finetuneVisionLayers) =>
set({ finetuneVisionLayers }),
setFinetuneLanguageLayers: (finetuneLanguageLayers) =>
set({ finetuneLanguageLayers }),
setFinetuneAttentionModules: (finetuneAttentionModules) =>
set({ finetuneAttentionModules }),
setFinetuneMLPModules: (finetuneMLPModules) => set({ finetuneMLPModules }),
setTargetModules: (targetModules) => set({ targetModules }),
canProceed: () => canProceedForStep(get()),
reset: () => set(initialState),
}),
};
return {
...initialState,
setStep: (step) => set({ currentStep: step }),
nextStep: () => set({ currentStep: clampStep(get().currentStep + 1) }),
prevStep: () => set({ currentStep: clampStep(get().currentStep - 1) }),
setModelType: (modelType) => {
_modelConfigController?.abort();
_modelConfigController = null;
set({
modelType,
selectedModel: null,
isCheckingVision: false,
isVisionModel: false,
isLoadingModelDefaults: false,
modelDefaultsError: null,
modelDefaultsAppliedFor: null,
});
},
setSelectedModel: (selectedModel) => {
const previousModel = get().selectedModel;
set({ selectedModel, modelDefaultsError: null });
if (!selectedModel) {
_modelConfigController?.abort();
_modelConfigController = null;
set({
isCheckingVision: false,
isVisionModel: false,
isLoadingModelDefaults: false,
modelDefaultsError: null,
modelDefaultsAppliedFor: null,
});
return;
}
const shouldLoadDefaults =
selectedModel !== previousModel ||
get().modelDefaultsAppliedFor !== selectedModel;
if (shouldLoadDefaults) {
void loadAndApplyModelDefaults(selectedModel);
}
},
ensureModelDefaultsLoaded: () => {
const state = get();
if (!state.selectedModel) return;
if (state.isLoadingModelDefaults) return;
if (state.modelDefaultsAppliedFor === state.selectedModel) return;
void loadAndApplyModelDefaults(state.selectedModel);
},
setTrainingMethod: (trainingMethod) => set({ trainingMethod }),
setHfToken: (hfToken) => set({ hfToken }),
setDatasetSource: (datasetSource) => set({ datasetSource }),
setDatasetFormat: (datasetFormat) => set({ datasetFormat }),
setDataset: (dataset) => {
_datasetCheckController?.abort();
_datasetCheckController = null;
set({
dataset,
datasetSubset: null,
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isCheckingDataset: false,
});
},
setDatasetSubset: (datasetSubset) => {
_datasetCheckController?.abort();
_datasetCheckController = null;
set({
datasetSubset,
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isCheckingDataset: false,
});
},
setDatasetSplit: (datasetSplit) => {
_datasetCheckController?.abort();
_datasetCheckController = null;
set({
datasetSplit,
datasetManualMapping: emptyManualMapping(),
isDatasetMultimodal: null,
isCheckingDataset: false,
});
const state = get();
const datasetName =
state.datasetSource === "huggingface"
? state.dataset
: state.uploadedFile;
if (!datasetName) return;
const controller = new AbortController();
_datasetCheckController = controller;
set({ isCheckingDataset: true });
checkDatasetFormat({
datasetName,
hfToken: state.hfToken.trim() || null,
subset: state.datasetSubset,
split: datasetSplit || "train",
})
.then((res) => {
if (controller.signal.aborted) return;
set({
isDatasetMultimodal: !!res.is_multimodal,
isCheckingDataset: false,
});
})
.catch(() => {
if (controller.signal.aborted) return;
set({ isDatasetMultimodal: null, isCheckingDataset: false });
});
},
setDatasetManualMapping: (datasetManualMapping) =>
set({ datasetManualMapping }),
setUploadedFile: (uploadedFile) => set({ uploadedFile }),
setEpochs: (epochs) => set({ epochs }),
setContextLength: (contextLength) => set({ contextLength }),
setLearningRate: (learningRate) => set({ learningRate }),
setLoraRank: (loraRank) => set({ loraRank }),
setLoraAlpha: (loraAlpha) => set({ loraAlpha }),
setLoraDropout: (loraDropout) => set({ loraDropout }),
setLoraVariant: (loraVariant) => set({ loraVariant }),
setBatchSize: (batchSize) => set({ batchSize }),
setGradientAccumulation: (gradientAccumulation) =>
set({ gradientAccumulation }),
setWeightDecay: (weightDecay) => set({ weightDecay }),
setWarmupSteps: (warmupSteps) => set({ warmupSteps }),
setMaxSteps: (maxSteps) => set({ maxSteps }),
setSaveSteps: (saveSteps) => set({ saveSteps }),
setEvalSteps: (evalSteps) => set({ evalSteps }),
setPacking: (packing) => set({ packing }),
setTrainOnCompletions: (trainOnCompletions) =>
set({ trainOnCompletions }),
setGradientCheckpointing: (gradientCheckpointing) =>
set({ gradientCheckpointing }),
setRandomSeed: (randomSeed) => set({ randomSeed }),
setEnableWandb: (enableWandb) => set({ enableWandb }),
setWandbToken: (wandbToken) => set({ wandbToken }),
setWandbProject: (wandbProject) => set({ wandbProject }),
setEnableTensorboard: (enableTensorboard) => set({ enableTensorboard }),
setTensorboardDir: (tensorboardDir) => set({ tensorboardDir }),
setLogFrequency: (logFrequency) => set({ logFrequency }),
setFinetuneVisionLayers: (finetuneVisionLayers) =>
set({ finetuneVisionLayers }),
setFinetuneLanguageLayers: (finetuneLanguageLayers) =>
set({ finetuneLanguageLayers }),
setFinetuneAttentionModules: (finetuneAttentionModules) =>
set({ finetuneAttentionModules }),
setFinetuneMLPModules: (finetuneMLPModules) =>
set({ finetuneMLPModules }),
setTargetModules: (targetModules) => set({ targetModules }),
canProceed: () => canProceedForStep(get()),
reset: () => set(initialState),
};
},
{
name: "unsloth_training_config_v1",
version: 2,
version: 3,
migrate: (persisted, version) => {
const s = persisted as Record<string, unknown>;
if (version >= 2) return s as unknown as TrainingConfigStore;
if (s.datasetSubset == null && s.datasetConfig != null) {
if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) {
s.datasetSubset = s.datasetConfig;
}
delete s.datasetConfig;
if (version < 3 && s.modelDefaultsAppliedFor == null) {
s.modelDefaultsAppliedFor = null;
}
return s as unknown as TrainingConfigStore;
},
partialize: (state) => {
const { modelType, isCheckingVision, isVisionModel, isCheckingDataset, isDatasetMultimodal, ...rest } = state;
return rest;
},
partialize: partializePersistedState,
},
),
);

View file

@ -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<TrainingRuntimeStore>()((set) => ({
@ -163,6 +178,7 @@ export const useTrainingRuntimeStore = create<TrainingRuntimeStore>()((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<TrainingRuntimeStore>()((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<TrainingRuntimeStore>()((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<TrainingRuntimeStore>()((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,
};
}),

View file

@ -53,6 +53,9 @@ export interface TrainingConfigState {
logFrequency: number;
isCheckingVision: boolean;
isVisionModel: boolean;
isLoadingModelDefaults: boolean;
modelDefaultsError: string | null;
modelDefaultsAppliedFor: string | null;
isCheckingDataset: boolean;
isDatasetMultimodal: boolean | null;
finetuneVisionLayers: boolean;
@ -68,6 +71,7 @@ export interface TrainingConfigActions {
prevStep: () => void;
setModelType: (type: ModelType) => void;
setSelectedModel: (model: string | null) => void;
ensureModelDefaultsLoaded: () => void;
setTrainingMethod: (method: TrainingMethod) => void;
setHfToken: (token: string) => void;
setDatasetSource: (source: DatasetSource) => void;

View file

@ -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;