diff --git a/studio/frontend/src/features/training/api/models-api.ts b/studio/frontend/src/features/training/api/models-api.ts index f40f29fda2..cbf203d05a 100644 --- a/studio/frontend/src/features/training/api/models-api.ts +++ b/studio/frontend/src/features/training/api/models-api.ts @@ -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; } 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 502889ee1e..30d03e7fc5 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -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 = new Set([ "modelType", "isCheckingVision", + "isAudioModel", "isLoadingModelDefaults", "modelDefaultsError", "modelDefaultsAppliedFor", @@ -127,6 +129,16 @@ export const useTrainingConfigStore = create()( 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()( ...patch, modelType: inferredModelType, isVisionModel: modelDetails.is_vision, + isAudioModel: isAudio, isLoadingModelDefaults: false, isCheckingVision: false, modelDefaultsError: null, @@ -148,6 +161,7 @@ export const useTrainingConfigStore = create()( set({ isLoadingModelDefaults: false, + isAudioModel: false, modelDefaultsError: error instanceof Error ? error.message @@ -161,12 +175,13 @@ export const useTrainingConfigStore = create()( 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()( 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()( selectedModel: null, isCheckingVision: false, isVisionModel: false, + isAudioModel: false, isDatasetAudio: false, isLoadingModelDefaults: false, modelDefaultsError: null, @@ -249,6 +273,7 @@ export const useTrainingConfigStore = create()( set({ isCheckingVision: false, isVisionModel: false, + isAudioModel: false, isDatasetAudio: false, isLoadingModelDefaults: false, modelDefaultsError: null, diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index bc8305430b..5a427dd153 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -60,6 +60,7 @@ export interface TrainingConfigState { logFrequency: number; isCheckingVision: boolean; isVisionModel: boolean; + isAudioModel: boolean; isLoadingModelDefaults: boolean; modelDefaultsError: string | null; modelDefaultsAppliedFor: string | null;