diff --git a/studio/frontend/src/config/training.ts b/studio/frontend/src/config/training.ts index 8840b9363c..da60328d40 100644 --- a/studio/frontend/src/config/training.ts +++ b/studio/frontend/src/config/training.ts @@ -73,10 +73,26 @@ 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 LR_SCHEDULER_OPTIONS: ReadonlyArray<{ value: string; label: string }> = [ + { value: "linear", label: "Linear" }, + { value: "cosine", label: "Cosine" }, +]; + export const DEFAULT_HYPERPARAMS = { epochs: 3, contextLength: 2048, learningRate: 2e-4, + optimizerType: "adamw_8bit", + lrSchedulerType: "linear", 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..de82fb0fb5 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -20,7 +20,12 @@ import { TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { CONTEXT_LENGTHS, TARGET_MODULES } from "@/config/training"; +import { + CONTEXT_LENGTHS, + LR_SCHEDULER_OPTIONS, + OPTIMIZER_OPTIONS, + TARGET_MODULES, +} from "@/config/training"; import { useTrainingConfigStore } from "@/features/training"; import type { GradientCheckpointing } from "@/types/training"; import { @@ -508,6 +513,78 @@ 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 + + + } + > + + + + How the learning rate changes over training. Linear decays + steadily; cosine decays in a curve.{" "} + + 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..ac93761b62 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -40,8 +40,8 @@ export function buildTrainingStartPayload( weight_decay: config.weightDecay, random_seed: config.randomSeed, packing: config.packing, - optim: "adamw_8bit", - lr_scheduler_type: "linear", + optim: config.optimizerType, + lr_scheduler_type: config.lrSchedulerType, use_lora: adapterMethod, lora_r: config.loraRank, lora_alpha: config.loraAlpha, diff --git a/studio/frontend/src/features/training/api/models-api.ts b/studio/frontend/src/features/training/api/models-api.ts index 6e6e3ff762..e22de5f581 100644 --- a/studio/frontend/src/features/training/api/models-api.ts +++ b/studio/frontend/src/features/training/api/models-api.ts @@ -9,6 +9,8 @@ interface BackendTrainingDefaults { max_seq_length?: number; num_epochs?: number; learning_rate?: number | string; + optim?: string; + lr_scheduler_type?: string; batch_size?: number; gradient_accumulation_steps?: number; warmup_steps?: number; diff --git a/studio/frontend/src/features/training/lib/model-defaults.ts b/studio/frontend/src/features/training/lib/model-defaults.ts index 07ffb2a422..35ce562dbf 100644 --- a/studio/frontend/src/features/training/lib/model-defaults.ts +++ b/studio/frontend/src/features/training/lib/model-defaults.ts @@ -7,6 +7,8 @@ type ModelDefaultsPatch = Partial< | "epochs" | "contextLength" | "learningRate" + | "optimizerType" + | "lrSchedulerType" | "loraRank" | "loraAlpha" | "loraDropout" @@ -86,6 +88,12 @@ export function mapBackendModelConfigToTrainingPatch( const learningRate = toNumber(training?.learning_rate); if (learningRate !== undefined) patch.learningRate = learningRate; + const optim = toStringValue(training?.optim); + if (optim !== undefined) patch.optimizerType = optim; + + const lrSchedulerType = toStringValue(training?.lr_scheduler_type); + if (lrSchedulerType !== undefined) patch.lrSchedulerType = lrSchedulerType; + const batchSize = toNumber(training?.batch_size); if (batchSize !== undefined) patch.batchSize = batchSize; 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..b6f9c01c42 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,8 @@ export const useTrainingConfigStore = create()( setEpochs: (epochs) => set({ epochs }), setContextLength: (contextLength) => set({ contextLength }), setLearningRate: (learningRate) => set({ learningRate }), + setOptimizerType: (optimizerType) => set({ optimizerType }), + setLrSchedulerType: (lrSchedulerType) => set({ lrSchedulerType }), setLoraRank: (loraRank) => set({ loraRank }), setLoraAlpha: (loraAlpha) => set({ loraAlpha }), setLoraDropout: (loraDropout) => set({ loraDropout }), @@ -346,7 +348,7 @@ export const useTrainingConfigStore = create()( }, { name: "unsloth_training_config_v1", - version: 3, + version: 5, migrate: (persisted, version) => { const s = persisted as Record; if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) { @@ -356,6 +358,12 @@ export const useTrainingConfigStore = create()( if (version < 3 && s.modelDefaultsAppliedFor == null) { s.modelDefaultsAppliedFor = null; } + if (version < 4 && s.optimizerType == null) { + s.optimizerType = DEFAULT_HYPERPARAMS.optimizerType; + } + if (version < 5 && s.lrSchedulerType == null) { + s.lrSchedulerType = DEFAULT_HYPERPARAMS.lrSchedulerType; + } 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..b268a08b79 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -30,6 +30,8 @@ export interface TrainingConfigState { epochs: number; contextLength: number; learningRate: number; + optimizerType: string; + lrSchedulerType: string; loraRank: number; loraAlpha: number; loraDropout: number; @@ -85,6 +87,8 @@ export interface TrainingConfigActions { setEpochs: (epochs: number) => void; setContextLength: (length: number) => void; setLearningRate: (rate: number) => void; + setOptimizerType: (value: string) => void; + setLrSchedulerType: (value: string) => void; setLoraRank: (rank: number) => void; setLoraAlpha: (alpha: number) => void; setLoraDropout: (dropout: number) => void;