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 65d802c77e..0275e5a6da 100644 --- a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx @@ -38,7 +38,10 @@ import { useInfiniteScroll, } from "@/hooks"; import { cn, formatCompact } from "@/lib/utils"; -import { useTrainingConfigStore } from "@/features/training"; +import { + HfDatasetSubsetSplitSelectors, + useTrainingConfigStore, +} from "@/features/training"; import type { DatasetFormat } from "@/types/training"; import { InformationCircleIcon, @@ -67,6 +70,10 @@ export function DatasetStep() { setDatasetFormat, dataset, setDataset, + datasetSubset, + setDatasetSubset, + datasetSplit, + setDatasetSplit, uploadedFile, setUploadedFile, } = useTrainingConfigStore( @@ -79,6 +86,10 @@ export function DatasetStep() { setDatasetFormat: s.setDatasetFormat, dataset: s.dataset, setDataset: s.setDataset, + datasetSubset: s.datasetSubset, + setDatasetSubset: s.setDatasetSubset, + datasetSplit: s.datasetSplit, + setDatasetSplit: s.setDatasetSplit, uploadedFile: s.uploadedFile, setUploadedFile: s.setUploadedFile, })), @@ -261,6 +272,17 @@ export function DatasetStep() { + + ) : ( <> diff --git a/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx b/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx index ae6edd230a..da2758ebfa 100644 --- a/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx @@ -23,6 +23,8 @@ export function SummaryStep() { datasetSource, datasetFormat, dataset, + datasetSubset, + datasetSplit, uploadedFile, epochs, contextLength, @@ -39,6 +41,8 @@ export function SummaryStep() { datasetSource, datasetFormat, dataset, + datasetSubset, + datasetSplit, uploadedFile, epochs, contextLength, @@ -53,6 +57,8 @@ export function SummaryStep() { datasetSource, datasetFormat, dataset, + datasetSubset, + datasetSplit, uploadedFile, epochs, contextLength, @@ -147,6 +153,18 @@ export function SummaryStep() { Source {datasetSource} + {datasetSubset && ( +
+ Subset + {datasetSubset} +
+ )} + {datasetSplit && ( +
+ Split + {datasetSplit} +
+ )}
Format {datasetFormat} diff --git a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx index 80fbcb65ee..f145b17f61 100644 --- a/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-preview-dialog.tsx @@ -40,15 +40,21 @@ type DatasetPreviewDialogProps = { onOpenChange: (open: boolean) => void; datasetName: string | null; hfToken: string | null; + datasetSubset?: string | null; + datasetSplit?: string | null; }; // --------------------------------------------------------------------------- // API -- uses existing /check-format endpoint // --------------------------------------------------------------------------- +// TODO(backend): Needs to accept `config` and `split` fields (see #37). +// The frontend already sends them in the request below. async function fetchCheckFormat( datasetName: string, hfToken: string | null, + subset?: string | null, + split?: string | null, ): Promise { const res = await fetch("/api/datasets/check-format", { method: "POST", @@ -56,7 +62,8 @@ async function fetchCheckFormat( body: JSON.stringify({ dataset_name: datasetName, hf_token: hfToken || undefined, - split: "train", + config: subset || undefined, + split: split || "train", }), }); if (!res.ok) { @@ -75,6 +82,8 @@ export function DatasetPreviewDialog({ onOpenChange, datasetName, hfToken, + datasetSubset, + datasetSplit, }: DatasetPreviewDialogProps) { const [data, setData] = useState(null); const [loading, setLoading] = useState(false); @@ -90,7 +99,7 @@ export function DatasetPreviewDialog({ setLoading(true); setError(null); - fetchCheckFormat(datasetName, hfToken) + fetchCheckFormat(datasetName, hfToken, datasetSubset, datasetSplit) .then((res) => { if (!cancelled) { setData(res); @@ -107,7 +116,7 @@ export function DatasetPreviewDialog({ return () => { cancelled = true; }; - }, [open, datasetName, hfToken]); + }, [open, datasetName, hfToken, datasetSubset, datasetSplit]); const rows = data?.preview_samples ?? []; const columns = data?.columns ?? []; @@ -115,9 +124,15 @@ export function DatasetPreviewDialog({ // Determine source label const sourceLabel = useMemo(() => { if (!datasetName) return ""; - if (datasetName.includes("/")) return `Hugging Face (${datasetName})`; + if (datasetName.includes("/")) { + let label = `Hugging Face (${datasetName}`; + if (datasetSubset) label += ` / ${datasetSubset}`; + if (datasetSplit) label += ` / ${datasetSplit}`; + label += ")"; + return label; + } return `Local Files (${datasetName})`; - }, [datasetName]); + }, [datasetName, datasetSubset, datasetSplit]); // Build TanStack Table columns from the column names const tableColumns = useMemo>[]>(() => { diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx index 0d325fbefd..c4d691d8c7 100644 --- a/studio/frontend/src/features/studio/sections/dataset-section.tsx +++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx @@ -28,7 +28,10 @@ import { useInfiniteScroll, } from "@/hooks"; import { formatCompact } from "@/lib/utils"; -import { useTrainingConfigStore } from "@/features/training"; +import { + HfDatasetSubsetSplitSelectors, + useTrainingConfigStore, +} from "@/features/training"; import { CloudUploadIcon, Database02Icon, @@ -42,25 +45,40 @@ import { useMemo, useRef, useState } from "react"; import { useShallow } from "zustand/react/shallow"; import { DatasetPreviewDialog } from "./dataset-preview-dialog"; +function isLikelyLocalDatasetRef(value: string) { + return ( + value.startsWith("/") || + value.startsWith("./") || + value.startsWith("../") || + value.includes("\\") || + /\.(jsonl|json|csv|parquet)$/i.test(value) + ); +} + export function DatasetSection() { - const { dataset, setDataset, datasetFormat, setDatasetFormat, hfToken } = - useTrainingConfigStore( - useShallow( - ({ - dataset, - setDataset, - datasetFormat, - setDatasetFormat, - hfToken, - }) => ({ - dataset, - setDataset, - datasetFormat, - setDatasetFormat, - hfToken, - }), - ), - ); + const { + dataset, + setDataset, + datasetFormat, + setDatasetFormat, + datasetSubset, + setDatasetSubset, + datasetSplit, + setDatasetSplit, + hfToken, + } = useTrainingConfigStore( + useShallow((s) => ({ + dataset: s.dataset, + setDataset: s.setDataset, + datasetFormat: s.datasetFormat, + setDatasetFormat: s.setDatasetFormat, + datasetSubset: s.datasetSubset, + setDatasetSubset: s.setDatasetSubset, + datasetSplit: s.datasetSplit, + setDatasetSplit: s.setDatasetSplit, + hfToken: s.hfToken, + })), + ); const [inputValue, setInputValue] = useState(""); const [previewOpen, setPreviewOpen] = useState(false); @@ -221,6 +239,17 @@ export function DatasetSection() {
+ +
Target Format @@ -280,6 +309,8 @@ export function DatasetSection() {

Hugging Face Dataset + {datasetSubset && ` / ${datasetSubset}`} + {datasetSplit && ` / ${datasetSplit}`}

@@ -321,6 +352,8 @@ export function DatasetSection() { onOpenChange={setPreviewOpen} datasetName={dataset} hfToken={hfToken} + datasetSubset={datasetSubset} + datasetSplit={datasetSplit} /> ); diff --git a/studio/frontend/src/features/training/api/mappers.ts b/studio/frontend/src/features/training/api/mappers.ts index 0150c62b41..990862a27e 100644 --- a/studio/frontend/src/features/training/api/mappers.ts +++ b/studio/frontend/src/features/training/api/mappers.ts @@ -22,6 +22,8 @@ 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, local_datasets: [], format_type: config.datasetFormat, num_epochs: config.epochs, 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 new file mode 100644 index 0000000000..00532ce860 --- /dev/null +++ b/studio/frontend/src/features/training/components/hf-dataset-subset-split-selectors.tsx @@ -0,0 +1,271 @@ +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { Spinner } from "@/components/ui/spinner"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/ui/tooltip"; +import { + Field, + FieldLabel, +} from "@/components/ui/field"; +import { useHfDatasetSplits } from "@/hooks"; +import { InformationCircleIcon } from "@hugeicons/core-free-icons"; +import { HugeiconsIcon } from "@hugeicons/react"; +import { useEffect } from "react"; + +type Props = { + variant: "wizard" | "studio"; + enabled: boolean; + datasetName: string | null; + accessToken?: string; + datasetSubset: string | null; + setDatasetSubset: (v: string | null) => void; + datasetSplit: string | null; + setDatasetSplit: (v: string | null) => void; +}; + +export function HfDatasetSubsetSplitSelectors({ + variant, + enabled, + datasetName, + accessToken, + datasetSubset, + setDatasetSubset, + datasetSplit, + setDatasetSplit, +}: Props) { + const { + subsets: hfSubsets, + splits: hfSplits, + hasMultipleSubsets, + hasMultipleSplits, + isLoading, + error, + } = useHfDatasetSplits(enabled ? datasetName : null, datasetSubset, { + accessToken, + }); + + useEffect(() => { + if (hfSubsets.length === 1 && datasetSubset !== hfSubsets[0]) { + setDatasetSubset(hfSubsets[0]); + } + }, [hfSubsets, datasetSubset, setDatasetSubset]); + + useEffect(() => { + if (hfSplits.length === 0) return; + if (hasMultipleSubsets && !datasetSubset) return; + if (hfSplits.length === 1 && datasetSplit !== hfSplits[0]) { + setDatasetSplit(hfSplits[0]); + } else if (!datasetSplit && hfSplits.includes("train")) { + setDatasetSplit("train"); + } else if (!datasetSplit) { + setDatasetSplit(hfSplits[0]); + } + }, [ + hfSplits, + hasMultipleSubsets, + datasetSubset, + datasetSplit, + setDatasetSplit, + ]); + + if (!enabled || !datasetName) return null; + + return ( + <> + {isLoading && ( +
+ + Loading dataset configs and splits... +
+ )} + + {error && ( +
+ Could not fetch dataset splits: {error} +
+ )} + + {!isLoading && !error && hasMultipleSubsets && ( + <> + {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. + + + + +
+ )} + + )} + + ); +} diff --git a/studio/frontend/src/features/training/index.ts b/studio/frontend/src/features/training/index.ts index 48b9079fd3..41d1aade75 100644 --- a/studio/frontend/src/features/training/index.ts +++ b/studio/frontend/src/features/training/index.ts @@ -5,4 +5,5 @@ export { } from "./stores/training-runtime-store"; export { useTrainingActions } from "./hooks/use-training-actions"; export { useTrainingRuntimeLifecycle } from "./hooks/use-training-runtime-lifecycle"; +export { HfDatasetSubsetSplitSelectors } from "./components/hf-dataset-subset-split-selectors"; export type { TrainingPhase } from "./types/runtime"; 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 a9d6d37f42..57d05db114 100644 --- a/studio/frontend/src/features/training/stores/training-config-store.ts +++ b/studio/frontend/src/features/training/stores/training-config-store.ts @@ -16,6 +16,8 @@ const initialState: TrainingConfigState = { datasetSource: "huggingface", datasetFormat: "auto", dataset: null, + datasetSubset: null, + datasetSplit: null, uploadedFile: null, ...DEFAULT_HYPERPARAMS, }; @@ -55,7 +57,11 @@ export const useTrainingConfigStore = create()( setHfToken: (hfToken) => set({ hfToken }), setDatasetSource: (datasetSource) => set({ datasetSource }), setDatasetFormat: (datasetFormat) => set({ datasetFormat }), - setDataset: (dataset) => set({ dataset }), + setDataset: (dataset) => + set({ dataset, datasetSubset: null, datasetSplit: null }), + setDatasetSubset: (datasetSubset) => + set({ datasetSubset, datasetSplit: null }), + setDatasetSplit: (datasetSplit) => set({ datasetSplit }), setUploadedFile: (uploadedFile) => set({ uploadedFile }), setEpochs: (epochs) => set({ epochs }), setContextLength: (contextLength) => set({ contextLength }), @@ -96,6 +102,16 @@ export const useTrainingConfigStore = create()( }), { name: "unsloth_training_config_v1", + version: 2, + migrate: (persisted, version) => { + const s = persisted as Record; + if (version >= 2) return s as unknown as TrainingConfigStore; + if (s.datasetSubset == null && s.datasetConfig != null) { + s.datasetSubset = s.datasetConfig; + } + delete s.datasetConfig; + return s as unknown as TrainingConfigStore; + }, partialize: (state) => { const { modelType, ...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 9cf789a702..6011ff6b96 100644 --- a/studio/frontend/src/features/training/types/api.ts +++ b/studio/frontend/src/features/training/types/api.ts @@ -5,6 +5,8 @@ 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; local_datasets: string[]; format_type: string; num_epochs: number; diff --git a/studio/frontend/src/features/training/types/config.ts b/studio/frontend/src/features/training/types/config.ts index c0d93f76a7..77706d60b2 100644 --- a/studio/frontend/src/features/training/types/config.ts +++ b/studio/frontend/src/features/training/types/config.ts @@ -18,6 +18,8 @@ export interface TrainingConfigState { datasetSource: DatasetSource; datasetFormat: DatasetFormat; dataset: string | null; + datasetSubset: string | null; + datasetSplit: string | null; uploadedFile: string | null; epochs: number; contextLength: number; @@ -60,6 +62,8 @@ export interface TrainingConfigActions { setDatasetSource: (source: DatasetSource) => void; setDatasetFormat: (format: DatasetFormat) => void; setDataset: (dataset: string | null) => void; + setDatasetSubset: (subset: string | null) => void; + setDatasetSplit: (split: string | null) => void; setUploadedFile: (file: string | null) => void; setEpochs: (epochs: number) => void; setContextLength: (length: number) => void; diff --git a/studio/frontend/src/hooks/index.ts b/studio/frontend/src/hooks/index.ts index 24f6faca79..1ada378651 100644 --- a/studio/frontend/src/hooks/index.ts +++ b/studio/frontend/src/hooks/index.ts @@ -2,4 +2,5 @@ export { useDebouncedValue } from "./use-debounced-value"; export { useGpuInfo } from "./use-gpu-info"; export { useHfModelSearch } from "./use-hf-model-search"; export { useHfDatasetSearch } from "./use-hf-dataset-search"; +export { useHfDatasetSplits } from "./use-hf-dataset-splits"; export { useInfiniteScroll } from "./use-infinite-scroll"; diff --git a/studio/frontend/src/hooks/use-hf-dataset-splits.ts b/studio/frontend/src/hooks/use-hf-dataset-splits.ts new file mode 100644 index 0000000000..7b8e4906ec --- /dev/null +++ b/studio/frontend/src/hooks/use-hf-dataset-splits.ts @@ -0,0 +1,139 @@ +import { useCallback, useEffect, useState } from "react"; + +// --------------------------------------------------------------------------- +// Types +// --------------------------------------------------------------------------- + +export interface HfSplitEntry { + dataset: string; + config: string; + split: string; +} + +export interface HfSplitsResponse { + splits: HfSplitEntry[]; + pending: unknown[]; + failed: unknown[]; +} + +export interface HfDatasetSplitsResult { + /** All unique subset names found in the dataset */ + subsets: string[]; + /** All split names available for the currently selected subset */ + splits: string[]; + /** Raw split entries from the API */ + entries: HfSplitEntry[]; + /** Whether the dataset has more than one subset */ + hasMultipleSubsets: boolean; + /** Whether the selected subset has more than one split */ + hasMultipleSplits: boolean; + /** True while the request is in-flight */ + isLoading: boolean; + /** Error message if the fetch failed */ + error: string | null; +} + +const HF_SPLITS_API = "https://datasets-server.huggingface.co/splits"; + +// --------------------------------------------------------------------------- +// Hook +// --------------------------------------------------------------------------- + +/** + * Fetches the available configs (subsets) and splits for a HuggingFace dataset + * using the datasets-server API. + * + * @param datasetName - HF dataset id (e.g. "ibm/duorc"), or null to skip. + * @param selectedSubset - Currently selected subset, used to filter splits. + * @param options.accessToken - Optional HF access token for gated datasets. + */ +export function useHfDatasetSplits( + datasetName: string | null, + selectedSubset: string | null, + options?: { accessToken?: string }, +): HfDatasetSplitsResult { + const [entries, setEntries] = useState([]); + const [isLoading, setIsLoading] = useState(false); + const [error, setError] = useState(null); + + const accessToken = options?.accessToken; + + const fetchSplits = useCallback( + async (dataset: string, signal: AbortSignal) => { + const url = `${HF_SPLITS_API}?dataset=${encodeURIComponent(dataset)}`; + const headers: Record = {}; + if (accessToken) { + headers.Authorization = `Bearer ${accessToken}`; + } + + const res = await fetch(url, { headers, signal }); + if (!res.ok) { + const body = await res.json().catch(() => null); + throw new Error( + body?.error || `Failed to fetch splits (${res.status})`, + ); + } + + const data: HfSplitsResponse = await res.json(); + return data.splits ?? []; + }, + [accessToken], + ); + + useEffect(() => { + if (!datasetName) { + setEntries([]); + setError(null); + setIsLoading(false); + return; + } + + const controller = new AbortController(); + setIsLoading(true); + setError(null); + + fetchSplits(datasetName, controller.signal) + .then((splits) => { + if (!controller.signal.aborted) { + setEntries(splits); + setError(null); + } + }) + .catch((err) => { + if (!controller.signal.aborted) { + setError(err.message || "Failed to fetch dataset splits"); + setEntries([]); + } + }) + .finally(() => { + if (!controller.signal.aborted) { + setIsLoading(false); + } + }); + + return () => controller.abort(); + }, [datasetName, fetchSplits]); + + // Derive unique subsets + const subsets = Array.from(new Set(entries.map((e) => e.config))); + + // Derive splits for the active subset. + // If dataset has >1 subset and none is selected yet, return no splits so UI + // doesn't auto-pick/show a split before subset is chosen. + const activeSubset = + selectedSubset ?? (subsets.length === 1 ? subsets[0] : null); + const filteredEntries = activeSubset + ? entries.filter((e) => e.config === activeSubset) + : []; + const splits = Array.from(new Set(filteredEntries.map((e) => e.split))); + + return { + subsets, + splits, + entries, + hasMultipleSubsets: subsets.length > 1, + hasMultipleSplits: activeSubset ? splits.length > 1 : false, + isLoading, + error, + }; +} diff --git a/studio/frontend/src/types/training.ts b/studio/frontend/src/types/training.ts index 49c1dddaa3..18c3bbe407 100644 --- a/studio/frontend/src/types/training.ts +++ b/studio/frontend/src/types/training.ts @@ -18,6 +18,8 @@ export interface WizardState { datasetSource: DatasetSource; datasetFormat: DatasetFormat; dataset: string | null; + datasetSubset: string | null; + datasetSplit: string | null; uploadedFile: string | null; epochs: number; contextLength: number; @@ -60,6 +62,8 @@ export interface WizardActions { setDatasetSource: (source: DatasetSource) => void; setDatasetFormat: (format: DatasetFormat) => void; setDataset: (dataset: string | null) => void; + setDatasetSubset: (subset: string | null) => void; + setDatasetSplit: (split: string | null) => void; setUploadedFile: (file: string | null) => void; setEpochs: (epochs: number) => void; setContextLength: (length: number) => void;