Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,14 @@
import type { TrainingViewData } from "@/features/training";
import { getTrainingRun, onTrainingRunUpdated } from "@/features/training";
import type { TrainingRunDetailResponse } from "@/features/training";
import { parseBackendTrainingMethod } from "@/features/training/lib/training-methods";
import { type ReactElement, useEffect, useState } from "react";
import { ChartsSection } from "./sections/charts-section";
import { ProgressSection } from "./sections/progress-section";
import { translate, useT } from "@/i18n";
import {
mapTrainingRunConfigOverride,
mapTrainingRunMethod,
} from "./lib/training-run-config";

type StudioT = ReturnType<typeof useT>;

Expand Down Expand Up @@ -79,10 +82,7 @@ function mapToViewData(
isTrainingRunning: false,
modelName: run.display_name ?? run.model_name,
projectName: run.project_name,
trainingMethod: parseBackendTrainingMethod(
detail.config?.training_type,
detail.config?.load_in_4bit,
),
trainingMethod: mapTrainingRunMethod(detail),
lossHistory,
lrHistory,
gradNormHistory,
Expand Down Expand Up @@ -147,25 +147,7 @@ export function HistoricalTrainingView({
}

const viewData = mapToViewData(detail, t);
const configOverride = detail.config
? {
epochs: detail.config.num_epochs as number | undefined,
batchSize: detail.config.batch_size as number | undefined,
learningRate: detail.config.learning_rate as string | undefined,
maxSteps: detail.config.max_steps as number | undefined,
contextLength: detail.config.max_seq_length as number | undefined,
warmupSteps: detail.config.warmup_steps as number | undefined,
optimizerType: detail.config.optim as string | undefined,
loraRank: detail.config.lora_r as number | undefined,
loraAlpha: detail.config.lora_alpha as number | undefined,
loraDropout: detail.config.lora_dropout as number | undefined,
loraVariant: detail.config.use_rslora
? "rslora"
: detail.config.use_loftq
? "loftq"
: "lora",
}
: undefined;
const configOverride = mapTrainingRunConfigOverride(detail);

return (
<div className="flex flex-col gap-6">
Expand Down
51 changes: 51 additions & 0 deletions studio/frontend/src/features/studio/lib/training-run-config.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0

import type { TrainingRunDetailResponse } from "@/features/training";
import { parseBackendTrainingMethod } from "@/features/training/lib/training-methods";
import type { TrainingMethod } from "@/types/training";

export type TrainingRunConfigOverride = {
epochs?: number;
batchSize?: number;
learningRate?: string;
maxSteps?: number | null;
contextLength?: number;
warmupSteps?: number | null;
optimizerType?: string;
loraRank?: number;
loraAlpha?: number;
loraDropout?: number;
loraVariant?: string;
};

export function mapTrainingRunConfigOverride(
detail: TrainingRunDetailResponse,
): TrainingRunConfigOverride | undefined {
if (!detail.config) return undefined;
const cfg = detail.config;
return {
epochs: cfg.num_epochs as number | undefined,
batchSize: cfg.batch_size as number | undefined,
learningRate: cfg.learning_rate as string | undefined,
maxSteps: cfg.max_steps as number | null | undefined,
contextLength: cfg.max_seq_length as number | undefined,
warmupSteps: cfg.warmup_steps as number | null | undefined,
optimizerType: cfg.optim as string | undefined,
loraRank: cfg.lora_r as number | undefined,
loraAlpha: cfg.lora_alpha as number | undefined,
loraDropout: cfg.lora_dropout as number | undefined,
loraVariant: cfg.use_rslora
? "rslora"
: cfg.use_loftq
? "loftq"
: "lora",
};
}

export function mapTrainingRunMethod(
detail: TrainingRunDetailResponse,
): TrainingMethod {
const cfg = detail.config;
return parseBackendTrainingMethod(cfg?.training_type, cfg?.load_in_4bit);
}
74 changes: 70 additions & 4 deletions studio/frontend/src/features/studio/live-training-view.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -5,10 +5,19 @@ import { cn } from "@/lib/utils";
import {
useTrainingConfigStore,
useTrainingRuntimeStore,
getTrainingRun,
TrainingRunRequestError,
} from "@/features/training";
import type { TrainingViewData } from "@/features/training";
import type { ReactElement } from "react";
import type {
TrainingViewData,
TrainingRunDetailResponse,
} from "@/features/training";
import { type ReactElement, useEffect, useState } from "react";
import { useShallow } from "zustand/react/shallow";
import {
mapTrainingRunConfigOverride,
mapTrainingRunMethod,
} from "./lib/training-run-config";
import { ChartsSection } from "./sections/charts-section";
import { ProgressSection } from "./sections/progress-section";
import { TrainingStartOverlay } from "./training-start-overlay";
Expand Down Expand Up @@ -52,6 +61,56 @@ export function LiveTrainingView(): ReactElement {
})),
);

const [activeRunDetail, setActiveRunDetail] =
useState<TrainingRunDetailResponse | null>(null);

useEffect(() => {
if (!runtime.jobId) {
setActiveRunDetail(null);
return;
}
const jobId = runtime.jobId;
const controller = new AbortController();
let active = true;
let retryTimer: number | null = null;
const loadActiveRunDetail = () => {
void getTrainingRun(jobId, controller.signal)
.then((detail) => {
if (!active || controller.signal.aborted) return;
setActiveRunDetail(detail);
})
.catch((err: unknown) => {
if (!active || controller.signal.aborted) return;
const status =
err instanceof TrainingRunRequestError ? err.status : null;
const retryNotFound =
status === 404 && (runtime.isStarting || runtime.isTrainingRunning);
if (
typeof status === "number" &&
status < 500 &&
!retryNotFound
) {
setActiveRunDetail(null);
return;
}
retryTimer = window.setTimeout(loadActiveRunDetail, 1000);
});
};
loadActiveRunDetail();
return () => {
active = false;
if (retryTimer !== null) {
window.clearTimeout(retryTimer);
}
controller.abort();
};
}, [runtime.isStarting, runtime.isTrainingRunning, runtime.jobId]);

const activeRunConfigOverride =
activeRunDetail?.run.id === runtime.jobId
? mapTrainingRunConfigOverride(activeRunDetail)
: undefined;

const activeProjectName =
runtime.startProjectName !== null
? runtime.startProjectName.trim() || null
Expand All @@ -76,7 +135,10 @@ export function LiveTrainingView(): ReactElement {
isTrainingRunning: runtime.isTrainingRunning,
modelName: runtime.startModelName ?? config.selectedModel ?? "",
projectName: activeProjectName,
trainingMethod: config.trainingMethod ?? "",
trainingMethod:
activeRunDetail?.run.id === runtime.jobId
? mapTrainingRunMethod(activeRunDetail)
: config.trainingMethod ?? "",
lossHistory: runtime.lossHistory,
lrHistory: runtime.lrHistory,
gradNormHistory: runtime.gradNormHistory,
Expand Down Expand Up @@ -105,7 +167,11 @@ export function LiveTrainingView(): ReactElement {
)}
>
<div data-tour="studio-training-progress">
<ProgressSection key={runtime.jobId ?? "no-job"} data={viewData} />
<ProgressSection
key={runtime.jobId ?? "no-job"}
data={viewData}
configOverride={activeRunConfigOverride}
/>
</div>
<ChartsSection
currentStep={viewData.currentStep}
Expand Down
88 changes: 64 additions & 24 deletions studio/frontend/src/features/studio/sections/progress-section.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ import type { TrainingViewData } from "@/features/training";
import { useGpuUtilization } from "@/hooks";
import type { GpuUtilization } from "@/hooks/use-gpu-utilization";
import { cn } from "@/lib/utils";
import type { TrainingRunConfigOverride } from "../lib/training-run-config";
import {
ChartAverageIcon,
DashboardSpeed01Icon,
Expand Down Expand Up @@ -78,22 +79,27 @@ function configRow(
return [label, value];
}

function resolveConfigValue<K extends keyof TrainingRunConfigOverride>(
configOverride: TrainingRunConfigOverride | undefined,
key: K,
fallback: Exclude<TrainingRunConfigOverride[K], undefined>,
): Exclude<TrainingRunConfigOverride[K], undefined> {
if (
configOverride &&
Object.prototype.hasOwnProperty.call(configOverride, key)
) {
const value = configOverride[key];
if (value !== undefined) {
return value as Exclude<TrainingRunConfigOverride[K], undefined>;
}
}
return fallback;
}

interface ProgressSectionProps {
data: TrainingViewData;
isHistorical?: boolean;
configOverride?: {
epochs?: number;
batchSize?: number;
learningRate?: string;
maxSteps?: number;
contextLength?: number;
warmupSteps?: number;
optimizerType?: string;
loraRank?: number;
loraAlpha?: number;
loraDropout?: number;
loraVariant?: string;
};
configOverride?: TrainingRunConfigOverride;
}

export function ProgressSection({
Expand Down Expand Up @@ -183,17 +189,51 @@ export function ProgressSection({
? data.currentGradNorm
: (lastValue(data.gradNormHistory) ?? data.currentGradNorm);

const cfgEpochs = isHistorical ? configOverride?.epochs : config.epochs;
const cfgBatchSize = isHistorical ? configOverride?.batchSize : config.batchSize;
const cfgLearningRate = isHistorical ? configOverride?.learningRate : config.learningRate;
const cfgMaxSteps = isHistorical ? configOverride?.maxSteps : config.maxSteps;
const cfgContextLength = isHistorical ? configOverride?.contextLength : config.contextLength;
const cfgWarmupSteps = isHistorical ? configOverride?.warmupSteps : config.warmupSteps;
const cfgOptimizerType = isHistorical ? configOverride?.optimizerType : config.optimizerType;
const cfgLoraRank = isHistorical ? configOverride?.loraRank : config.loraRank;
const cfgLoraAlpha = isHistorical ? configOverride?.loraAlpha : config.loraAlpha;
const cfgLoraDropout = isHistorical ? configOverride?.loraDropout : config.loraDropout;
const cfgLoraVariant = isHistorical ? configOverride?.loraVariant : config.loraVariant;
const cfgEpochs = isHistorical
? configOverride?.epochs
: resolveConfigValue(configOverride, "epochs", config.epochs);
const cfgBatchSize = isHistorical
? configOverride?.batchSize
: resolveConfigValue(configOverride, "batchSize", config.batchSize);
const cfgLearningRate = isHistorical
? configOverride?.learningRate
: resolveConfigValue(
configOverride,
"learningRate",
String(config.learningRate),
);
const cfgMaxSteps = isHistorical
? configOverride?.maxSteps
: resolveConfigValue(configOverride, "maxSteps", config.maxSteps);
const cfgContextLength = isHistorical
? configOverride?.contextLength
: resolveConfigValue(
configOverride,
"contextLength",
config.contextLength,
);
const cfgWarmupSteps = isHistorical
? configOverride?.warmupSteps
: resolveConfigValue(configOverride, "warmupSteps", config.warmupSteps);
const cfgOptimizerType = isHistorical
? configOverride?.optimizerType
: resolveConfigValue(
configOverride,
"optimizerType",
config.optimizerType,
);
const cfgLoraRank = isHistorical
? configOverride?.loraRank
: resolveConfigValue(configOverride, "loraRank", config.loraRank);
const cfgLoraAlpha = isHistorical
? configOverride?.loraAlpha
: resolveConfigValue(configOverride, "loraAlpha", config.loraAlpha);
const cfgLoraDropout = isHistorical
? configOverride?.loraDropout
: resolveConfigValue(configOverride, "loraDropout", config.loraDropout);
const cfgLoraVariant = isHistorical
? configOverride?.loraVariant
: resolveConfigValue(configOverride, "loraVariant", config.loraVariant);

const optimizerLabel =
OPTIMIZER_OPTIONS.find((o) => o.value === cfgOptimizerType)?.label ??
Expand Down
15 changes: 14 additions & 1 deletion studio/frontend/src/features/training/api/history-api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -12,9 +12,22 @@ import type {

const readError = (r: Response): Promise<string> => readFastApiError(r);

export class TrainingRunRequestError extends Error {
status: number | null;

constructor(message: string, status: number | null) {
super(message);
this.name = "TrainingRunRequestError";
this.status = status;
}
}

async function parseJson<T>(response: Response): Promise<T> {
if (!response.ok) {
throw new Error(await readError(response));
throw new TrainingRunRequestError(
await readError(response),
response.status,
);
}
return (await response.json()) as T;
}
Expand Down
1 change: 1 addition & 0 deletions studio/frontend/src/features/training/index.ts
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ export {
getTrainingRun,
deleteTrainingRun,
renameTrainingRun,
TrainingRunRequestError,
} from "./api/history-api";
export {
onTrainingRunUpdated,
Expand Down
Loading