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 };
+}