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 579296cc7f..0281f9ffae 100644
--- a/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx
+++ b/studio/frontend/src/features/onboarding/components/steps/summary-step.tsx
@@ -1,14 +1,12 @@
import { Badge } from "@/components/ui/badge";
import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card";
import { Separator } from "@/components/ui/separator";
-import { DATASETS, findModelById } from "@/config/training";
import { useWizardStore } from "@/stores/training";
import { isAdapterMethod } from "@/types/training";
import { GpuIcon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { useShallow } from "zustand/react/shallow";
-// Mock system info - like checking via code
const SYSTEM_INFO = {
gpu: "NVIDIA RTX 4090",
vram: "24 GB",
@@ -33,34 +31,24 @@ export function SummaryStep() {
loraAlpha,
loraDropout,
} = useWizardStore(
- useShallow((s) => ({
- modelType: s.modelType,
- selectedModel: s.selectedModel,
- trainingMethod: s.trainingMethod,
- datasetSource: s.datasetSource,
- datasetFormat: s.datasetFormat,
- dataset: s.dataset,
- uploadedFile: s.uploadedFile,
- epochs: s.epochs,
- contextLength: s.contextLength,
- learningRate: s.learningRate,
- loraRank: s.loraRank,
- loraAlpha: s.loraAlpha,
- loraDropout: s.loraDropout,
+ useShallow(({
+ modelType, selectedModel, trainingMethod, datasetSource, datasetFormat,
+ dataset, uploadedFile, epochs, contextLength, learningRate,
+ loraRank, loraAlpha, loraDropout,
+ }) => ({
+ modelType, selectedModel, trainingMethod, datasetSource, datasetFormat,
+ dataset, uploadedFile, epochs, contextLength, learningRate,
+ loraRank, loraAlpha, loraDropout,
})),
);
- const modelData = findModelById(selectedModel);
- const datasetData = DATASETS.find((d) => d.id === dataset);
const showLoraParams = isAdapterMethod(trainingMethod);
const datasetName =
- datasetSource === "upload" ? uploadedFile : datasetData?.name;
- const datasetDesc =
- datasetSource === "upload" ? "Uploaded file" : datasetData?.description;
+ datasetSource === "upload" ? uploadedFile : dataset;
return (
-
+
System
@@ -103,10 +91,9 @@ export function SummaryStep() {
-
- {modelData?.name ?? "—"}
+
+ {selectedModel ?? "—"}
- {modelData?.params}
@@ -129,14 +116,8 @@ export function SummaryStep() {
- {datasetName ?? "—"}
-
- {datasetDesc}
-
+ {datasetName ?? "—"}
- {datasetData?.size && (
-
{datasetData.size}
- )}
diff --git a/studio/frontend/src/features/studio/sections/dataset-section.tsx b/studio/frontend/src/features/studio/sections/dataset-section.tsx
index 466565a479..0b2c2652ff 100644
--- a/studio/frontend/src/features/studio/sections/dataset-section.tsx
+++ b/studio/frontend/src/features/studio/sections/dataset-section.tsx
@@ -44,18 +44,27 @@ import { useShallow } from "zustand/react/shallow";
export function DatasetSection() {
const { dataset, setDataset, datasetFormat, setDatasetFormat, hfToken } =
useWizardStore(
- useShallow((s) => ({
- dataset: s.dataset,
- setDataset: s.setDataset,
- datasetFormat: s.datasetFormat,
- setDatasetFormat: s.setDatasetFormat,
- hfToken: s.hfToken,
+ useShallow(({ dataset, setDataset, datasetFormat, setDatasetFormat, hfToken }) => ({
+ dataset, setDataset, datasetFormat, setDatasetFormat, hfToken,
})),
);
const [inputValue, setInputValue] = useState("");
const selectingRef = useRef(false);
const debouncedQuery = useDebouncedValue(inputValue);
+
+ function handleDatasetSelect(id: string | null) {
+ selectingRef.current = true;
+ setDataset(id);
+ }
+
+ function handleInputChange(val: string) {
+ if (selectingRef.current) {
+ selectingRef.current = false;
+ return;
+ }
+ setInputValue(val);
+ }
const {
results: hfResults,
isLoading,
@@ -65,7 +74,13 @@ export function DatasetSection() {
accessToken: hfToken || undefined,
});
- const resultIds = useMemo(() => hfResults.map((r) => r.id), [hfResults]);
+ const resultIds = useMemo(() => {
+ const ids = hfResults.map((r) => r.id);
+ if (dataset && !ids.includes(dataset)) {
+ ids.unshift(dataset);
+ }
+ return ids;
+ }, [hfResults, dataset]);
const comboboxAnchorRef = useRef
(null);
const { scrollRef, sentinelRef } = useInfiniteScroll(
@@ -82,7 +97,6 @@ export function DatasetSection() {
className="lg:col-span-4 min-h-[450px]"
>
- {/* Load from Hub */}
Load from Hub
@@ -118,8 +132,8 @@ export function DatasetSection() {
filteredItems={resultIds}
filter={null}
value={dataset}
- onValueChange={(id) => { selectingRef.current = true; setDataset(id); }}
- onInputValueChange={(val) => { if (selectingRef.current) { selectingRef.current = false; return; } setInputValue(val); }}
+ onValueChange={handleDatasetSelect}
+ onInputValueChange={handleInputChange}
itemToStringValue={(id) => id}
autoHighlight={true}
>
@@ -145,10 +159,14 @@ export function DatasetSection() {
>
{(id: string) => {
- const r = hfResults.find((r) => r.id === id);
+ const r = hfResults.find((ds) => ds.id === id);
const detail = r?.totalExamples
? `${formatCompact(r.totalExamples)} rows`
- : (r?.sizeCategory ?? null);
+ : r?.sizeCategory
+ ? r.sizeCategory
+ : r?.downloads != null
+ ? `↓${formatCompact(r.downloads)}`
+ : null;
return (
- {detail ? (
+ {detail && (
{detail}
- ) : r?.downloads != null ? (
-
- ↓{formatCompact(r.downloads)}
-
- ) : null}
+ )}
);
}}
@@ -193,7 +207,6 @@ export function DatasetSection() {
- {/* Format */}
Dataset Format
@@ -239,7 +252,6 @@ export function DatasetSection() {
- {/* Active dataset display */}
{dataset ? (
@@ -269,7 +281,6 @@ export function DatasetSection() {
)}
- {/* Action buttons */}