diff --git a/studio/frontend/src/features/studio/sections/progress-section.tsx b/studio/frontend/src/features/studio/sections/progress-section.tsx index a04273f568..10a547b82d 100644 --- a/studio/frontend/src/features/studio/sections/progress-section.tsx +++ b/studio/frontend/src/features/studio/sections/progress-section.tsx @@ -237,6 +237,7 @@ export function ProgressSection(): ReactElement { onClick={() => { setStopRequested(true); setStopDialogOpen(false); + useTrainingRuntimeStore.getState().setStopRequested(true); void stopTrainingRun(false).then((ok) => { if (!ok) setStopRequested(false); }); @@ -248,6 +249,7 @@ export function ProgressSection(): ReactElement { onClick={() => { setStopRequested(true); setStopDialogOpen(false); + useTrainingRuntimeStore.getState().setStopRequested(true); void stopTrainingRun(true).then((ok) => { if (!ok) setStopRequested(false); }); diff --git a/studio/frontend/src/features/studio/studio-page.tsx b/studio/frontend/src/features/studio/studio-page.tsx index 26353b9b7e..69abc1969e 100644 --- a/studio/frontend/src/features/studio/studio-page.tsx +++ b/studio/frontend/src/features/studio/studio-page.tsx @@ -45,11 +45,16 @@ export function StudioPage(): ReactElement { const dialogInitial = useDatasetPreviewDialogStore((s) => s.initialData); const closeDialog = useDatasetPreviewDialogStore((s) => s.close); + const stopRequested = useTrainingRuntimeStore((state) => state.stopRequested); const canGoBack = showTrainingView && - !isTrainingRunning && !isHydratingRuntime && - (runtimePhase === "stopped" || runtimePhase === "error" || runtimePhase === "completed" || runtimePhase === "idle"); + (stopRequested || + (!isTrainingRunning && + (runtimePhase === "stopped" || + runtimePhase === "error" || + runtimePhase === "completed" || + runtimePhase === "idle"))); const tourEnabled = hasHydratedRuntime && !isHydratingRuntime; const isConfigTour = !showTrainingView; const tourSteps = showTrainingView ? studioTrainingTourSteps : studioTourSteps; diff --git a/studio/frontend/src/features/studio/training-start-overlay.tsx b/studio/frontend/src/features/studio/training-start-overlay.tsx index b6c413b400..f23cf529df 100644 --- a/studio/frontend/src/features/studio/training-start-overlay.tsx +++ b/studio/frontend/src/features/studio/training-start-overlay.tsx @@ -1,9 +1,23 @@ +import { + AlertDialog, + AlertDialogAction, + AlertDialogCancel, + AlertDialogContent, + AlertDialogDescription, + AlertDialogFooter, + AlertDialogHeader, + AlertDialogTitle, +} from "@/components/ui/alert-dialog"; +import { Button } from "@/components/ui/button"; import { AnimatedSpan, Terminal, TypingAnimation, -} from "@/components/ui/terminal" -import type { ReactElement } from "react" +} from "@/components/ui/terminal"; +import { useTrainingActions, useTrainingRuntimeStore } from "@/features/training"; +import { StopIcon } from "@hugeicons/core-free-icons"; +import { HugeiconsIcon } from "@hugeicons/react"; +import { useEffect, useState, type ReactElement } from "react"; type TrainingStartOverlayProps = { message: string @@ -14,9 +28,58 @@ export function TrainingStartOverlay({ message, currentStep, }: TrainingStartOverlayProps): ReactElement { + const { stopTrainingRun } = useTrainingActions(); + const isStarting = useTrainingRuntimeStore((s) => s.isStarting); + const [cancelDialogOpen, setCancelDialogOpen] = useState(false); + const [cancelRequested, setCancelRequested] = useState(false); + + useEffect(() => { + if (!isStarting) { + setCancelRequested(false); + } + }, [isStarting]); + return ( -
+
+
+ + + + + Cancel Training + + Do you want to cancel the current training run? + + + + Continue Training + { + setCancelRequested(true); + setCancelDialogOpen(false); + useTrainingRuntimeStore.getState().setStopRequested(true); + void stopTrainingRun(false).then((ok) => { + if (!ok) setCancelRequested(false); + }); + }} + > + Cancel Training + + + + +
Unsloth mascot()((set) => ({ ...initialState, + setStopRequested: (value) => set({ stopRequested: value }), setHydrating: (value) => set({ isHydrating: value }), setHasHydrated: (value) => set({ hasHydrated: value }), setStarting: (value) => set({ isStarting: value }), @@ -173,12 +175,15 @@ export const useTrainingRuntimeStore = create()((set) => ( const detailLoss = payload.details?.loss; const detailLr = payload.details?.learning_rate; const detailEpoch = payload.details?.epoch; + const stopRequested = + payload.is_training_running ? state.stopRequested : false; return { ...state, jobId: payload.job_id || state.jobId, phase: payload.phase, isTrainingRunning: payload.is_training_running, + stopRequested, evalEnabled: payload.eval_enabled ?? state.evalEnabled, message: payload.message, error: payload.error, diff --git a/studio/frontend/src/features/training/types/runtime.ts b/studio/frontend/src/features/training/types/runtime.ts index 7ebf09518d..b6e417c963 100644 --- a/studio/frontend/src/features/training/types/runtime.ts +++ b/studio/frontend/src/features/training/types/runtime.ts @@ -95,9 +95,11 @@ export interface TrainingRuntimeState { gradNormHistory: TrainingSeriesPoint[]; evalLossHistory: TrainingSeriesPoint[]; resetGeneration: number; + stopRequested: boolean; } export interface TrainingRuntimeActions { + setStopRequested: (value: boolean) => void; setHydrating: (value: boolean) => void; setHasHydrated: (value: boolean) => void; setStarting: (value: boolean) => void;