From e1fbccfc57d3ab490ba9cf99aa4523ec89cdccc2 Mon Sep 17 00:00:00 2001 From: imagineer99 Date: Tue, 17 Feb 2026 05:53:21 +0000 Subject: [PATCH 1/6] fix: subset param name --- studio/frontend/src/features/training/api/datasets-api.ts | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/studio/frontend/src/features/training/api/datasets-api.ts b/studio/frontend/src/features/training/api/datasets-api.ts index 37f3744701..7bba75ca38 100644 --- a/studio/frontend/src/features/training/api/datasets-api.ts +++ b/studio/frontend/src/features/training/api/datasets-api.ts @@ -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, }), From e803e13d3e262cd74901206e762999d474fb4190 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 17 Feb 2026 14:21:18 +0000 Subject: [PATCH 2/6] add full dependency chain for unsloth + unsloth-extras --- .gitignore | 1 + requirements/base.txt | 3 ++ requirements/extras-no-deps.txt | 13 ++++++++ requirements/extras.txt | 56 +++++++++++++++++++++++++++++++++ requirements/overrides.txt | 7 +++++ requirements/studio.txt | 13 ++++++++ requirements/triton-kernels.txt | 2 ++ setup.sh | 16 ++++++++-- studio/backend/requirements.txt | 7 ----- 9 files changed, 109 insertions(+), 9 deletions(-) create mode 100644 requirements/base.txt create mode 100644 requirements/extras-no-deps.txt create mode 100644 requirements/extras.txt create mode 100644 requirements/overrides.txt create mode 100644 requirements/studio.txt create mode 100644 requirements/triton-kernels.txt delete mode 100644 studio/backend/requirements.txt diff --git a/.gitignore b/.gitignore index e07aa496b8..25a5ba54a4 100755 --- a/.gitignore +++ b/.gitignore @@ -18,6 +18,7 @@ unsloth_compiled_cache/ # ML artifacts (large files) outputs/ +exports/ *.gguf *.safetensors diff --git a/requirements/base.txt b/requirements/base.txt new file mode 100644 index 0000000000..407ae01b52 --- /dev/null +++ b/requirements/base.txt @@ -0,0 +1,3 @@ +# Core unsloth packages +unsloth-zoo +unsloth diff --git a/requirements/extras-no-deps.txt b/requirements/extras-no-deps.txt new file mode 100644 index 0000000000..c2a9c6bad8 --- /dev/null +++ b/requirements/extras-no-deps.txt @@ -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 diff --git a/requirements/extras.txt b/requirements/extras.txt new file mode 100644 index 0000000000..3ed20faa8b --- /dev/null +++ b/requirements/extras.txt @@ -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 diff --git a/requirements/overrides.txt b/requirements/overrides.txt new file mode 100644 index 0000000000..02770f3953 --- /dev/null +++ b/requirements/overrides.txt @@ -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 diff --git a/requirements/studio.txt b/requirements/studio.txt new file mode 100644 index 0000000000..fd7da2626d --- /dev/null +++ b/requirements/studio.txt @@ -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 diff --git a/requirements/triton-kernels.txt b/requirements/triton-kernels.txt new file mode 100644 index 0000000000..17e265b35e --- /dev/null +++ b/requirements/triton-kernels.txt @@ -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 diff --git a/setup.sh b/setup.sh index bb5da805a8..3a24a17e61 100755 --- a/setup.sh +++ b/setup.sh @@ -132,15 +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 +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/rolandtannous/unsloth-zoo/fix/use-sys-executable-for-pip-and-python/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. diff --git a/studio/backend/requirements.txt b/studio/backend/requirements.txt deleted file mode 100644 index 3f97cca66c..0000000000 --- a/studio/backend/requirements.txt +++ /dev/null @@ -1,7 +0,0 @@ -fastapi>=0.100.0 -uvicorn>=0.27.0 -pydantic>=2.0 -torch -psutil -nest-asyncio>=1.5.8 - From 3badf6649cebe7a97cfcc0b758ca042488b73b4d Mon Sep 17 00:00:00 2001 From: Shine1i Date: Tue, 17 Feb 2026 17:58:36 +0100 Subject: [PATCH 3/6] feat: add default model configuration mapping and auto-apply logic - Implemented backend model configuration mapping to training state. - Added auto-apply logic for default configurations when models are selected. - Introduced utilities for type conversion and validation within training configuration. --- .../components/steps/model-selection-step.tsx | 8 +- .../src/features/studio/studio-page.tsx | 15 +- .../src/features/training/api/models-api.ts | 85 +++- .../features/training/lib/model-defaults.ts | 180 +++++++++ .../training/stores/training-config-store.ts | 379 +++++++++++------- .../src/features/training/types/config.ts | 4 + 6 files changed, 516 insertions(+), 155 deletions(-) create mode 100644 studio/frontend/src/features/training/lib/model-defaults.ts diff --git a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx index d11ff6bac0..3f63775ff9 100644 --- a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx @@ -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 ( diff --git a/studio/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx index f013934530..4677e5a80c 100644 --- a/studio/frontend/src/features/studio/studio-page.tsx +++ b/studio/frontend/src/features/studio/studio-page.tsx @@ -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 (
diff --git a/studio/frontend/src/features/training/api/models-api.ts b/studio/frontend/src/features/training/api/models-api.ts index acb05977ba..44a703b8ca 100644 --- a/studio/frontend/src/features/training/api/models-api.ts +++ b/studio/frontend/src/features/training/api/models-api.ts @@ -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 { - 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 { + 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; } diff --git a/studio/frontend/src/features/training/lib/model-defaults.ts b/studio/frontend/src/features/training/lib/model-defaults.ts new file mode 100644 index 0000000000..07ffb2a422 --- /dev/null +++ b/studio/frontend/src/features/training/lib/model-defaults.ts @@ -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; +} diff --git a/studio/frontend/src/features/training/stores/training-config-store.ts b/studio/frontend/src/features/training/stores/training-config-store.ts index 6a792f5fba..2935c0deaa 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -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,6 +29,9 @@ const initialState: TrainingConfigState = { uploadedFile: null, isCheckingVision: false, isVisionModel: false, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: null, isCheckingDataset: false, isDatasetMultimodal: null, ...DEFAULT_HYPERPARAMS, @@ -40,6 +44,9 @@ 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; + function clampStep(step: number): StepNumber { return Math.min(MAX_STEP, Math.max(MIN_STEP, step)) as StepNumber; } @@ -64,165 +71,255 @@ function canProceedForStep(state: TrainingConfigState): boolean { export const useTrainingConfigStore = create()( 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, + 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({ + ...mapBackendModelConfigToTrainingPatch(modelDetails.config), + isVisionModel: modelDetails.is_vision, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: modelName, + }); + }) + .catch((error) => { + if (controller.signal.aborted) return; + if (get().selectedModel !== modelName) return; + + set({ + isLoadingModelDefaults: false, + modelDefaultsError: + error instanceof Error + ? error.message + : "Failed to load model defaults", + }); + }); + }; + + return { + ...initialState, + setStep: (step) => set({ currentStep: step }), + nextStep: () => set({ currentStep: clampStep(get().currentStep + 1) }), + prevStep: () => set({ currentStep: clampStep(get().currentStep - 1) }), + setModelType: (modelType) => { + _visionCheckController?.abort(); + _visionCheckController = null; + _modelConfigController?.abort(); + _modelConfigController = null; + + set({ + modelType, + selectedModel: null, + isCheckingVision: false, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: null, + }); + }, + setSelectedModel: (selectedModel) => { + const previousModel = get().selectedModel; + set({ selectedModel, modelDefaultsError: null }); + + _visionCheckController?.abort(); + _visionCheckController = null; + + if (!selectedModel) { + _modelConfigController?.abort(); + _modelConfigController = null; set({ - isVisionModel: isVision, isCheckingVision: false, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: null, }); - }) - .catch(() => { - 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; + return; + } - const controller = new AbortController(); - _datasetCheckController = controller; - set({ isCheckingDataset: true }); + const shouldLoadDefaults = + selectedModel !== previousModel || + get().modelDefaultsAppliedFor !== selectedModel; + if (shouldLoadDefaults) { + void loadAndApplyModelDefaults(selectedModel); + } - 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, + // Fire async backend check to determine if model is vision + const controller = new AbortController(); + _visionCheckController = controller; + set({ isCheckingVision: true }); + + checkVisionModel(selectedModel) + .then((isVision) => { + if (controller.signal.aborted) return; + set({ + isVisionModel: isVision, + isCheckingVision: false, + }); + }) + .catch(() => { + if (controller.signal.aborted) return; + set({ isCheckingVision: false }); }); - }) - .catch(() => { - if (controller.signal.aborted) return; - set({ isDatasetMultimodal: null, isCheckingDataset: false }); + }, + 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, }); - }, - 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), - }), + }, + 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; - 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; + const { + modelType, + isCheckingVision, + isVisionModel, + isLoadingModelDefaults, + modelDefaultsError, + isCheckingDataset, + isDatasetMultimodal, + ...rest + } = state; + void modelType; + void isCheckingVision; + void isVisionModel; + void isLoadingModelDefaults; + void modelDefaultsError; + void isCheckingDataset; + void isDatasetMultimodal; return rest; }, }, diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index ad883bc357..67f1d00edc 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -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; From 7203766fac8151bebd9abf942c565092d1b9c3eb Mon Sep 17 00:00:00 2001 From: Shine1i Date: Tue, 17 Feb 2026 18:01:45 +0100 Subject: [PATCH 4/6] refactor: streamline vision model detection and improve state persistence logic - Removed redundant vision-check controllers. - Added `NON_PERSISTED_STATE_KEYS` to manage persisted training state. - Introduced `partializePersistedState` for cleaner state filtering. --- .../training/stores/training-config-store.ts | 87 +++++++++---------- 1 file changed, 40 insertions(+), 47 deletions(-) diff --git a/studio/frontend/src/features/training/stores/training-config-store.ts b/studio/frontend/src/features/training/stores/training-config-store.ts index 2935c0deaa..3d7d352817 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -37,16 +37,33 @@ const initialState: TrainingConfigState = { ...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 = new Set([ + "modelType", + "isCheckingVision", + "isVisionModel", + "isLoadingModelDefaults", + "modelDefaultsError", + "isCheckingDataset", + "isDatasetMultimodal", +]); + +function partializePersistedState( + state: TrainingConfigStore, +): Partial { + return Object.fromEntries( + Object.entries(state).filter(([key]) => { + const stateKey = key as keyof TrainingConfigState; + return !NON_PERSISTED_STATE_KEYS.has(stateKey); + }), + ) as Partial; +} + function clampStep(step: number): StepNumber { return Math.min(MAX_STEP, Math.max(MIN_STEP, step)) as StepNumber; } @@ -78,6 +95,7 @@ export const useTrainingConfigStore = create()( _modelConfigController = controller; set({ isLoadingModelDefaults: true, + isCheckingVision: true, modelDefaultsError: null, }); @@ -90,6 +108,7 @@ export const useTrainingConfigStore = create()( ...mapBackendModelConfigToTrainingPatch(modelDetails.config), isVisionModel: modelDetails.is_vision, isLoadingModelDefaults: false, + isCheckingVision: false, modelDefaultsError: null, modelDefaultsAppliedFor: modelName, }); @@ -105,6 +124,20 @@ export const useTrainingConfigStore = create()( ? error.message : "Failed to load model defaults", }); + + // 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 }); + }); }); }; @@ -114,8 +147,6 @@ export const useTrainingConfigStore = create()( nextStep: () => set({ currentStep: clampStep(get().currentStep + 1) }), prevStep: () => set({ currentStep: clampStep(get().currentStep - 1) }), setModelType: (modelType) => { - _visionCheckController?.abort(); - _visionCheckController = null; _modelConfigController?.abort(); _modelConfigController = null; @@ -123,6 +154,7 @@ export const useTrainingConfigStore = create()( modelType, selectedModel: null, isCheckingVision: false, + isVisionModel: false, isLoadingModelDefaults: false, modelDefaultsError: null, modelDefaultsAppliedFor: null, @@ -132,14 +164,12 @@ export const useTrainingConfigStore = create()( const previousModel = get().selectedModel; set({ selectedModel, modelDefaultsError: null }); - _visionCheckController?.abort(); - _visionCheckController = null; - if (!selectedModel) { _modelConfigController?.abort(); _modelConfigController = null; set({ isCheckingVision: false, + isVisionModel: false, isLoadingModelDefaults: false, modelDefaultsError: null, modelDefaultsAppliedFor: null, @@ -153,24 +183,6 @@ export const useTrainingConfigStore = create()( if (shouldLoadDefaults) { void loadAndApplyModelDefaults(selectedModel); } - - // Fire async backend check to determine if model is vision - const controller = new AbortController(); - _visionCheckController = controller; - set({ isCheckingVision: true }); - - checkVisionModel(selectedModel) - .then((isVision) => { - if (controller.signal.aborted) return; - set({ - isVisionModel: isVision, - isCheckingVision: false, - }); - }) - .catch(() => { - if (controller.signal.aborted) return; - set({ isCheckingVision: false }); - }); }, ensureModelDefaultsLoaded: () => { const state = get(); @@ -302,26 +314,7 @@ export const useTrainingConfigStore = create()( } return s as unknown as TrainingConfigStore; }, - partialize: (state) => { - const { - modelType, - isCheckingVision, - isVisionModel, - isLoadingModelDefaults, - modelDefaultsError, - isCheckingDataset, - isDatasetMultimodal, - ...rest - } = state; - void modelType; - void isCheckingVision; - void isVisionModel; - void isLoadingModelDefaults; - void modelDefaultsError; - void isCheckingDataset; - void isDatasetMultimodal; - return rest; - }, + partialize: partializePersistedState, }, ), ); From 0be3e6f5252e797d1e7b4ff1a69ff10b7fda5715 Mon Sep 17 00:00:00 2001 From: Shine1i Date: Tue, 17 Feb 2026 18:26:59 +0100 Subject: [PATCH 5/6] feat: integrate gradient norm tracking in training runtime and metrics - Enhanced chart logic to filter and visualize finite gradient norm values. --- studio/backend/core/training/training.py | 13 +++++ studio/backend/models/responses.py | 2 + studio/backend/models/training.py | 3 +- studio/backend/routes/training.py | 35 ++++++++++-- .../studio/sections/charts-content.tsx | 26 ++++++--- .../training/stores/training-runtime-store.ts | 53 ++++++++++++++----- .../src/features/training/types/runtime.ts | 4 ++ 7 files changed, 111 insertions(+), 25 deletions(-) diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index 51d411be78..9b589d215d 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -4,6 +4,7 @@ Training backend for FastAPI integration import matplotlib.pyplot as plt from typing import Any, Generator, Tuple import logging +import math from .trainer import get_trainer, TrainingProgress from utils.hardware import clear_gpu_cache @@ -28,6 +29,8 @@ class TrainingBackend: self.loss_history = [] self.lr_history = [] self.step_history = [] + self.grad_norm_history = [] + self.grad_norm_step_history = [] self.eval_loss_history = [] self.eval_step_history = [] self.eval_enabled = False @@ -43,6 +46,14 @@ class TrainingBackend: self.loss_history.append(progress.loss) self.lr_history.append(progress.learning_rate) self.step_history.append(progress.step) + if progress.step >= 0 and progress.grad_norm is not None: + try: + grad_norm = float(progress.grad_norm) + except (TypeError, ValueError): + grad_norm = None + if grad_norm is not None and math.isfinite(grad_norm): + self.grad_norm_history.append(grad_norm) + self.grad_norm_step_history.append(progress.step) if progress.eval_loss is not None: self.eval_loss_history.append(progress.eval_loss) self.eval_step_history.append(progress.step) @@ -144,6 +155,8 @@ class TrainingBackend: self.loss_history = [] self.lr_history = [] self.step_history = [] + self.grad_norm_history = [] + self.grad_norm_step_history = [] self.eval_loss_history = [] self.eval_step_history = [] self.eval_enabled = False diff --git a/studio/backend/models/responses.py b/studio/backend/models/responses.py index 2aa798c5c9..c72dfc54dd 100644 --- a/studio/backend/models/responses.py +++ b/studio/backend/models/responses.py @@ -19,6 +19,8 @@ class TrainingMetricsResponse(BaseModel): loss_history: List[float] = Field(default_factory=list, description="Loss values per step") lr_history: List[float] = Field(default_factory=list, description="Learning rate per step") step_history: List[int] = Field(default_factory=list, description="Step numbers") + grad_norm_history: List[float] = Field(default_factory=list, description="Gradient norm values") + grad_norm_step_history: List[int] = Field(default_factory=list, description="Step numbers for gradient norm values") current_loss: Optional[float] = Field(None, description="Most recent loss value") current_lr: Optional[float] = Field(None, description="Most recent learning rate") current_step: Optional[int] = Field(None, description="Most recent step number") diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index fd04baf3a0..2b989e6a82 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -104,7 +104,7 @@ class TrainingStatus(BaseModel): metric_history: Optional[dict] = Field( None, description="Full metric history arrays for chart recovery after SSE reconnection. " - "Keys: 'steps', 'loss', 'lr' — each a list of numeric values.", + "Keys: 'steps', 'loss', 'lr', 'grad_norm', 'grad_norm_steps' — each a list of numeric values.", ) @@ -122,4 +122,3 @@ class TrainingProgress(BaseModel): grad_norm: Optional[float] = Field(None, description="L2 norm of gradients, computed before gradient clipping") num_tokens: Optional[int] = Field(None, description="Total number of tokens processed so far") eval_loss: Optional[float] = Field(None, description="Eval loss from the most recent evaluation step") - diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index c09b814fa1..1770999284 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -325,6 +325,8 @@ async def reset_training( backend.loss_history = [] backend.lr_history = [] backend.step_history = [] + backend.grad_norm_history = [] + backend.grad_norm_step_history = [] return {"status": "ok"} except Exception as e: logger.error(f"Error resetting training: {e}", exc_info=True) @@ -408,6 +410,8 @@ async def get_training_status( "steps": list(backend.step_history), "loss": list(backend.loss_history), "lr": list(backend.lr_history), + "grad_norm": list(getattr(backend, "grad_norm_history", [])), + "grad_norm_steps": list(getattr(backend, "grad_norm_step_history", [])), "eval_loss": list(backend.eval_loss_history), "eval_steps": list(backend.eval_step_history), } @@ -445,6 +449,8 @@ async def get_training_metrics( loss_history = backend.loss_history lr_history = backend.lr_history step_history = backend.step_history + grad_norm_history = getattr(backend, "grad_norm_history", []) + grad_norm_step_history = getattr(backend, "grad_norm_step_history", []) # Get current values current_loss = loss_history[-1] if loss_history else None @@ -455,6 +461,8 @@ async def get_training_metrics( loss_history=loss_history, lr_history=lr_history, step_history=step_history, + grad_norm_history=grad_norm_history, + grad_norm_step_history=grad_norm_step_history, current_loss=current_loss, current_lr=current_lr, current_step=current_step, @@ -505,6 +513,8 @@ async def stream_training_progress( total_steps: int, epoch: Optional[float] = None, progress: Optional[Any] = None, + grad_norm_override: Optional[float] = None, + eval_loss_override: Optional[float] = None, ) -> TrainingProgress: total = max(total_steps, 0) if step < 0 or total == 0: @@ -517,9 +527,13 @@ async def stream_training_progress( # Get actual values from progress object if available elapsed_seconds = getattr(progress, 'elapsed_seconds', None) if progress else None eta_seconds = getattr(progress, 'eta_seconds', None) if progress else None - grad_norm = getattr(progress, 'grad_norm', None) if progress else None + grad_norm = grad_norm_override + if grad_norm is None and progress: + grad_norm = getattr(progress, 'grad_norm', None) num_tokens = getattr(progress, 'num_tokens', None) if progress else None - eval_loss = getattr(progress, 'eval_loss', None) if progress else None + eval_loss = eval_loss_override + if eval_loss is None and progress: + eval_loss = getattr(progress, 'eval_loss', None) return TrainingProgress( job_id=job_id, @@ -558,6 +572,13 @@ async def stream_training_progress( # ── Replay missed steps on reconnect ───────────────────── if resume_from_step is not None and backend.step_history: replayed = 0 + grad_norm_by_step = { + step_val: grad_val + for step_val, grad_val in zip( + getattr(backend, "grad_norm_step_history", []), + getattr(backend, "grad_norm_history", []), + ) + } for i, step_val in enumerate(backend.step_history): if step_val > resume_from_step: loss_val = backend.loss_history[i] if i < len(backend.loss_history) else 0.0 @@ -567,7 +588,15 @@ async def stream_training_progress( ) total_replay = getattr(tp_replay, "total_steps", step_val) if tp_replay else step_val epoch_replay = getattr(tp_replay, "epoch", None) if tp_replay else None - payload = build_progress(step_val, loss_val, lr_val, total_replay, epoch_replay, progress=tp_replay) + payload = build_progress( + step_val, + loss_val, + lr_val, + total_replay, + epoch_replay, + progress=tp_replay, + grad_norm_override=grad_norm_by_step.get(step_val), + ) yield format_sse(payload.model_dump_json(), event="progress", event_id=step_val) replayed += 1 if replayed: diff --git a/studio/frontend/src/features/studio/sections/charts-content.tsx b/studio/frontend/src/features/studio/sections/charts-content.tsx index 61cbb0a3ea..bb7a7a131c 100644 --- a/studio/frontend/src/features/studio/sections/charts-content.tsx +++ b/studio/frontend/src/features/studio/sections/charts-content.tsx @@ -117,12 +117,13 @@ function buildStepTicks(min: number, max: number, targetCount = 6): number[] { } function buildYDomain(values: number[]): [number, number] { - if (values.length === 0) { + const finiteValues = values.filter((value) => Number.isFinite(value)); + if (finiteValues.length === 0) { return [0, 1]; } - const min = Math.min(...values); - const max = Math.max(...values); + const min = Math.min(...finiteValues); + const max = Math.max(...finiteValues); if (min === max) { const base = Math.abs(min); @@ -240,7 +241,8 @@ export function ChartsContent({ (point) => point.step >= visibleStepDomain[0] && point.step <= visibleStepDomain[1], ) - .map((point) => point.gradNorm), + .map((point) => point.gradNorm) + .filter((value) => Number.isFinite(value)), [reducedGradNormData, visibleStepDomain], ); @@ -251,7 +253,8 @@ export function ChartsContent({ (point) => point.step >= visibleStepDomain[0] && point.step <= visibleStepDomain[1], ) - .map((point) => point.lr), + .map((point) => point.lr) + .filter((value) => Number.isFinite(value)), [reducedLrData, visibleStepDomain], ); @@ -345,6 +348,7 @@ export function ChartsContent({ @@ -444,6 +448,7 @@ export function ChartsContent({ @@ -509,6 +514,7 @@ export function ChartsContent({ @@ -536,7 +542,10 @@ export function ChartsContent({ tickMargin={4} fontSize={10} width={52} - tickFormatter={(value) => Number(value).toExponential(0)} + tickFormatter={(value) => { + const num = Number(value); + return Number.isFinite(num) ? num.toExponential(0) : "0e+0"; + }} /> `Step ${payload?.[0]?.payload?.step ?? ""}` } - formatter={(value) => [Number(value).toExponential(3), "LR"]} + formatter={(value) => { + const num = Number(value); + return [Number.isFinite(num) ? num.toExponential(3) : "0e+0", "LR"]; + }} /> } /> diff --git a/studio/frontend/src/features/training/stores/training-runtime-store.ts b/studio/frontend/src/features/training/stores/training-runtime-store.ts index fd0d3d5733..6a3756a8ca 100644 --- a/studio/frontend/src/features/training/stores/training-runtime-store.ts +++ b/studio/frontend/src/features/training/stores/training-runtime-store.ts @@ -56,6 +56,11 @@ function toSeries(steps: number[], values: number[]): TrainingSeriesPoint[] { return sortSeries(points); } +function toFiniteNumber(value: unknown): number | null { + if (typeof value !== "number") return null; + return Number.isFinite(value) ? value : null; +} + function upsertPoint( points: TrainingSeriesPoint[], step: number, @@ -74,22 +79,32 @@ function upsertPoint( function applyMetricHistoryFromStatus(payload: TrainingStatusResponse): { lossHistory: TrainingSeriesPoint[] | null; lrHistory: TrainingSeriesPoint[] | null; + gradNormHistory: TrainingSeriesPoint[] | null; evalLossHistory: TrainingSeriesPoint[] | null; } { const history = payload.metric_history; if (!history || !history.steps?.length) { - return { lossHistory: null, lrHistory: null, evalLossHistory: null }; + return { + lossHistory: null, + lrHistory: null, + gradNormHistory: null, + evalLossHistory: null, + }; } const steps = history.steps; const lossHistory = history.loss ? toSeries(steps, history.loss) : null; const lrHistory = history.lr ? toSeries(steps, history.lr) : null; + const gradNormHistory = + history.grad_norm && history.grad_norm_steps + ? toSeries(history.grad_norm_steps, history.grad_norm) + : null; const evalLossHistory = history.eval_loss && history.eval_steps ? toSeries(history.eval_steps, history.eval_loss) : null; - return { lossHistory, lrHistory, evalLossHistory }; + return { lossHistory, lrHistory, gradNormHistory, evalLossHistory }; } export const useTrainingRuntimeStore = create()((set) => ({ @@ -163,6 +178,7 @@ export const useTrainingRuntimeStore = create()((set) => ( typeof detailEpoch === "number" ? detailEpoch : state.currentEpoch, lossHistory: metricHistory.lossHistory ?? state.lossHistory, lrHistory: metricHistory.lrHistory ?? state.lrHistory, + gradNormHistory: metricHistory.gradNormHistory ?? state.gradNormHistory, evalLossHistory: metricHistory.evalLossHistory ?? state.evalLossHistory, }; }), @@ -171,6 +187,10 @@ export const useTrainingRuntimeStore = create()((set) => ( set((state) => { const lossHistory = toSeries(payload.step_history, payload.loss_history); const lrHistory = toSeries(payload.step_history, payload.lr_history); + const gradNormHistory = toSeries( + payload.grad_norm_step_history, + payload.grad_norm_history, + ); const latestStep = payload.current_step ?? (payload.step_history.length > 0 @@ -181,6 +201,8 @@ export const useTrainingRuntimeStore = create()((set) => ( ...state, lossHistory: lossHistory.length > 0 ? lossHistory : state.lossHistory, lrHistory: lrHistory.length > 0 ? lrHistory : state.lrHistory, + gradNormHistory: + gradNormHistory.length > 0 ? gradNormHistory : state.gradNormHistory, currentStep: typeof latestStep === "number" ? Math.max(latestStep, state.currentStep) @@ -199,36 +221,41 @@ export const useTrainingRuntimeStore = create()((set) => ( applyProgress: (payload: TrainingProgressPayload, eventId?: number) => set((state) => { const step = Math.max(payload.step, 0); + const currentLoss = toFiniteNumber(payload.loss); + const currentLearningRate = toFiniteNumber(payload.learning_rate); + const currentGradNorm = toFiniteNumber(payload.grad_norm); + const evalLoss = toFiniteNumber(payload.eval_loss); + return { ...state, jobId: payload.job_id || state.jobId, currentStep: step, totalSteps: Math.max(payload.total_steps, state.totalSteps), - currentLoss: payload.loss, - currentLearningRate: payload.learning_rate, + currentLoss: currentLoss ?? state.currentLoss, + currentLearningRate: currentLearningRate ?? state.currentLearningRate, progressPercent: payload.progress_percent, currentEpoch: payload.epoch ?? state.currentEpoch, elapsedSeconds: payload.elapsed_seconds, etaSeconds: payload.eta_seconds, - currentGradNorm: payload.grad_norm, + currentGradNorm, currentNumTokens: payload.num_tokens, firstStepReceived: state.firstStepReceived || step > 0, lastEventId: typeof eventId === "number" ? eventId : state.lastEventId, lossHistory: - step > 0 - ? upsertPoint(state.lossHistory, step, payload.loss) + step > 0 && currentLoss !== null + ? upsertPoint(state.lossHistory, step, currentLoss) : state.lossHistory, lrHistory: - step > 0 - ? upsertPoint(state.lrHistory, step, payload.learning_rate) + step > 0 && currentLearningRate !== null + ? upsertPoint(state.lrHistory, step, currentLearningRate) : state.lrHistory, gradNormHistory: - step > 0 && typeof payload.grad_norm === "number" - ? upsertPoint(state.gradNormHistory, step, payload.grad_norm) + step > 0 && currentGradNorm !== null + ? upsertPoint(state.gradNormHistory, step, currentGradNorm) : state.gradNormHistory, evalLossHistory: - step > 0 && typeof payload.eval_loss === "number" - ? upsertPoint(state.evalLossHistory, step, payload.eval_loss) + step > 0 && evalLoss !== null + ? upsertPoint(state.evalLossHistory, step, evalLoss) : state.evalLossHistory, }; }), diff --git a/studio/frontend/src/features/training/types/runtime.ts b/studio/frontend/src/features/training/types/runtime.ts index 8e0b43bd52..df24c020ca 100644 --- a/studio/frontend/src/features/training/types/runtime.ts +++ b/studio/frontend/src/features/training/types/runtime.ts @@ -26,6 +26,8 @@ export interface TrainingStatusResponse { steps?: number[]; loss?: number[]; lr?: number[]; + grad_norm?: number[]; + grad_norm_steps?: number[]; eval_loss?: number[]; eval_steps?: number[]; } | null; @@ -35,6 +37,8 @@ export interface TrainingMetricsResponse { loss_history: number[]; lr_history: number[]; step_history: number[]; + grad_norm_history: number[]; + grad_norm_step_history: number[]; current_loss: number | null; current_lr: number | null; current_step: number | null; From caa28e5d46a87393eefa7c3225ecdf570ec639e8 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Tue, 17 Feb 2026 17:55:05 +0000 Subject: [PATCH 6/6] replace with patch from merged PR in unsloth-zoo --- setup.sh | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/setup.sh b/setup.sh index 3a24a17e61..e5ea68c407 100755 --- a/setup.sh +++ b/setup.sh @@ -146,7 +146,7 @@ run_quiet "pip install torchao+transformers" pip install --force-reinstall --no- 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/rolandtannous/unsloth-zoo/fix/use-sys-executable-for-pip-and-python/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 studio" pip install -r "$SCRIPT_DIR/requirements/studio.txt"