@@ -98,7 +101,7 @@ export function TrainingSection() {
data-tour="studio-start"
className="w-full cursor-pointer bg-gradient-to-r from-emerald-500 to-teal-500 text-white hover:from-emerald-600 hover:to-teal-600"
onClick={() => void startTrainingRun()}
- disabled={isStarting}
+ disabled={isStarting || isIncompatible}
>
{isStarting ? "Starting..." : "Start Training"}
@@ -106,6 +109,11 @@ export function TrainingSection() {
{startError && (
{startError}
)}
+ {isIncompatible && (
+
+ Text model is not compatible with a multimodal dataset. Switch to a vision model or choose a text-only dataset.
+
+ )}
{/* Save / Clear */}
diff --git a/studio/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx
index ebf4e3881b..f013934530 100644
--- a/studio/frontend/src/features/studio/studio-page.tsx
+++ b/studio/frontend/src/features/studio/studio-page.tsx
@@ -70,7 +70,7 @@ export function StudioPage(): ReactElement {
datasetSplit={config.datasetSplit}
mode={dialogMode}
initialData={dialogInitial}
- isVlm={config.modelType === "vision"}
+ isVlm={config.isVisionModel && config.isDatasetMultimodal === true}
/>
{canGoBack && (
diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts
index 05fdd98959..ac02218e64 100644
--- a/studio/frontend/src/features/training/api/mappers.ts
+++ b/studio/frontend/src/features/training/api/mappers.ts
@@ -54,6 +54,7 @@ export function buildTrainingStartPayload(
finetune_language_layers: config.finetuneLanguageLayers,
finetune_attention_modules: config.finetuneAttentionModules,
finetune_mlp_modules: config.finetuneMLPModules,
+ is_dataset_multimodal: !!config.isDatasetMultimodal,
enable_wandb: config.enableWandb,
wandb_token: config.enableWandb ? config.wandbToken.trim() || null : null,
wandb_project: config.enableWandb
@@ -72,7 +73,7 @@ function buildCustomFormatMapping(
const { input, output } = config.datasetManualMapping;
if (!input || !output) return undefined;
- if (config.modelType === "vision") {
+ if (config.isVisionModel && config.isDatasetMultimodal) {
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..1a2a1aef8c 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 && config.isDatasetMultimodal === true;
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..09f453694d 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,8 @@ 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";
const MIN_STEP: StepNumber = 1;
const MAX_STEP: StepNumber = STEPS.length as StepNumber;
@@ -24,9 +26,20 @@ const initialState: TrainingConfigState = {
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
uploadedFile: null,
+ isCheckingVision: false,
+ isVisionModel: false,
+ isCheckingDataset: false,
+ isDatasetMultimodal: null,
...DEFAULT_HYPERPARAMS,
};
+// AbortController for in-flight vision checks so rapid model changes
+// cancel stale requests.
+let _visionCheckController: AbortController | null = null;
+
+// AbortController for in-flight dataset multimodal checks.
+let _datasetCheckController: AbortController | null = null;
+
function clampStep(step: number): StepNumber {
return Math.min(MAX_STEP, Math.max(MIN_STEP, step)) as StepNumber;
}
@@ -57,26 +70,104 @@ 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 }),
setDatasetFormat: (datasetFormat) => set({ datasetFormat }),
- setDataset: (dataset) =>
+ setDataset: (dataset) => {
+ // Cancel any in-flight dataset check
+ _datasetCheckController?.abort();
+ _datasetCheckController = null;
set({
dataset,
datasetSubset: null,
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
- }),
- setDatasetSubset: (datasetSubset) =>
+ isDatasetMultimodal: null,
+ isCheckingDataset: false,
+ });
+ },
+ setDatasetSubset: (datasetSubset) => {
+ _datasetCheckController?.abort();
+ _datasetCheckController = null;
set({
datasetSubset,
datasetSplit: null,
datasetManualMapping: emptyManualMapping(),
- }),
- setDatasetSplit: (datasetSplit) =>
- set({ datasetSplit, 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;
+
+ 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 }),
@@ -130,7 +221,7 @@ export const useTrainingConfigStore = create()(
return s as unknown as TrainingConfigStore;
},
partialize: (state) => {
- const { modelType, ...rest } = state;
+ const { modelType, isCheckingVision, isVisionModel, isCheckingDataset, isDatasetMultimodal, ...rest } = state;
return rest;
},
},
diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts
index f6c43616d9..95de92cfdb 100644
--- a/studio/frontend/src/features/training/types/api.ts
+++ b/studio/frontend/src/features/training/types/api.ts
@@ -36,6 +36,7 @@ export interface TrainingStartRequest {
finetune_language_layers: boolean;
finetune_attention_modules: boolean;
finetune_mlp_modules: boolean;
+ is_dataset_multimodal: boolean;
enable_wandb: boolean;
wandb_token: string | null;
wandb_project: string | null;
diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts
index 6a28500b7e..551d947617 100644
--- a/studio/frontend/src/features/training/types/config.ts
+++ b/studio/frontend/src/features/training/types/config.ts
@@ -50,6 +50,10 @@ export interface TrainingConfigState {
enableTensorboard: boolean;
tensorboardDir: string;
logFrequency: number;
+ isCheckingVision: boolean;
+ isVisionModel: boolean;
+ isCheckingDataset: boolean;
+ isDatasetMultimodal: boolean | null;
finetuneVisionLayers: boolean;
finetuneLanguageLayers: boolean;
finetuneAttentionModules: boolean;