diff --git a/studio/backend/utils/datasets/dataset_utils.py b/studio/backend/utils/datasets/dataset_utils.py
index 72298dc245..e92e145269 100644
--- a/studio/backend/utils/datasets/dataset_utils.py
+++ b/studio/backend/utils/datasets/dataset_utils.py
@@ -507,7 +507,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)
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..7b42cfbdc4 100644
--- a/studio/frontend/src/features/studio/sections/dataset-section.tsx
+++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx
@@ -66,6 +66,8 @@ export function DatasetSection() {
setDatasetSubset,
datasetSplit,
setDatasetSplit,
+ datasetEvalSplit,
+ setDatasetEvalSplit,
hfToken,
} = useTrainingConfigStore(
useShallow((s) => ({
@@ -77,6 +79,8 @@ export function DatasetSection() {
setDatasetSubset: s.setDatasetSubset,
datasetSplit: s.datasetSplit,
setDatasetSplit: s.setDatasetSplit,
+ datasetEvalSplit: s.datasetEvalSplit,
+ setDatasetEvalSplit: s.setDatasetEvalSplit,
hfToken: s.hfToken,
})),
);
@@ -282,6 +286,8 @@ export function DatasetSection() {
setDatasetSubset={setDatasetSubset}
datasetSplit={datasetSplit}
setDatasetSplit={setDatasetSplit}
+ datasetEvalSplit={datasetEvalSplit}
+ setDatasetEvalSplit={setDatasetEvalSplit}
/>
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..897eae206a 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,144 @@ 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.
-
-
-
-
-
- )}
- >
- )}
-
- {!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;