Merge branch 'nightly' into integrate/exports-page
This commit is contained in:
commit
b8171a86ac
22 changed files with 733 additions and 196 deletions
3
requirements/base.txt
Normal file
3
requirements/base.txt
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
# Core unsloth packages
|
||||
unsloth-zoo
|
||||
unsloth
|
||||
13
requirements/extras-no-deps.txt
Normal file
13
requirements/extras-no-deps.txt
Normal 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
56
requirements/extras.txt
Normal 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
|
||||
7
requirements/overrides.txt
Normal file
7
requirements/overrides.txt
Normal 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
13
requirements/studio.txt
Normal 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
|
||||
2
requirements/triton-kernels.txt
Normal file
2
requirements/triton-kernels.txt
Normal 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
|
||||
18
setup.sh
18
setup.sh
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -1,7 +0,0 @@
|
|||
fastapi>=0.100.0
|
||||
uvicorn>=0.27.0
|
||||
pydantic>=2.0
|
||||
torch
|
||||
psutil
|
||||
nest-asyncio>=1.5.8
|
||||
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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"];
|
||||
}}
|
||||
/>
|
||||
}
|
||||
/>
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}),
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
180
studio/frontend/src/features/training/lib/model-defaults.ts
Normal file
180
studio/frontend/src/features/training/lib/model-defaults.ts
Normal 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;
|
||||
}
|
||||
|
|
@ -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,
|
||||
},
|
||||
),
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
};
|
||||
}),
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue