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:
commit
a9ec5d894c
13 changed files with 554 additions and 26 deletions
|
|
@ -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}
|
||||
/>
|
||||
</>
|
||||
) : (
|
||||
<>
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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>>[]>(() => {
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
)}
|
||||
</>
|
||||
)}
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
139
studio/frontend/src/hooks/use-hf-dataset-splits.ts
Normal file
139
studio/frontend/src/hooks/use-hf-dataset-splits.ts
Normal 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,
|
||||
};
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue