Merge pull request #98 from unslothai/feat/dataset-config-splits

feat: check dataset configs and splits before hitting check-format
This commit is contained in:
Wasim Yousef Said 2026-02-15 12:17:22 -08:00 committed by GitHub
commit a9ec5d894c
13 changed files with 554 additions and 26 deletions

View file

@ -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() {
</Combobox>
</div>
</Field>
<HfDatasetSubsetSplitSelectors
variant="wizard"
enabled={datasetSource === "huggingface"}
datasetName={dataset}
accessToken={hfToken || undefined}
datasetSubset={datasetSubset}
setDatasetSubset={setDatasetSubset}
datasetSplit={datasetSplit}
setDatasetSplit={setDatasetSplit}
/>
</>
) : (
<>

View file

@ -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() {
<span className="text-muted-foreground">Source</span>
<span className="capitalize">{datasetSource}</span>
</div>
{datasetSubset && (
<div className="flex items-center justify-between text-sm">
<span className="text-muted-foreground">Subset</span>
<span className="font-mono text-xs">{datasetSubset}</span>
</div>
)}
{datasetSplit && (
<div className="flex items-center justify-between text-sm">
<span className="text-muted-foreground">Split</span>
<span className="font-mono text-xs">{datasetSplit}</span>
</div>
)}
<div className="flex items-center justify-between text-sm">
<span className="text-muted-foreground">Format</span>
<span className="capitalize">{datasetFormat}</span>

View file

@ -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<CheckFormatResponse> {
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<CheckFormatResponse | null>(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<ColumnDef<Record<string, unknown>>[]>(() => {

View file

@ -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() {
</div>
</div>
<HfDatasetSubsetSplitSelectors
variant="studio"
enabled={!!dataset && !isLikelyLocalDatasetRef(dataset)}
datasetName={dataset}
accessToken={hfToken || undefined}
datasetSubset={datasetSubset}
setDatasetSubset={setDatasetSubset}
datasetSplit={datasetSplit}
setDatasetSplit={setDatasetSplit}
/>
<div className="flex flex-col gap-2">
<span className="flex items-center gap-1.5 text-xs font-medium text-muted-foreground">
Target Format
@ -280,6 +309,8 @@ export function DatasetSection() {
</p>
<p className="text-[10px] text-muted-foreground">
Hugging Face Dataset
{datasetSubset && ` / ${datasetSubset}`}
{datasetSplit && ` / ${datasetSplit}`}
</p>
</div>
</div>
@ -321,6 +352,8 @@ export function DatasetSection() {
onOpenChange={setPreviewOpen}
datasetName={dataset}
hfToken={hfToken}
datasetSubset={datasetSubset}
datasetSplit={datasetSplit}
/>
</SectionCard>
);

View file

@ -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,

View file

@ -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 && (
<div
className={
variant === "wizard"
? "flex items-center gap-2 text-xs text-muted-foreground py-1"
: "flex items-center gap-2 rounded-lg border bg-muted/20 px-3.5 py-3 text-xs text-muted-foreground"
}
>
<Spinner className="size-3.5" />
Loading dataset configs and splits...
</div>
)}
{error && (
<div
className={
variant === "wizard"
? "rounded-lg border border-amber-200 bg-amber-50 px-3 py-2 text-xs text-amber-700 dark:border-amber-800 dark:bg-amber-950 dark:text-amber-400"
: "rounded-lg border border-amber-200 bg-amber-50 px-3.5 py-2.5 text-xs text-amber-700 dark:border-amber-800 dark:bg-amber-950 dark:text-amber-400"
}
>
Could not fetch dataset splits: {error}
</div>
)}
{!isLoading && !error && hasMultipleSubsets && (
<>
{variant === "wizard" ? (
<Field>
<FieldLabel className="flex items-center gap-1.5">
Subset
<Tooltip>
<TooltipTrigger asChild={true}>
<button
type="button"
className="text-muted-foreground/50 hover:text-muted-foreground"
>
<HugeiconsIcon
icon={InformationCircleIcon}
className="size-3.5"
/>
</button>
</TooltipTrigger>
<TooltipContent className="max-w-xs">
This dataset has multiple subsets. Select which one to use
for training.
</TooltipContent>
</Tooltip>
</FieldLabel>
<Select
value={datasetSubset ?? ""}
onValueChange={(v) => setDatasetSubset(v || null)}
>
<SelectTrigger className="w-full">
<SelectValue placeholder="Select a subset..." />
</SelectTrigger>
<SelectContent>
{hfSubsets.map((subset) => (
<SelectItem key={subset} value={subset}>
{subset}
</SelectItem>
))}
</SelectContent>
</Select>
</Field>
) : (
<div className="flex flex-col gap-1.5">
<span className="flex items-center gap-1.5 text-xs font-medium text-muted-foreground">
Subset
<Tooltip>
<TooltipTrigger asChild={true}>
<button
type="button"
className="text-foreground/70 hover:text-foreground"
>
<HugeiconsIcon
icon={InformationCircleIcon}
className="size-3"
/>
</button>
</TooltipTrigger>
<TooltipContent>
This dataset has multiple subsets. Select which one to use
for training.
</TooltipContent>
</Tooltip>
</span>
<Select
value={datasetSubset ?? ""}
onValueChange={(v) => setDatasetSubset(v || null)}
>
<SelectTrigger className="w-full">
<SelectValue placeholder="Select a subset..." />
</SelectTrigger>
<SelectContent>
{hfSubsets.map((subset) => (
<SelectItem key={subset} value={subset}>
{subset}
</SelectItem>
))}
</SelectContent>
</Select>
</div>
)}
</>
)}
{!isLoading && !error && hasMultipleSplits && (
<>
{variant === "wizard" ? (
<Field>
<FieldLabel className="flex items-center gap-1.5">
Split
<Tooltip>
<TooltipTrigger asChild={true}>
<button
type="button"
className="text-muted-foreground/50 hover:text-muted-foreground"
>
<HugeiconsIcon
icon={InformationCircleIcon}
className="size-3.5"
/>
</button>
</TooltipTrigger>
<TooltipContent className="max-w-xs">
Select which split of the dataset to use for training.
</TooltipContent>
</Tooltip>
</FieldLabel>
<Select
value={datasetSplit ?? ""}
onValueChange={(v) => setDatasetSplit(v || null)}
>
<SelectTrigger className="w-full">
<SelectValue placeholder="Select a split..." />
</SelectTrigger>
<SelectContent>
{hfSplits.map((split) => (
<SelectItem key={split} value={split}>
{split}
</SelectItem>
))}
</SelectContent>
</Select>
</Field>
) : (
<div className="flex flex-col gap-1.5">
<span className="flex items-center gap-1.5 text-xs font-medium text-muted-foreground">
Split
<Tooltip>
<TooltipTrigger asChild={true}>
<button
type="button"
className="text-foreground/70 hover:text-foreground"
>
<HugeiconsIcon
icon={InformationCircleIcon}
className="size-3"
/>
</button>
</TooltipTrigger>
<TooltipContent>
Select which split of the dataset to use for training.
</TooltipContent>
</Tooltip>
</span>
<Select
value={datasetSplit ?? ""}
onValueChange={(v) => setDatasetSplit(v || null)}
>
<SelectTrigger className="w-full">
<SelectValue placeholder="Select a split..." />
</SelectTrigger>
<SelectContent>
{hfSplits.map((split) => (
<SelectItem key={split} value={split}>
{split}
</SelectItem>
))}
</SelectContent>
</Select>
</div>
)}
</>
)}
</>
);
}

View file

@ -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";

View file

@ -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<TrainingConfigStore>()(
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<TrainingConfigStore>()(
}),
{
name: "unsloth_training_config_v1",
version: 2,
migrate: (persisted, version) => {
const s = persisted as Record<string, unknown>;
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;

View file

@ -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;

View file

@ -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;

View file

@ -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";

View file

@ -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<HfSplitEntry[]>([]);
const [isLoading, setIsLoading] = useState(false);
const [error, setError] = useState<string | null>(null);
const accessToken = options?.accessToken;
const fetchSplits = useCallback(
async (dataset: string, signal: AbortSignal) => {
const url = `${HF_SPLITS_API}?dataset=${encodeURIComponent(dataset)}`;
const headers: Record<string, string> = {};
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,
};
}

View file

@ -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;