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;