Merge pull request #354 from unslothai/fix/audio-train-completions
fix: uncheck train_on_completions for audio models
This commit is contained in:
commit
b91cdda2b9
3 changed files with 29 additions and 3 deletions
|
|
@ -61,8 +61,8 @@ export interface ModelConfigResponse {
|
|||
model_name?: string | null;
|
||||
config?: BackendModelConfig | null;
|
||||
is_vision: boolean;
|
||||
is_audio: boolean;
|
||||
is_lora: boolean;
|
||||
is_audio?: boolean;
|
||||
base_model?: string | null;
|
||||
model_type?: "text" | "vision" | "audio" | "embeddings" | null;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -36,6 +36,7 @@ const initialState: TrainingConfigState = {
|
|||
uploadedFile: null,
|
||||
isCheckingVision: false,
|
||||
isVisionModel: false,
|
||||
isAudioModel: false,
|
||||
isLoadingModelDefaults: false,
|
||||
modelDefaultsError: null,
|
||||
modelDefaultsAppliedFor: null,
|
||||
|
|
@ -58,6 +59,7 @@ let _trainOnCompletionsManuallySet = false;
|
|||
const NON_PERSISTED_STATE_KEYS: ReadonlySet<keyof TrainingConfigState> = new Set([
|
||||
"modelType",
|
||||
"isCheckingVision",
|
||||
"isAudioModel",
|
||||
"isLoadingModelDefaults",
|
||||
"modelDefaultsError",
|
||||
"modelDefaultsAppliedFor",
|
||||
|
|
@ -127,6 +129,16 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
patch.trainOnCompletions = false;
|
||||
}
|
||||
|
||||
const isAudio = !!modelDetails.is_audio;
|
||||
// Pure audio model → always uncheck trainOnCompletions.
|
||||
if (isAudio && !modelDetails.is_vision) {
|
||||
patch.trainOnCompletions = false;
|
||||
}
|
||||
// Audio-capable vision model (e.g. gemma3n) + audio dataset → uncheck.
|
||||
if (isAudio && modelDetails.is_vision && get().isDatasetAudio) {
|
||||
patch.trainOnCompletions = false;
|
||||
}
|
||||
|
||||
// Use backend-provided model_type when available, otherwise
|
||||
// infer from is_vision (temporary until backend ships model_type).
|
||||
const inferredModelType: ModelType = modelDetails.model_type
|
||||
|
|
@ -136,6 +148,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
...patch,
|
||||
modelType: inferredModelType,
|
||||
isVisionModel: modelDetails.is_vision,
|
||||
isAudioModel: isAudio,
|
||||
isLoadingModelDefaults: false,
|
||||
isCheckingVision: false,
|
||||
modelDefaultsError: null,
|
||||
|
|
@ -148,6 +161,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
|
||||
set({
|
||||
isLoadingModelDefaults: false,
|
||||
isAudioModel: false,
|
||||
modelDefaultsError:
|
||||
error instanceof Error
|
||||
? error.message
|
||||
|
|
@ -161,12 +175,13 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
set({
|
||||
modelType: isVision ? "vision" : "text",
|
||||
isVisionModel: isVision,
|
||||
isAudioModel: false,
|
||||
isCheckingVision: false,
|
||||
});
|
||||
})
|
||||
.catch(() => {
|
||||
if (get().selectedModel !== modelName) return;
|
||||
set({ isCheckingVision: false });
|
||||
set({ isCheckingVision: false, isAudioModel: false });
|
||||
});
|
||||
});
|
||||
};
|
||||
|
|
@ -194,10 +209,18 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
isCheckingDataset: false,
|
||||
};
|
||||
if (!_trainOnCompletionsManuallySet) {
|
||||
const { isVisionModel } = get();
|
||||
const { isVisionModel, isAudioModel } = get();
|
||||
if (isVisionModel && isImage) {
|
||||
updates.trainOnCompletions = false;
|
||||
}
|
||||
// Pure audio model → always uncheck regardless of dataset.
|
||||
if (isAudioModel && !isVisionModel) {
|
||||
updates.trainOnCompletions = false;
|
||||
}
|
||||
// Audio-capable vision model (e.g. gemma3n) + audio dataset → uncheck.
|
||||
if (isAudioModel && isVisionModel && isAudio) {
|
||||
updates.trainOnCompletions = false;
|
||||
}
|
||||
}
|
||||
set(updates);
|
||||
})
|
||||
|
|
@ -233,6 +256,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
selectedModel: null,
|
||||
isCheckingVision: false,
|
||||
isVisionModel: false,
|
||||
isAudioModel: false,
|
||||
isDatasetAudio: false,
|
||||
isLoadingModelDefaults: false,
|
||||
modelDefaultsError: null,
|
||||
|
|
@ -249,6 +273,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
|
|||
set({
|
||||
isCheckingVision: false,
|
||||
isVisionModel: false,
|
||||
isAudioModel: false,
|
||||
isDatasetAudio: false,
|
||||
isLoadingModelDefaults: false,
|
||||
modelDefaultsError: null,
|
||||
|
|
|
|||
|
|
@ -60,6 +60,7 @@ export interface TrainingConfigState {
|
|||
logFrequency: number;
|
||||
isCheckingVision: boolean;
|
||||
isVisionModel: boolean;
|
||||
isAudioModel: boolean;
|
||||
isLoadingModelDefaults: boolean;
|
||||
modelDefaultsError: string | null;
|
||||
modelDefaultsAppliedFor: string | null;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue