Compare commits

...
Sign in to create a new pull request.

2 commits

Author SHA1 Message Date
Manan Shah
1dbf7d085c
Merge branch 'main' into fix/checking-dataset-logic 2026-03-17 01:16:50 -05:00
Manan17
fedcfc45ea Revamping the checking model+dataset logic 2026-03-17 06:14:21 +00:00
3 changed files with 11 additions and 12 deletions

View file

@ -80,9 +80,6 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
multimodal_info = detect_multimodal_dataset(dataset) multimodal_info = detect_multimodal_dataset(dataset)
is_audio = multimodal_info.get("is_audio", False) is_audio = multimodal_info.get("is_audio", False)
if multimodal_info["is_image"]:
is_vlm = True # Route to VLM detection for image datasets
# Common audio fields for all return paths # Common audio fields for all return paths
audio_fields = { audio_fields = {
"is_audio": is_audio, "is_audio": is_audio,
@ -153,8 +150,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": heuristic_mapping, "suggested_mapping": heuristic_mapping,
"detected_image_column": None, "detected_image_column": None,
"detected_text_column": None, "detected_text_column": None,
"is_image": False, "is_image": multimodal_info["is_image"],
"multimodal_columns": None, "multimodal_columns": multimodal_info.get("multimodal_columns"),
**audio_fields, **audio_fields,
} }
else: else:
@ -166,8 +163,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": None, "suggested_mapping": None,
"detected_image_column": None, "detected_image_column": None,
"detected_text_column": None, "detected_text_column": None,
"is_image": False, "is_image": multimodal_info["is_image"],
"multimodal_columns": None, "multimodal_columns": multimodal_info.get("multimodal_columns"),
"warning": ( "warning": (
f"Could not auto-detect column roles for columns: {columns}. " f"Could not auto-detect column roles for columns: {columns}. "
"Please assign roles manually, or use AI Assist." "Please assign roles manually, or use AI Assist."
@ -183,8 +180,8 @@ def check_dataset_format(dataset, is_vlm: bool = False) -> dict:
"suggested_mapping": None, "suggested_mapping": None,
"detected_image_column": None, "detected_image_column": None,
"detected_text_column": None, "detected_text_column": None,
"is_image": False, "is_image": multimodal_info["is_image"],
"multimodal_columns": None, "multimodal_columns": multimodal_info.get("multimodal_columns"),
**audio_fields, **audio_fields,
} }

View file

@ -46,7 +46,8 @@ export function TrainingSection() {
const store = useTrainingConfigStore(); const store = useTrainingConfigStore();
const { isStarting, startError, startTrainingRun } = useTrainingActions(); const { isStarting, startError, startTrainingRun } = useTrainingActions();
const isIncompatible = const isIncompatible =
!store.isVisionModel && store.isDatasetImage === true; (!store.isVisionModel && store.isDatasetImage === true) ||
(!store.isAudioModel && store.isDatasetAudio === true);
const configValidation = validateTrainingConfig(store); const configValidation = validateTrainingConfig(store);
const fileInputRef = useRef<HTMLInputElement>(null); const fileInputRef = useRef<HTMLInputElement>(null);
@ -155,10 +156,10 @@ export function TrainingSection() {
data-tour="studio-start" 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" 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()} onClick={() => void startTrainingRun()}
disabled={isStarting || isIncompatible || !configValidation.ok} disabled={isStarting || isIncompatible || store.isCheckingDataset || !configValidation.ok}
> >
<HugeiconsIcon icon={Rocket01Icon} className="size-4" /> <HugeiconsIcon icon={Rocket01Icon} className="size-4" />
{isStarting ? "Starting..." : "Start Training"} {isStarting ? "Starting..." : store.isCheckingDataset ? "Checking dataset..." : "Start Training"}
</Button> </Button>
{startError && ( {startError && (
<p className="text-xs text-red-500 leading-relaxed">{startError}</p> <p className="text-xs text-red-500 leading-relaxed">{startError}</p>

View file

@ -210,6 +210,7 @@ export const useTrainingConfigStore = create<TrainingConfigStore>()(
hfToken: state.hfToken.trim() || null, hfToken: state.hfToken.trim() || null,
subset: state.datasetSubset, subset: state.datasetSubset,
split, split,
isVlm: state.isVisionModel,
}) })
.then((res) => { .then((res) => {
if (controller.signal.aborted) return; if (controller.signal.aborted) return;