Merge pull request #354 from unslothai/fix/audio-train-completions

fix: uncheck train_on_completions for audio models
This commit is contained in:
Roland Tannous 2026-03-11 00:05:51 +04:00 committed by GitHub
commit b91cdda2b9
3 changed files with 29 additions and 3 deletions

View file

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

View file

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

View file

@ -60,6 +60,7 @@ export interface TrainingConfigState {
logFrequency: number;
isCheckingVision: boolean;
isVisionModel: boolean;
isAudioModel: boolean;
isLoadingModelDefaults: boolean;
modelDefaultsError: string | null;
modelDefaultsAppliedFor: string | null;