From 93b31f0db2bfcdc2e73bb458d8c1c65035cead6b Mon Sep 17 00:00:00 2001 From: samit Date: Thu, 19 Feb 2026 15:03:23 -0800 Subject: [PATCH] added optim in the frontend --- studio/frontend/src/config/training.ts | 10 +++++ .../studio/sections/params-section.tsx | 43 ++++++++++++++++++- .../studio/sections/progress-section.tsx | 7 +++ .../src/features/training/api/mappers.ts | 2 +- .../training/stores/training-config-store.ts | 6 ++- .../src/features/training/types/config.ts | 2 + 6 files changed, 67 insertions(+), 3 deletions(-) diff --git a/studio/frontend/src/config/training.ts b/studio/frontend/src/config/training.ts index 8840b9363c..9ad09ba813 100644 --- a/studio/frontend/src/config/training.ts +++ b/studio/frontend/src/config/training.ts @@ -73,10 +73,20 @@ export const TARGET_MODULES = [ "down_proj", ]; +export const OPTIMIZER_OPTIONS: ReadonlyArray<{ value: string; label: string }> = [ + { value: "adamw_8bit", label: "AdamW 8-bit" }, + { value: "paged_adamw_8bit", label: "Paged AdamW 8-bit" }, + { value: "adamw_bnb_8bit", label: "AdamW BNB 8-bit" }, + { value: "paged_adamw_32bit", label: "Paged AdamW 32-bit" }, + { value: "adamw_torch", label: "AdamW (PyTorch)" }, + { value: "adamw_torch_fused", label: "AdamW (PyTorch Fused)" }, +]; + export const DEFAULT_HYPERPARAMS = { epochs: 3, contextLength: 2048, learningRate: 2e-4, + optimizerType: "adamw_8bit", loraRank: 16, loraAlpha: 32, loraDropout: 0.05, diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 2069da7a90..cf25d5ba54 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -20,7 +20,11 @@ import { TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { CONTEXT_LENGTHS, TARGET_MODULES } from "@/config/training"; +import { + CONTEXT_LENGTHS, + OPTIMIZER_OPTIONS, + TARGET_MODULES, +} from "@/config/training"; import { useTrainingConfigStore } from "@/features/training"; import type { GradientCheckpointing } from "@/types/training"; import { @@ -508,6 +512,43 @@ export function ParamsSection(): ReactElement { value="optimization" className="mt-3 flex flex-col gap-3" > + + Optimization algorithm. 8-bit variants reduce memory usage. + Fused is recommended for vision models.{" "} + + Read more + + + } + > + + o.value === config.optimizerType)?.label ?? + config.optimizerType; + const configItems = [ { section: "Hyperparams", @@ -133,6 +139,7 @@ export function ProgressSection(): ReactElement { ["Epochs", config.epochs], ["Batch size", config.batchSize], ["Learning rate", config.learningRate], + ["Optimizer", optimizerLabel], ["Max steps", config.maxSteps], ["Context length", config.contextLength], ["Warmup steps", config.warmupSteps], diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 510cc542b8..78386f197b 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -40,7 +40,7 @@ export function buildTrainingStartPayload( weight_decay: config.weightDecay, random_seed: config.randomSeed, packing: config.packing, - optim: "adamw_8bit", + optim: config.optimizerType, lr_scheduler_type: "linear", use_lora: adapterMethod, lora_r: config.loraRank, 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 67317dbc09..62d4808b2c 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -305,6 +305,7 @@ export const useTrainingConfigStore = create()( setEpochs: (epochs) => set({ epochs }), setContextLength: (contextLength) => set({ contextLength }), setLearningRate: (learningRate) => set({ learningRate }), + setOptimizerType: (optimizerType) => set({ optimizerType }), setLoraRank: (loraRank) => set({ loraRank }), setLoraAlpha: (loraAlpha) => set({ loraAlpha }), setLoraDropout: (loraDropout) => set({ loraDropout }), @@ -346,7 +347,7 @@ export const useTrainingConfigStore = create()( }, { name: "unsloth_training_config_v1", - version: 3, + version: 4, migrate: (persisted, version) => { const s = persisted as Record; if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) { @@ -356,6 +357,9 @@ export const useTrainingConfigStore = create()( if (version < 3 && s.modelDefaultsAppliedFor == null) { s.modelDefaultsAppliedFor = null; } + if (version < 4 && s.optimizerType == null) { + s.optimizerType = DEFAULT_HYPERPARAMS.optimizerType; + } return s as unknown as TrainingConfigStore; }, partialize: partializePersistedState, diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index 53beabd383..8557739b06 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -30,6 +30,7 @@ export interface TrainingConfigState { epochs: number; contextLength: number; learningRate: number; + optimizerType: string; loraRank: number; loraAlpha: number; loraDropout: number; @@ -85,6 +86,7 @@ export interface TrainingConfigActions { setEpochs: (epochs: number) => void; setContextLength: (length: number) => void; setLearningRate: (rate: number) => void; + setOptimizerType: (value: string) => void; setLoraRank: (rank: number) => void; setLoraAlpha: (alpha: number) => void; setLoraDropout: (dropout: number) => void;