From 6f9dd90d56f3f7d36300adee8901311bb9da9103 Mon Sep 17 00:00:00 2001 From: Manan17 Date: Tue, 17 Feb 2026 00:34:48 +0000 Subject: [PATCH] Integration of the api with the EXPORT page with UI changes --- .../src/features/export/api/export-api.ts | 127 ++++ .../export/components/export-dialog.tsx | 80 ++- .../src/features/export/export-page.tsx | 566 ++++++++++++------ 3 files changed, 595 insertions(+), 178 deletions(-) create mode 100644 studio/frontend/src/features/export/api/export-api.ts diff --git a/studio/frontend/src/features/export/api/export-api.ts b/studio/frontend/src/features/export/api/export-api.ts new file mode 100644 index 0000000000..cbaa4a8145 --- /dev/null +++ b/studio/frontend/src/features/export/api/export-api.ts @@ -0,0 +1,127 @@ +import { authFetch } from "@/features/auth"; + +async function readError(response: Response): Promise { + try { + const payload = (await response.json()) as { detail?: string; message?: string }; + return payload.detail || payload.message || `Request failed (${response.status})`; + } catch { + return `Request failed (${response.status})`; + } +} + +async function parseJson(response: Response): Promise { + if (!response.ok) { + throw new Error(await readError(response)); + } + return (await response.json()) as T; +} + +export interface CheckpointInfo { + display_name: string; + path: string; + loss?: number | null; +} + +export interface ModelCheckpoints { + name: string; + checkpoints: CheckpointInfo[]; + base_model?: string | null; + peft_type?: string | null; + lora_rank?: number | null; +} + +export interface CheckpointListResponse { + outputs_dir: string; + models: ModelCheckpoints[]; +} + +export interface ExportOperationResponse { + success: boolean; + message: string; + details?: Record | null; +} + +export async function fetchCheckpoints(): Promise { + const response = await authFetch("/api/models/checkpoints"); + return parseJson(response); +} + +export async function loadCheckpoint(params: { + checkpoint_path: string; + max_seq_length?: number; + load_in_4bit?: boolean; +}): Promise { + const response = await authFetch("/api/export/load-checkpoint", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(params), + }); + return parseJson(response); +} + +export async function exportMerged(params: { + save_directory: string; + format_type?: string; + push_to_hub?: boolean; + repo_id?: string | null; + hf_token?: string | null; + private?: boolean; +}): Promise { + const response = await authFetch("/api/export/export/merged", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(params), + }); + return parseJson(response); +} + +export async function exportBase(params: { + save_directory: string; + push_to_hub?: boolean; + repo_id?: string | null; + hf_token?: string | null; + private?: boolean; + base_model_id?: string | null; +}): Promise { + const response = await authFetch("/api/export/export/base", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(params), + }); + return parseJson(response); +} + +export async function exportGGUF(params: { + save_directory: string; + quantization_method: string; + push_to_hub?: boolean; + repo_id?: string | null; + hf_token?: string | null; +}): Promise { + const response = await authFetch("/api/export/export/gguf", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(params), + }); + return parseJson(response); +} + +export async function exportLoRA(params: { + save_directory: string; + push_to_hub?: boolean; + repo_id?: string | null; + hf_token?: string | null; + private?: boolean; +}): Promise { + const response = await authFetch("/api/export/export/lora", { + method: "POST", + headers: { "Content-Type": "application/json" }, + body: JSON.stringify(params), + }); + return parseJson(response); +} + +export async function cleanupExport(): Promise { + const response = await authFetch("/api/export/cleanup", { method: "POST" }); + return parseJson(response); +} diff --git a/studio/frontend/src/features/export/components/export-dialog.tsx b/studio/frontend/src/features/export/components/export-dialog.tsx index 4f66048270..fedbcca76a 100644 --- a/studio/frontend/src/features/export/components/export-dialog.tsx +++ b/studio/frontend/src/features/export/components/export-dialog.tsx @@ -13,8 +13,9 @@ import { InputGroupAddon, InputGroupInput, } from "@/components/ui/input-group"; +import { Spinner } from "@/components/ui/spinner"; import { Switch } from "@/components/ui/switch"; -import { ArrowRight01Icon, Key01Icon } from "@hugeicons/core-free-icons"; +import { AlertCircleIcon, ArrowRight01Icon, CheckmarkCircle02Icon, Key01Icon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { AnimatePresence, motion } from "motion/react"; import { collapseAnim } from "../anim"; @@ -41,6 +42,10 @@ interface ExportDialogProps { onHfTokenChange: (v: string) => void; privateRepo: boolean; onPrivateRepoChange: (v: boolean) => void; + onExport: () => void; + exporting: boolean; + exportError: string | null; + exportSuccess: boolean; } export function ExportDialog({ @@ -62,10 +67,41 @@ export function ExportDialog({ onHfTokenChange, privateRepo, onPrivateRepoChange, + onExport, + exporting, + exportError, + exportSuccess, }: ExportDialogProps) { return ( - - + { + if (exporting) return; + onOpenChange(v); + }} + > + { if (exporting) e.preventDefault(); }}> + {exportSuccess ? ( + <> +
+
+ +
+
+

Export Complete

+

+ {destination === "hub" + ? "Model successfully pushed to Hugging Face Hub." + : "Model saved locally."} +

+
+
+ + + + + ) : ( + <> Export Model @@ -77,6 +113,7 @@ export function ExportDialog({ - + + + )}
); diff --git a/studio/frontend/src/features/export/export-page.tsx b/studio/frontend/src/features/export/export-page.tsx index 96ca3589eb..5c11a5914b 100644 --- a/studio/frontend/src/features/export/export-page.tsx +++ b/studio/frontend/src/features/export/export-page.tsx @@ -8,80 +8,56 @@ import { SelectValue, } from "@/components/ui/select"; import { Separator } from "@/components/ui/separator"; +import { Spinner } from "@/components/ui/spinner"; import { Tooltip, TooltipContent, TooltipTrigger, } from "@/components/ui/tooltip"; -import { useTrainingRuntimeStore } from "@/features/training"; import { useTrainingConfigStore } from "@/features/training"; -import { isAdapterMethod } from "@/types/training"; -import { InformationCircleIcon, PackageIcon } from "@hugeicons/core-free-icons"; +import { AlertCircleIcon, InformationCircleIcon, PackageIcon } from "@hugeicons/core-free-icons"; import { HugeiconsIcon } from "@hugeicons/react"; import { AnimatePresence, motion } from "motion/react"; -import { useMemo, useState } from "react"; +import { useCallback, useEffect, useMemo, useState } from "react"; import { useShallow } from "zustand/react/shallow"; import { collapseAnim } from "./anim"; +import type { ModelCheckpoints } from "./api/export-api"; +import { + cleanupExport, + exportBase, + exportGGUF, + exportLoRA, + exportMerged, + fetchCheckpoints, + loadCheckpoint, +} from "./api/export-api"; 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"; import { GuidedTour, useGuidedTourController } from "@/features/tour"; import { exportTourSteps } from "./tour"; export function ExportPage() { - const { - trainingMethod, - selectedModel, - saveSteps, - epochs, - loraRank, - hfToken, - setHfToken, - } = useTrainingConfigStore( + const { hfToken, setHfToken } = useTrainingConfigStore( useShallow((s) => ({ - trainingMethod: s.trainingMethod, - selectedModel: s.selectedModel, - saveSteps: s.saveSteps, - epochs: s.epochs, - loraRank: s.loraRank, hfToken: s.hfToken, setHfToken: s.setHfToken, })), ); - const totalSteps = useTrainingRuntimeStore((state) => state.totalSteps); - const isAdapter = isAdapterMethod(trainingMethod); - const checkpoints = useMemo(() => { - if (isAdapter) { - const interval = saveSteps > 0 ? saveSteps : 100; - const total = totalSteps > 0 ? 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).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, saveSteps, totalSteps]); + // ---- API-driven checkpoint state ---- + const [models, setModels] = useState([]); + const [loadingCheckpoints, setLoadingCheckpoints] = useState(true); + const [checkpointError, setCheckpointError] = useState(null); + const [selectedModelIdx, setSelectedModelIdx] = useState(null); const [checkpoint, setCheckpoint] = useState(null); + const [exportMethod, setExportMethod] = useState(null); const [quantLevels, setQuantLevels] = useState([]); const [dialogOpen, setDialogOpen] = useState(false); @@ -91,11 +67,68 @@ export function ExportPage() { const [modelName, setModelName] = useState(""); const [privateRepo, setPrivateRepo] = useState(false); + const [exporting, setExporting] = useState(false); + const [exportError, setExportError] = useState(null); + const [exportSuccess, setExportSuccess] = useState(false); + const tour = useGuidedTourController({ id: "export", steps: exportTourSteps, }); + // ---- Fetch checkpoints on mount ---- + useEffect(() => { + let cancelled = false; + setLoadingCheckpoints(true); + setCheckpointError(null); + fetchCheckpoints() + .then((data) => { + if (!cancelled) { + setModels(data.models); + } + }) + .catch((err) => { + if (!cancelled) { + setCheckpointError( + err instanceof Error ? err.message : "Failed to load checkpoints", + ); + } + }) + .finally(() => { + if (!cancelled) setLoadingCheckpoints(false); + }); + return () => { + cancelled = true; + }; + }, []); + + // ---- Derived state ---- + const selectedModelData = useMemo( + () => + selectedModelIdx != null + ? models.find((m) => m.name === selectedModelIdx) ?? null + : null, + [models, selectedModelIdx], + ); + + const checkpointsForModel = useMemo( + () => selectedModelData?.checkpoints ?? [], + [selectedModelData], + ); + + // Derive training info from selected model's API metadata + const baseModelName = selectedModelData?.base_model ?? "—"; + const isAdapter = !!selectedModelData?.peft_type; + const loraRank = selectedModelData?.lora_rank ?? null; + const trainingMethodLabel = selectedModelData?.peft_type + ? "LoRA / QLoRA" + : "Full Fine-tune"; + + // Reset checkpoint when the selected model changes + useEffect(() => { + setCheckpoint(null); + }, [selectedModelIdx]); + const handleMethodChange = (method: ExportMethod) => { setExportMethod(method); if (method !== "gguf") { @@ -108,8 +141,100 @@ export function ExportPage() { checkpoint && exportMethod && (exportMethod !== "gguf" || quantLevels.length > 0); - const baseModelName = selectedModel ?? "—"; + // ---- Export handler ---- + const handleExport = useCallback(async () => { + if (!checkpoint) return; + + const selectedCp = checkpointsForModel.find( + (cp) => cp.display_name === checkpoint, + ); + if (!selectedCp) return; + + setExporting(true); + setExportError(null); + setExportSuccess(false); + + const saveDir = `./exports/${selectedModelIdx ?? "model"}/${checkpoint}`; + const pushToHub = destination === "hub"; + const repoId = pushToHub && hfUsername && modelName + ? `${hfUsername}/${modelName}` + : undefined; + const token = pushToHub && hfToken ? hfToken : undefined; + + try { + // 1. Load checkpoint + await loadCheckpoint({ checkpoint_path: selectedCp.path }); + + // 2. Run export based on method + if (exportMethod === "merged") { + if (isAdapter) { + await exportMerged({ + save_directory: saveDir, + push_to_hub: pushToHub, + repo_id: repoId, + hf_token: token, + private: privateRepo, + }); + } else { + await exportBase({ + save_directory: saveDir, + push_to_hub: pushToHub, + repo_id: repoId, + hf_token: token, + private: privateRepo, + base_model_id: selectedModelData?.base_model, + }); + } + } else if (exportMethod === "gguf") { + for (const quant of quantLevels) { + await exportGGUF({ + save_directory: saveDir, + quantization_method: quant, + push_to_hub: pushToHub, + repo_id: repoId, + hf_token: token, + }); + } + } else if (exportMethod === "lora") { + await exportLoRA({ + save_directory: saveDir, + push_to_hub: pushToHub, + repo_id: repoId, + hf_token: token, + private: privateRepo, + }); + } + + setExportSuccess(true); + } catch (err) { + setExportError( + err instanceof Error ? err.message : "Export failed", + ); + } finally { + try { + await cleanupExport(); + } catch { + // cleanup is best-effort + } + setExporting(false); + } + }, [ + checkpoint, + checkpointsForModel, + selectedModelIdx, + selectedModelData, + exportMethod, + isAdapter, + quantLevels, + destination, + hfUsername, + modelName, + hfToken, + privateRepo, + ]); + + // ---- Render ---- return (
@@ -132,141 +257,236 @@ export function ExportPage() { featured={true} className="shadow-border ring-1 ring-border" > - {/* Top row: Checkpoint + metadata | Guide */} -
-
-
- - -
+ {/* Loading / error states */} + {loadingCheckpoints && ( +
+ + Loading checkpoints… +
+ )} -
- - Training Info - -
-
- Base Model - {baseModelName} + {checkpointError && ( +
+ + {checkpointError} +
+ )} + + {!loadingCheckpoints && !checkpointError && ( + <> + {/* Top row: Dropdowns + metadata | Guide */} +
+
+ {/* Training run dropdown */} +
+ +
-
- Method - - {METHOD_LABELS[trainingMethod] ?? trainingMethod} + + {/* Checkpoint dropdown */} +
+ + +
+ +
+ + Training Info -
-
- Checkpoints - {checkpoints.length} -
-
- Epochs - {epochs} -
- {isAdapter && ( -
- LoRA Rank - {loraRank} +
+
+ Base Model + {baseModelName} +
+
+ Method + + {trainingMethodLabel} + +
+
+ Checkpoints + + {checkpointsForModel.length} + +
+ {isAdapter && ( +
+ LoRA Rank + {loraRank} +
+ )}
- )} +
+
+ +
+ + Quick Guide + +
    + {GUIDE_STEPS.map((step, i) => ( +
  1. + + {i + 1} + + {step} +
  2. + ))} +
-
-
- - Quick Guide - -
    - {GUIDE_STEPS.map((step, i) => ( -
  1. - - {i + 1} - - {step} -
  2. - ))} -
-
-
+ - + + {exportMethod === "gguf" && ( + + + + )} + - - {exportMethod === "gguf" && ( - - - - )} - - - -
-
- - Est. size: {estimatedSize} · Free disk space: 120 GB -
- -
+ +
+ {/* TODO: unhide once estimated size comes from the backend API */} + {/*
+ + Est. size: {estimatedSize} · Free disk space: 120 GB +
*/} + +
+ + )}
@@ -289,6 +509,10 @@ export function ExportPage() { onHfTokenChange={setHfToken} privateRepo={privateRepo} onPrivateRepoChange={setPrivateRepo} + onExport={handleExport} + exporting={exporting} + exportError={exportError} + exportSuccess={exportSuccess} />
);