diff --git a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx index b90cdd0e3b..fc579dd32a 100644 --- a/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx +++ b/studio/frontend/src/components/assistant-ui/model-selector/pickers.tsx @@ -1,7 +1,20 @@ import { Input } from "@/components/ui/input"; import { Spinner } from "@/components/ui/spinner"; -import { useDebouncedValue, useHfModelSearch, useInfiniteScroll } from "@/hooks"; +import { + Tooltip, + TooltipContent, + TooltipTrigger, +} from "@/components/ui/tooltip"; +import { + useDebouncedValue, + useGpuInfo, + useHfModelSearch, + useInfiniteScroll, + useRecommendedModelVram, +} from "@/hooks"; import { cn, formatCompact } from "@/lib/utils"; +import type { VramFitStatus } from "@/lib/vram"; +import { checkVramFit, estimateLoadingVram } from "@/lib/vram"; import { Search01Icon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { useMemo, useState, type ReactNode } from "react"; @@ -28,27 +41,77 @@ function ModelRow({ meta, selected, onClick, + vramStatus, + vramEst, + gpuGb, }: { label: string; meta?: string; selected?: boolean; onClick: () => void; + vramStatus?: VramFitStatus | null; + vramEst?: number; + gpuGb?: number; }) { - return ( + const exceeds = vramStatus === "exceeds"; + const showVramTooltip = + vramEst != null && vramEst > 0 && gpuGb != null && gpuGb > 0; + const vramTooltipText = + showVramTooltip && vramStatus + ? exceeds + ? `Needs ~${vramEst}GB VRAM (GPU: ${gpuGb}GB)` + : vramStatus === "tight" + ? `~${vramEst}GB VRAM (tight fit on ${gpuGb}GB)` + : `~${vramEst}GB VRAM` + : null; + + const content = ( ); + + if (vramTooltipText) { + return ( + + {content} + + {label} + {vramTooltipText} + + + ); + } + return content; } export function HubModelPicker({ @@ -60,6 +123,7 @@ export function HubModelPicker({ value?: string; onSelect: (id: string, meta: ModelSelectorChangeMeta) => void; }) { + const gpu = useGpuInfo(); const [query, setQuery] = useState(""); const debouncedQuery = useDebouncedValue(query); const { results, isLoading, isLoadingMore, fetchMore } = useHfModelSearch( @@ -71,6 +135,9 @@ export function HubModelPicker({ [models, value], ); + const { paramCountById: recommendedParamCountById } = + useRecommendedModelVram(recommendedIds); + const showHfSection = debouncedQuery.trim().length > 0; const recommendedSet = useMemo(() => new Set(recommendedIds), [recommendedIds]); @@ -94,6 +161,49 @@ export function HubModelPicker({ [results], ); + const vramMap = useMemo(() => { + const map = new Map< + string, + { est: number; status: VramFitStatus | null; detail: string | null } + >(); + for (const r of results) { + const detail = r.totalParams + ? formatCompact(r.totalParams) + : r.downloads != null + ? `↓${formatCompact(r.downloads)}` + : null; + if (r.totalParams) { + const est = estimateLoadingVram(r.totalParams, "qlora"); + 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; + }, [results, gpu]); + + const recommendedVramMap = useMemo(() => { + const map = new Map< + string, + { est: number; status: VramFitStatus | null; detail: string | null } + >(); + for (const id of recommendedIds) { + const totalParams = recommendedParamCountById.get(id); + if (totalParams) { + const est = estimateLoadingVram(totalParams, "qlora"); + const status = gpu.available + ? checkVramFit(est, gpu.memoryTotalGb) + : null; + const detail = formatCompact(totalParams); + map.set(id, { est, status, detail }); + } + } + return map; + }, [recommendedIds, recommendedParamCountById, gpu]); + const { scrollRef, sentinelRef } = useInfiniteScroll(fetchMore, results.length); return ( @@ -124,14 +234,23 @@ export function HubModelPicker({ No default models. ) : ( - recommendedIds.map((id) => ( - onSelect(id, { source: "hub", isLora: false })} - /> - )) + recommendedIds.map((id) => { + const vram = recommendedVramMap.get(id); + return ( + + onSelect(id, { source: "hub", isLora: false }) + } + vramStatus={vram?.status ?? null} + vramEst={vram?.est} + gpuGb={gpu.available ? gpu.memoryTotalGb : undefined} + /> + ); + }) )} ) : null} @@ -144,15 +263,23 @@ export function HubModelPicker({ No matching models. ) : ( - hfIds.map((id) => ( - onSelect(id, { source: "hub", isLora: false })} - /> - )) + hfIds.map((id) => { + const vram = vramMap.get(id); + return ( + + onSelect(id, { source: "hub", isLora: false }) + } + vramStatus={vram?.status ?? null} + vramEst={vram?.est} + gpuGb={gpu.available ? gpu.memoryTotalGb : undefined} + /> + ); + }) )}
{isLoadingMore ? ( diff --git a/studio/frontend/src/hooks/index.ts b/studio/frontend/src/hooks/index.ts index b6c4fcad14..3932baf088 100644 --- a/studio/frontend/src/hooks/index.ts +++ b/studio/frontend/src/hooks/index.ts @@ -3,6 +3,7 @@ export { useGpuInfo } from "./use-gpu-info"; export { useGpuUtilization } from "./use-gpu-utilization"; export { useHardwareInfo } from "./use-hardware-info"; export { useHfModelSearch } from "./use-hf-model-search"; +export { useRecommendedModelVram } from "./use-recommended-model-vram"; export { useHfDatasetSearch } from "./use-hf-dataset-search"; export { useHfDatasetSplits } from "./use-hf-dataset-splits"; export { useHfTokenValidation } from "./use-hf-token-validation"; diff --git a/studio/frontend/src/hooks/use-recommended-model-vram.ts b/studio/frontend/src/hooks/use-recommended-model-vram.ts new file mode 100644 index 0000000000..69692ff416 --- /dev/null +++ b/studio/frontend/src/hooks/use-recommended-model-vram.ts @@ -0,0 +1,57 @@ +import { modelInfo } from "@huggingface/hub"; +import { useEffect, useState } from "react"; + +/** + * Fetches Hugging Face model info (safetensors total param count) for a list of + * model IDs. Used to show VRAM fit (FIT / TIGHT / OOM) for recommended/default + * models in the chat model dropdown. + */ +export function useRecommendedModelVram(ids: string[]) { + const [paramCountById, setParamCountById] = useState< + Map + >(new Map()); + const [isLoading, setIsLoading] = useState(false); + + const stableKey = [...ids].filter(Boolean).sort().join(","); + + useEffect(() => { + const stableIds = stableKey ? stableKey.split(",") : []; + if (stableIds.length === 0) { + setParamCountById(new Map()); + setIsLoading(false); + return; + } + let canceled = false; + void (async () => { + setIsLoading(true); + const next = new Map(); + await Promise.all( + stableIds.map(async (id) => { + if (canceled) return; + try { + const info = await modelInfo({ + name: id, + additionalFields: ["safetensors"], + }); + const raw = info as { safetensors?: { total?: number } }; + const total = raw.safetensors?.total; + if (typeof total === "number" && total > 0) { + next.set(id, total); + } + } catch { + // Model not on HF or no safetensors; skip + } + }), + ); + if (!canceled) { + setParamCountById(next); + setIsLoading(false); + } + })(); + return () => { + canceled = true; + }; + }, [stableKey]); + + return { paramCountById, isLoading }; +}