+ {canGoBack && (
+
+ )}
+
{/* Header */}
diff --git a/studio/frontend/src/features/training/api/train-api.ts b/studio/frontend/src/features/training/api/train-api.ts
index e3c53b0bc2..d2f589c298 100644
--- a/studio/frontend/src/features/training/api/train-api.ts
+++ b/studio/frontend/src/features/training/api/train-api.ts
@@ -41,11 +41,22 @@ export async function startTraining(
return parseJson(response);
}
-export async function stopTraining(): Promise {
- const response = await authFetch("/api/train/stop", { method: "POST" });
+export async function stopTraining(save = true): Promise {
+ const response = await authFetch("/api/train/stop", {
+ method: "POST",
+ headers: { "Content-Type": "application/json" },
+ body: JSON.stringify({ save }),
+ });
return parseJson(response);
}
+export async function resetTraining(): Promise {
+ const response = await authFetch("/api/train/reset", { method: "POST" });
+ if (!response.ok) {
+ throw new Error(await readError(response));
+ }
+}
+
export async function getTrainingStatus(): Promise {
const response = await authFetch("/api/train/status");
return parseJson(response);
diff --git a/studio/frontend/src/features/training/hooks/use-training-actions.ts b/studio/frontend/src/features/training/hooks/use-training-actions.ts
index 510c45ed4a..4ace9dd5ed 100644
--- a/studio/frontend/src/features/training/hooks/use-training-actions.ts
+++ b/studio/frontend/src/features/training/hooks/use-training-actions.ts
@@ -1,7 +1,7 @@
import { useCallback } from "react";
import { useTrainingConfigStore } from "../stores/training-config-store";
import { useTrainingRuntimeStore } from "../stores/training-runtime-store";
-import { startTraining, stopTraining } from "../api/train-api";
+import { startTraining, stopTraining, resetTraining } from "../api/train-api";
import { buildTrainingStartPayload } from "../api/mappers";
import { syncTrainingRuntimeFromBackend } from "../lib/sync-runtime";
import { validateTrainingConfig } from "../lib/validation";
@@ -45,12 +45,12 @@ export function useTrainingActions() {
}
}, []);
- const stopTrainingRun = useCallback(async (): Promise => {
+ const stopTrainingRun = useCallback(async (save = true): Promise => {
const runtimeStore = useTrainingRuntimeStore.getState();
runtimeStore.setStartError(null);
try {
- await stopTraining();
+ await stopTraining(save);
await syncTrainingRuntimeFromBackend();
return true;
} catch (error) {
@@ -61,10 +61,20 @@ export function useTrainingActions() {
}
}, []);
+ const dismissTrainingRun = useCallback(async (): Promise => {
+ useTrainingRuntimeStore.getState().resetRuntime();
+ try {
+ await resetTraining();
+ } catch {
+ // Frontend already reset; backend will catch up on next poll
+ }
+ }, []);
+
return {
isStarting,
startError,
startTrainingRun,
stopTrainingRun,
+ dismissTrainingRun,
};
}
diff --git a/studio/frontend/src/features/training/hooks/use-training-runtime-lifecycle.ts b/studio/frontend/src/features/training/hooks/use-training-runtime-lifecycle.ts
index 3baff3a62d..4bf329a33e 100644
--- a/studio/frontend/src/features/training/hooks/use-training-runtime-lifecycle.ts
+++ b/studio/frontend/src/features/training/hooks/use-training-runtime-lifecycle.ts
@@ -1,3 +1,4 @@
+import { hasAuthToken } from "@/features/auth";
import { useEffect } from "react";
import {
getTrainingMetrics,
@@ -48,23 +49,27 @@ export function useTrainingRuntimeLifecycle(): void {
};
const pollMetrics = async () => {
+ if (!hasAuthToken()) return;
+ const gen = runtimeStore.getState().resetGeneration;
try {
const metrics = await getTrainingMetrics();
- if (disposed) {
+ if (disposed || runtimeStore.getState().resetGeneration !== gen) {
return;
}
runtimeStore.getState().applyMetrics(metrics);
} catch (error) {
- if (!isAbortError(error) && !disposed) {
+ if (!isAbortError(error) && !disposed && hasAuthToken()) {
runtimeStore.getState().setSseConnected(false);
}
}
};
const pollStatus = async () => {
+ if (!hasAuthToken()) return;
+ const gen = runtimeStore.getState().resetGeneration;
try {
const status = await getTrainingStatus();
- if (disposed) {
+ if (disposed || runtimeStore.getState().resetGeneration !== gen) {
return;
}
@@ -77,7 +82,7 @@ export function useTrainingRuntimeLifecycle(): void {
stopStream();
}
} catch (error) {
- if (!isAbortError(error) && !disposed) {
+ if (!isAbortError(error) && !disposed && hasAuthToken()) {
runtimeStore.getState().setSseConnected(false);
}
}
diff --git a/studio/frontend/src/features/training/lib/sync-runtime.ts b/studio/frontend/src/features/training/lib/sync-runtime.ts
index b5fbd0bafb..bf255f58c6 100644
--- a/studio/frontend/src/features/training/lib/sync-runtime.ts
+++ b/studio/frontend/src/features/training/lib/sync-runtime.ts
@@ -6,12 +6,17 @@ import { useTrainingRuntimeStore } from "../stores/training-runtime-store";
import type { TrainingStatusResponse } from "../types/runtime";
export async function syncTrainingRuntimeFromBackend(): Promise {
+ const gen = useTrainingRuntimeStore.getState().resetGeneration;
+
const [status, metrics] = await Promise.all([
getTrainingStatus(),
getTrainingMetrics(),
]);
const runtimeStore = useTrainingRuntimeStore.getState();
+ if (runtimeStore.resetGeneration !== gen) {
+ return status;
+ }
runtimeStore.applyStatus(status);
runtimeStore.applyMetrics(metrics);
diff --git a/studio/frontend/src/features/training/stores/training-runtime-store.ts b/studio/frontend/src/features/training/stores/training-runtime-store.ts
index bc40019346..94db0f28f5 100644
--- a/studio/frontend/src/features/training/stores/training-runtime-store.ts
+++ b/studio/frontend/src/features/training/stores/training-runtime-store.ts
@@ -34,6 +34,7 @@ const initialState: TrainingRuntimeState = {
lossHistory: [],
lrHistory: [],
gradNormHistory: [],
+ resetGeneration: 0,
};
function sortSeries(points: TrainingSeriesPoint[]): TrainingSeriesPoint[] {
@@ -95,12 +96,13 @@ export const useTrainingRuntimeStore = create()((set) => (
setLastEventId: (value) => set({ lastEventId: value }),
resetRuntime: () =>
- set({
+ set((state) => ({
...initialState,
lossHistory: [],
lrHistory: [],
gradNormHistory: [],
- }),
+ resetGeneration: state.resetGeneration + 1,
+ })),
setStartQueued: (jobId, message) =>
set({
diff --git a/studio/frontend/src/features/training/types/runtime.ts b/studio/frontend/src/features/training/types/runtime.ts
index fe2afbd36d..389418680a 100644
--- a/studio/frontend/src/features/training/types/runtime.ts
+++ b/studio/frontend/src/features/training/types/runtime.ts
@@ -82,6 +82,7 @@ export interface TrainingRuntimeState {
lossHistory: TrainingSeriesPoint[];
lrHistory: TrainingSeriesPoint[];
gradNormHistory: TrainingSeriesPoint[];
+ resetGeneration: number;
}
export interface TrainingRuntimeActions {