From 09f3b6bce5e84fbc54353d6a4cf72dabf19cb3cb Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Mon, 16 Feb 2026 18:18:28 +0000 Subject: [PATCH 1/2] feat(frontend): auto-detect vision models via backend, separate search filter from model classification --- .../studio/sections/params-section.tsx | 4 +- .../src/features/studio/studio-page.tsx | 2 +- .../src/features/training/api/mappers.ts | 2 +- .../src/features/training/api/models-api.ts | 21 ++++++++++ .../training/hooks/use-training-actions.ts | 2 +- .../training/stores/training-config-store.ts | 42 ++++++++++++++++++- .../src/features/training/types/config.ts | 2 + 7 files changed, 68 insertions(+), 7 deletions(-) create mode 100644 studio/frontend/src/features/training/api/models-api.ts 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; From fa0ca59215cfa041f7ab82bc89b05796c81c3572 Mon Sep 17 00:00:00 2001 From: Roland Tannous Date: Mon, 16 Feb 2026 19:18:49 +0000 Subject: [PATCH 2/2] feat: auto-detect model+dataset compatibility to select VLM vs LLM training path --- studio/backend/core/training/trainer.py | 12 ++-- studio/backend/core/training/training.py | 6 +- studio/backend/models/training.py | 1 + studio/backend/routes/training.py | 1 + .../studio/sections/params-section.tsx | 8 +-- .../studio/sections/training-section.tsx | 10 ++- .../src/features/studio/studio-page.tsx | 2 +- .../src/features/training/api/mappers.ts | 3 +- .../training/hooks/use-training-actions.ts | 2 +- .../training/stores/training-config-store.ts | 67 +++++++++++++++++-- .../src/features/training/types/api.ts | 1 + .../src/features/training/types/config.ts | 2 + 12 files changed, 94 insertions(+), 21 deletions(-) diff --git a/studio/backend/core/training/trainer.py b/studio/backend/core/training/trainer.py index b3116c5c04..8a164193ea 100644 --- a/studio/backend/core/training/trainer.py +++ b/studio/backend/core/training/trainer.py @@ -110,17 +110,21 @@ class UnslothTrainer: model_name: str, max_seq_length: int = 2048, load_in_4bit: bool = True, - hf_token: Optional[str] = None) -> bool: + hf_token: Optional[str] = None, + is_dataset_multimodal: bool = False) -> bool: """Load model for training (supports both text and vision models)""" try: print("\nClearing GPU memory before training...") clear_gpu_cache() - # Detect if this is a vision model first - self.is_vlm = is_vision_model(model_name) + # Detect if this is a vision model AND dataset is multimodal + # A vision-capable model with a text-only dataset should use FastLanguageModel + self.is_vlm = is_vision_model(model_name) and is_dataset_multimodal self.model_name = model_name - logger.info(f"Model type detected: {'Vision' if self.is_vlm else 'Text'}") + logger.info(f"Model architecture is vision: {is_vision_model(model_name)}") + logger.info(f"Dataset is multimodal: {is_dataset_multimodal}") + logger.info(f"Using VLM path: {self.is_vlm}") # Reset training state for new run self._update_progress( diff --git a/studio/backend/core/training/training.py b/studio/backend/core/training/training.py index d9aaa8ca0e..251b246f35 100644 --- a/studio/backend/core/training/training.py +++ b/studio/backend/core/training/training.py @@ -96,7 +96,8 @@ class TrainingBackend: # Optional parameters custom_format_mapping: dict = None, subset: str = None, - split: str = "train") -> bool: + split: str = "train", + is_dataset_multimodal: bool = False) -> bool: """ Start training. @@ -150,7 +151,8 @@ class TrainingBackend: model_name=model_name, max_seq_length=max_seq_length, load_in_4bit=load_in_4bit if use_lora_actual else False, # Only 4bit for LoRA - hf_token=hf_token if hf_token.strip() else None + hf_token=hf_token if hf_token.strip() else None, + is_dataset_multimodal=is_dataset_multimodal, ) if not success or self.trainer.should_stop: diff --git a/studio/backend/models/training.py b/studio/backend/models/training.py index e0839ae485..dcc6aff0ae 100644 --- a/studio/backend/models/training.py +++ b/studio/backend/models/training.py @@ -55,6 +55,7 @@ class TrainingStartRequest(BaseModel): finetune_language_layers: bool = Field(False, description="Finetune language layers") finetune_attention_modules: bool = Field(False, description="Finetune attention modules") finetune_mlp_modules: bool = Field(False, description="Finetune MLP modules") + is_dataset_multimodal: bool = Field(False, description="Whether the dataset contains multimodal (image) data") # Logging parameters enable_wandb: bool = Field(False, description="Enable Weights & Biases logging") diff --git a/studio/backend/routes/training.py b/studio/backend/routes/training.py index 5bf59d85d2..13f7135ba4 100644 --- a/studio/backend/routes/training.py +++ b/studio/backend/routes/training.py @@ -176,6 +176,7 @@ async def start_training( "finetune_language_layers": request.finetune_language_layers, "finetune_attention_modules": request.finetune_attention_modules, "finetune_mlp_modules": request.finetune_mlp_modules, + "is_dataset_multimodal": request.is_dataset_multimodal, "enable_wandb": request.enable_wandb, "wandb_token": request.wandb_token or "", "wandb_project": request.wandb_project or "", diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 0b85cafe7e..9a3495120d 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.isVisionModel; + const showVisionLora = store.isVisionModel && store.isDatasetMultimodal === true; const [loraOpen, setLoraOpen] = useState(false); const [hyperOpen, setHyperOpen] = useState(false); @@ -350,7 +350,7 @@ export function ParamsSection(): ReactElement { /> {/* Vision checkboxes */} - {isVision && ( + {showVisionLora && (
{( [ @@ -400,7 +400,7 @@ export function ParamsSection(): ReactElement { )} {/* Text target modules */} - {!isVision && ( + {!showVisionLora && (
Target Modules @@ -693,7 +693,7 @@ export function ParamsSection(): ReactElement { - {!store.isVisionModel && ( + {!showVisionLora && (
@@ -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 c4da41db3d..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.isVisionModel} + 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 fc61d70269..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.isVisionModel) { + if (config.isVisionModel && config.isDatasetMultimodal) { return { [input]: "image", [output]: "text" }; } 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 b7336a6c3f..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.isVisionModel; + 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 8a686ff414..09f453694d 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -4,6 +4,7 @@ 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; @@ -27,6 +28,8 @@ const initialState: TrainingConfigState = { uploadedFile: null, isCheckingVision: false, isVisionModel: false, + isCheckingDataset: false, + isDatasetMultimodal: null, ...DEFAULT_HYPERPARAMS, }; @@ -34,6 +37,9 @@ const initialState: TrainingConfigState = { // 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; } @@ -100,21 +106,68 @@ export const useTrainingConfigStore = create()( 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 }), @@ -168,7 +221,7 @@ export const useTrainingConfigStore = create()( return s as unknown as TrainingConfigStore; }, partialize: (state) => { - const { modelType, isCheckingVision, isVisionModel, ...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 476cfc267d..551d947617 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -52,6 +52,8 @@ export interface TrainingConfigState { logFrequency: number; isCheckingVision: boolean; isVisionModel: boolean; + isCheckingDataset: boolean; + isDatasetMultimodal: boolean | null; finetuneVisionLayers: boolean; finetuneLanguageLayers: boolean; finetuneAttentionModules: boolean;