diff --git a/studio/frontend/bun.lock b/studio/frontend/bun.lock index f438f4e3ea..27bc43b946 100644 --- a/studio/frontend/bun.lock +++ b/studio/frontend/bun.lock @@ -9,6 +9,8 @@ "@assistant-ui/react-markdown": "^0.12.1", "@assistant-ui/react-streamdown": "^0.1.0", "@base-ui/react": "^1.1.0", + "@dagrejs/dagre": "^2.0.4", + "@dagrejs/graphlib": "^3.0.4", "@fontsource-variable/figtree": "^5.2.10", "@fontsource-variable/inter": "^5.2.8", "@fontsource-variable/space-grotesk": "^5.2.10", @@ -189,6 +191,10 @@ "@chevrotain/utils": ["@chevrotain/utils@11.0.3", "", {}, "sha512-YslZMgtJUyuMbZ+aKvfF3x1f5liK4mWNxghFRv7jqRR9C3R3fAOGTTKvxXDa2Y1s9zSbcpuO0cAxDYsc9SrXoQ=="], + "@dagrejs/dagre": ["@dagrejs/dagre@2.0.4", "", { "dependencies": { "@dagrejs/graphlib": "3.0.4" } }, "sha512-J6vCWTNpicHF4zFlZG1cS5DkGzMr9941gddYkakjrg3ZNev4bbqEgLHFTWiFrcJm7UCRu7olO3K6IRDd9gSGhA=="], + + "@dagrejs/graphlib": ["@dagrejs/graphlib@3.0.4", "", {}, "sha512-HxZ7fCvAwTLCWCO0WjDkzAFQze8LdC6iOpKbetDKHIuDfIgMlIzYzqZ4nxwLlclQX+3ZVeZ1K2OuaOE2WWcyOg=="], + "@date-fns/tz": ["@date-fns/tz@1.4.1", "", {}, "sha512-P5LUNhtbj6YfI3iJjw5EL9eUAG6OitD0W3fWQcpQjDRc/QIsL0tRNuO1PcDvPccWL1fSTXXdE1ds+l95DV/OFA=="], "@dotenvx/dotenvx": ["@dotenvx/dotenvx@1.52.0", "", { "dependencies": { "commander": "^11.1.0", "dotenv": "^17.2.1", "eciesjs": "^0.4.10", "execa": "^5.1.1", "fdir": "^6.2.0", "ignore": "^5.3.0", "object-treeify": "1.1.33", "picomatch": "^4.0.2", "which": "^4.0.0" }, "bin": { "dotenvx": "src/cli/dotenvx.js" } }, "sha512-CaQcc8JvtzQhUSm9877b6V4Tb7HCotkcyud9X2YwdqtQKwgljkMRwU96fVYKnzN3V0Hj74oP7Es+vZ0mS+Aa1w=="], diff --git a/studio/frontend/package.json b/studio/frontend/package.json index ac61c22969..0d1352efb3 100644 --- a/studio/frontend/package.json +++ b/studio/frontend/package.json @@ -72,12 +72,12 @@ "zustand": "^5.0.10" }, "devDependencies": { + "@biomejs/biome": "^1.9.4", "@eslint/js": "^9.39.1", "@types/node": "^24.10.1", "@types/react": "^19.2.5", "@types/react-dom": "^19.2.3", "@vitejs/plugin-react": "^5.1.1", - "@biomejs/biome": "^1.9.4", "eslint": "^9.39.1", "eslint-plugin-react-hooks": "^7.0.1", "eslint-plugin-react-refresh": "^0.4.24", diff --git a/studio/frontend/src/features/studio/sections/model-section.tsx b/studio/frontend/src/features/studio/sections/model-section.tsx index fcaa319e1d..3fe304a01f 100644 --- a/studio/frontend/src/features/studio/sections/model-section.tsx +++ b/studio/frontend/src/features/studio/sections/model-section.tsx @@ -28,10 +28,17 @@ import { import { MODEL_TYPE_TO_HF_TASK } from "@/config/training"; import { useDebouncedValue, + useGpuInfo, useHfModelSearch, useInfiniteScroll, } from "@/hooks"; import { formatCompact } from "@/lib/utils"; +import { + type TrainingMethod as VramTrainingMethod, + type VramFitStatus, + checkVramFit, + estimateLoadingVram, +} from "@/lib/vram"; import { useTrainingConfigStore } from "@/features/training"; import type { TrainingMethod } from "@/types/training"; import { @@ -57,6 +64,8 @@ const DARK_CONTENT = "bg-foreground text-background shadow-xl border-background/10 [--accent:rgba(255,255,255,0.1)] [--accent-foreground:white] [&_[data-slot=select-item]]:text-white/70 [&_[data-slot=select-scroll-up-button]]:bg-foreground [&_[data-slot=select-scroll-down-button]]:bg-foreground"; export function ModelSection() { + const gpu = useGpuInfo(); + const { modelType, selectedModel, @@ -122,6 +131,37 @@ export function ModelSection() { return ids; }, [hfResults, selectedModel]); + // Pre-compute VRAM fit status for every model in the current result set. + // Keyed by model id so the render callback is a simple O(1) lookup. + // + // Pre-compute VRAM fit status for every model in the current result set. + // Keyed by model id so the render callback is a simple O(1) lookup. + // Re-computes when the training method changes (QLoRA=4-bit vs LoRA/Full=fp16). + const vramMap = useMemo(() => { + const method = trainingMethod as VramTrainingMethod; + const map = new Map< + string, + { est: number; status: VramFitStatus | null; detail: string | null } + >(); + for (const r of hfResults) { + const detail = r.totalParams + ? formatCompact(r.totalParams) + : r.downloads != null + ? `\u2193${formatCompact(r.downloads)}` + : null; + if (r.totalParams) { + const est = estimateLoadingVram(r.totalParams, method); + const status = gpu.available + ? checkVramFit(est, gpu.memoryTotalGb) + : null; + map.set(r.id, { est, status, detail }); + } else { + map.set(r.id, { est: 0, status: null, detail }); + } + } + return map; + }, [hfResults, gpu, trainingMethod]); + const comboboxAnchorRef = useRef(null); const { scrollRef, sentinelRef } = useInfiniteScroll( fetchMore, @@ -225,21 +265,21 @@ export function ModelSection() { > {(id: string) => { - const r = hfResults.find((m) => m.id === id); - const detail = r?.totalParams - ? formatCompact(r.totalParams) - : r?.downloads != null - ? `↓${formatCompact(r.downloads)}` - : null; + const entry = vramMap.get(id); + const detail = entry?.detail ?? null; + const fitStatus = entry?.status ?? null; + const vramEst = entry?.est ?? null; + const exceeds = fitStatus === "exceeds"; + return ( - + {id} @@ -248,13 +288,34 @@ export function ModelSection() { className="max-w-xs break-all" > {id} + {vramEst != null && vramEst > 0 && gpu.available && ( + + {exceeds + ? `Needs ~${vramEst}GB VRAM (GPU: ${gpu.memoryTotalGb}GB)` + : fitStatus === "tight" + ? `~${vramEst}GB VRAM (tight fit on ${gpu.memoryTotalGb}GB)` + : `~${vramEst}GB VRAM`} + + )} - {detail && ( - - {detail} - - )} + + {fitStatus === "exceeds" && ( + + OOM + + )} + {fitStatus === "tight" && ( + + TIGHT + + )} + {detail && ( + + {detail} + + )} + ); }} diff --git a/studio/frontend/src/hooks/index.ts b/studio/frontend/src/hooks/index.ts index 4e338169e4..24f6faca79 100644 --- a/studio/frontend/src/hooks/index.ts +++ b/studio/frontend/src/hooks/index.ts @@ -1,4 +1,5 @@ export { useDebouncedValue } from "./use-debounced-value"; +export { useGpuInfo } from "./use-gpu-info"; export { useHfModelSearch } from "./use-hf-model-search"; export { useHfDatasetSearch } from "./use-hf-dataset-search"; export { useInfiniteScroll } from "./use-infinite-scroll"; diff --git a/studio/frontend/src/hooks/use-gpu-info.ts b/studio/frontend/src/hooks/use-gpu-info.ts new file mode 100644 index 0000000000..98ca056050 --- /dev/null +++ b/studio/frontend/src/hooks/use-gpu-info.ts @@ -0,0 +1,68 @@ +import { useEffect, useState } from "react"; + +export interface GpuInfo { + available: boolean; + name: string; + memoryTotalGb: number; +} + +const DEFAULT_GPU: GpuInfo = { + available: false, + name: "Unknown", + memoryTotalGb: 0, +}; + +// Module-level cache so multiple components share one fetch. +let cachedGpu: GpuInfo | null = null; +let fetchPromise: Promise | null = null; + +async function fetchGpuOnce(): Promise { + if (cachedGpu) return cachedGpu; + if (fetchPromise) return fetchPromise; + + fetchPromise = (async () => { + try { + const res = await fetch("/api/system"); + if (!res.ok) throw new Error(`HTTP ${res.status}`); + const data = await res.json(); + const gpuData = data?.gpu; + if (!gpuData?.available || !gpuData.devices?.length) return DEFAULT_GPU; + const dev = gpuData.devices[0]; + const info: GpuInfo = { + available: true, + name: dev.name ?? "Unknown", + memoryTotalGb: dev.memory_total_gb ?? 0, + }; + cachedGpu = info; + return info; + } catch { + // Reset promise so subsequent calls retry (e.g. backend wasn't ready) + fetchPromise = null; + return DEFAULT_GPU; + } + })(); + + return fetchPromise; +} + +/** + * Fetch GPU info from the backend /api/system endpoint. + * + * The result is cached at module level -- only one network request is made + * regardless of how many components call this hook. + */ +export function useGpuInfo(): GpuInfo { + const [gpu, setGpu] = useState(cachedGpu ?? DEFAULT_GPU); + + useEffect(() => { + if (cachedGpu) return; + + let cancelled = false; + fetchGpuOnce().then((info) => { + if (!cancelled) setGpu(info); + }); + return () => { cancelled = true; }; + }, []); + + return gpu; +} diff --git a/studio/frontend/src/lib/vram.ts b/studio/frontend/src/lib/vram.ts new file mode 100644 index 0000000000..7aea44e850 --- /dev/null +++ b/studio/frontend/src/lib/vram.ts @@ -0,0 +1,95 @@ +/** + * VRAM estimation for model loading (4-bit quantization via bitsandbytes). + * + * Estimates the total driver-level VRAM (what nvidia-smi reports) needed to + * load a model in 4-bit with Unsloth / bitsandbytes. This determines + * whether a model will fit on the user's GPU before any training begins. + * + * Formula: totalParams * 0.90 + 1.4 GB + * + * Calibrated against isolated Unsloth model loads on RTX 5070 Ti (2026.2): + * Qwen2.5-0.5B (0.49B) : est 1.8 vs actual 1.86 GB (-3%) + * Llama-3.2-1B (1.24B) : est 2.5 vs actual 2.54 GB (-1%) + * Llama-3.2-3B (3.21B) : est 4.3 vs actual 4.40 GB (-2%) + * Llama-3.1-8B (8.03B) : est 8.6 vs actual 8.14 GB (+6%) + * + * Accuracy: within 3% for 0.5B-3B models, within 6% for 8B. + */ + +// --------------------------------------------------------------------------- +// Constants (exported for testing) +// --------------------------------------------------------------------------- + +/** + * Effective bytes per parameter for 4-bit model weights at driver level. + * + * Raw bnb 4-bit is ~0.5 bytes/param, but embedding and lm_head layers remain + * in fp16 and bnb adds per-block quantization metadata, bringing the + * effective rate to ~0.84-0.93 across tested architectures. 0.9 is the + * calibrated middle ground. + */ +export const BNB_4BIT_LOADING_BYTES = 0.9; + +/** + * Fixed overhead (GB) for the CUDA driver context and PyTorch runtime. + * + * This is independent of model size -- it is the baseline GPU memory consumed + * before any model weights are loaded. Measured at 1.34-1.46 GB across + * tested models; we use 1.4 as the default. + */ +export const LOADING_OVERHEAD_GB = 1.4; + +// --------------------------------------------------------------------------- +// Estimation +// --------------------------------------------------------------------------- + +export type VramFitStatus = "fits" | "tight" | "exceeds"; + +/** + * Bytes per parameter when loading a model at fp16/bf16 (LoRA, full FT). + * + * This is the theoretical value (2 bytes = 16 bits). Not yet calibrated + * against actual measurements -- the real driver-level usage may be slightly + * higher due to buffers and metadata, similar to how 4-bit is 0.9 vs 0.5. + */ +export const FP16_LOADING_BYTES = 2.0; + +export type TrainingMethod = "qlora" | "lora" | "full"; + +/** + * Estimate VRAM (GB) needed to load a model with Unsloth. + * + * The bytes-per-param rate depends on the training method: + * - QLoRA : 4-bit quantized via bnb -> 0.90 bytes/param (calibrated) + * - LoRA : fp16 -> 2.0 bytes/param (theoretical) + * - Full : fp16 -> 2.0 bytes/param (theoretical) + * + * Formula: totalParams * bytesPerParam + 1.4 GB overhead + */ +export function estimateLoadingVram( + totalParams: number, + method: TrainingMethod = "qlora", +): number { + const bytesPerParam = + method === "qlora" ? BNB_4BIT_LOADING_BYTES : FP16_LOADING_BYTES; + const gb = (totalParams / 1e9) * bytesPerParam + LOADING_OVERHEAD_GB; + return Math.round(gb * 10) / 10; +} + +/** + * Check whether a model fits in the available GPU VRAM. + * + * fits - uses <= 75% of available + * tight - uses 75-100% of available + * exceeds - uses > 100% of available + */ +export function checkVramFit( + requiredGb: number, + availableGb: number, +): VramFitStatus { + if (availableGb <= 0) return requiredGb <= 0 ? "fits" : "exceeds"; + const ratio = requiredGb / availableGb; + if (ratio <= 0.75) return "fits"; + if (ratio <= 1.0) return "tight"; + return "exceeds"; +}