feat: add default model configuration mapping and auto-apply logic

- Implemented backend model configuration mapping to training state.
- Added auto-apply logic for default configurations when models are selected.
- Introduced utilities for type conversion and validation within training configuration.
This commit is contained in:
Shine1i 2026-02-17 17:58:36 +01:00
commit 3badf6649c
6 changed files with 514 additions and 153 deletions

View file

@ -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 (
<FieldGroup>
<Field>

View file

@ -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 (
<div className="min-h-screen bg-background">

View file

@ -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<boolean> {
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<ModelConfigResponse> {
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;
}

View file

@ -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;
}

View file

@ -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<TrainingConfigStore>()(
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<string, unknown>;
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;
},
},

View file

@ -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;