diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py index 72298dc245..3a4d54f93f 100644 --- a/studio/backend/utils/datasets/dataset_utils.py +++ b/studio/backend/utils/datasets/dataset_utils.py @@ -357,48 +357,25 @@ def format_dataset( conversations = [] num_examples = len(examples[list(examples.keys())[0]]) - # NEW: Check if this is user-provided or auto-detected - is_user_provided = custom_format_mapping is not None # Passed explicitly - - # Preserve non-mapped columns ONLY if auto-detected - preserved_columns = {} - if not is_user_provided: # Only preserve for auto-detection - all_columns = set(examples.keys()) - mapped_columns = set(custom_mapping.keys()) - non_mapped_columns = all_columns - mapped_columns - - for col in non_mapped_columns: - preserved_columns[col] = examples[col] + # Preserve non-mapped columns + all_columns = set(examples.keys()) + mapped_columns = set(custom_mapping.keys()) + preserved_columns = { + col: examples[col] + for col in all_columns - mapped_columns + } for i in range(num_examples): convo = [] - - # Enforce standard role order - role_order = ['system', 'user', 'assistant'] - - for target_role in role_order: + for target_role in ['system', 'user', 'assistant']: for col_name, role in custom_mapping.items(): if role == target_role and col_name in examples: content = examples[col_name][i] - - # NEW: Different behavior based on mapping source - if is_user_provided: - # User explicitly mapped this - always include even if empty - convo.append({"role": role, "content": str(content) if content else ""}) - else: - # Auto-detected - skip empty (original behavior) - if content and str(content).strip(): - convo.append({"role": role, "content": str(content)}) - + if content and str(content).strip(): + convo.append({"role": role, "content": str(content)}) conversations.append(convo) - result = {"conversations": conversations} - - # Only add preserved columns if auto-detected - if not is_user_provided: - result.update(preserved_columns) - - return result + return {"conversations": conversations, **preserved_columns} try: @@ -507,7 +484,7 @@ def format_dataset( } # CHATML MODE: Convert to ChatML - elif format_type in ["chatml", "conversational"]: + elif format_type in ["chatml", "conversational", "sharegpt"]: if detected["format"] == "alpaca": converted = convert_alpaca_to_chatml(dataset, batch_size, num_proc) @@ -556,36 +533,38 @@ def format_dataset( else: warnings.append(f"Unknown format, attempting standardization") - try: - standardized = standardize_chat_format( - dataset, tokenizer, aliases_for_system, - aliases_for_user, aliases_for_assistant, - batch_size, num_proc - ) - return { - "dataset": standardized, - "detected_format": "unknown", - "final_format": f"chatml_{detected['chat_column']}", - "chat_column": detected["chat_column"], - "is_standardized": True, - "requires_manual_mapping": False, - "is_multimodal": multimodal_info["is_multimodal"], - "multimodal_info": multimodal_info, - "warnings": warnings - } - except Exception as e: - warnings.append(f"Standardization failed: {e}") - return { - "dataset": dataset, - "detected_format": "unknown", - "final_format": "unknown", - "chat_column": detected["chat_column"], - "is_standardized": False, - "requires_manual_mapping": True, - "is_multimodal": multimodal_info["is_multimodal"], - "multimodal_info": multimodal_info, - "warnings": warnings - } + if detected["chat_column"]: + try: + standardized = standardize_chat_format( + dataset, tokenizer, aliases_for_system, + aliases_for_user, aliases_for_assistant, + batch_size, num_proc + ) + return { + "dataset": standardized, + "detected_format": "unknown", + "final_format": f"chatml_{detected['chat_column']}", + "chat_column": detected["chat_column"], + "is_standardized": True, + "requires_manual_mapping": False, + "is_multimodal": multimodal_info["is_multimodal"], + "multimodal_info": multimodal_info, + "warnings": warnings + } + except Exception as e: + warnings.append(f"Standardization failed: {e}") + + return { + "dataset": dataset, + "detected_format": "unknown", + "final_format": "unknown", + "chat_column": detected["chat_column"], + "is_standardized": False, + "requires_manual_mapping": True, + "is_multimodal": multimodal_info["is_multimodal"], + "multimodal_info": multimodal_info, + "warnings": warnings + } else: raise ValueError(f"Unknown format_type: {format_type}") @@ -816,8 +795,10 @@ def format_and_template_dataset( ) # Step 2: Apply chat template - if "gemma" in model_name.lower() and not dataset_info["is_multimodal"] and (format_type != "alpaca" or (format_type == "auto" and dataset_info["detected_format"] != "alpaca")): - print("remove_bos_prefix is true") + # Gemma emits a leading that must be stripped for text-only chatml/sharegpt. + is_alpaca = format_type == "alpaca" or (format_type == "auto" and dataset_info["detected_format"] == "alpaca") + is_gemma = "gemma" in model_name.lower() + if is_gemma and not dataset_info["is_multimodal"] and not is_alpaca: remove_bos_prefix = True template_result = apply_chat_template_to_dataset( dataset_info=dataset_info, @@ -839,14 +820,24 @@ def format_and_template_dataset( all_warnings = dataset_info.get("warnings", []) + template_result.get("warnings", []) all_errors = template_result.get("errors", []) + # If format_dataset returned "unknown" but apply_chat_template rescued + # it via heuristic detection, update final_format to reflect reality. + final_format = dataset_info["final_format"] + requires_manual = dataset_info.get("requires_manual_mapping", False) + if final_format == "unknown" and template_result["success"]: + out_ds = template_result["dataset"] + if hasattr(out_ds, "column_names") and "text" in out_ds.column_names: + final_format = "chatml_conversations" + requires_manual = False + return { "dataset": template_result["dataset"], "detected_format": dataset_info["detected_format"], - "final_format": dataset_info["final_format"], + "final_format": final_format, "chat_column": dataset_info.get("chat_column"), "is_vlm": False, # This is LLM flow "success": template_result["success"], - "requires_manual_mapping": dataset_info.get("requires_manual_mapping", False), + "requires_manual_mapping": requires_manual, "warnings": all_warnings, "errors": all_errors, "summary": summary, diff --git a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx index 484b10ddb4..1ad4fad8f0 100644 --- a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx @@ -75,6 +75,8 @@ export function DatasetStep() { setDatasetSubset, datasetSplit, setDatasetSplit, + datasetEvalSplit, + setDatasetEvalSplit, uploadedFile, setUploadedFile, } = useTrainingConfigStore( @@ -91,6 +93,8 @@ export function DatasetStep() { setDatasetSubset: s.setDatasetSubset, datasetSplit: s.datasetSplit, setDatasetSplit: s.setDatasetSplit, + datasetEvalSplit: s.datasetEvalSplit, + setDatasetEvalSplit: s.setDatasetEvalSplit, uploadedFile: s.uploadedFile, setUploadedFile: s.setUploadedFile, })), @@ -304,6 +308,8 @@ export function DatasetStep() { setDatasetSubset={setDatasetSubset} datasetSplit={datasetSplit} setDatasetSplit={setDatasetSplit} + datasetEvalSplit={datasetEvalSplit} + setDatasetEvalSplit={setDatasetEvalSplit} /> ) : ( diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index 1dfce7d21a..6f6a6a401e 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -1,5 +1,10 @@ import { SectionCard } from "@/components/section-card"; import { Button } from "@/components/ui/button"; +import { + Collapsible, + CollapsibleContent, + CollapsibleTrigger, +} from "@/components/ui/collapsible"; import { Combobox, ComboboxContent, @@ -35,6 +40,7 @@ import { useTrainingConfigStore, } from "@/features/training"; import { + ArrowDown01Icon, CloudUploadIcon, Database02Icon, FileAttachmentIcon, @@ -66,6 +72,8 @@ export function DatasetSection() { setDatasetSubset, datasetSplit, setDatasetSplit, + datasetEvalSplit, + setDatasetEvalSplit, hfToken, } = useTrainingConfigStore( useShallow((s) => ({ @@ -77,11 +85,14 @@ export function DatasetSection() { setDatasetSubset: s.setDatasetSubset, datasetSplit: s.datasetSplit, setDatasetSplit: s.setDatasetSplit, + datasetEvalSplit: s.datasetEvalSplit, + setDatasetEvalSplit: s.setDatasetEvalSplit, hfToken: s.hfToken, })), ); const [inputValue, setInputValue] = useState(""); + const [advancedOpen, setAdvancedOpen] = useState(false); const openPreview = useDatasetPreviewDialogStore((s) => s.openPreview); const selectingRef = useRef(false); const debouncedQuery = useDebouncedValue(inputValue); @@ -132,79 +143,82 @@ export function DatasetSection() { title="Dataset" description="Select or upload training data" accent="indigo" - className="md:min-h-[450px] dark:shadow-border" + className="md:min-h-[470px] dark:shadow-border" >
-
- - Load from Hub - - - - - - Search Hugging Face datasets or enter a path like - 'username/dataset-name'.{" "} - - Read more - - - - -
{ - if (event.key !== "Enter") return; - if (!(event.target instanceof HTMLInputElement)) return; - event.preventDefault(); - if (hfResults.length > 0) { - handleDatasetSelect(hfResults[0].id); - } else { - const text = event.target.value.trim(); - if (text) handleDatasetSelect(text); - } - }} - > - id} - autoHighlight={true} +
+ + Load from Hub + + + + + + Search Hugging Face datasets or enter a path like + 'username/dataset-name'.{" "} + + Read more + + + + +
{ + if (event.key !== "Enter") return; + if (!(event.target instanceof HTMLInputElement)) return; + event.preventDefault(); + if (hfResults.length > 0) { + handleDatasetSelect(hfResults[0].id); + } else { + const text = event.target.value.trim(); + if (text) handleDatasetSelect(text); + } + }} > - - - - - - - {isLoading ? ( -
- Searching... -
- ) : ( - No datasets found - )} -
id} + autoHighlight={true} + > + + + + + + + {isLoading ? ( +
+ Searching... +
+ ) : ( + No datasets found + )} +
{(id: string) => { const r = hfResults.find((ds) => ds.id === id); @@ -220,58 +234,58 @@ export function DatasetSection() { - - - + className="justify-between" + > + + + + {id} + + + {id} + + + {detail && ( + + {detail} - - - {id} - - - {detail && ( - - {detail} - - )} - - ); - }} - -
- {isLoadingMore && ( -
- -
- )} -
- - + )} + + ); + }} + +
+ {isLoadingMore && ( +
+ +
+ )} +
+ + +
+ {(tokenValidationError ?? hfSearchError) && ( +

+ {tokenValidationError ?? hfSearchError} + {" — "} + + Get or update token + +

+ )} + {isCheckingToken && ( +

Checking token…

+ )}
- {(tokenValidationError ?? hfSearchError) && ( -

- {tokenValidationError ?? hfSearchError} - {" — "} - - Get or update token - -

- )} - {isCheckingToken && ( -

Checking token…

- )} -
-
- - Target Format - - - - - - Format of your training data. Auto-detect works for most - datasets.{" "} - - Read more - - - - - -
+ + + + Advanced + + +
+ + Target Format + + + + + + Format of your training data. Auto-detect works for most + datasets.{" "} + + Read more + + + + + +
+
+
{dataset ? (
diff --git a/studio/frontend/src/features/studio/sections/params-section.tsx b/studio/frontend/src/features/studio/sections/params-section.tsx index 144a1f34d7..b7ad7aec0b 100644 --- a/studio/frontend/src/features/studio/sections/params-section.tsx +++ b/studio/frontend/src/features/studio/sections/params-section.tsx @@ -127,7 +127,7 @@ export function ParamsSection(): ReactElement { title="Parameters" description="Configure training hyperparameters" accent="orange" - className="md:min-h-[450px]" + className="md:min-h-[470px]" >
{/* Max Steps */} diff --git a/studio/frontend/src/features/studio/sections/training-section.tsx b/studio/frontend/src/features/studio/sections/training-section.tsx index 3f342ad68c..af6a16fc44 100644 --- a/studio/frontend/src/features/studio/sections/training-section.tsx +++ b/studio/frontend/src/features/studio/sections/training-section.tsx @@ -98,7 +98,7 @@ export function TrainingSection() { title="Training" description="Monitor and control training" accent="blue" - className="md:min-h-[450px]" + className="md:min-h-[470px]" >
{/* Loss chart */} diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 5f98913158..1adfbd8b6d 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -24,8 +24,9 @@ export function buildTrainingStartPayload( load_in_4bit: adapterMethod ? isQlorMethod : false, max_seq_length: config.contextLength, hf_dataset: hfDataset, - hf_dataset_config: hfDataset ? config.datasetSubset : null, - hf_dataset_split: hfDataset ? config.datasetSplit : null, + subset: hfDataset ? config.datasetSubset : null, + train_split: hfDataset ? config.datasetSplit : null, + eval_split: hfDataset ? config.datasetEvalSplit : null, local_datasets: [], format_type: config.datasetFormat, custom_format_mapping: customFormatMapping, diff --git a/studio/frontend/src/features/training/components/hf-dataset-subset-split-selectors.tsx b/studio/frontend/src/features/training/components/hf-dataset-subset-split-selectors.tsx index 00532ce860..d21fd55ff4 100644 --- a/studio/frontend/src/features/training/components/hf-dataset-subset-split-selectors.tsx +++ b/studio/frontend/src/features/training/components/hf-dataset-subset-split-selectors.tsx @@ -29,6 +29,8 @@ type Props = { setDatasetSubset: (v: string | null) => void; datasetSplit: string | null; setDatasetSplit: (v: string | null) => void; + datasetEvalSplit: string | null; + setDatasetEvalSplit: (v: string | null) => void; }; export function HfDatasetSubsetSplitSelectors({ @@ -40,12 +42,13 @@ export function HfDatasetSubsetSplitSelectors({ setDatasetSubset, datasetSplit, setDatasetSplit, + datasetEvalSplit, + setDatasetEvalSplit, }: Props) { const { subsets: hfSubsets, splits: hfSplits, hasMultipleSubsets, - hasMultipleSplits, isLoading, error, } = useHfDatasetSplits(enabled ? datasetName : null, datasetSubset, { @@ -78,6 +81,8 @@ export function HfDatasetSubsetSplitSelectors({ if (!enabled || !datasetName) return null; + const showDropdowns = !isLoading && !error && hfSubsets.length > 0; + return ( <> {isLoading && ( @@ -105,167 +110,169 @@ export function HfDatasetSubsetSplitSelectors({
)} - {!isLoading && !error && hasMultipleSubsets && ( + {showDropdowns && ( <> - {variant === "wizard" ? ( - - - Subset - - - - - - This dataset has multiple subsets. Select which one to use - for training. - - - - - - ) : ( -
- - Subset - - - - - - This dataset has multiple subsets. Select which one to use - for training. - - - - + {variant === "studio" ? ( +
+ +
- )} - - )} - - {!isLoading && !error && hasMultipleSplits && ( - <> - {variant === "wizard" ? ( - - - Split - - - - - - Select which split of the dataset to use for training. - - - - - ) : ( -
- - Split - - - - - - Select which split of the dataset to use for training. - - - - -
+ <> + + + )} + )} ); } + +function SelectorDropdown({ + variant, + label, + tooltip, + value, + onChange, + options, + placeholder, + allowNone = false, +}: { + variant: "wizard" | "studio"; + label: string; + tooltip: string; + value: string | null; + onChange: (v: string | null) => void; + options: string[]; + placeholder: string; + allowNone?: boolean; +}) { + if (variant === "wizard") { + return ( + + + {label} + + + + + + {tooltip} + + + + + + ); + } + + return ( +
+ + {label} + + + + + + {tooltip} + + + + +
+ ); +} 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 12f0c2ab93..b2d1858716 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -26,6 +26,7 @@ const initialState: TrainingConfigState = { dataset: null, datasetSubset: null, datasetSplit: null, + datasetEvalSplit: null, datasetManualMapping: emptyManualMapping(), uploadedFile: null, isCheckingVision: false, @@ -252,6 +253,7 @@ export const useTrainingConfigStore = create()( dataset, datasetSubset: null, datasetSplit: null, + datasetEvalSplit: null, datasetManualMapping: emptyManualMapping(), isDatasetMultimodal: null, isCheckingDataset: false, @@ -264,6 +266,7 @@ export const useTrainingConfigStore = create()( set({ datasetSubset, datasetSplit: null, + datasetEvalSplit: null, datasetManualMapping: emptyManualMapping(), isDatasetMultimodal: null, isCheckingDataset: false, @@ -300,6 +303,12 @@ export const useTrainingConfigStore = create()( const split = state.datasetSplit || "train"; runDatasetCheck(datasetName, split); }, + setDatasetEvalSplit: (datasetEvalSplit) => { + set({ + datasetEvalSplit, + evalSteps: datasetEvalSplit ? 0.1 : 0, + }); + }, setDatasetManualMapping: (datasetManualMapping) => set({ datasetManualMapping }), setUploadedFile: (uploadedFile) => set({ uploadedFile }), @@ -359,7 +368,7 @@ export const useTrainingConfigStore = create()( }, { name: "unsloth_training_config_v1", - version: 5, + version: 6, migrate: (persisted, version) => { const s = persisted as Record; if (version < 2 && s.datasetSubset == null && s.datasetConfig != null) { @@ -375,6 +384,9 @@ export const useTrainingConfigStore = create()( if (version < 5 && s.lrSchedulerType == null) { s.lrSchedulerType = DEFAULT_HYPERPARAMS.lrSchedulerType; } + if (version < 6 && s.datasetEvalSplit == null) { + s.datasetEvalSplit = null; + } return s as unknown as TrainingConfigStore; }, partialize: partializePersistedState, diff --git a/studio/frontend/src/features/training/types/api.ts b/studio/frontend/src/features/training/types/api.ts index fce7e18b0b..22f02e4331 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -5,8 +5,9 @@ export interface TrainingStartRequest { load_in_4bit: boolean; max_seq_length: number; hf_dataset: string | null; - hf_dataset_config: string | null; - hf_dataset_split: string | null; + subset: string | null; + train_split: string | null; + eval_split: string | null; local_datasets: string[]; format_type: string; custom_format_mapping?: Record | null; diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index f64c5d5aa8..6c2feec172 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -24,6 +24,7 @@ export interface TrainingConfigState { dataset: string | null; datasetSubset: string | null; datasetSplit: string | null; + datasetEvalSplit: string | null; datasetManualMapping: DatasetManualMapping; uploadedFile: string | null; epochs: number; @@ -81,6 +82,7 @@ export interface TrainingConfigActions { setDataset: (dataset: string | null) => void; setDatasetSubset: (subset: string | null) => void; setDatasetSplit: (split: string | null) => void; + setDatasetEvalSplit: (split: string | null) => void; setDatasetManualMapping: (mapping: DatasetManualMapping) => void; setUploadedFile: (file: string | null) => void; setEpochs: (epochs: number) => void;