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;