Skip to content

Commit 5516781

Browse files
committed
feat(h05): checkpoint save/load path inputs in TopBar → stage_train G12
TopBar gains two text inputs in the Train dropdown menu: - `train-checkpoint-save-path` → opts.checkpoint_save_path - `train-checkpoint-load-path` → opts.checkpoint_load_path App.handleRunPipeline forwards both into stage_options.train when non-empty so backend stage_train (V5-G12) can safetensors.save_file / safetensors.load_file the model weights. Backend extras already emit `checkpoint.saved_path` and `checkpoint.loaded_path`. Tests: - vbgui vitest TopBar.test.tsx: two H05 specs cover (a) filled inputs forward as expected and (b) empty inputs omit fields — 20/20 TopBar tests passing. - vbgui vitest full suite: 200/200 passing. - vbgui e2e 59_checkpoint_save_load.spec.ts: H05.A: save-path → extras.checkpoint.saved_path matches; file lands on disk after the train run. H05.B: depends on H05.A having written the file; load-path → extras.checkpoint.loaded_path matches. Strict loss-continuation (H05.6 in plan) is deferred to H19 (cppmega-mlx-56o — strict identical-loss-continuation) where the spec/data pinning also lands. - vbgui e2e V6 suite (55–59): 7/7 passing — no regression. - backend pytest stage_train + warm + checkpoint: 44/44 passing.
1 parent 98a925a commit 5516781

4 files changed

Lines changed: 129 additions & 3 deletions

File tree

Lines changed: 58 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,58 @@
1+
// H05: TopBar checkpoint save/load path inputs forward into stage_train
2+
// G12. The two tests run sequentially in this file (--workers=1):
3+
// 1. Save-path → extras.checkpoint.saved_path matches and the
4+
// safetensors file lands on disk for the next test to load.
5+
// 2. Load-path with that file → extras.checkpoint.loaded_path matches.
6+
//
7+
// The strict loss-continuation round-trip (H05.6) is deferred to H19
8+
// (strict identical-loss-continuation) where the spec/data also gets
9+
// pinned to a deterministic seed. H05 closes the path-forwarding gap.
10+
11+
import { test, expect } from "@playwright/test";
12+
import { unlinkSync, existsSync } from "node:fs";
13+
import { gotoApp, selectPreset, closeModal } from "../fixtures";
14+
15+
const SAVE = "/tmp/vbgui_h05_ckpt.safetensors";
16+
17+
test.beforeAll(() => {
18+
if (existsSync(SAVE)) unlinkSync(SAVE);
19+
});
20+
21+
test("H05.A: save-path forwards → checkpoint.saved_path matches",
22+
async ({ page }) => {
23+
test.setTimeout(120_000);
24+
await gotoApp(page);
25+
await selectPreset(page, "llama3_8b");
26+
await page.getByTestId("run-pipeline-toggle").click();
27+
await page.getByTestId("train-num-steps").fill("2");
28+
await page.getByTestId("train-checkpoint-save-path").fill(SAVE);
29+
await page.getByTestId("run-pipeline-train").click();
30+
await page.getByTestId("run-result-modal").waitFor({ timeout: 60_000 });
31+
await page.getByTestId("run-result-expand-train").click();
32+
await page.getByTestId("run-result-extras-row-train").waitFor();
33+
const saved = await page.getByTestId(
34+
"run-result-extras-train-checkpoint-saved_path").textContent();
35+
expect(saved?.trim()).toBe(SAVE);
36+
await closeModal(page);
37+
expect(existsSync(SAVE)).toBe(true);
38+
});
39+
40+
test("H05.B: load-path forwards → checkpoint.loaded_path matches",
41+
async ({ page }) => {
42+
test.setTimeout(120_000);
43+
// Guard: depends on file written by H05.A.
44+
expect(existsSync(SAVE)).toBe(true);
45+
await gotoApp(page);
46+
await selectPreset(page, "llama3_8b");
47+
await page.getByTestId("run-pipeline-toggle").click();
48+
await page.getByTestId("train-num-steps").fill("2");
49+
await page.getByTestId("train-checkpoint-load-path").fill(SAVE);
50+
await page.getByTestId("run-pipeline-train").click();
51+
await page.getByTestId("run-result-modal").waitFor({ timeout: 60_000 });
52+
await page.getByTestId("run-result-expand-train").click();
53+
await page.getByTestId("run-result-extras-row-train").waitFor();
54+
const loaded = await page.getByTestId(
55+
"run-result-extras-train-checkpoint-loaded_path").textContent();
56+
expect(loaded?.trim()).toBe(SAVE);
57+
await closeModal(page);
58+
});

vbgui/src/App.tsx

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -293,7 +293,10 @@ export function App(): JSX.Element {
293293

294294
const handleRunPipeline = useCallback(async (
295295
mode: RunMode,
296-
opts?: { num_steps?: number; warm_start?: boolean },
296+
opts?: { num_steps?: number; warm_start?: boolean;
297+
checkpoint_save_path?: string;
298+
checkpoint_load_path?: string;
299+
},
297300
) => {
298301
const snap = wireSpecRef.current;
299302
if (snap.nodes.length === 0) {
@@ -322,6 +325,13 @@ export function App(): JSX.Element {
322325
if (opts?.warm_start && lastTrainRunId) {
323326
trainOpts.continue_from_run_id = lastTrainRunId;
324327
}
328+
// H05: forward checkpoint save/load paths to stage_train G12.
329+
if (opts?.checkpoint_save_path) {
330+
trainOpts.checkpoint_save_path = opts.checkpoint_save_path;
331+
}
332+
if (opts?.checkpoint_load_path) {
333+
trainOpts.checkpoint_load_path = opts.checkpoint_load_path;
334+
}
325335
if (trainParquetPath) trainOpts.parquet_path = trainParquetPath;
326336
if (trainTokenizerPath) trainOpts.tokenizer_path = trainTokenizerPath;
327337
// Forward SideChannelsTab train selection as synthetic int lists for the

vbgui/src/components/TopBar.tsx

Lines changed: 33 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,10 @@ export interface TopBarProps {
1414
onTopologyChange: (t: TopologyFactory) => void;
1515
onCompileModeChange: (m: SpecState["sharding"]["compile_mode"]) => void;
1616
onRunPipeline: (mode: RunMode,
17-
opts?: { num_steps?: number; warm_start?: boolean }) => void;
17+
opts?: { num_steps?: number; warm_start?: boolean;
18+
checkpoint_save_path?: string;
19+
checkpoint_load_path?: string;
20+
}) => void;
1821
/** H02: toggle callbacks. */
1922
onMixedPrecisionChange?: (enabled: boolean) => void;
2023
onFp8EnabledChange?: (enabled: boolean) => void;
@@ -40,6 +43,8 @@ export function TopBar(p: TopBarProps): JSX.Element {
4043
const [open, setOpen] = useState(false);
4144
const [trainNumSteps, setTrainNumSteps] = useState<number>(2);
4245
const [warmStart, setWarmStart] = useState<boolean>(false);
46+
const [ckptSavePath, setCkptSavePath] = useState<string>("");
47+
const [ckptLoadPath, setCkptLoadPath] = useState<string>("");
4348
return (
4449
<header data-testid="top-bar"
4550
style={{ height: 56, display: "flex", alignItems: "center",
@@ -174,11 +179,37 @@ export function TopBar(p: TopBarProps): JSX.Element {
174179
onChange={(e) => setWarmStart(e.target.checked)} />
175180
warm-start (continue from last run)
176181
</label>
182+
<div style={{ padding: "6px 12px", display: "flex",
183+
flexDirection: "column", gap: 4, fontSize: 11,
184+
color: "#374151" }}>
185+
<label style={{ display: "flex", alignItems: "center",
186+
gap: 6 }}>
187+
<span style={{ width: 78, color: "#6b7280" }}>ckpt save:</span>
188+
<input data-testid="train-checkpoint-save-path" type="text"
189+
placeholder="/tmp/ckpt.safetensors"
190+
value={ckptSavePath}
191+
onChange={(e) => setCkptSavePath(e.target.value)}
192+
style={{ width: 200 }} />
193+
</label>
194+
<label style={{ display: "flex", alignItems: "center",
195+
gap: 6 }}>
196+
<span style={{ width: 78, color: "#6b7280" }}>ckpt load:</span>
197+
<input data-testid="train-checkpoint-load-path" type="text"
198+
placeholder="/tmp/prev.safetensors"
199+
value={ckptLoadPath}
200+
onChange={(e) => setCkptLoadPath(e.target.value)}
201+
style={{ width: 200 }} />
202+
</label>
203+
</div>
177204
<button data-testid="run-pipeline-train"
178205
onClick={() => { setOpen(false);
179206
p.onRunPipeline("train",
180207
{ num_steps: trainNumSteps,
181-
warm_start: warmStart }); }}
208+
warm_start: warmStart,
209+
checkpoint_save_path:
210+
ckptSavePath || undefined,
211+
checkpoint_load_path:
212+
ckptLoadPath || undefined }); }}
182213
disabled={!!p.trainDisabled}
183214
title={p.trainDisabled?.reason ?? ""}
184215
style={{ ...menuItem,

vbgui/tests/TopBar.test.tsx

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,33 @@ describe("TopBar", () => {
131131
expect.objectContaining({ num_steps: expect.any(Number) }));
132132
});
133133

134+
it("H05: train-checkpoint-save/load-path inputs forward into opts",
135+
() => {
136+
const onRunPipeline = vi.fn();
137+
render(<TopBar {...defaultTopProps({ onRunPipeline })} />);
138+
fireEvent.click(screen.getByTestId("run-pipeline-toggle"));
139+
fireEvent.change(screen.getByTestId("train-checkpoint-save-path"),
140+
{ target: { value: "/tmp/save.safetensors" } });
141+
fireEvent.change(screen.getByTestId("train-checkpoint-load-path"),
142+
{ target: { value: "/tmp/load.safetensors" } });
143+
fireEvent.click(screen.getByTestId("run-pipeline-train"));
144+
expect(onRunPipeline).toHaveBeenLastCalledWith("train",
145+
expect.objectContaining({
146+
checkpoint_save_path: "/tmp/save.safetensors",
147+
checkpoint_load_path: "/tmp/load.safetensors",
148+
}));
149+
});
150+
151+
it("H05: empty checkpoint inputs omit fields", () => {
152+
const onRunPipeline = vi.fn();
153+
render(<TopBar {...defaultTopProps({ onRunPipeline })} />);
154+
fireEvent.click(screen.getByTestId("run-pipeline-toggle"));
155+
fireEvent.click(screen.getByTestId("run-pipeline-train"));
156+
const opts = onRunPipeline.mock.calls.at(-1)?.[1];
157+
expect(opts.checkpoint_save_path).toBeUndefined();
158+
expect(opts.checkpoint_load_path).toBeUndefined();
159+
});
160+
134161
it("H04: train-warm-start checkbox forwards warm_start flag", () => {
135162
const onRunPipeline = vi.fn();
136163
render(<TopBar {...defaultTopProps({ onRunPipeline })} />);

0 commit comments

Comments
 (0)