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:
parent
bbd7d6d122
commit
3badf6649c
6 changed files with 514 additions and 153 deletions
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
180
studio/frontend/src/features/training/lib/model-defaults.ts
Normal file
180
studio/frontend/src/features/training/lib/model-defaults.ts
Normal 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;
|
||||
}
|
||||
|
|
@ -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;
|
||||
},
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue