diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx
index 415aca3fac..0b85cafe7e 100644
--- a/studio/frontend/src/features/studio/sections/params-section.tsx
+++ b/studio/frontend/src/features/studio/sections/params-section.tsx
@@ -109,7 +109,7 @@ function SliderRow({
export function ParamsSection(): ReactElement {
const store = useTrainingConfigStore();
const isLora = store.trainingMethod !== "full";
- const isVision = store.modelType === "vision";
+ const isVision = store.isVisionModel;
const [loraOpen, setLoraOpen] = useState(false);
const [hyperOpen, setHyperOpen] = useState(false);
@@ -693,7 +693,7 @@ export function ParamsSection(): ReactElement {
- {store.modelType !== "vision" && (
+ {!store.isVisionModel && (
{canGoBack && (
diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts
index 05fdd98959..fc61d70269 100644
--- a/studio/frontend/src/features/training/api/mappers.ts
+++ b/studio/frontend/src/features/training/api/mappers.ts
@@ -72,7 +72,7 @@ function buildCustomFormatMapping(
const { input, output } = config.datasetManualMapping;
if (!input || !output) return undefined;
- if (config.modelType === "vision") {
+ if (config.isVisionModel) {
return { [input]: "image", [output]: "text" };
}
diff --git a/studio/frontend/src/features/training/api/models-api.ts b/studio/frontend/src/features/training/api/models-api.ts
new file mode 100644
index 0000000000..acb05977ba
--- /dev/null
+++ b/studio/frontend/src/features/training/api/models-api.ts
@@ -0,0 +1,21 @@
+import { authFetch } from "@/features/auth";
+
+interface VisionCheckResponse {
+ model_name: string;
+ is_vision: boolean;
+}
+
+/**
+ * Check whether a model is a vision model by asking the backend.
+ * Calls GET /api/models/check-vision/{model_name}.
+ */
+export async function checkVisionModel(modelName: string): Promise {
+ 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;
+}
diff --git a/studio/frontend/src/features/training/hooks/use-training-actions.ts b/studio/frontend/src/features/training/hooks/use-training-actions.ts
index e3600432e9..b7336a6c3f 100644
--- a/studio/frontend/src/features/training/hooks/use-training-actions.ts
+++ b/studio/frontend/src/features/training/hooks/use-training-actions.ts
@@ -29,7 +29,7 @@ export function useTrainingActions() {
try {
const datasetName = getDatasetName(config);
- const isVlm = config.modelType === "vision";
+ const isVlm = config.isVisionModel;
if (datasetName) {
const check = await checkDatasetFormat({
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 731ded9d80..8a686ff414 100644
--- a/studio/frontend/src/features/training/stores/training-config-store.ts
+++ b/studio/frontend/src/features/training/stores/training-config-store.ts
@@ -3,6 +3,7 @@ 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";
const MIN_STEP: StepNumber = 1;
const MAX_STEP: StepNumber = STEPS.length as StepNumber;
@@ -24,9 +25,15 @@ const initialState: TrainingConfigState = {
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
uploadedFile: null,
+ isCheckingVision: false,
+ isVisionModel: false,
...DEFAULT_HYPERPARAMS,
};
+// AbortController for in-flight vision checks so rapid model changes
+// cancel stale requests.
+let _visionCheckController: AbortController | null = null;
+
function clampStep(step: number): StepNumber {
return Math.min(MAX_STEP, Math.max(MIN_STEP, step)) as StepNumber;
}
@@ -57,7 +64,38 @@ export const useTrainingConfigStore = create()(
nextStep: () => set({ currentStep: clampStep(get().currentStep + 1) }),
prevStep: () => set({ currentStep: clampStep(get().currentStep - 1) }),
setModelType: (modelType) => set({ modelType, selectedModel: null }),
- setSelectedModel: (selectedModel) => set({ selectedModel }),
+ 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
+ const controller = new AbortController();
+ _visionCheckController = controller;
+ set({ isCheckingVision: true });
+
+ checkVisionModel(selectedModel)
+ .then((isVision) => {
+ // Only apply if this is still the active check
+ if (controller.signal.aborted) return;
+ set({
+ isVisionModel: isVision,
+ isCheckingVision: false,
+ });
+ })
+ .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 }),
@@ -130,7 +168,7 @@ export const useTrainingConfigStore = create()(
return s as unknown as TrainingConfigStore;
},
partialize: (state) => {
- const { modelType, ...rest } = state;
+ const { modelType, isCheckingVision, isVisionModel, ...rest } = state;
return rest;
},
},
diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts
index 6a28500b7e..476cfc267d 100644
--- a/studio/frontend/src/features/training/types/config.ts
+++ b/studio/frontend/src/features/training/types/config.ts
@@ -50,6 +50,8 @@ export interface TrainingConfigState {
enableTensorboard: boolean;
tensorboardDir: string;
logFrequency: number;
+ isCheckingVision: boolean;
+ isVisionModel: boolean;
finetuneVisionLayers: boolean;
finetuneLanguageLayers: boolean;
finetuneAttentionModules: boolean;