From 51f5b5f262d9aecf360e32cbd775a9b43b109d88 Mon Sep 17 00:00:00 2001 From: shine1i Date: Mon, 2 Feb 2026 11:08:31 +0100 Subject: [PATCH 01/12] feat: replace config summary with model export feature, including export methods, quantization options, and new UI components --- .claude/settings.local.json | 15 ++ studio/frontend/bun.lock | 31 ++- studio/frontend/package.json | 1 + studio/frontend/src/app/router.tsx | 2 + studio/frontend/src/app/routes/export.tsx | 15 ++ studio/frontend/src/components/navbar.tsx | 2 +- .../frontend/src/components/section-card.tsx | 2 +- studio/frontend/src/config/training.ts | 17 +- .../export/components/export-dialog.tsx | 178 ++++++++++++++ .../export/components/method-picker.tsx | 87 +++++++ .../export/components/quant-picker.tsx | 84 +++++++ .../frontend/src/features/export/constants.ts | 56 +++++ .../src/features/export/export-page.tsx | 230 ++++++++++++++++++ studio/frontend/src/features/export/index.ts | 1 + .../sections/config-summary-section.tsx | 88 ------- .../studio/sections/model-section.tsx | 15 +- studio/frontend/src/hooks/index.ts | 5 +- .../frontend/src/hooks/use-debounced-value.ts | 10 + .../src/hooks/use-hf-dataset-search.ts | 72 ++++++ .../frontend/src/hooks/use-hf-model-search.ts | 74 ++++++ studio/frontend/src/types/training.ts | 1 + 21 files changed, 862 insertions(+), 124 deletions(-) create mode 100644 .claude/settings.local.json create mode 100644 studio/frontend/src/app/routes/export.tsx create mode 100644 studio/frontend/src/features/export/components/export-dialog.tsx create mode 100644 studio/frontend/src/features/export/components/method-picker.tsx create mode 100644 studio/frontend/src/features/export/components/quant-picker.tsx create mode 100644 studio/frontend/src/features/export/constants.ts create mode 100644 studio/frontend/src/features/export/export-page.tsx create mode 100644 studio/frontend/src/features/export/index.ts delete mode 100644 studio/frontend/src/features/studio/sections/config-summary-section.tsx create mode 100644 studio/frontend/src/hooks/use-debounced-value.ts create mode 100644 studio/frontend/src/hooks/use-hf-dataset-search.ts create mode 100644 studio/frontend/src/hooks/use-hf-model-search.ts diff --git a/.claude/settings.local.json b/.claude/settings.local.json new file mode 100644 index 0000000000..7802dc0461 --- /dev/null +++ b/.claude/settings.local.json @@ -0,0 +1,15 @@ +{ + "permissions": { + "allow": [ + "Bash(tree:*)", + "Bash(findstr:*)", + "Bash(bun run typecheck:*)", + "mcp__plugin_serena_serena__list_dir", + "Bash(bun x tsc:*)", + "mcp__plugin_perplexity_perplexity__perplexity_ask", + "WebSearch", + "WebFetch(domain:www.npmjs.com)", + "WebFetch(domain:github.com)" + ] + } +} diff --git a/studio/frontend/bun.lock b/studio/frontend/bun.lock index e33edd872d..1417bea8db 100644 --- a/studio/frontend/bun.lock +++ b/studio/frontend/bun.lock @@ -14,6 +14,7 @@ "@fontsource-variable/space-grotesk": "^5.2.10", "@hugeicons/core-free-icons": "^3.1.1", "@hugeicons/react": "^1.1.4", + "@huggingface/hub": "^2.8.0", "@radix-ui/react-select": "^2.2.6", "@radix-ui/react-slot": "^1.2.3", "@streamdown/cjk": "^1.0.1", @@ -276,6 +277,10 @@ "@hugeicons/react": ["@hugeicons/react@1.1.4", "", { "peerDependencies": { "react": ">=16.0.0" } }, "sha512-gsc3eZyd2fGqRUThW9+lfjxxsOkz6KNVmRXRgJjP32GL0OnnLJnl3hytKt47CBbiQj2xE2kCw+rnP3UQCThcKw=="], + "@huggingface/hub": ["@huggingface/hub@2.8.0", "", { "dependencies": { "@huggingface/tasks": "^0.19.80" }, "optionalDependencies": { "cli-progress": "^3.12.0" }, "bin": { "hfjs": "dist/cli.js" } }, "sha512-eh7lXCrZeNor2YE+2jn2F75/GEzq+TAh81jLTskiqBNcKuBes0I7TIGP1qiDC+rupP765C5vx3XMNxXQVN3N9w=="], + + "@huggingface/tasks": ["@huggingface/tasks@0.19.82", "", {}, "sha512-i8TzJb6Zk7KnYRL8unnYRuh/tW7ku3hJtQw972q/ZXvjkr/YIwWAAVxEfXuaj0vtCUbUljamJji6oWVG4+0PLQ=="], + "@humanfs/core": ["@humanfs/core@0.19.1", "", {}, "sha512-5DyQ4+1JEUzejeK1JGICcideyfUbGixgS9jNgex5nqkW+cY7WZhxBigmieN5Qnw9ZosSNVC9KQKyb+GUaGyKUA=="], "@humanfs/node": ["@humanfs/node@0.16.7", "", { "dependencies": { "@humanfs/core": "^0.19.1", "@humanwhocodes/retry": "^0.4.0" } }, "sha512-/zUx+yOsIrG4Y43Eh2peDeKCxlRt/gET6aHfaKpuq267qXdYDFViVHfMaLyygZOnl0kGWxFIgsBy8QFuTLUXEQ=="], @@ -786,6 +791,8 @@ "cli-cursor": ["cli-cursor@5.0.0", "", { "dependencies": { "restore-cursor": "^5.0.0" } }, "sha512-aCj4O5wKyszjMmDT4tZj93kxyydN/K5zPWSCe6/0AV/AA1pqe5ZBIw0a2ZfPQV7lL5/yb5HsUreJ6UFAF1tEQw=="], + "cli-progress": ["cli-progress@3.12.0", "", { "dependencies": { "string-width": "^4.2.3" } }, "sha512-tRkV3HJ1ASwm19THiiLIXLO7Im7wlTuKnvkYaTkyoAPefqjNg7W7DHKUlGRxy9vxDvbyCYQkQozvptuMkGCg8A=="], + "cli-spinners": ["cli-spinners@2.9.2", "", {}, "sha512-ywqV+5MmyL4E7ybXgKys4DugZbX0FC6LnwrhjuykIjnK9k8OQacQ7axGKnjDXWNhns0xot3bZI5h55H8yo9cJg=="], "cli-width": ["cli-width@4.1.0", "", {}, "sha512-ouuZd4/dm2Sw5Gmqy6bGyNNNe1qt9RpmxveLSO7KcgsTnU7RXfsw+/bukWGo1abgBiMAic068rclZsO4IWmmxQ=="], @@ -962,7 +969,7 @@ "electron-to-chromium": ["electron-to-chromium@1.5.278", "", {}, "sha512-dQ0tM1svDRQOwxnXxm+twlGTjr9Upvt8UFWAgmLsxEzFQxhbti4VwxmMjsDxVC51Zo84swW7FVCXEV+VAkhuPw=="], - "emoji-regex": ["emoji-regex@10.6.0", "", {}, "sha512-toUI84YS5YmxW219erniWD0CIVOo46xGKColeNQRgOzDorgBi1v4D71/OFzgD9GO2UGKIv1C3Sp8DAn0+j5w7A=="], + "emoji-regex": ["emoji-regex@8.0.0", "", {}, "sha512-MSjYzcWNOA0ewAHpz0MxpYFvwg6yjy1NG3xteoqz644VCo/RPgnr1/GGt+ic3iJTzQ8Eu3TdM14SawnVUmGE6A=="], "encodeurl": ["encodeurl@2.0.0", "", {}, "sha512-Q0n9HRi4m6JuGIV1eFlmvJB7ZEVxu93IrMyiMsGC0lrMJMWzRgx6WGquyfQgZVb31vhGgXnfmPNNXmxnOkRBrg=="], @@ -1700,7 +1707,7 @@ "strict-event-emitter": ["strict-event-emitter@0.5.1", "", {}, "sha512-vMgjE/GGEPEFnhFub6pa4FmJBRBVOLpIII2hvCZ8Kzb7K0hlHo7mQv6xYrBvCL2LtAIBwFUK8wvuJgTVSQ5MFQ=="], - "string-width": ["string-width@7.2.0", "", { "dependencies": { "emoji-regex": "^10.3.0", "get-east-asian-width": "^1.0.0", "strip-ansi": "^7.1.0" } }, "sha512-tsaTIkKW9b4N+AEj+SVA+WhJzV7/zMhcSu78mLKWSk7cXMOSHsBKFWUs0fWwq8QyK3MgJBQRX6Gbi4kYbdvGkQ=="], + "string-width": ["string-width@4.2.3", "", { "dependencies": { "emoji-regex": "^8.0.0", "is-fullwidth-code-point": "^3.0.0", "strip-ansi": "^6.0.1" } }, "sha512-wKyQRQpjJ0sIp62ErSZdGsjMJWsap5oRNihHhu6G7JVO/9jIB6UyevL+tXuOqrng8j/cxKTWyWUwvSTriiZz/g=="], "stringify-entities": ["stringify-entities@4.0.4", "", { "dependencies": { "character-entities-html4": "^2.0.0", "character-entities-legacy": "^3.0.0" } }, "sha512-IwfBptatlO+QCJUo19AqvrPNqlVMpW9YEL2LIVY+Rpv2qsjCGxaDLNRgeGsQWJhfItebuJhsGSLjaBbNSQ+ieg=="], @@ -2086,8 +2093,6 @@ "chevrotain/lodash-es": ["lodash-es@4.17.21", "", {}, "sha512-mKnC+QJ9pWVzv+C4/U3rRsHapFfHvQFoFB92e52xeyGMcX6/OlIl78je1u8vePzYZSkkogMPJ2yjxxsb89cxyw=="], - "cliui/string-width": ["string-width@4.2.3", "", { "dependencies": { "emoji-regex": "^8.0.0", "is-fullwidth-code-point": "^3.0.0", "strip-ansi": "^6.0.1" } }, "sha512-wKyQRQpjJ0sIp62ErSZdGsjMJWsap5oRNihHhu6G7JVO/9jIB6UyevL+tXuOqrng8j/cxKTWyWUwvSTriiZz/g=="], - "cliui/strip-ansi": ["strip-ansi@6.0.1", "", { "dependencies": { "ansi-regex": "^5.0.1" } }, "sha512-Y38VPSHcqkFrCpFnQ9vuSXmquuv5oXOKpGeT6aGrr3o3Gc9AlVa6JBfUSOCnbxGGZF+/0ooI7KrPuUSztUdU5A=="], "cliui/wrap-ansi": ["wrap-ansi@7.0.0", "", { "dependencies": { "ansi-styles": "^4.0.0", "string-width": "^4.1.0", "strip-ansi": "^6.0.0" } }, "sha512-YVGIj2kamLSTxw6NsZjoBxfSwsn0ycdesmc4p+Q21c5zPuZ1pl+NfxVdxPtdHvmNVOQ6XSYG4AUtyt/Fi7D16Q=="], @@ -2124,6 +2129,8 @@ "ora/chalk": ["chalk@5.6.2", "", {}, "sha512-7NzBL0rN6fMUW+f7A6Io4h40qQlG+xGmtMxfbnH/K7TAtt8JQWVQK+6g0UXKMeVJoyV5EkkNsErQ8pVD3bLHbA=="], + "ora/string-width": ["string-width@7.2.0", "", { "dependencies": { "emoji-regex": "^10.3.0", "get-east-asian-width": "^1.0.0", "strip-ansi": "^7.1.0" } }, "sha512-tsaTIkKW9b4N+AEj+SVA+WhJzV7/zMhcSu78mLKWSk7cXMOSHsBKFWUs0fWwq8QyK3MgJBQRX6Gbi4kYbdvGkQ=="], + "parse-entities/@types/unist": ["@types/unist@2.0.11", "", {}, "sha512-CmBKiL6NNo/OqgmMn95Fk9Whlp2mtvIv+KNpQKN2F4SjvrEesubTRWGYSg+BnWZOnlCaSTU1sMpsBOzgbYhnsA=="], "postcss/nanoid": ["nanoid@3.3.11", "", { "bin": { "nanoid": "bin/nanoid.cjs" } }, "sha512-N8SpfPUnUp1bK+PMYW8qSWdl9U+wwNWI4QKxOYDy9JAro3WMX7p2OeVRF9v+347pnakNevPmiHhNmZ2HbFA76w=="], @@ -2146,12 +2153,10 @@ "shadcn/zod": ["zod@3.25.76", "", {}, "sha512-gzUt/qt81nXsFGKIFcC3YnfEAx5NkunCfnDlvuBSSFS02bcXu4Lmea0AFIUwbLWxWPx3d9p8S5QoaujKcNQxcQ=="], - "wrap-ansi/string-width": ["string-width@4.2.3", "", { "dependencies": { "emoji-regex": "^8.0.0", "is-fullwidth-code-point": "^3.0.0", "strip-ansi": "^6.0.1" } }, "sha512-wKyQRQpjJ0sIp62ErSZdGsjMJWsap5oRNihHhu6G7JVO/9jIB6UyevL+tXuOqrng8j/cxKTWyWUwvSTriiZz/g=="], + "string-width/strip-ansi": ["strip-ansi@6.0.1", "", { "dependencies": { "ansi-regex": "^5.0.1" } }, "sha512-Y38VPSHcqkFrCpFnQ9vuSXmquuv5oXOKpGeT6aGrr3o3Gc9AlVa6JBfUSOCnbxGGZF+/0ooI7KrPuUSztUdU5A=="], "wrap-ansi/strip-ansi": ["strip-ansi@6.0.1", "", { "dependencies": { "ansi-regex": "^5.0.1" } }, "sha512-Y38VPSHcqkFrCpFnQ9vuSXmquuv5oXOKpGeT6aGrr3o3Gc9AlVa6JBfUSOCnbxGGZF+/0ooI7KrPuUSztUdU5A=="], - "yargs/string-width": ["string-width@4.2.3", "", { "dependencies": { "emoji-regex": "^8.0.0", "is-fullwidth-code-point": "^3.0.0", "strip-ansi": "^6.0.1" } }, "sha512-wKyQRQpjJ0sIp62ErSZdGsjMJWsap5oRNihHhu6G7JVO/9jIB6UyevL+tXuOqrng8j/cxKTWyWUwvSTriiZz/g=="], - "@dotenvx/dotenvx/execa/get-stream": ["get-stream@6.0.1", "", {}, "sha512-ts6Wi+2j3jQjqi70w5AlN8DFnkSwC+MqmxEzdEALB2qXZYV3X/b1CTfgPLGJNMeAWxdPfU8FO1ms3NUfaHCPYg=="], "@dotenvx/dotenvx/execa/human-signals": ["human-signals@2.1.0", "", {}, "sha512-B4FFZ6q/T2jhhksgkbEW3HBvWIfDW85snkQgawt07S7J5QXTk6BkNV+0yAeZrM5QpMAdYlocGoljn0sJ/WQkFw=="], @@ -2238,8 +2243,6 @@ "ajv-formats/ajv/json-schema-traverse": ["json-schema-traverse@1.0.0", "", {}, "sha512-NM8/P9n3XjXhIZn1lLhkFaACTOURQXjWhV4BA/RnOv8xvgqtqpAX9IO4mRQxSx1Rlo4tqzeqb0sOlruaOy3dug=="], - "cliui/string-width/emoji-regex": ["emoji-regex@8.0.0", "", {}, "sha512-MSjYzcWNOA0ewAHpz0MxpYFvwg6yjy1NG3xteoqz644VCo/RPgnr1/GGt+ic3iJTzQ8Eu3TdM14SawnVUmGE6A=="], - "cliui/strip-ansi/ansi-regex": ["ansi-regex@5.0.1", "", {}, "sha512-quJQXlTSUGL2LH9SUXo8VwsY4soanhgo6LNSm84E1LBcE8s3O0wpdiRzyR9z/ZZJMlMWv37qOOb9pdJlMUEKFQ=="], "cmdk/@radix-ui/react-primitive/@radix-ui/react-slot": ["@radix-ui/react-slot@1.2.3", "", { "dependencies": { "@radix-ui/react-compose-refs": "1.1.2" }, "peerDependencies": { "@types/react": "*", "react": "^16.8 || ^17.0 || ^18.0 || ^19.0 || ^19.0.0-rc" }, "optionalPeers": ["@types/react"] }, "sha512-aeNmHnBxbi2St0au6VBVC7JXFlhLlOnvIIlePNniyUNAClzmtAUEY8/pBiK3iHjufOlwA+c20/8jngo7xcrg8A=="], @@ -2254,14 +2257,10 @@ "motion/framer-motion/motion-utils": ["motion-utils@12.29.2", "", {}, "sha512-G3kc34H2cX2gI63RqU+cZq+zWRRPSsNIOjpdl9TN4AQwC4sgwYPl/Q/Obf/d53nOm569T0fYK+tcoSV50BWx8A=="], - "wrap-ansi/string-width/emoji-regex": ["emoji-regex@8.0.0", "", {}, "sha512-MSjYzcWNOA0ewAHpz0MxpYFvwg6yjy1NG3xteoqz644VCo/RPgnr1/GGt+ic3iJTzQ8Eu3TdM14SawnVUmGE6A=="], + "ora/string-width/emoji-regex": ["emoji-regex@10.6.0", "", {}, "sha512-toUI84YS5YmxW219erniWD0CIVOo46xGKColeNQRgOzDorgBi1v4D71/OFzgD9GO2UGKIv1C3Sp8DAn0+j5w7A=="], + + "string-width/strip-ansi/ansi-regex": ["ansi-regex@5.0.1", "", {}, "sha512-quJQXlTSUGL2LH9SUXo8VwsY4soanhgo6LNSm84E1LBcE8s3O0wpdiRzyR9z/ZZJMlMWv37qOOb9pdJlMUEKFQ=="], "wrap-ansi/strip-ansi/ansi-regex": ["ansi-regex@5.0.1", "", {}, "sha512-quJQXlTSUGL2LH9SUXo8VwsY4soanhgo6LNSm84E1LBcE8s3O0wpdiRzyR9z/ZZJMlMWv37qOOb9pdJlMUEKFQ=="], - - "yargs/string-width/emoji-regex": ["emoji-regex@8.0.0", "", {}, "sha512-MSjYzcWNOA0ewAHpz0MxpYFvwg6yjy1NG3xteoqz644VCo/RPgnr1/GGt+ic3iJTzQ8Eu3TdM14SawnVUmGE6A=="], - - "yargs/string-width/strip-ansi": ["strip-ansi@6.0.1", "", { "dependencies": { "ansi-regex": "^5.0.1" } }, "sha512-Y38VPSHcqkFrCpFnQ9vuSXmquuv5oXOKpGeT6aGrr3o3Gc9AlVa6JBfUSOCnbxGGZF+/0ooI7KrPuUSztUdU5A=="], - - "yargs/string-width/strip-ansi/ansi-regex": ["ansi-regex@5.0.1", "", {}, "sha512-quJQXlTSUGL2LH9SUXo8VwsY4soanhgo6LNSm84E1LBcE8s3O0wpdiRzyR9z/ZZJMlMWv37qOOb9pdJlMUEKFQ=="], } } diff --git a/studio/frontend/package.json b/studio/frontend/package.json index 5a2f5a3063..e09c04695b 100644 --- a/studio/frontend/package.json +++ b/studio/frontend/package.json @@ -22,6 +22,7 @@ "@fontsource-variable/space-grotesk": "^5.2.10", "@hugeicons/core-free-icons": "^3.1.1", "@hugeicons/react": "^1.1.4", + "@huggingface/hub": "^2.8.0", "@radix-ui/react-select": "^2.2.6", "@radix-ui/react-slot": "^1.2.3", "@streamdown/cjk": "^1.0.1", diff --git a/studio/frontend/src/app/router.tsx b/studio/frontend/src/app/router.tsx index 95012e1a83..e2771919da 100644 --- a/studio/frontend/src/app/router.tsx +++ b/studio/frontend/src/app/router.tsx @@ -4,6 +4,7 @@ import { Route as chatRoute } from "./routes/chat"; import { Route as gridTestRoute } from "./routes/grid-test"; import { Route as homeRoute } from "./routes/home"; import { Route as onboardingRoute } from "./routes/onboarding"; +import { Route as exportRoute } from "./routes/export"; import { Route as studioRoute } from "./routes/studio"; const routeTree = rootRoute.addChildren([ @@ -12,6 +13,7 @@ const routeTree = rootRoute.addChildren([ gridTestRoute, studioRoute, chatRoute, + exportRoute, ]); export const router = createRouter({ routeTree }); diff --git a/studio/frontend/src/app/routes/export.tsx b/studio/frontend/src/app/routes/export.tsx new file mode 100644 index 0000000000..93bb83cd98 --- /dev/null +++ b/studio/frontend/src/app/routes/export.tsx @@ -0,0 +1,15 @@ +import { createRoute } from "@tanstack/react-router"; +import { lazy } from "react"; +import { Route as rootRoute } from "./__root"; + +const ExportPage = lazy(() => + import("@/features/export/export-page").then((m) => ({ + default: m.ExportPage, + })), +); + +export const Route = createRoute({ + getParentRoute: () => rootRoute, + path: "/export", + component: ExportPage, +}); diff --git a/studio/frontend/src/components/navbar.tsx b/studio/frontend/src/components/navbar.tsx index 3ccb8ef579..6153a74cc3 100644 --- a/studio/frontend/src/components/navbar.tsx +++ b/studio/frontend/src/components/navbar.tsx @@ -17,7 +17,7 @@ import { useState } from "react"; const NAV_ITEMS = [ { label: "Studio", href: "/studio", icon: ZapIcon, enabled: true }, { label: "Evaluate", href: "/evaluate", enabled: false }, - { label: "Export", href: "/export", enabled: false }, + { label: "Export", href: "/export", enabled: true }, { label: "Chat", href: "/chat", enabled: true }, ]; diff --git a/studio/frontend/src/components/section-card.tsx b/studio/frontend/src/components/section-card.tsx index aee6fdfa05..1548fd57db 100644 --- a/studio/frontend/src/components/section-card.tsx +++ b/studio/frontend/src/components/section-card.tsx @@ -52,7 +52,7 @@ export function SectionCard({ return (
= { + text: "text-generation", + vision: "image-text-to-text", + tts: "text-to-speech", + embeddings: "feature-extraction", +}; diff --git a/studio/frontend/src/features/export/components/export-dialog.tsx b/studio/frontend/src/features/export/components/export-dialog.tsx new file mode 100644 index 0000000000..1e6ff9962b --- /dev/null +++ b/studio/frontend/src/features/export/components/export-dialog.tsx @@ -0,0 +1,178 @@ +import { Button } from "@/components/ui/button"; +import { + Dialog, + DialogContent, + DialogDescription, + DialogFooter, + DialogHeader, + DialogTitle, +} from "@/components/ui/dialog"; +import { Input } from "@/components/ui/input"; +import { + InputGroup, + InputGroupAddon, + InputGroupInput, +} from "@/components/ui/input-group"; +import { Switch } from "@/components/ui/switch"; +import { ArrowRight01Icon, Key01Icon } from "@hugeicons/core-free-icons"; +import { HugeiconsIcon } from "@hugeicons/react"; +import { AnimatePresence, motion } from "motion/react"; +import { EXPORT_METHODS, type ExportMethod } from "../constants"; + +type Destination = "local" | "hub"; + +const anim = { + initial: { height: 0, opacity: 0 }, + animate: { height: "auto" as const, opacity: 1 }, + exit: { height: 0, opacity: 0 }, + transition: { duration: 0.3, ease: [0.25, 0.1, 0.25, 1] as const }, +}; + +interface ExportDialogProps { + open: boolean; + onOpenChange: (open: boolean) => void; + checkpoint: string | null; + exportMethod: ExportMethod | null; + quantLevels: string[]; + estimatedSize: string; + baseModelName: string; + isAdapter: boolean; + destination: Destination; + onDestinationChange: (v: Destination) => void; + hfUsername: string; + onHfUsernameChange: (v: string) => void; + modelName: string; + onModelNameChange: (v: string) => void; + hfToken: string; + onHfTokenChange: (v: string) => void; + privateRepo: boolean; + onPrivateRepoChange: (v: boolean) => void; +} + +export function ExportDialog({ + open, + onOpenChange, + checkpoint, + exportMethod, + quantLevels, + estimatedSize, + baseModelName, + isAdapter, + destination, + onDestinationChange, + hfUsername, + onHfUsernameChange, + modelName, + onModelNameChange, + hfToken, + onHfTokenChange, + privateRepo, + onPrivateRepoChange, +}: ExportDialogProps) { + return ( + + + + Export Model + Choose where to save your exported model. + + +
+ + +
+ + + {destination === "hub" && ( + +
+
+
+ + onHfUsernameChange(e.target.value)} /> +
+
+ + onModelNameChange(e.target.value)} /> +
+
+ +
+
+ + + Get token + + +
+ + + + + onHfTokenChange(e.target.value)} /> + +

Leave empty if already logged in via CLI.

+
+ +
+ + +
+
+
+ )} +
+ + {/* Summary */} +
+
+ Base Model + {baseModelName} +
+
+ {isAdapter ? "Checkpoint" : "Model"} + {checkpoint} +
+
+ Export Method + + {EXPORT_METHODS.find((m) => m.value === exportMethod)?.title} + +
+ {exportMethod === "gguf" && quantLevels.length > 0 && ( +
+ Quantizations + {quantLevels.join(", ")} +
+ )} +
+ Est. size + {estimatedSize} +
+
+ + + + + +
+
+ ); +} diff --git a/studio/frontend/src/features/export/components/method-picker.tsx b/studio/frontend/src/features/export/components/method-picker.tsx new file mode 100644 index 0000000000..b0052f3069 --- /dev/null +++ b/studio/frontend/src/features/export/components/method-picker.tsx @@ -0,0 +1,87 @@ +import { Badge } from "@/components/ui/badge"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/ui/tooltip"; +import { cn } from "@/lib/utils"; +import { CheckmarkCircle01Icon, InformationCircleIcon } from "@hugeicons/core-free-icons"; +import { HugeiconsIcon } from "@hugeicons/react"; +import { EXPORT_METHODS, type ExportMethod } from "../constants"; + +interface MethodPickerProps { + value: ExportMethod | null; + onChange: (v: ExportMethod) => void; +} + +export function MethodPicker({ value, onChange }: MethodPickerProps) { + return ( +
+ + Export Method + + + + + + How your model is packaged for deployment.{" "} + Read more + + + +
+ {EXPORT_METHODS.map((m) => { + const selected = value === m.value; + return ( + + ); + })} +
+
+ ); +} diff --git a/studio/frontend/src/features/export/components/quant-picker.tsx b/studio/frontend/src/features/export/components/quant-picker.tsx new file mode 100644 index 0000000000..e16dbd5494 --- /dev/null +++ b/studio/frontend/src/features/export/components/quant-picker.tsx @@ -0,0 +1,84 @@ +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/ui/tooltip"; +import { cn } from "@/lib/utils"; +import { CheckmarkCircle01Icon, InformationCircleIcon, LayersIcon } from "@hugeicons/core-free-icons"; +import { HugeiconsIcon } from "@hugeicons/react"; +import { QUANT_OPTIONS } from "../constants"; + +interface QuantPickerProps { + value: string[]; + onChange: (v: string[]) => void; +} + +export function QuantPicker({ value, onChange }: QuantPickerProps) { + const toggle = (qv: string) => { + onChange( + value.includes(qv) ? value.filter((q) => q !== qv) : [...value, qv], + ); + }; + + return ( +
+
+ + Quantization Levels + + + + + + Lower quantization (Q2, Q3) = smaller files but reduced quality. Q4–Q5 is a good balance.{" "} + Read more + + + — select one or more +
+
+ {QUANT_OPTIONS.map((q) => { + const active = value.includes(q.value); + return ( + + ); + })} +
+ {value.length > 0 && ( +
+ + {value.length} selected + + +
+ )} +
+ ); +} diff --git a/studio/frontend/src/features/export/constants.ts b/studio/frontend/src/features/export/constants.ts new file mode 100644 index 0000000000..b289372f05 --- /dev/null +++ b/studio/frontend/src/features/export/constants.ts @@ -0,0 +1,56 @@ +import type { TrainingMethod } from "@/types/training"; + +export type ExportMethod = "merged" | "lora" | "gguf"; + +export const EXPORT_METHODS: { + value: ExportMethod; + title: string; + description: string; + tooltip: string; + badge?: string; +}[] = [ + { value: "merged", title: "Merged Model", description: "Full 16-bit model ready for inference.", tooltip: "Merges adapter weights into the base model. Best for direct deployment with vLLM or TGI." }, + { value: "lora", title: "LoRA Only", description: "Lightweight adapter files (~100 MB). Needs base model.", tooltip: "Exports only the trained adapter. Pair with the base model at inference time to save storage." }, + { value: "gguf", title: "GGUF / Llama.cpp", description: "Quantized formats for local AI runners.", tooltip: "Converts to GGUF for llama.cpp, Ollama, and other local runners. Pick a quantization level below." }, +]; + +export const QUANT_OPTIONS = [ + { value: "q2_k", label: "Q2_K", size: "~2.5 GB" }, + { value: "iq3_m", label: "IQ3_M", size: "~3.1 GB" }, + { value: "q3_k_m", label: "Q3_K_M", size: "~3.5 GB" }, + { value: "q4_0", label: "Q4_0", size: "~4.1 GB" }, + { value: "q4_k_m", label: "Q4_K_M", size: "~4.8 GB", recommended: true }, + { value: "q5_0", label: "Q5_0", size: "~5.0 GB" }, + { value: "q5_k_m", label: "Q5_K_M", size: "~5.6 GB" }, + { value: "q6_k", label: "Q6_K", size: "~6.6 GB" }, + { value: "q8_0", label: "Q8_0", size: "~8.2 GB" }, + { value: "f16", label: "F16", size: "~14.2 GB" }, +]; + +export function getEstimatedSize(method: ExportMethod | null, quantLevels: string[]) { + const sizeOf = (v: string) => QUANT_OPTIONS.find((q) => q.value === v)?.size ?? "—"; + if (method === "gguf" && quantLevels.length > 0) { + if (quantLevels.length === 1) return sizeOf(quantLevels[0]); + const total = quantLevels + .map((q) => Number.parseFloat(sizeOf(q).replace(/[^0-9.]/g, ""))) + .reduce((a, b) => a + b, 0); + return `~${total.toFixed(1)} GB (${quantLevels.length} files)`; + } + if (method === "merged") return "~14.2 GB"; + if (method === "lora") return "~100 MB"; + return "—"; +} + +export const METHOD_LABELS: Record = { + qlora: "QLoRA", + lora: "LoRA", + full: "Full Fine-tune", +}; + +export const GUIDE_STEPS = [ + "Select a training checkpoint to export from", + "Choose an export method based on your use case", + "Pick quantization levels if using GGUF", + "Click Export and choose your destination", + "Test your model and compare outputs in Chat", +]; diff --git a/studio/frontend/src/features/export/export-page.tsx b/studio/frontend/src/features/export/export-page.tsx new file mode 100644 index 0000000000..0ba0c817a8 --- /dev/null +++ b/studio/frontend/src/features/export/export-page.tsx @@ -0,0 +1,230 @@ +import { Button } from "@/components/ui/button"; +import { + Select, + SelectContent, + SelectItem, + SelectTrigger, + SelectValue, +} from "@/components/ui/select"; +import { Separator } from "@/components/ui/separator"; +import { SectionCard } from "@/components/section-card"; +import { MODELS } from "@/config/training"; +import { useWizardStore } from "@/stores/training"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/ui/tooltip"; +import { InformationCircleIcon, PackageIcon } from "@hugeicons/core-free-icons"; +import { HugeiconsIcon } from "@hugeicons/react"; +import { AnimatePresence, motion } from "motion/react"; +import { useMemo, useState } from "react"; +import { ExportDialog } from "./components/export-dialog"; +import { MethodPicker } from "./components/method-picker"; +import { QuantPicker } from "./components/quant-picker"; +import { + type ExportMethod, + GUIDE_STEPS, + METHOD_LABELS, + getEstimatedSize, +} from "./constants"; + +const anim = { + initial: { height: 0, opacity: 0 }, + animate: { height: "auto" as const, opacity: 1 }, + exit: { height: 0, opacity: 0 }, + transition: { duration: 0.3, ease: [0.25, 0.1, 0.25, 1] as const }, +}; + +export function ExportPage() { + const store = useWizardStore(); + const isAdapter = store.trainingMethod === "lora" || store.trainingMethod === "qlora"; + const modelInfo = useMemo( + () => MODELS.find((m) => m.id === store.selectedModel), + [store.selectedModel], + ); + + const checkpoints = useMemo(() => { + if (isAdapter) { + const interval = store.saveSteps > 0 ? store.saveSteps : 100; + const total = store.trainingMetrics?.totalSteps ?? 500; + const entries: { value: string; label: string; detail: string }[] = []; + for (let step = interval; step <= total; step += interval) { + const loss = (1.5 - (step / total) * 0.7 + Math.random() * 0.05).toFixed(2); + entries.push({ + value: `checkpoint-${step}`, + label: `checkpoint-${step}`, + detail: step === total ? `Best Loss: ${loss}` : `Loss: ${loss}`, + }); + } + return entries.reverse(); + } + return [{ value: "final-model", label: "Final Model", detail: "Full fine-tuned weights" }]; + }, [isAdapter, store.saveSteps, store.trainingMetrics?.totalSteps]); + + const [checkpoint, setCheckpoint] = useState(null); + const [exportMethod, setExportMethod] = useState(null); + const [quantLevels, setQuantLevels] = useState([]); + const [dialogOpen, setDialogOpen] = useState(false); + + const [destination, setDestination] = useState<"local" | "hub">("local"); + const [hfUsername, setHfUsername] = useState(""); + const [modelName, setModelName] = useState(""); + const [privateRepo, setPrivateRepo] = useState(false); + + const handleMethodChange = (method: ExportMethod) => { + setExportMethod(method); + if (method !== "gguf") setQuantLevels([]); + }; + + const estimatedSize = getEstimatedSize(exportMethod, quantLevels); + const canExport = checkpoint && exportMethod && (exportMethod !== "gguf" || quantLevels.length > 0); + const baseModelName = modelInfo?.name ?? store.selectedModel ?? "—"; + + return ( +
+
+
+

Export Model

+

Export your fine-tuned model for deployment

+
+ + } + title="Export Configuration" + description="Select checkpoint, method, and quantization" + accent="emerald" + featured + className="shadow-border ring-1 ring-border" + > + {/* Top row: Checkpoint + metadata | Guide */} +
+
+
+ + +
+ +
+ Training Info +
+
+ Base Model + {baseModelName} +
+
+ Method + {METHOD_LABELS[store.trainingMethod] ?? store.trainingMethod} +
+
+ Checkpoints + {checkpoints.length} +
+
+ Epochs + {store.epochs} +
+ {isAdapter && ( +
+ LoRA Rank + {store.loraRank} +
+ )} + {modelInfo?.params && ( +
+ Params + {modelInfo.params} +
+ )} +
+
+
+ +
+ Quick Guide +
    + {GUIDE_STEPS.map((step, i) => ( +
  1. + + {i + 1} + + {step} +
  2. + ))} +
+
+
+ + + + + {exportMethod === "gguf" && ( + + + + )} + + + +
+
+ + Est. size: {estimatedSize} · Free disk space: 120 GB +
+ +
+
+
+ + +
+ ); +} diff --git a/studio/frontend/src/features/export/index.ts b/studio/frontend/src/features/export/index.ts new file mode 100644 index 0000000000..320f8b6d2e --- /dev/null +++ b/studio/frontend/src/features/export/index.ts @@ -0,0 +1 @@ +export { ExportPage } from "./export-page"; diff --git a/studio/frontend/src/features/studio/sections/config-summary-section.tsx b/studio/frontend/src/features/studio/sections/config-summary-section.tsx deleted file mode 100644 index db85fdc11d..0000000000 --- a/studio/frontend/src/features/studio/sections/config-summary-section.tsx +++ /dev/null @@ -1,88 +0,0 @@ -import { SectionCard } from "@/components/section-card"; -import { Button } from "@/components/ui/button"; -import { useWizardStore } from "@/stores/training"; -import { Settings02Icon, StopIcon } from "@hugeicons/core-free-icons"; -import { HugeiconsIcon } from "@hugeicons/react"; - -export function ConfigSummarySection() { - const store = useWizardStore(); - - const items = [ - { - section: "Model", - rows: [ - ["Model", store.selectedModel ?? "—"], - ["Type", store.modelType ?? "—"], - ["Method", store.trainingMethod], - ], - }, - { - section: "Dataset", - rows: [ - ["Source", store.datasetSource], - ["Dataset", store.dataset ?? store.uploadedFile ?? "—"], - ["Format", store.datasetFormat], - ], - }, - { - section: "Hyperparams", - rows: [ - ["Epochs", store.epochs], - ["Batch size", store.batchSize], - ["Learning rate", store.learningRate], - ["Max steps", store.maxSteps], - ["Context length", store.contextLength], - ["Warmup steps", store.warmupSteps], - ], - }, - ...(store.trainingMethod !== "full" - ? [ - { - section: "LoRA", - rows: [ - ["Rank", store.loraRank], - ["Alpha", store.loraAlpha], - ["Dropout", store.loraDropout], - ["Variant", store.loraVariant], - ], - }, - ] - : []), - ]; - - return ( - } - title="Config" - description="Training configuration" - accent="indigo" - className="lg:col-span-4" - > -
- {items.map((group) => ( -
-

- {group.section} -

- {group.rows.map(([label, value]) => ( -
- {String(label)} - - {String(value)} - -
- ))} -
- ))} - - -
-
- ); -} diff --git a/studio/frontend/src/features/studio/sections/model-section.tsx b/studio/frontend/src/features/studio/sections/model-section.tsx index 0fff945c8f..f9e5018779 100644 --- a/studio/frontend/src/features/studio/sections/model-section.tsx +++ b/studio/frontend/src/features/studio/sections/model-section.tsx @@ -29,19 +29,6 @@ import { HugeiconsIcon } from "@hugeicons/react"; import { useMemo } from "react"; import { useShallow } from "zustand/react/shallow"; -const HF_REPO_MAP: Record = { - "llava-1.6-7b": "unsloth/llava-v1.6-mistral-7b", - "llava-1.6-13b": "unsloth/llava-v1.6-vicuna-13b", - "qwen-vl-7b": "Qwen/Qwen-VL-Chat", - bark: "suno/bark", - "xtts-v2": "coqui/XTTS-v2", - "gemma-3-27b": "unsloth/gemma-3-27b", - "llama-3.1-8b": "unsloth/Llama-3.1-8B", - "mistral-7b": "unsloth/mistral-7b-v0.3", - "phi-4": "unsloth/phi-4", - "qwen-2.5-7b": "Qwen/Qwen2.5-7B", -}; - const DOT_COLORS = [ "bg-amber-400", "bg-blue-400", @@ -180,7 +167,7 @@ export function ModelSection() { placeholder="unsloth/gemma-3-27b" value={ selectedModel - ? (HF_REPO_MAP[selectedModel] ?? selectedModel) + ? (MODELS.find((m) => m.id === selectedModel)?.hfRepo ?? selectedModel) : "" } onChange={(e) => setSelectedModel(e.target.value || null)} diff --git a/studio/frontend/src/hooks/index.ts b/studio/frontend/src/hooks/index.ts index f3adbbdac9..0e981fab4f 100644 --- a/studio/frontend/src/hooks/index.ts +++ b/studio/frontend/src/hooks/index.ts @@ -1,2 +1,3 @@ -// Shared hooks -export {}; +export { useDebouncedValue } from "./use-debounced-value"; +export { useHfModelSearch } from "./use-hf-model-search"; +export { useHfDatasetSearch } from "./use-hf-dataset-search"; diff --git a/studio/frontend/src/hooks/use-debounced-value.ts b/studio/frontend/src/hooks/use-debounced-value.ts new file mode 100644 index 0000000000..860d5181a4 --- /dev/null +++ b/studio/frontend/src/hooks/use-debounced-value.ts @@ -0,0 +1,10 @@ +import { useEffect, useState } from "react"; + +export function useDebouncedValue(value: T, delayMs = 300): T { + const [debounced, setDebounced] = useState(value); + useEffect(() => { + const id = setTimeout(() => setDebounced(value), delayMs); + return () => clearTimeout(id); + }, [value, delayMs]); + return debounced; +} diff --git a/studio/frontend/src/hooks/use-hf-dataset-search.ts b/studio/frontend/src/hooks/use-hf-dataset-search.ts new file mode 100644 index 0000000000..727b6c1442 --- /dev/null +++ b/studio/frontend/src/hooks/use-hf-dataset-search.ts @@ -0,0 +1,72 @@ +import { listDatasets } from "@huggingface/hub"; +import { useEffect, useState } from "react"; + +export interface HfDatasetResult { + id: string; + downloads: number; + likes: number; +} + +interface HfSearchState { + results: HfDatasetResult[]; + isLoading: boolean; + error: string | null; +} + +export function useHfDatasetSearch( + query: string, + options?: { limit?: number; accessToken?: string }, +): HfSearchState { + const { limit = 20, accessToken } = options ?? {}; + const [state, setState] = useState({ + results: [], + isLoading: false, + error: null, + }); + + useEffect(() => { + if (!query.trim()) { + setState({ results: [], isLoading: false, error: null }); + return; + } + + let cancelled = false; + setState((prev) => ({ ...prev, isLoading: true, error: null })); + + (async () => { + try { + const results: HfDatasetResult[] = []; + const iter = listDatasets({ + search: { query }, + limit, + ...(accessToken ? { credentials: { accessToken } } : {}), + }); + for await (const ds of iter) { + if (cancelled) return; + results.push({ + id: ds.id, + downloads: ds.downloads, + likes: ds.likes, + }); + } + if (!cancelled) { + setState({ results, isLoading: false, error: null }); + } + } catch (err) { + if (!cancelled) { + setState({ + results: [], + isLoading: false, + error: err instanceof Error ? err.message : "Search failed", + }); + } + } + })(); + + return () => { + cancelled = true; + }; + }, [query, limit, accessToken]); + + return state; +} diff --git a/studio/frontend/src/hooks/use-hf-model-search.ts b/studio/frontend/src/hooks/use-hf-model-search.ts new file mode 100644 index 0000000000..fe0cc358e7 --- /dev/null +++ b/studio/frontend/src/hooks/use-hf-model-search.ts @@ -0,0 +1,74 @@ +import { listModels } from "@huggingface/hub"; +import { useEffect, useState } from "react"; + +export interface HfModelResult { + id: string; + downloads: number; + likes: number; + task?: string; +} + +interface HfSearchState { + results: HfModelResult[]; + isLoading: boolean; + error: string | null; +} + +export function useHfModelSearch( + query: string, + options?: { task?: string; limit?: number; accessToken?: string }, +): HfSearchState { + const { task, limit = 20, accessToken } = options ?? {}; + const [state, setState] = useState({ + results: [], + isLoading: false, + error: null, + }); + + useEffect(() => { + if (!query.trim()) { + setState({ results: [], isLoading: false, error: null }); + return; + } + + let cancelled = false; + setState((prev) => ({ ...prev, isLoading: true, error: null })); + + (async () => { + try { + const results: HfModelResult[] = []; + const iter = listModels({ + search: { query, ...(task ? { task } : {}) }, + limit, + ...(accessToken ? { credentials: { accessToken } } : {}), + }); + for await (const model of iter) { + if (cancelled) return; + results.push({ + id: model.id, + downloads: model.downloads, + likes: model.likes, + task: model.task, + }); + } + if (!cancelled) { + setState({ results, isLoading: false, error: null }); + } + } catch (err) { + if (!cancelled) { + setState({ + results: [], + isLoading: false, + error: err instanceof Error ? err.message : "Search failed", + }); + } + } + })(); + + return () => { + cancelled = true; + }; + }, [query, task, limit, accessToken]); + + return state; +} diff --git a/studio/frontend/src/types/training.ts b/studio/frontend/src/types/training.ts index dfee530f96..9688ec1ed4 100644 --- a/studio/frontend/src/types/training.ts +++ b/studio/frontend/src/types/training.ts @@ -128,6 +128,7 @@ export interface ModelOption { params: string; vram?: string; context?: string; + hfRepo?: string; recommended?: boolean; } From bbf8b4ffbee4c82392930e04680a974142b7f1d4 Mon Sep 17 00:00:00 2001 From: shine1i Date: Mon, 2 Feb 2026 12:14:18 +0100 Subject: [PATCH 02/12] feat: add Hugging Face search integration for datasets and models, extend infinite scroll support, and improve UI components with animations and tooltips --- .claude/settings.local.json | 15 -- studio/frontend/src/components/ui/spinner.tsx | 11 + studio/frontend/src/config/training.ts | 8 +- studio/frontend/src/features/export/anim.ts | 6 + .../export/components/export-dialog.tsx | 10 +- .../src/features/export/export-page.tsx | 53 ++-- .../components/steps/dataset-step.tsx | 146 +++++++---- .../components/steps/model-selection-step.tsx | 135 ++++++++--- .../components/steps/summary-step.tsx | 8 +- .../studio/sections/dataset-section.tsx | 140 +++++++++-- .../studio/sections/model-section.tsx | 228 ++++++++++++------ studio/frontend/src/hooks/index.ts | 1 + .../src/hooks/use-hf-dataset-search.ts | 127 +++++----- .../frontend/src/hooks/use-hf-model-search.ts | 89 +++---- .../src/hooks/use-hf-paginated-search.ts | 116 +++++++++ .../frontend/src/hooks/use-infinite-scroll.ts | 21 ++ studio/frontend/src/lib/utils.ts | 7 + studio/frontend/src/types/training.ts | 4 + 18 files changed, 789 insertions(+), 336 deletions(-) delete mode 100644 .claude/settings.local.json create mode 100644 studio/frontend/src/components/ui/spinner.tsx create mode 100644 studio/frontend/src/features/export/anim.ts create mode 100644 studio/frontend/src/hooks/use-hf-paginated-search.ts create mode 100644 studio/frontend/src/hooks/use-infinite-scroll.ts diff --git a/.claude/settings.local.json b/.claude/settings.local.json deleted file mode 100644 index 7802dc0461..0000000000 --- a/.claude/settings.local.json +++ /dev/null @@ -1,15 +0,0 @@ -{ - "permissions": { - "allow": [ - "Bash(tree:*)", - "Bash(findstr:*)", - "Bash(bun run typecheck:*)", - "mcp__plugin_serena_serena__list_dir", - "Bash(bun x tsc:*)", - "mcp__plugin_perplexity_perplexity__perplexity_ask", - "WebSearch", - "WebFetch(domain:www.npmjs.com)", - "WebFetch(domain:github.com)" - ] - } -} diff --git a/studio/frontend/src/components/ui/spinner.tsx b/studio/frontend/src/components/ui/spinner.tsx new file mode 100644 index 0000000000..3030726f0f --- /dev/null +++ b/studio/frontend/src/components/ui/spinner.tsx @@ -0,0 +1,11 @@ +import { cn } from "@/lib/utils" +import { HugeiconsIcon } from "@hugeicons/react" +import { Loading03Icon } from "@hugeicons/core-free-icons" + +function Spinner({ className }: { className?: string }) { + return ( + + ) +} + +export { Spinner } diff --git a/studio/frontend/src/config/training.ts b/studio/frontend/src/config/training.ts index 7b39b38e88..fd2b13e87c 100644 --- a/studio/frontend/src/config/training.ts +++ b/studio/frontend/src/config/training.ts @@ -4,6 +4,7 @@ import type { ModelType, StepConfig, } from "@/types/training"; +import type { PipelineType } from "@huggingface/hub"; export const STEPS: StepConfig[] = [ { @@ -257,7 +258,12 @@ export const DEFAULT_HYPERPARAMS = { targetModules: TARGET_MODULES, }; -export const MODEL_TYPE_TO_HF_TASK: Record = { +export function findModelById(id: string | null): ModelOption | undefined { + if (!id) return undefined; + return MODELS.find((m) => m.id === id || m.hfRepo === id); +} + +export const MODEL_TYPE_TO_HF_TASK: Record = { text: "text-generation", vision: "image-text-to-text", tts: "text-to-speech", diff --git a/studio/frontend/src/features/export/anim.ts b/studio/frontend/src/features/export/anim.ts new file mode 100644 index 0000000000..0bfc5aaa8a --- /dev/null +++ b/studio/frontend/src/features/export/anim.ts @@ -0,0 +1,6 @@ +export const collapseAnim = { + initial: { height: 0, opacity: 0 }, + animate: { height: "auto" as const, opacity: 1 }, + exit: { height: 0, opacity: 0 }, + transition: { duration: 0.3, ease: [0.25, 0.1, 0.25, 1] as const }, +}; diff --git a/studio/frontend/src/features/export/components/export-dialog.tsx b/studio/frontend/src/features/export/components/export-dialog.tsx index 1e6ff9962b..dcbd785fd0 100644 --- a/studio/frontend/src/features/export/components/export-dialog.tsx +++ b/studio/frontend/src/features/export/components/export-dialog.tsx @@ -17,17 +17,11 @@ import { Switch } from "@/components/ui/switch"; import { ArrowRight01Icon, Key01Icon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { AnimatePresence, motion } from "motion/react"; +import { collapseAnim } from "../anim"; import { EXPORT_METHODS, type ExportMethod } from "../constants"; type Destination = "local" | "hub"; -const anim = { - initial: { height: 0, opacity: 0 }, - animate: { height: "auto" as const, opacity: 1 }, - exit: { height: 0, opacity: 0 }, - transition: { duration: 0.3, ease: [0.25, 0.1, 0.25, 1] as const }, -}; - interface ExportDialogProps { open: boolean; onOpenChange: (open: boolean) => void; @@ -96,7 +90,7 @@ export function ExportDialog({ {destination === "hub" && ( - +
diff --git a/studio/frontend/src/features/export/export-page.tsx b/studio/frontend/src/features/export/export-page.tsx index 0ba0c817a8..09c6514b65 100644 --- a/studio/frontend/src/features/export/export-page.tsx +++ b/studio/frontend/src/features/export/export-page.tsx @@ -8,8 +8,9 @@ import { } from "@/components/ui/select"; import { Separator } from "@/components/ui/separator"; import { SectionCard } from "@/components/section-card"; -import { MODELS } from "@/config/training"; +import { findModelById } from "@/config/training"; import { useWizardStore } from "@/stores/training"; +import { isAdapterMethod } from "@/types/training"; import { Tooltip, TooltipContent, @@ -19,6 +20,8 @@ import { InformationCircleIcon, PackageIcon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { AnimatePresence, motion } from "motion/react"; import { useMemo, useState } from "react"; +import { useShallow } from "zustand/react/shallow"; +import { collapseAnim } from "./anim"; import { ExportDialog } from "./components/export-dialog"; import { MethodPicker } from "./components/method-picker"; import { QuantPicker } from "./components/quant-picker"; @@ -29,25 +32,27 @@ import { getEstimatedSize, } from "./constants"; -const anim = { - initial: { height: 0, opacity: 0 }, - animate: { height: "auto" as const, opacity: 1 }, - exit: { height: 0, opacity: 0 }, - transition: { duration: 0.3, ease: [0.25, 0.1, 0.25, 1] as const }, -}; - export function ExportPage() { - const store = useWizardStore(); - const isAdapter = store.trainingMethod === "lora" || store.trainingMethod === "qlora"; - const modelInfo = useMemo( - () => MODELS.find((m) => m.id === store.selectedModel), - [store.selectedModel], - ); + const { trainingMethod, selectedModel, saveSteps, trainingMetrics, epochs, loraRank, hfToken, setHfToken } = + useWizardStore( + useShallow((s) => ({ + trainingMethod: s.trainingMethod, + selectedModel: s.selectedModel, + saveSteps: s.saveSteps, + trainingMetrics: s.trainingMetrics, + epochs: s.epochs, + loraRank: s.loraRank, + hfToken: s.hfToken, + setHfToken: s.setHfToken, + })), + ); + const isAdapter = isAdapterMethod(trainingMethod); + const modelInfo = useMemo(() => findModelById(selectedModel), [selectedModel]); const checkpoints = useMemo(() => { if (isAdapter) { - const interval = store.saveSteps > 0 ? store.saveSteps : 100; - const total = store.trainingMetrics?.totalSteps ?? 500; + const interval = saveSteps > 0 ? saveSteps : 100; + const total = trainingMetrics?.totalSteps ?? 500; const entries: { value: string; label: string; detail: string }[] = []; for (let step = interval; step <= total; step += interval) { const loss = (1.5 - (step / total) * 0.7 + Math.random() * 0.05).toFixed(2); @@ -60,7 +65,7 @@ export function ExportPage() { return entries.reverse(); } return [{ value: "final-model", label: "Final Model", detail: "Full fine-tuned weights" }]; - }, [isAdapter, store.saveSteps, store.trainingMetrics?.totalSteps]); + }, [isAdapter, saveSteps, trainingMetrics?.totalSteps]); const [checkpoint, setCheckpoint] = useState(null); const [exportMethod, setExportMethod] = useState(null); @@ -79,7 +84,7 @@ export function ExportPage() { const estimatedSize = getEstimatedSize(exportMethod, quantLevels); const canExport = checkpoint && exportMethod && (exportMethod !== "gguf" || quantLevels.length > 0); - const baseModelName = modelInfo?.name ?? store.selectedModel ?? "—"; + const baseModelName = modelInfo?.name ?? selectedModel ?? "—"; return (
@@ -141,7 +146,7 @@ export function ExportPage() {
Method - {METHOD_LABELS[store.trainingMethod] ?? store.trainingMethod} + {METHOD_LABELS[trainingMethod] ?? trainingMethod}
Checkpoints @@ -149,12 +154,12 @@ export function ExportPage() {
Epochs - {store.epochs} + {epochs}
{isAdapter && (
LoRA Rank - {store.loraRank} + {loraRank}
)} {modelInfo?.params && ( @@ -186,7 +191,7 @@ export function ExportPage() { {exportMethod === "gguf" && ( - + )} @@ -220,8 +225,8 @@ export function ExportPage() { onHfUsernameChange={setHfUsername} modelName={modelName} onModelNameChange={setModelName} - hfToken={store.hfToken} - onHfTokenChange={store.setHfToken} + hfToken={hfToken} + onHfTokenChange={setHfToken} privateRepo={privateRepo} onPrivateRepoChange={setPrivateRepo} /> 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 a5d07c638e..291e065249 100644 --- a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx @@ -26,13 +26,15 @@ import { SelectTrigger, SelectValue, } from "@/components/ui/select"; +import { Spinner } from "@/components/ui/spinner"; import { Tooltip, TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; import { DATASETS } from "@/config/training"; -import { cn } from "@/lib/utils"; +import { useDebouncedValue, useHfDatasetSearch, useInfiniteScroll } from "@/hooks"; +import { cn, formatCompact } from "@/lib/utils"; import { useWizardStore } from "@/stores/training"; import type { DatasetFormat } from "@/types/training"; import { @@ -42,7 +44,7 @@ import { Upload04Icon, } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; -import { useMemo, useRef } from "react"; +import { useMemo, useRef, useState } from "react"; import { useShallow } from "zustand/react/shallow"; const FORMAT_OPTIONS: { value: DatasetFormat; label: string }[] = [ @@ -79,26 +81,56 @@ export function DatasetStep() { })), ); - const sortedDatasets = useMemo( - () => - // Sort recommended first - [...DATASETS].sort( - (a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0), - ), + const [inputValue, setInputValue] = useState(""); + const debouncedQuery = useDebouncedValue(inputValue); + const { results: hfResults, isLoading, isLoadingMore, hasMore, fetchMore } = useHfDatasetSearch(debouncedQuery, { + accessToken: hfToken || undefined, + }); + + const curatedDatasets = useMemo( + () => [...DATASETS].sort((a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0)), [], ); - const selectedDatasetData = DATASETS.find((d) => d.id === dataset); + const datasetMap = useMemo(() => { + const map = new Map(); + for (const d of curatedDatasets) { + map.set(d.id, { label: d.name, description: d.description, size: d.size, recommended: d.recommended }); + } + for (const r of hfResults) { + if (!map.has(r.id)) { + map.set(r.id, { label: r.id, downloads: r.downloads, totalExamples: r.totalExamples, sizeCategory: r.sizeCategory }); + } + } + return map; + }, [curatedDatasets, hfResults]); + + const displayIds = useMemo(() => { + if (!debouncedQuery.trim()) { + return curatedDatasets.map((d) => d.id); + } + const q = debouncedQuery.toLowerCase(); + const curatedIds = curatedDatasets + .filter((d) => d.name.toLowerCase().includes(q) || d.id.toLowerCase().includes(q)) + .map((d) => d.id); + const liveIds = hfResults.map((r) => r.id).filter((id) => !curatedIds.includes(id)); + return [...curatedIds, ...liveIds]; + }, [debouncedQuery, curatedDatasets, hfResults]); + + const allIds = useMemo( + () => [...new Set([...curatedDatasets.map((d) => d.id), ...hfResults.map((r) => r.id)])], + [curatedDatasets, hfResults], + ); + const comboboxAnchorRef = useRef(null); + const { scrollRef, sentinelRef } = useInfiniteScroll(fetchMore); const handleFileUpload = () => { - // Mock file upload setUploadedFile("my_dataset.jsonl"); }; return ( - {/* Source Toggle */} Source
@@ -163,48 +195,71 @@ export function DatasetStep() { Search datasets
d.name)} - value={selectedDatasetData?.name ?? null} - onValueChange={(name) => { - const ds = sortedDatasets.find((d) => d.name === name); - if (ds) { - setDataset(ds.id); - } - }} + items={allIds} + filteredItems={displayIds} + filter={null} + value={dataset} + onValueChange={(id) => setDataset(id)} + onInputValueChange={(val) => setInputValue(val)} + itemToStringValue={(id) => datasetMap.get(id)?.label ?? id} autoHighlight={true} > - + - No datasets found - - {(name: string) => { - const ds = sortedDatasets.find((d) => d.name === name); - return ( - -
- {name} - {ds && ( - - {ds.description} + {isLoading ? ( +
Searching…
+ ) : ( + No datasets found + )} +
+ + {(id: string) => { + const meta = datasetMap.get(id); + const label = meta?.label ?? id; + const rowLabel = meta?.size ?? (meta?.totalExamples ? `${formatCompact(meta.totalExamples)} rows` : null); + return ( + + + +
+ {label} + {meta?.description && ( + {meta.description} + )} +
+
+ + {label} + +
+ {rowLabel ? ( + + {rowLabel} + + ) : meta?.sizeCategory ? ( + + {meta.sizeCategory} - )} -
- {ds && ( - - {ds.size} - - )} - - ); - }} - + ) : meta?.downloads != null ? ( + + ↓{formatCompact(meta.downloads)} + + ) : null} + + ); + }} + + {hasMore &&
} + {isLoadingMore && ( +
+ +
+ )} +
@@ -250,7 +305,6 @@ export function DatasetStep() { )} - {/* Format Selection */}
diff --git a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx index 16a0c1e967..66217fb822 100644 --- a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx @@ -25,12 +25,15 @@ import { SelectTrigger, SelectValue, } from "@/components/ui/select"; +import { Spinner } from "@/components/ui/spinner"; import { Tooltip, TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { MODELS } from "@/config/training"; +import { MODEL_TYPE_TO_HF_TASK, MODELS } from "@/config/training"; +import { useDebouncedValue, useHfModelSearch, useInfiniteScroll } from "@/hooks"; +import { formatCompact } from "@/lib/utils"; import { useWizardStore } from "@/stores/training"; import type { TrainingMethod } from "@/types/training"; import { @@ -39,7 +42,7 @@ import { Search01Icon, } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; -import { useMemo, useRef } from "react"; +import { useMemo, useRef, useState } from "react"; import { useShallow } from "zustand/react/shallow"; export function ModelSelectionStep() { @@ -63,18 +66,51 @@ export function ModelSelectionStep() { })), ); - const filteredModels = useMemo(() => { - if (!modelType) { - return []; - } - // Sort recommended first + const [inputValue, setInputValue] = useState(""); + const debouncedQuery = useDebouncedValue(inputValue); + const task = modelType ? MODEL_TYPE_TO_HF_TASK[modelType] : undefined; + const { results: hfResults, isLoading, isLoadingMore, hasMore, fetchMore } = useHfModelSearch(debouncedQuery, { + task, + accessToken: hfToken || undefined, + }); + + const curatedModels = useMemo(() => { + if (!modelType) return []; return MODELS.filter((m) => m.type === modelType).sort( (a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0), ); }, [modelType]); - const selectedModelData = MODELS.find((m) => m.id === selectedModel); + const modelMap = useMemo(() => { + const map = new Map(); + for (const m of curatedModels) { + map.set(m.hfRepo ?? m.id, { label: m.name, params: m.params, recommended: m.recommended }); + } + for (const r of hfResults) { + if (!map.has(r.id)) { + map.set(r.id, { label: r.id, downloads: r.downloads, totalParams: r.totalParams }); + } + } + return map; + }, [curatedModels, hfResults]); + + const displayIds = useMemo(() => { + if (!debouncedQuery.trim()) { + return curatedModels.map((m) => m.hfRepo ?? m.id); + } + const q = debouncedQuery.toLowerCase(); + const curatedIds = curatedModels + .filter((m) => m.name.toLowerCase().includes(q) || m.id.toLowerCase().includes(q) || m.hfRepo?.toLowerCase().includes(q)) + .map((m) => m.hfRepo ?? m.id); + const liveIds = hfResults.map((r) => r.id).filter((id) => !curatedIds.includes(id)); + return [...curatedIds, ...liveIds]; + }, [debouncedQuery, curatedModels, hfResults]); + + const allIds = useMemo(() => [...new Set([...curatedModels.map((m) => m.hfRepo ?? m.id), ...hfResults.map((r) => r.id)])], [curatedModels, hfResults]); + + const selectedModelData = MODELS.find((m) => m.id === selectedModel || m.hfRepo === selectedModel); const comboboxAnchorRef = useRef(null); + const { scrollRef, sentinelRef } = useInfiniteScroll(fetchMore); return ( @@ -123,7 +159,7 @@ export function ModelSelectionStep() { - Search from our curated list of optimized models.{" "} + Search Hugging Face models or pick from our recommended list.{" "}
m.name)} - value={selectedModelData?.name ?? null} - onValueChange={(name) => { - const model = filteredModels.find((m) => m.name === name); - if (model) { - setSelectedModel(model.id); - } - }} + items={allIds} + filteredItems={displayIds} + filter={null} + value={selectedModel} + onValueChange={(id) => setSelectedModel(id)} + onInputValueChange={(val) => setInputValue(val)} + itemToStringValue={(id) => modelMap.get(id)?.label ?? id} autoHighlight={true} > - + - No models found - - {(name: string) => { - const model = filteredModels.find((m) => m.name === name); - return ( - - {name} - {model && ( - - {model.params} - - )} - - ); - }} - + {isLoading ? ( +
Searching…
+ ) : ( + No models found + )} +
+ + {(id: string) => { + const meta = modelMap.get(id); + const label = meta?.label ?? id; + const sizeLabel = meta?.params ?? (meta?.totalParams ? formatCompact(meta.totalParams) : null); + return ( + + + + {label} + + + {label} + + + + {meta?.recommended && ( + + Recommended + + )} + {sizeLabel ? ( + {sizeLabel} + ) : meta?.downloads != null ? ( + ↓{formatCompact(meta.downloads)} + ) : null} + + + ); + }} + + {hasMore &&
} + {isLoadingMore && ( +
+ +
+ )} +
- {selectedModelData && ( + {(selectedModelData || selectedModel) && (
@@ -208,7 +271,7 @@ export function ModelSelectionStep() { - Choose how to fine-tune {selectedModelData.name} + Choose how to fine-tune {selectedModelData?.name ?? selectedModel}
- - - - + setSelectedModel(id)} + onInputValueChange={(val) => setInputValue(val)} + itemToStringValue={(id) => modelMap.get(id)?.label ?? id} + autoHighlight={true} > - {filteredModels.map((m, i) => ( - - - - {m.name} - - {m.params} - - - - ))} - - -
- - {/* HF Repo */} -
- - Hugging Face Repo - - - - - - m.id === selectedModel)?.hfRepo ?? selectedModel) - : "" - } - onChange={(e) => setSelectedModel(e.target.value || null)} - /> - + + + + + + + {isLoading ? ( +
+ Searching… +
+ ) : ( + No models found + )} +
+ + {(id: string) => { + const meta = modelMap.get(id); + const label = meta?.label ?? id; + const sizeLabel = meta?.params ?? (meta?.totalParams ? formatCompact(meta.totalParams) : null); + return ( + + + + {label} + + + {label} + + + {sizeLabel ? ( + + {sizeLabel} + + ) : meta?.downloads != null ? ( + + ↓{formatCompact(meta.downloads)} + + ) : null} + + ); + }} + + {hasMore &&
} + {isLoadingMore && ( +
+ +
+ )} +
+ + +
{/* Training Method */} diff --git a/studio/frontend/src/hooks/index.ts b/studio/frontend/src/hooks/index.ts index 0e981fab4f..4e338169e4 100644 --- a/studio/frontend/src/hooks/index.ts +++ b/studio/frontend/src/hooks/index.ts @@ -1,3 +1,4 @@ export { useDebouncedValue } from "./use-debounced-value"; export { useHfModelSearch } from "./use-hf-model-search"; export { useHfDatasetSearch } from "./use-hf-dataset-search"; +export { useInfiniteScroll } from "./use-infinite-scroll"; diff --git a/studio/frontend/src/hooks/use-hf-dataset-search.ts b/studio/frontend/src/hooks/use-hf-dataset-search.ts index 727b6c1442..748dff6ba2 100644 --- a/studio/frontend/src/hooks/use-hf-dataset-search.ts +++ b/studio/frontend/src/hooks/use-hf-dataset-search.ts @@ -1,72 +1,85 @@ import { listDatasets } from "@huggingface/hub"; -import { useEffect, useState } from "react"; +import { useCallback } from "react"; +import { useHfPaginatedSearch } from "./use-hf-paginated-search"; + +interface DatasetInfoSplit { + name: string; + num_bytes: number; + num_examples: number; +} + +interface CardDataWithInfo { + size_categories?: string[]; + pretty_name?: string; + dataset_info?: + | { + splits?: DatasetInfoSplit[]; + download_size?: number; + dataset_size?: number; + } + | Array<{ splits?: DatasetInfoSplit[] }>; +} + +function extractTotalExamples( + cardData: CardDataWithInfo | undefined, +): number | undefined { + if (!cardData?.dataset_info) return undefined; + const infos = Array.isArray(cardData.dataset_info) + ? cardData.dataset_info + : [cardData.dataset_info]; + let total = 0; + let found = false; + for (const info of infos) { + for (const split of info.splits ?? []) { + if (typeof split.num_examples === "number") { + total += split.num_examples; + found = true; + } + } + } + return found ? total : undefined; +} export interface HfDatasetResult { id: string; downloads: number; likes: number; + totalExamples?: number; + sizeCategory?: string; } -interface HfSearchState { - results: HfDatasetResult[]; - isLoading: boolean; - error: string | null; +function mapDataset(raw: unknown): HfDatasetResult { + const ds = raw as { + name: string; + downloads: number; + likes: number; + cardData?: unknown; + }; + const card = ds.cardData as CardDataWithInfo | undefined; + return { + id: ds.name, + downloads: ds.downloads, + likes: ds.likes, + totalExamples: extractTotalExamples(card), + sizeCategory: card?.size_categories?.[0], + }; } export function useHfDatasetSearch( query: string, - options?: { limit?: number; accessToken?: string }, -): HfSearchState { - const { limit = 20, accessToken } = options ?? {}; - const [state, setState] = useState({ - results: [], - isLoading: false, - error: null, - }); + options?: { accessToken?: string }, +) { + const { accessToken } = options ?? {}; - useEffect(() => { - if (!query.trim()) { - setState({ results: [], isLoading: false, error: null }); - return; - } + const createIter = useCallback( + () => + listDatasets({ + search: { query }, + additionalFields: ["cardData"], + ...(accessToken ? { credentials: { accessToken } } : {}), + }) as AsyncGenerator, + [query, accessToken], + ); - let cancelled = false; - setState((prev) => ({ ...prev, isLoading: true, error: null })); - - (async () => { - try { - const results: HfDatasetResult[] = []; - const iter = listDatasets({ - search: { query }, - limit, - ...(accessToken ? { credentials: { accessToken } } : {}), - }); - for await (const ds of iter) { - if (cancelled) return; - results.push({ - id: ds.id, - downloads: ds.downloads, - likes: ds.likes, - }); - } - if (!cancelled) { - setState({ results, isLoading: false, error: null }); - } - } catch (err) { - if (!cancelled) { - setState({ - results: [], - isLoading: false, - error: err instanceof Error ? err.message : "Search failed", - }); - } - } - })(); - - return () => { - cancelled = true; - }; - }, [query, limit, accessToken]); - - return state; + return useHfPaginatedSearch(query, createIter, mapDataset); } diff --git a/studio/frontend/src/hooks/use-hf-model-search.ts b/studio/frontend/src/hooks/use-hf-model-search.ts index fe0cc358e7..7b1ef0ee14 100644 --- a/studio/frontend/src/hooks/use-hf-model-search.ts +++ b/studio/frontend/src/hooks/use-hf-model-search.ts @@ -1,74 +1,45 @@ +import type { PipelineType } from "@huggingface/hub"; import { listModels } from "@huggingface/hub"; -import { useEffect, useState } from "react"; +import { useCallback } from "react"; +import { useHfPaginatedSearch } from "./use-hf-paginated-search"; export interface HfModelResult { id: string; downloads: number; likes: number; - task?: string; + totalParams?: number; } -interface HfSearchState { - results: HfModelResult[]; - isLoading: boolean; - error: string | null; +function mapModel(raw: unknown): HfModelResult { + const m = raw as { + name: string; + downloads: number; + likes: number; + safetensors?: { total: number }; + }; + return { + id: m.name, + downloads: m.downloads, + likes: m.likes, + totalParams: m.safetensors?.total, + }; } export function useHfModelSearch( query: string, - options?: { task?: string; limit?: number; accessToken?: string }, -): HfSearchState { - const { task, limit = 20, accessToken } = options ?? {}; - const [state, setState] = useState({ - results: [], - isLoading: false, - error: null, - }); + options?: { task?: PipelineType; accessToken?: string }, +) { + const { task, accessToken } = options ?? {}; - useEffect(() => { - if (!query.trim()) { - setState({ results: [], isLoading: false, error: null }); - return; - } + const createIter = useCallback( + () => + listModels({ + search: { query, ...(task ? { task } : {}) }, + additionalFields: ["safetensors"], + ...(accessToken ? { credentials: { accessToken } } : {}), + }) as AsyncGenerator, + [query, task, accessToken], + ); - let cancelled = false; - setState((prev) => ({ ...prev, isLoading: true, error: null })); - - (async () => { - try { - const results: HfModelResult[] = []; - const iter = listModels({ - search: { query, ...(task ? { task } : {}) }, - limit, - ...(accessToken ? { credentials: { accessToken } } : {}), - }); - for await (const model of iter) { - if (cancelled) return; - results.push({ - id: model.id, - downloads: model.downloads, - likes: model.likes, - task: model.task, - }); - } - if (!cancelled) { - setState({ results, isLoading: false, error: null }); - } - } catch (err) { - if (!cancelled) { - setState({ - results: [], - isLoading: false, - error: err instanceof Error ? err.message : "Search failed", - }); - } - } - })(); - - return () => { - cancelled = true; - }; - }, [query, task, limit, accessToken]); - - return state; + return useHfPaginatedSearch(query, createIter, mapModel); } diff --git a/studio/frontend/src/hooks/use-hf-paginated-search.ts b/studio/frontend/src/hooks/use-hf-paginated-search.ts new file mode 100644 index 0000000000..7beeb1d1b5 --- /dev/null +++ b/studio/frontend/src/hooks/use-hf-paginated-search.ts @@ -0,0 +1,116 @@ +import { useCallback, useEffect, useRef, useState } from "react"; + +interface HfPaginatedState { + results: T[]; + isLoading: boolean; + isLoadingMore: boolean; + hasMore: boolean; + error: string | null; +} + +const INITIAL: HfPaginatedState = { + results: [], + isLoading: false, + isLoadingMore: false, + hasMore: false, + error: null, +}; +const BATCH = 20; + +async function pullBatch( + iter: AsyncGenerator, + mapItem: (raw: unknown) => T, + size: number, +) { + const items: T[] = []; + for (let i = 0; i < size; i++) { + const result = await iter.next(); + if (result.done) return { items, done: true }; + items.push(mapItem(result.value)); + } + return { items, done: false }; +} + +export function useHfPaginatedSearch( + query: string, + createIter: () => AsyncGenerator, + mapItem: (raw: unknown) => T, +): HfPaginatedState & { fetchMore: () => void } { + const [state, setState] = useState>( + INITIAL as HfPaginatedState, + ); + const stateRef = useRef(state); + stateRef.current = state; + + const iterRef = useRef | null>(null); + const versionRef = useRef(0); + + useEffect(() => { + const v = ++versionRef.current; + iterRef.current = null; + + if (!query.trim()) { + setState(INITIAL as HfPaginatedState); + return; + } + + setState((prev) => ({ + ...prev, + results: [], + isLoading: true, + error: null, + hasMore: false, + })); + + const iter = createIter(); + iterRef.current = iter; + + pullBatch(iter, mapItem, BATCH) + .then(({ items, done }) => { + if (versionRef.current !== v) return; + setState({ + results: items, + isLoading: false, + isLoadingMore: false, + hasMore: !done, + error: null, + }); + }) + .catch((err) => { + if (versionRef.current !== v) return; + setState({ + results: [], + isLoading: false, + isLoadingMore: false, + hasMore: false, + error: err instanceof Error ? err.message : "Search failed", + }); + }); + }, [query, createIter, mapItem]); + + const fetchMore = useCallback(() => { + const iter = iterRef.current; + const { isLoading, isLoadingMore, hasMore } = stateRef.current; + if (!iter || isLoading || isLoadingMore || !hasMore) return; + + const v = versionRef.current; + setState((prev) => ({ ...prev, isLoadingMore: true })); + + pullBatch(iter, mapItem, BATCH) + .then(({ items, done }) => { + if (versionRef.current !== v) return; + setState((prev) => ({ + ...prev, + results: [...prev.results, ...items], + isLoadingMore: false, + hasMore: !done, + })); + }) + .catch(() => { + if (versionRef.current !== v) return; + setState((prev) => ({ ...prev, isLoadingMore: false, hasMore: false })); + }); + }, [mapItem]); + + return { ...state, fetchMore }; +} diff --git a/studio/frontend/src/hooks/use-infinite-scroll.ts b/studio/frontend/src/hooks/use-infinite-scroll.ts new file mode 100644 index 0000000000..488739604f --- /dev/null +++ b/studio/frontend/src/hooks/use-infinite-scroll.ts @@ -0,0 +1,21 @@ +import { useEffect, useRef } from "react"; + +export function useInfiniteScroll(fetchMore: () => void) { + const scrollRef = useRef(null); + const sentinelRef = useRef(null); + + useEffect(() => { + const el = sentinelRef.current; + if (!el) return; + const obs = new IntersectionObserver( + ([e]) => { + if (e.isIntersecting) fetchMore(); + }, + { threshold: 0, root: scrollRef.current }, + ); + obs.observe(el); + return () => obs.disconnect(); + }, [fetchMore]); + + return { scrollRef, sentinelRef }; +} diff --git a/studio/frontend/src/lib/utils.ts b/studio/frontend/src/lib/utils.ts index a70ebb68c7..3f05e80d09 100644 --- a/studio/frontend/src/lib/utils.ts +++ b/studio/frontend/src/lib/utils.ts @@ -4,3 +4,10 @@ import { twMerge } from "tailwind-merge"; export function cn(...inputs: ClassValue[]): string { return twMerge(clsx(inputs)); } + +export function formatCompact(n: number): string { + if (n >= 1_000_000_000) return `${(n / 1_000_000_000).toFixed(1)}B`; + if (n >= 1_000_000) return `${(n / 1_000_000).toFixed(1)}M`; + if (n >= 1_000) return `${(n / 1_000).toFixed(1)}K`; + return String(n); +} diff --git a/studio/frontend/src/types/training.ts b/studio/frontend/src/types/training.ts index 9688ec1ed4..0a286571e1 100644 --- a/studio/frontend/src/types/training.ts +++ b/studio/frontend/src/types/training.ts @@ -1,5 +1,9 @@ export type ModelType = "vision" | "tts" | "embeddings" | "text"; export type TrainingMethod = "qlora" | "lora" | "full"; + +export function isAdapterMethod(method: TrainingMethod): boolean { + return method === "lora" || method === "qlora"; +} export type StepNumber = 1 | 2 | 3 | 4 | 5; export type DatasetSource = "huggingface" | "upload"; export type DatasetFormat = "auto" | "alpaca" | "chatml" | "sharegpt"; From e423174c0b57e066777675e513b90658b51f8c58 Mon Sep 17 00:00:00 2001 From: shine1i Date: Mon, 2 Feb 2026 12:51:04 +0100 Subject: [PATCH 03/12] refactor: format and clean up imports, hooks, and UI components for consistent structure and readability across models and datasets sections --- .gitignore | 1 + studio/frontend/src/config/training.ts | 57 +++----- .../export/components/export-dialog.tsx | 66 +++++++-- .../export/components/method-picker.tsx | 64 +++++++-- .../export/components/quant-picker.tsx | 45 ++++-- .../frontend/src/features/export/constants.ts | 44 ++++-- .../src/features/export/export-page.tsx | 136 +++++++++++++----- .../components/steps/dataset-step.tsx | 99 ++++++++++--- .../components/steps/model-selection-step.tsx | 112 ++++++++++++--- .../components/steps/summary-step.tsx | 2 +- .../studio/sections/dataset-section.tsx | 106 +++++++++++--- .../studio/sections/model-section.tsx | 94 +++++++++--- .../src/hooks/use-hf-dataset-search.ts | 4 +- .../src/hooks/use-hf-paginated-search.ts | 24 +++- .../frontend/src/hooks/use-infinite-scroll.ts | 8 +- studio/frontend/src/hooks/use-mobile.ts | 24 ++-- 16 files changed, 670 insertions(+), 216 deletions(-) diff --git a/.gitignore b/.gitignore index 1453f625c2..3bb5a62a82 100755 --- a/.gitignore +++ b/.gitignore @@ -25,6 +25,7 @@ models/ # IDE / Editors .vscode/ .idea/ +.claude/ *.swp *.swo diff --git a/studio/frontend/src/config/training.ts b/studio/frontend/src/config/training.ts index fd2b13e87c..362766e04b 100644 --- a/studio/frontend/src/config/training.ts +++ b/studio/frontend/src/config/training.ts @@ -98,8 +98,21 @@ export const MODELS: ModelOption[] = [ hfRepo: "Qwen/Qwen-VL-Chat", }, // TTS models - { id: "bark", name: "Bark", type: "tts", params: "1B", hfRepo: "suno/bark", recommended: true }, - { id: "xtts-v2", name: "XTTS v2", type: "tts", params: "500M", hfRepo: "coqui/XTTS-v2" }, + { + id: "bark", + name: "Bark", + type: "tts", + params: "1B", + hfRepo: "suno/bark", + recommended: true, + }, + { + id: "xtts-v2", + name: "XTTS v2", + type: "tts", + params: "500M", + hfRepo: "coqui/XTTS-v2", + }, // Embedding models { id: "bge-large", @@ -111,24 +124,6 @@ export const MODELS: ModelOption[] = [ { id: "e5-large", name: "E5 Large", type: "embeddings", params: "335M" }, { id: "gte-large", name: "GTE Large", type: "embeddings", params: "335M" }, // Text models - { - id: "llama-3.1-8b", - name: "Llama 3.1 8B", - type: "text", - params: "8B", - vram: "~6GB", - context: "128K", - hfRepo: "unsloth/Llama-3.1-8B", - recommended: true, - }, - { - id: "llama-3.1-70b", - name: "Llama 3.1 70B", - type: "text", - params: "70B", - vram: "~40GB", - context: "128K", - }, { id: "mistral-7b", name: "Mistral 7B", @@ -154,24 +149,6 @@ export const MODELS: ModelOption[] = [ vram: "~3GB", context: "128K", }, - { - id: "gemma-2-9b", - name: "Gemma 2 9B", - type: "text", - params: "9B", - vram: "~7GB", - context: "8K", - }, - { - id: "gemma-3-27b", - name: "Gemma 3 27B", - type: "text", - params: "27B", - vram: "~18GB", - context: "128K", - hfRepo: "unsloth/gemma-3-27b", - recommended: true, - }, ]; export const DATASETS: DatasetOption[] = [ @@ -259,7 +236,9 @@ export const DEFAULT_HYPERPARAMS = { }; export function findModelById(id: string | null): ModelOption | undefined { - if (!id) return undefined; + if (!id) { + return undefined; + } return MODELS.find((m) => m.id === id || m.hfRepo === id); } diff --git a/studio/frontend/src/features/export/components/export-dialog.tsx b/studio/frontend/src/features/export/components/export-dialog.tsx index dcbd785fd0..4f66048270 100644 --- a/studio/frontend/src/features/export/components/export-dialog.tsx +++ b/studio/frontend/src/features/export/components/export-dialog.tsx @@ -68,7 +68,9 @@ export function ExportDialog({ Export Model - Choose where to save your exported model. + + Choose where to save your exported model. +
@@ -94,18 +96,32 @@ export function ExportDialog({
- - onHfUsernameChange(e.target.value)} /> + + onHfUsernameChange(e.target.value)} + />
- - onModelNameChange(e.target.value)} /> + + onModelNameChange(e.target.value)} + />
- onHfTokenChange(e.target.value)} /> + onHfTokenChange(e.target.value)} + /> -

Leave empty if already logged in via CLI.

+

+ Leave empty if already logged in via CLI. +

- - + +
@@ -153,7 +189,9 @@ export function ExportDialog({ {exportMethod === "gguf" && quantLevels.length > 0 && (
Quantizations - {quantLevels.join(", ")} + + {quantLevels.join(", ")} +
)}
@@ -163,7 +201,9 @@ export function ExportDialog({
- + diff --git a/studio/frontend/src/features/export/components/method-picker.tsx b/studio/frontend/src/features/export/components/method-picker.tsx index b0052f3069..ba0ca50543 100644 --- a/studio/frontend/src/features/export/components/method-picker.tsx +++ b/studio/frontend/src/features/export/components/method-picker.tsx @@ -5,7 +5,10 @@ import { TooltipTrigger, } from "@/components/ui/tooltip"; import { cn } from "@/lib/utils"; -import { CheckmarkCircle01Icon, InformationCircleIcon } from "@hugeicons/core-free-icons"; +import { + CheckmarkCircle01Icon, + InformationCircleIcon, +} from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { EXPORT_METHODS, type ExportMethod } from "../constants"; @@ -20,14 +23,24 @@ export function MethodPicker({ value, onChange }: MethodPickerProps) { Export Method - - How your model is packaged for deployment.{" "} - Read more + + Read more + @@ -49,34 +62,61 @@ export function MethodPicker({ value, onChange }: MethodPickerProps) {
{selected && ( - + )}
{m.title} - - e.stopPropagation()}> - + + e.stopPropagation()} + > + {m.tooltip}{" "} - Read more + + Read more + {m.badge && ( - + {m.badge} )}
- {m.description} + + {m.description} +
); diff --git a/studio/frontend/src/features/export/components/quant-picker.tsx b/studio/frontend/src/features/export/components/quant-picker.tsx index e16dbd5494..9b6c070453 100644 --- a/studio/frontend/src/features/export/components/quant-picker.tsx +++ b/studio/frontend/src/features/export/components/quant-picker.tsx @@ -4,7 +4,11 @@ import { TooltipTrigger, } from "@/components/ui/tooltip"; import { cn } from "@/lib/utils"; -import { CheckmarkCircle01Icon, InformationCircleIcon, LayersIcon } from "@hugeicons/core-free-icons"; +import { + CheckmarkCircle01Icon, + InformationCircleIcon, + LayersIcon, +} from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { QUANT_OPTIONS } from "../constants"; @@ -23,20 +27,38 @@ export function QuantPicker({ value, onChange }: QuantPickerProps) { return (
- - Quantization Levels + + + Quantization Levels + - - - Lower quantization (Q2, Q3) = smaller files but reduced quality. Q4–Q5 is a good balance.{" "} - Read more + Lower quantization (Q2, Q3) = smaller files but reduced quality. + Q4–Q5 is a good balance.{" "} + + Read more + - — select one or more + + — select one or more +
{QUANT_OPTIONS.map((q) => { @@ -53,7 +75,12 @@ export function QuantPicker({ value, onChange }: QuantPickerProps) { : "ring-border text-muted-foreground hover:text-foreground hover:ring-foreground/20", )} > - {active && } + {active && ( + + )} {q.label} {q.size} {q.recommended && !active && ( diff --git a/studio/frontend/src/features/export/constants.ts b/studio/frontend/src/features/export/constants.ts index b289372f05..154058982d 100644 --- a/studio/frontend/src/features/export/constants.ts +++ b/studio/frontend/src/features/export/constants.ts @@ -9,9 +9,27 @@ export const EXPORT_METHODS: { tooltip: string; badge?: string; }[] = [ - { value: "merged", title: "Merged Model", description: "Full 16-bit model ready for inference.", tooltip: "Merges adapter weights into the base model. Best for direct deployment with vLLM or TGI." }, - { value: "lora", title: "LoRA Only", description: "Lightweight adapter files (~100 MB). Needs base model.", tooltip: "Exports only the trained adapter. Pair with the base model at inference time to save storage." }, - { value: "gguf", title: "GGUF / Llama.cpp", description: "Quantized formats for local AI runners.", tooltip: "Converts to GGUF for llama.cpp, Ollama, and other local runners. Pick a quantization level below." }, + { + value: "merged", + title: "Merged Model", + description: "Full 16-bit model ready for inference.", + tooltip: + "Merges adapter weights into the base model. Best for direct deployment with vLLM or TGI.", + }, + { + value: "lora", + title: "LoRA Only", + description: "Lightweight adapter files (~100 MB). Needs base model.", + tooltip: + "Exports only the trained adapter. Pair with the base model at inference time to save storage.", + }, + { + value: "gguf", + title: "GGUF / Llama.cpp", + description: "Quantized formats for local AI runners.", + tooltip: + "Converts to GGUF for llama.cpp, Ollama, and other local runners. Pick a quantization level below.", + }, ]; export const QUANT_OPTIONS = [ @@ -27,17 +45,27 @@ export const QUANT_OPTIONS = [ { value: "f16", label: "F16", size: "~14.2 GB" }, ]; -export function getEstimatedSize(method: ExportMethod | null, quantLevels: string[]) { - const sizeOf = (v: string) => QUANT_OPTIONS.find((q) => q.value === v)?.size ?? "—"; +export function getEstimatedSize( + method: ExportMethod | null, + quantLevels: string[], +) { + const sizeOf = (v: string) => + QUANT_OPTIONS.find((q) => q.value === v)?.size ?? "—"; if (method === "gguf" && quantLevels.length > 0) { - if (quantLevels.length === 1) return sizeOf(quantLevels[0]); + if (quantLevels.length === 1) { + return sizeOf(quantLevels[0]); + } const total = quantLevels .map((q) => Number.parseFloat(sizeOf(q).replace(/[^0-9.]/g, ""))) .reduce((a, b) => a + b, 0); return `~${total.toFixed(1)} GB (${quantLevels.length} files)`; } - if (method === "merged") return "~14.2 GB"; - if (method === "lora") return "~100 MB"; + if (method === "merged") { + return "~14.2 GB"; + } + if (method === "lora") { + return "~100 MB"; + } return "—"; } diff --git a/studio/frontend/src/features/export/export-page.tsx b/studio/frontend/src/features/export/export-page.tsx index 09c6514b65..8c242941cb 100644 --- a/studio/frontend/src/features/export/export-page.tsx +++ b/studio/frontend/src/features/export/export-page.tsx @@ -1,3 +1,4 @@ +import { SectionCard } from "@/components/section-card"; import { Button } from "@/components/ui/button"; import { Select, @@ -7,15 +8,14 @@ import { SelectValue, } from "@/components/ui/select"; import { Separator } from "@/components/ui/separator"; -import { SectionCard } from "@/components/section-card"; -import { findModelById } from "@/config/training"; -import { useWizardStore } from "@/stores/training"; -import { isAdapterMethod } from "@/types/training"; import { Tooltip, TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; +import { findModelById } from "@/config/training"; +import { useWizardStore } from "@/stores/training"; +import { isAdapterMethod } from "@/types/training"; import { InformationCircleIcon, PackageIcon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { AnimatePresence, motion } from "motion/react"; @@ -33,21 +33,32 @@ import { } from "./constants"; export function ExportPage() { - const { trainingMethod, selectedModel, saveSteps, trainingMetrics, epochs, loraRank, hfToken, setHfToken } = - useWizardStore( - useShallow((s) => ({ - trainingMethod: s.trainingMethod, - selectedModel: s.selectedModel, - saveSteps: s.saveSteps, - trainingMetrics: s.trainingMetrics, - epochs: s.epochs, - loraRank: s.loraRank, - hfToken: s.hfToken, - setHfToken: s.setHfToken, - })), - ); + const { + trainingMethod, + selectedModel, + saveSteps, + trainingMetrics, + epochs, + loraRank, + hfToken, + setHfToken, + } = useWizardStore( + useShallow((s) => ({ + trainingMethod: s.trainingMethod, + selectedModel: s.selectedModel, + saveSteps: s.saveSteps, + trainingMetrics: s.trainingMetrics, + epochs: s.epochs, + loraRank: s.loraRank, + hfToken: s.hfToken, + setHfToken: s.setHfToken, + })), + ); const isAdapter = isAdapterMethod(trainingMethod); - const modelInfo = useMemo(() => findModelById(selectedModel), [selectedModel]); + const modelInfo = useMemo( + () => findModelById(selectedModel), + [selectedModel], + ); const checkpoints = useMemo(() => { if (isAdapter) { @@ -55,7 +66,11 @@ export function ExportPage() { const total = trainingMetrics?.totalSteps ?? 500; const entries: { value: string; label: string; detail: string }[] = []; for (let step = interval; step <= total; step += interval) { - const loss = (1.5 - (step / total) * 0.7 + Math.random() * 0.05).toFixed(2); + const loss = ( + 1.5 - + (step / total) * 0.7 + + Math.random() * 0.05 + ).toFixed(2); entries.push({ value: `checkpoint-${step}`, label: `checkpoint-${step}`, @@ -64,7 +79,13 @@ export function ExportPage() { } return entries.reverse(); } - return [{ value: "final-model", label: "Final Model", detail: "Full fine-tuned weights" }]; + return [ + { + value: "final-model", + label: "Final Model", + detail: "Full fine-tuned weights", + }, + ]; }, [isAdapter, saveSteps, trainingMetrics?.totalSteps]); const [checkpoint, setCheckpoint] = useState(null); @@ -79,19 +100,28 @@ export function ExportPage() { const handleMethodChange = (method: ExportMethod) => { setExportMethod(method); - if (method !== "gguf") setQuantLevels([]); + if (method !== "gguf") { + setQuantLevels([]); + } }; const estimatedSize = getEstimatedSize(exportMethod, quantLevels); - const canExport = checkpoint && exportMethod && (exportMethod !== "gguf" || quantLevels.length > 0); + const canExport = + checkpoint && + exportMethod && + (exportMethod !== "gguf" || quantLevels.length > 0); const baseModelName = modelInfo?.name ?? selectedModel ?? "—"; return (
-

Export Model

-

Export your fine-tuned model for deployment

+

+ Export Model +

+

+ Export your fine-tuned model for deployment +

{/* Top row: Checkpoint + metadata | Guide */} @@ -109,27 +139,47 @@ export function ExportPage() { [...DATASETS].sort((a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0)), + () => + [...DATASETS].sort( + (a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0), + ), [], ); const datasetMap = useMemo(() => { - const map = new Map(); + const map = new Map< + string, + { + label: string; + description?: string; + size?: string; + totalExamples?: number; + sizeCategory?: string; + downloads?: number; + } + >(); for (const d of curatedDatasets) { - map.set(d.id, { label: d.name, description: d.description, size: d.size }); + map.set(d.id, { + label: d.name, + description: d.description, + size: d.size, + }); } for (const r of hfResults) { if (!map.has(r.id)) { - map.set(r.id, { label: r.id, downloads: r.downloads, totalExamples: r.totalExamples, sizeCategory: r.sizeCategory }); + map.set(r.id, { + label: r.id, + downloads: r.downloads, + totalExamples: r.totalExamples, + sizeCategory: r.sizeCategory, + }); } } return map; @@ -89,14 +119,24 @@ export function DatasetSection() { } const q = debouncedQuery.toLowerCase(); const curatedIds = curatedDatasets - .filter((d) => d.name.toLowerCase().includes(q) || d.id.toLowerCase().includes(q)) + .filter( + (d) => + d.name.toLowerCase().includes(q) || d.id.toLowerCase().includes(q), + ) .map((d) => d.id); - const liveIds = hfResults.map((r) => r.id).filter((id) => !curatedIds.includes(id)); + const liveIds = hfResults + .map((r) => r.id) + .filter((id) => !curatedIds.includes(id)); return [...curatedIds, ...liveIds]; }, [debouncedQuery, curatedDatasets, hfResults]); const allIds = useMemo( - () => [...new Set([...curatedDatasets.map((d) => d.id), ...hfResults.map((r) => r.id)])], + () => [ + ...new Set([ + ...curatedDatasets.map((d) => d.id), + ...hfResults.map((r) => r.id), + ]), + ], [curatedDatasets, hfResults], ); @@ -129,7 +169,8 @@ export function DatasetSection() { - Search Hugging Face datasets or enter a path like 'username/dataset-name'.{" "} + Search Hugging Face datasets or enter a path like + 'username/dataset-name'.{" "} datasetMap.get(id)?.label ?? id} autoHighlight={true} > - + {isLoading ? ( -
Searching…
+
+ Searching… +
) : ( No datasets found )} -
+
{(id: string) => { const meta = datasetMap.get(id); const label = meta?.label ?? id; - const rowLabel = meta?.size ?? (meta?.totalExamples ? `${formatCompact(meta.totalExamples)} rows` : null); + const rowLabel = + meta?.size ?? + (meta?.totalExamples + ? `${formatCompact(meta.totalExamples)} rows` + : null); return ( - + - - {label} + + + {label} + - + {label} diff --git a/studio/frontend/src/features/studio/sections/model-section.tsx b/studio/frontend/src/features/studio/sections/model-section.tsx index f8edb3e148..e6f8849ff0 100644 --- a/studio/frontend/src/features/studio/sections/model-section.tsx +++ b/studio/frontend/src/features/studio/sections/model-section.tsx @@ -25,8 +25,12 @@ import { TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { MODEL_TYPE_TO_HF_TASK, MODELS } from "@/config/training"; -import { useDebouncedValue, useHfModelSearch, useInfiniteScroll } from "@/hooks"; +import { MODELS, MODEL_TYPE_TO_HF_TASK } from "@/config/training"; +import { + useDebouncedValue, + useHfModelSearch, + useInfiniteScroll, +} from "@/hooks"; import { formatCompact } from "@/lib/utils"; import { useWizardStore } from "@/stores/training"; import type { TrainingMethod } from "@/types/training"; @@ -76,26 +80,51 @@ export function ModelSection() { const [inputValue, setInputValue] = useState(""); const debouncedQuery = useDebouncedValue(inputValue); const task = modelType ? MODEL_TYPE_TO_HF_TASK[modelType] : undefined; - const { results: hfResults, isLoading, isLoadingMore, hasMore, fetchMore } = useHfModelSearch(debouncedQuery, { + const { + results: hfResults, + isLoading, + isLoadingMore, + hasMore, + fetchMore, + } = useHfModelSearch(debouncedQuery, { task, accessToken: hfToken || undefined, }); const curatedModels = useMemo(() => { - if (!modelType) return MODELS; + if (!modelType) { + return MODELS; + } return MODELS.filter((m) => m.type === modelType).sort( (a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0), ); }, [modelType]); const modelMap = useMemo(() => { - const map = new Map(); + const map = new Map< + string, + { + label: string; + params?: string; + totalParams?: number; + downloads?: number; + recommended?: boolean; + } + >(); for (const m of curatedModels) { - map.set(m.hfRepo ?? m.id, { label: m.name, params: m.params, recommended: m.recommended }); + map.set(m.hfRepo ?? m.id, { + label: m.name, + params: m.params, + recommended: m.recommended, + }); } for (const r of hfResults) { if (!map.has(r.id)) { - map.set(r.id, { label: r.id, downloads: r.downloads, totalParams: r.totalParams }); + map.set(r.id, { + label: r.id, + downloads: r.downloads, + totalParams: r.totalParams, + }); } } return map; @@ -107,14 +136,26 @@ export function ModelSection() { } const q = debouncedQuery.toLowerCase(); const curatedIds = curatedModels - .filter((m) => m.name.toLowerCase().includes(q) || m.id.toLowerCase().includes(q) || m.hfRepo?.toLowerCase().includes(q)) + .filter( + (m) => + m.name.toLowerCase().includes(q) || + m.id.toLowerCase().includes(q) || + m.hfRepo?.toLowerCase().includes(q), + ) .map((m) => m.hfRepo ?? m.id); - const liveIds = hfResults.map((r) => r.id).filter((id) => !curatedIds.includes(id)); + const liveIds = hfResults + .map((r) => r.id) + .filter((id) => !curatedIds.includes(id)); return [...curatedIds, ...liveIds]; }, [debouncedQuery, curatedModels, hfResults]); const allIds = useMemo( - () => [...new Set([...curatedModels.map((m) => m.hfRepo ?? m.id), ...hfResults.map((r) => r.id)])], + () => [ + ...new Set([ + ...curatedModels.map((m) => m.hfRepo ?? m.id), + ...hfResults.map((r) => r.id), + ]), + ], [curatedModels, hfResults], ); @@ -161,7 +202,10 @@ export function ModelSection() { placeholder="./models/my-model" value={ selectedModel - ? (MODELS.find((m) => m.id === selectedModel || m.hfRepo === selectedModel)?.hfRepo ?? selectedModel) + ? (MODELS.find( + (m) => + m.id === selectedModel || m.hfRepo === selectedModel, + )?.hfRepo ?? selectedModel) : "" } onChange={(e) => setSelectedModel(e.target.value || null)} @@ -222,19 +266,35 @@ export function ModelSection() { ) : ( No models found )} -
+
{(id: string) => { const meta = modelMap.get(id); const label = meta?.label ?? id; - const sizeLabel = meta?.params ?? (meta?.totalParams ? formatCompact(meta.totalParams) : null); + const sizeLabel = + meta?.params ?? + (meta?.totalParams + ? formatCompact(meta.totalParams) + : null); return ( - + - - {label} + + + {label} + - + {label} diff --git a/studio/frontend/src/hooks/use-hf-dataset-search.ts b/studio/frontend/src/hooks/use-hf-dataset-search.ts index 748dff6ba2..d10b7d11e9 100644 --- a/studio/frontend/src/hooks/use-hf-dataset-search.ts +++ b/studio/frontend/src/hooks/use-hf-dataset-search.ts @@ -23,7 +23,9 @@ interface CardDataWithInfo { function extractTotalExamples( cardData: CardDataWithInfo | undefined, ): number | undefined { - if (!cardData?.dataset_info) return undefined; + if (!cardData?.dataset_info) { + return undefined; + } const infos = Array.isArray(cardData.dataset_info) ? cardData.dataset_info : [cardData.dataset_info]; diff --git a/studio/frontend/src/hooks/use-hf-paginated-search.ts b/studio/frontend/src/hooks/use-hf-paginated-search.ts index 7beeb1d1b5..bc76b96613 100644 --- a/studio/frontend/src/hooks/use-hf-paginated-search.ts +++ b/studio/frontend/src/hooks/use-hf-paginated-search.ts @@ -25,7 +25,9 @@ async function pullBatch( const items: T[] = []; for (let i = 0; i < size; i++) { const result = await iter.next(); - if (result.done) return { items, done: true }; + if (result.done) { + return { items, done: true }; + } items.push(mapItem(result.value)); } return { items, done: false }; @@ -67,7 +69,9 @@ export function useHfPaginatedSearch( pullBatch(iter, mapItem, BATCH) .then(({ items, done }) => { - if (versionRef.current !== v) return; + if (versionRef.current !== v) { + return; + } setState({ results: items, isLoading: false, @@ -77,7 +81,9 @@ export function useHfPaginatedSearch( }); }) .catch((err) => { - if (versionRef.current !== v) return; + if (versionRef.current !== v) { + return; + } setState({ results: [], isLoading: false, @@ -91,14 +97,18 @@ export function useHfPaginatedSearch( const fetchMore = useCallback(() => { const iter = iterRef.current; const { isLoading, isLoadingMore, hasMore } = stateRef.current; - if (!iter || isLoading || isLoadingMore || !hasMore) return; + if (!iter || isLoading || isLoadingMore || !hasMore) { + return; + } const v = versionRef.current; setState((prev) => ({ ...prev, isLoadingMore: true })); pullBatch(iter, mapItem, BATCH) .then(({ items, done }) => { - if (versionRef.current !== v) return; + if (versionRef.current !== v) { + return; + } setState((prev) => ({ ...prev, results: [...prev.results, ...items], @@ -107,7 +117,9 @@ export function useHfPaginatedSearch( })); }) .catch(() => { - if (versionRef.current !== v) return; + if (versionRef.current !== v) { + return; + } setState((prev) => ({ ...prev, isLoadingMore: false, hasMore: false })); }); }, [mapItem]); diff --git a/studio/frontend/src/hooks/use-infinite-scroll.ts b/studio/frontend/src/hooks/use-infinite-scroll.ts index 488739604f..14c80f4810 100644 --- a/studio/frontend/src/hooks/use-infinite-scroll.ts +++ b/studio/frontend/src/hooks/use-infinite-scroll.ts @@ -6,10 +6,14 @@ export function useInfiniteScroll(fetchMore: () => void) { useEffect(() => { const el = sentinelRef.current; - if (!el) return; + if (!el) { + return; + } const obs = new IntersectionObserver( ([e]) => { - if (e.isIntersecting) fetchMore(); + if (e.isIntersecting) { + fetchMore(); + } }, { threshold: 0, root: scrollRef.current }, ); diff --git a/studio/frontend/src/hooks/use-mobile.ts b/studio/frontend/src/hooks/use-mobile.ts index 2b0fe1dfef..a93d583938 100644 --- a/studio/frontend/src/hooks/use-mobile.ts +++ b/studio/frontend/src/hooks/use-mobile.ts @@ -1,19 +1,21 @@ -import * as React from "react" +import * as React from "react"; -const MOBILE_BREAKPOINT = 768 +const MOBILE_BREAKPOINT = 768; export function useIsMobile() { - const [isMobile, setIsMobile] = React.useState(undefined) + const [isMobile, setIsMobile] = React.useState( + undefined, + ); React.useEffect(() => { - const mql = window.matchMedia(`(max-width: ${MOBILE_BREAKPOINT - 1}px)`) + const mql = window.matchMedia(`(max-width: ${MOBILE_BREAKPOINT - 1}px)`); const onChange = () => { - setIsMobile(window.innerWidth < MOBILE_BREAKPOINT) - } - mql.addEventListener("change", onChange) - setIsMobile(window.innerWidth < MOBILE_BREAKPOINT) - return () => mql.removeEventListener("change", onChange) - }, []) + setIsMobile(window.innerWidth < MOBILE_BREAKPOINT); + }; + mql.addEventListener("change", onChange); + setIsMobile(window.innerWidth < MOBILE_BREAKPOINT); + return () => mql.removeEventListener("change", onChange); + }, []); - return !!isMobile + return !!isMobile; } From 45df407b7879494ac856b7c77b8f752eba116266 Mon Sep 17 00:00:00 2001 From: shine1i Date: Mon, 2 Feb 2026 13:16:08 +0100 Subject: [PATCH 04/12] refactor: simplify model and dataset combobox logic, remove curated items, and streamline search handling across components --- studio/frontend/src/components/ui/tooltip.tsx | 2 +- .../components/steps/dataset-step.tsx | 129 +++----------- .../components/steps/model-selection-step.tsx | 139 ++++----------- .../studio/sections/dataset-section.tsx | 158 +++--------------- .../studio/sections/model-section.tsx | 117 +++---------- .../src/hooks/use-hf-dataset-search.ts | 4 +- .../frontend/src/hooks/use-hf-model-search.ts | 29 +++- .../src/hooks/use-hf-paginated-search.ts | 19 +-- .../frontend/src/hooks/use-infinite-scroll.ts | 4 +- 9 files changed, 139 insertions(+), 462 deletions(-) diff --git a/studio/frontend/src/components/ui/tooltip.tsx b/studio/frontend/src/components/ui/tooltip.tsx index 516d2b8d5f..6da22c3b65 100644 --- a/studio/frontend/src/components/ui/tooltip.tsx +++ b/studio/frontend/src/components/ui/tooltip.tsx @@ -4,7 +4,7 @@ import type * as React from "react"; import { cn } from "@/lib/utils"; function TooltipProvider({ - delayDuration = 0, + delayDuration = 400, ...props }: React.ComponentProps) { return ( 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 24f6564e66..39664f290f 100644 --- a/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/dataset-step.tsx @@ -32,7 +32,6 @@ import { TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { DATASETS } from "@/config/training"; import { useDebouncedValue, useHfDatasetSearch, @@ -86,88 +85,24 @@ export function DatasetStep() { ); const [inputValue, setInputValue] = useState(""); + const selectingRef = useRef(false); const debouncedQuery = useDebouncedValue(inputValue); const { results: hfResults, isLoading, isLoadingMore, - hasMore, fetchMore, } = useHfDatasetSearch(debouncedQuery, { accessToken: hfToken || undefined, }); - const curatedDatasets = useMemo( - () => - [...DATASETS].sort( - (a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0), - ), - [], - ); - - const datasetMap = useMemo(() => { - const map = new Map< - string, - { - label: string; - description?: string; - size?: string; - totalExamples?: number; - sizeCategory?: string; - downloads?: number; - recommended?: boolean; - } - >(); - for (const d of curatedDatasets) { - map.set(d.id, { - label: d.name, - description: d.description, - size: d.size, - recommended: d.recommended, - }); - } - for (const r of hfResults) { - if (!map.has(r.id)) { - map.set(r.id, { - label: r.id, - downloads: r.downloads, - totalExamples: r.totalExamples, - sizeCategory: r.sizeCategory, - }); - } - } - return map; - }, [curatedDatasets, hfResults]); - - const displayIds = useMemo(() => { - if (!debouncedQuery.trim()) { - return curatedDatasets.map((d) => d.id); - } - const q = debouncedQuery.toLowerCase(); - const curatedIds = curatedDatasets - .filter( - (d) => - d.name.toLowerCase().includes(q) || d.id.toLowerCase().includes(q), - ) - .map((d) => d.id); - const liveIds = hfResults - .map((r) => r.id) - .filter((id) => !curatedIds.includes(id)); - return [...curatedIds, ...liveIds]; - }, [debouncedQuery, curatedDatasets, hfResults]); - - const allIds = useMemo( - () => [ - ...new Set([ - ...curatedDatasets.map((d) => d.id), - ...hfResults.map((r) => r.id), - ]), - ], - [curatedDatasets, hfResults], - ); + const resultIds = useMemo(() => hfResults.map((r) => r.id), [hfResults]); const comboboxAnchorRef = useRef(null); - const { scrollRef, sentinelRef } = useInfiniteScroll(fetchMore); + const { scrollRef, sentinelRef } = useInfiniteScroll( + fetchMore, + hfResults.length, + ); const handleFileUpload = () => { setUploadedFile("my_dataset.jsonl"); @@ -239,13 +174,13 @@ export function DatasetStep() { Search datasets
setDataset(id)} - onInputValueChange={(val) => setInputValue(val)} - itemToStringValue={(id) => datasetMap.get(id)?.label ?? id} + onValueChange={(id) => { selectingRef.current = true; setDataset(id); }} + onInputValueChange={(val) => { if (selectingRef.current) { selectingRef.current = false; return; } setInputValue(val); }} + itemToStringValue={(id) => id} autoHighlight={true} > {isLoading ? (
- Searching… + Searching...
) : ( No datasets found @@ -270,13 +205,10 @@ export function DatasetStep() { > {(id: string) => { - const meta = datasetMap.get(id); - const label = meta?.label ?? id; - const rowLabel = - meta?.size ?? - (meta?.totalExamples - ? `${formatCompact(meta.totalExamples)} rows` - : null); + const r = hfResults.find((r) => r.id === id); + const detail = r?.totalExamples + ? `${formatCompact(r.totalExamples)} rows` + : (r?.sizeCategory ?? null); return ( - -
- {label} - {meta?.description && ( - - {meta.description} - - )} -
+ + + {id} + - {label} + {id}
- {rowLabel ? ( - - {rowLabel} - - ) : meta?.sizeCategory ? ( + {detail ? ( - {meta.sizeCategory} + {detail} - ) : meta?.downloads != null ? ( + ) : r?.downloads != null ? ( - ↓{formatCompact(meta.downloads)} + ↓{formatCompact(r.downloads)} ) : null}
); }}
- {hasMore &&
} +
{isLoadingMore && (
diff --git a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx index dd71ebc219..7f8346f2ff 100644 --- a/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx +++ b/studio/frontend/src/features/onboarding/components/steps/model-selection-step.tsx @@ -1,4 +1,3 @@ -import { Badge } from "@/components/ui/badge"; import { Combobox, ComboboxContent, @@ -31,7 +30,7 @@ import { TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { MODELS, MODEL_TYPE_TO_HF_TASK } from "@/config/training"; +import { MODEL_TYPE_TO_HF_TASK } from "@/config/training"; import { useDebouncedValue, useHfModelSearch, @@ -71,92 +70,26 @@ export function ModelSelectionStep() { ); const [inputValue, setInputValue] = useState(""); + const selectingRef = useRef(false); const debouncedQuery = useDebouncedValue(inputValue); const task = modelType ? MODEL_TYPE_TO_HF_TASK[modelType] : undefined; const { results: hfResults, isLoading, isLoadingMore, - hasMore, fetchMore, } = useHfModelSearch(debouncedQuery, { task, accessToken: hfToken || undefined, }); - const curatedModels = useMemo(() => { - if (!modelType) { - return []; - } - return MODELS.filter((m) => m.type === modelType).sort( - (a, b) => (b.recommended ? 1 : 0) - (a.recommended ? 1 : 0), - ); - }, [modelType]); + const resultIds = useMemo(() => hfResults.map((r) => r.id), [hfResults]); - const modelMap = useMemo(() => { - const map = new Map< - string, - { - label: string; - params?: string; - totalParams?: number; - downloads?: number; - recommended?: boolean; - } - >(); - for (const m of curatedModels) { - map.set(m.hfRepo ?? m.id, { - label: m.name, - params: m.params, - recommended: m.recommended, - }); - } - for (const r of hfResults) { - if (!map.has(r.id)) { - map.set(r.id, { - label: r.id, - downloads: r.downloads, - totalParams: r.totalParams, - }); - } - } - return map; - }, [curatedModels, hfResults]); - - const displayIds = useMemo(() => { - if (!debouncedQuery.trim()) { - return curatedModels.map((m) => m.hfRepo ?? m.id); - } - const q = debouncedQuery.toLowerCase(); - const curatedIds = curatedModels - .filter( - (m) => - m.name.toLowerCase().includes(q) || - m.id.toLowerCase().includes(q) || - m.hfRepo?.toLowerCase().includes(q), - ) - .map((m) => m.hfRepo ?? m.id); - const liveIds = hfResults - .map((r) => r.id) - .filter((id) => !curatedIds.includes(id)); - return [...curatedIds, ...liveIds]; - }, [debouncedQuery, curatedModels, hfResults]); - - const allIds = useMemo( - () => [ - ...new Set([ - ...curatedModels.map((m) => m.hfRepo ?? m.id), - ...hfResults.map((r) => r.id), - ]), - ], - [curatedModels, hfResults], - ); - - const selectedModelData = MODELS.find( - (m) => m.id === selectedModel || m.hfRepo === selectedModel, - ); const comboboxAnchorRef = useRef(null); - const { scrollRef, sentinelRef } = useInfiniteScroll(fetchMore); + const { scrollRef, sentinelRef } = useInfiniteScroll( + fetchMore, + hfResults.length, + ); return ( @@ -219,13 +152,13 @@ export function ModelSelectionStep() {
setSelectedModel(id)} - onInputValueChange={(val) => setInputValue(val)} - itemToStringValue={(id) => modelMap.get(id)?.label ?? id} + onValueChange={(id) => { selectingRef.current = true; setSelectedModel(id); }} + onInputValueChange={(val) => { if (selectingRef.current) { selectingRef.current = false; return; } setInputValue(val); }} + itemToStringValue={(id) => id} autoHighlight={true} > @@ -247,13 +180,10 @@ export function ModelSelectionStep() { > {(id: string) => { - const meta = modelMap.get(id); - const label = meta?.label ?? id; - const sizeLabel = - meta?.params ?? - (meta?.totalParams - ? formatCompact(meta.totalParams) - : null); + const r = hfResults.find((r) => r.id === id); + const sizeLabel = r?.totalParams + ? formatCompact(r.totalParams) + : null; return ( - {label} + {id} - {label} + {id} - - {meta?.recommended && ( - - Recommended - - )} - {sizeLabel ? ( - {sizeLabel} - ) : meta?.downloads != null ? ( - - ↓{formatCompact(meta.downloads)} - - ) : null} - + {sizeLabel ? ( + + {sizeLabel} + + ) : r?.downloads != null ? ( + + ↓{formatCompact(r.downloads)} + + ) : null} ); }} - {hasMore &&
} +
{isLoadingMore && (
@@ -306,7 +228,7 @@ export function ModelSelectionStep() {
- {(selectedModelData || selectedModel) && ( + {selectedModel && (
@@ -340,8 +262,7 @@ export function ModelSelectionStep() { - Choose how to fine-tune{" "} - {selectedModelData?.name ?? selectedModel} + Choose how to fine-tune {selectedModel}
- {/* Active dataset display */} {dataset ? (
@@ -269,7 +281,6 @@ export function DatasetSection() {
)} - {/* Action buttons */}
diff --git a/studio/frontend/src/types/training.ts b/studio/frontend/src/types/training.ts index 0a286571e1..9d67439258 100644 --- a/studio/frontend/src/types/training.ts +++ b/studio/frontend/src/types/training.ts @@ -123,23 +123,4 @@ export interface StepConfig { title: string; subtitle: string; description: string; -} - -export interface ModelOption { - id: string; - name: string; - type: ModelType; - params: string; - vram?: string; - context?: string; - hfRepo?: string; - recommended?: boolean; -} - -export interface DatasetOption { - id: string; - name: string; - description: string; - size: string; - recommended?: boolean; -} +} \ No newline at end of file From 9d926395acca4a7a50cc0767a04cf429fba009c7 Mon Sep 17 00:00:00 2001 From: shine1i Date: Mon, 2 Feb 2026 14:52:58 +0100 Subject: [PATCH 07/12] refactor: remove unused components, mock data, and redundant logic across chat features; streamline settings and runtime handling for better maintainability --- .../frontend/src/features/chat/chat-page.tsx | 36 ++--- .../src/features/chat/chat-settings-sheet.tsx | 37 ++--- .../src/features/chat/chat-top-bar.tsx | 151 ------------------ .../src/features/chat/runtime-provider.tsx | 46 +----- .../src/features/chat/thread-sidebar.tsx | 1 - studio/frontend/src/features/chat/types.ts | 4 - 6 files changed, 30 insertions(+), 245 deletions(-) delete mode 100644 studio/frontend/src/features/chat/chat-top-bar.tsx diff --git a/studio/frontend/src/features/chat/chat-page.tsx b/studio/frontend/src/features/chat/chat-page.tsx index cc1fdfbe6a..6de42ed671 100644 --- a/studio/frontend/src/features/chat/chat-page.tsx +++ b/studio/frontend/src/features/chat/chat-page.tsx @@ -18,7 +18,6 @@ import { memo, useCallback, useEffect, - useMemo, useRef, useState, } from "react"; @@ -38,7 +37,6 @@ import { import { ThreadSidebar } from "./thread-sidebar"; import type { ChatView } from "./types"; -// TODO: fetch from API at runtime const LORA_MODELS: ModelOption[] = [ { id: "outputs/llama-3.1-8b-instruct-lora", @@ -75,17 +73,9 @@ const GGUF_MODELS: ModelOption[] = [ }, ]; -type SingleContentProps = { - threadId?: string; -}; - -type CompareContentProps = { - pairId: string; -}; - const SingleContent = memo(function SingleContent({ threadId, -}: SingleContentProps): ReactElement { +}: { threadId?: string }): ReactElement { return (
@@ -97,7 +87,7 @@ const SingleContent = memo(function SingleContent({ const CompareContent = memo(function CompareContent({ pairId, -}: CompareContentProps): ReactElement { +}: { pairId: string }): ReactElement { const handlesRef = useRef>({}); const [baseThreadId, setBaseThreadId] = useState(); const [loraThreadId, setLoraThreadId] = useState(); @@ -123,9 +113,9 @@ const CompareContent = memo(function CompareContent({ return (
-
+
-
+
Base Model @@ -142,8 +132,8 @@ const CompareContent = memo(function CompareContent({
-
- +
+ Fine-tuned (LoRA)
@@ -215,13 +205,10 @@ export function ChatPage(): ReactElement { [], ); - const models = useMemo( - () => - inferenceParams.inferenceEngine === "llama-cpp" - ? GGUF_MODELS - : LORA_MODELS, - [inferenceParams.inferenceEngine], - ); + const models = + inferenceParams.inferenceEngine === "llama-cpp" + ? GGUF_MODELS + : LORA_MODELS; return ( - {/* main chat area */}
- {/* top bar */}
@@ -275,7 +260,6 @@ export function ChatPage(): ReactElement { )}
- {/* inline settings panel on right */} (BUILTIN_PRESETS); const [activePreset, setActivePreset] = useState("Default"); - const set = (key: keyof InferenceParams) => (v: number | string) => - onParamsChange({ ...params, [key]: v }); + function set(key: K) { + return (v: InferenceParams[K]) => onParamsChange({ ...params, [key]: v }); + } function applyPreset(name: string) { const p = presets.find((pr) => pr.name === name); @@ -218,10 +211,9 @@ export function ChatSettingsPanel({ return (