diff --git a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx index d11ff6bac0..3f63775ff9 100644 --- a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx @@ -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 ( diff --git a/studio/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx index f013934530..4677e5a80c 100644 --- a/studio/frontend/src/features/studio/studio-page.tsx +++ b/studio/frontend/src/features/studio/studio-page.tsx @@ -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 (
diff --git a/studio/frontend/src/features/training/api/models-api.ts b/studio/frontend/src/features/training/api/models-api.ts index acb05977ba..44a703b8ca 100644 --- a/studio/frontend/src/features/training/api/models-api.ts +++ b/studio/frontend/src/features/training/api/models-api.ts @@ -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 { - 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 { + 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; } diff --git a/studio/frontend/src/features/training/lib/model-defaults.ts b/studio/frontend/src/features/training/lib/model-defaults.ts new file mode 100644 index 0000000000..07ffb2a422 --- /dev/null +++ b/studio/frontend/src/features/training/lib/model-defaults.ts @@ -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; +} 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 6a792f5fba..2935c0deaa 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -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,6 +29,9 @@ const initialState: TrainingConfigState = { uploadedFile: null, isCheckingVision: false, isVisionModel: false, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: null, isCheckingDataset: false, isDatasetMultimodal: null, ...DEFAULT_HYPERPARAMS, @@ -40,6 +44,9 @@ 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; + function clampStep(step: number): StepNumber { return Math.min(MAX_STEP, Math.max(MIN_STEP, step)) as StepNumber; } @@ -64,165 +71,255 @@ function canProceedForStep(state: TrainingConfigState): boolean { export const useTrainingConfigStore = create()( 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, + 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({ + ...mapBackendModelConfigToTrainingPatch(modelDetails.config), + isVisionModel: modelDetails.is_vision, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: modelName, + }); + }) + .catch((error) => { + if (controller.signal.aborted) return; + if (get().selectedModel !== modelName) return; + + set({ + isLoadingModelDefaults: false, + modelDefaultsError: + error instanceof Error + ? error.message + : "Failed to load model defaults", + }); + }); + }; + + return { + ...initialState, + setStep: (step) => set({ currentStep: step }), + nextStep: () => set({ currentStep: clampStep(get().currentStep + 1) }), + prevStep: () => set({ currentStep: clampStep(get().currentStep - 1) }), + setModelType: (modelType) => { + _visionCheckController?.abort(); + _visionCheckController = null; + _modelConfigController?.abort(); + _modelConfigController = null; + + set({ + modelType, + selectedModel: null, + isCheckingVision: false, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: null, + }); + }, + setSelectedModel: (selectedModel) => { + const previousModel = get().selectedModel; + set({ selectedModel, modelDefaultsError: null }); + + _visionCheckController?.abort(); + _visionCheckController = null; + + if (!selectedModel) { + _modelConfigController?.abort(); + _modelConfigController = null; set({ - isVisionModel: isVision, isCheckingVision: false, + isLoadingModelDefaults: false, + modelDefaultsError: null, + modelDefaultsAppliedFor: null, }); - }) - .catch(() => { - 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; + return; + } - const controller = new AbortController(); - _datasetCheckController = controller; - set({ isCheckingDataset: true }); + const shouldLoadDefaults = + selectedModel !== previousModel || + get().modelDefaultsAppliedFor !== selectedModel; + if (shouldLoadDefaults) { + void loadAndApplyModelDefaults(selectedModel); + } - 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, + // Fire async backend check to determine if model is vision + const controller = new AbortController(); + _visionCheckController = controller; + set({ isCheckingVision: true }); + + checkVisionModel(selectedModel) + .then((isVision) => { + if (controller.signal.aborted) return; + set({ + isVisionModel: isVision, + isCheckingVision: false, + }); + }) + .catch(() => { + if (controller.signal.aborted) return; + set({ isCheckingVision: false }); }); - }) - .catch(() => { - if (controller.signal.aborted) return; - set({ isDatasetMultimodal: null, isCheckingDataset: false }); + }, + 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, }); - }, - 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), - }), + }, + 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; - 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; + const { + modelType, + isCheckingVision, + isVisionModel, + isLoadingModelDefaults, + modelDefaultsError, + isCheckingDataset, + isDatasetMultimodal, + ...rest + } = state; + void modelType; + void isCheckingVision; + void isVisionModel; + void isLoadingModelDefaults; + void modelDefaultsError; + void isCheckingDataset; + void isDatasetMultimodal; return rest; }, }, diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index ad883bc357..67f1d00edc 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -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;