Skip to content

Commit 558f0b1

Browse files
committed
feat: move train side-channel selection into sidebar
1 parent e05848d commit 558f0b1

7 files changed

Lines changed: 110 additions & 41 deletions

File tree

vbgui/e2e/scenarios/24_side_channels.spec.ts

Lines changed: 9 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
// V4-10/H17: side_channels toggle in train dropdown reaches stage_train
1+
// V4-10/H10/H17: SideChannelsTab train selection reaches stage_train
22
// and doc_ids has a real forward effect through the attention mask.
33

44
import { test, expect } from "@playwright/test";
@@ -10,8 +10,12 @@ test("V4-10: doc_ids toggle reaches stage_train side_channels_observed",
1010
await gotoApp(page);
1111
await selectPreset(page, "llama3_8b");
1212

13+
await page.getByTestId("sidebar-tab-side_channels").click();
14+
await page.getByTestId("side-channel-train-doc_ids").check();
15+
await expect(page.getByTestId("train-side-channel-doc_ids"))
16+
.toHaveCount(0);
17+
1318
await page.getByTestId("run-pipeline-toggle").click();
14-
await page.getByTestId("train-side-channel-doc_ids").check();
1519
await page.getByTestId("run-pipeline-train").click();
1620

1721
const modal = page.getByTestId("run-result-modal");
@@ -39,9 +43,10 @@ test("V4-10: both toggles enabled → both observed", async ({ page }) => {
3943
await gotoApp(page);
4044
await selectPreset(page, "llama3_8b");
4145

46+
await page.getByTestId("sidebar-tab-side_channels").click();
47+
await page.getByTestId("side-channel-train-doc_ids").check();
48+
await page.getByTestId("side-channel-train-token_ids").check();
4249
await page.getByTestId("run-pipeline-toggle").click();
43-
await page.getByTestId("train-side-channel-doc_ids").check();
44-
await page.getByTestId("train-side-channel-token_ids").check();
4550
await page.getByTestId("run-pipeline-train").click();
4651

4752
const modal = page.getByTestId("run-result-modal");

vbgui/src/App.tsx

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -107,6 +107,7 @@ export function App(): JSX.Element {
107107
useState<string | null>(null);
108108
const [availableSideChannels, setAvailableSideChannels] =
109109
useState<string[]>(["doc_ids", "token_ids"]);
110+
const [trainSideChannels, setTrainSideChannels] = useState<string[]>([]);
110111

111112
const rpc = useRpc({
112113
baseUrl: (import.meta.env.VITE_BACKEND_URL as string | undefined)
@@ -126,6 +127,11 @@ export function App(): JSX.Element {
126127
wireSpecRef.current = { nodes, edges, spec, availableSideChannels };
127128
}, [nodes, edges, spec, availableSideChannels]);
128129

130+
useEffect(() => {
131+
const available = new Set(availableSideChannels);
132+
setTrainSideChannels((prev) => prev.filter((name) => available.has(name)));
133+
}, [availableSideChannels]);
134+
129135
const runVerify = useCallback(async () => {
130136
const snap = wireSpecRef.current;
131137
if (snap.nodes.length === 0) return;
@@ -283,7 +289,7 @@ export function App(): JSX.Element {
283289

284290
const handleRunPipeline = useCallback(async (
285291
mode: RunMode,
286-
opts?: { num_steps?: number; side_channels?: string[] },
292+
opts?: { num_steps?: number },
287293
) => {
288294
const snap = wireSpecRef.current;
289295
if (snap.nodes.length === 0) {
@@ -309,11 +315,11 @@ export function App(): JSX.Element {
309315
}
310316
if (trainParquetPath) trainOpts.parquet_path = trainParquetPath;
311317
if (trainTokenizerPath) trainOpts.tokenizer_path = trainTokenizerPath;
312-
// Forward side-channel selection as synthetic int lists for the
318+
// Forward SideChannelsTab train selection as synthetic int lists for the
313319
// stage_train G17 math-effect smoke path.
314-
if (opts?.side_channels && opts.side_channels.length > 0) {
320+
if (trainSideChannels.length > 0) {
315321
const sc: Record<string, number[]> = {};
316-
for (const name of opts.side_channels) {
322+
for (const name of trainSideChannels) {
317323
sc[name] = [0, 1, 2, 3, 4, 5, 6, 7]; // synthetic 8-token sample
318324
}
319325
trainOpts.side_channels = sc;
@@ -338,7 +344,7 @@ export function App(): JSX.Element {
338344
setTrainRunId(null);
339345
}
340346
}
341-
}, [rpc, trainParquetPath, trainTokenizerPath]);
347+
}, [rpc, trainParquetPath, trainSideChannels, trainTokenizerPath]);
342348

343349
const handleCancelTrain = useCallback(async () => {
344350
const runId = trainRunId;
@@ -474,6 +480,7 @@ export function App(): JSX.Element {
474480
rewriters={spec.rewriters}
475481
sideChannels={spec.side_channels}
476482
availableSideChannels={availableSideChannels}
483+
selectedTrainSideChannels={trainSideChannels}
477484
sharding={spec.sharding}
478485
gotchas={spec.gotchas}
479486
proposals={proposals}
@@ -493,6 +500,7 @@ export function App(): JSX.Element {
493500
onRewriterApply={() => void scheduleVerify()}
494501
onSideChannelsApply={(s) =>
495502
dispatch({ type: "side_channels.set", side_channels: s })}
503+
onTrainSideChannelsChange={setTrainSideChannels}
496504
onShardingChange={(s) =>
497505
dispatch({ type: "sharding.set", sharding: s })}
498506
onShardingAccept={handleShardingAccept}

vbgui/src/components/Sidebar.tsx

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,7 @@ export interface SidebarProps {
2323
rewriters: RewriterState[];
2424
sideChannels: SideChannelState;
2525
availableSideChannels: string[];
26+
selectedTrainSideChannels: string[];
2627
sharding: ShardingState;
2728
gotchas: GotchaState[];
2829
proposals: ShardingProposalView[];
@@ -33,6 +34,7 @@ export interface SidebarProps {
3334
onRewriterReorder: (from: number, to: number) => void;
3435
onRewriterApply?: () => void;
3536
onSideChannelsApply: (s: SideChannelState) => void;
37+
onTrainSideChannelsChange: (channels: string[]) => void;
3638
onShardingChange: (s: ShardingState) => void;
3739
onShardingAccept: (idx: number) => void;
3840
onGotchaAutoFix?: (id: string) => void;
@@ -103,8 +105,11 @@ export function Sidebar(p: SidebarProps): JSX.Element {
103105
{active === "side_channels" && (
104106
<SideChannelsTab sideChannels={p.sideChannels}
105107
availableChannels={p.availableSideChannels}
108+
selectedTrainChannels={p.selectedTrainSideChannels}
106109
gotchas={p.gotchas}
107-
onApply={p.onSideChannelsApply} />
110+
onApply={p.onSideChannelsApply}
111+
onTrainChannelsChange={
112+
p.onTrainSideChannelsChange} />
108113
)}
109114
{active === "sharding" && (
110115
<ShardingTab sharding={p.sharding} proposals={p.proposals}

vbgui/src/components/TopBar.tsx

Lines changed: 2 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@ 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; side_channels?: string[] }) => void;
17+
opts?: { num_steps?: number }) => void;
1818
/** H02: toggle callbacks. */
1919
onMixedPrecisionChange?: (enabled: boolean) => void;
2020
onFp8EnabledChange?: (enabled: boolean) => void;
@@ -39,11 +39,6 @@ export interface TopBarProps {
3939
export function TopBar(p: TopBarProps): JSX.Element {
4040
const [open, setOpen] = useState(false);
4141
const [trainNumSteps, setTrainNumSteps] = useState<number>(2);
42-
// V4-10: side-channel toggles for the train run. Off by default;
43-
// when on, App.handleRunPipeline forwards a synthetic int list to
44-
// backend opts.side_channels so stage_train can record observation.
45-
const [scDocIds, setScDocIds] = useState<boolean>(false);
46-
const [scTokenIds, setScTokenIds] = useState<boolean>(false);
4742
return (
4843
<header data-testid="top-bar"
4944
style={{ height: 56, display: "flex", alignItems: "center",
@@ -170,30 +165,10 @@ export function TopBar(p: TopBarProps): JSX.Element {
170165
e.target.value || "1", 10)))}
171166
style={{ width: 50 }} />
172167
</div>
173-
<div style={{ padding: "0 12px 6px", display: "flex",
174-
alignItems: "center", gap: 10, fontSize: 11 }}>
175-
<span style={{ color: "#6b7280" }}>side-channels:</span>
176-
<label style={{ display: "flex", gap: 3, alignItems: "center" }}>
177-
<input data-testid="train-side-channel-doc_ids"
178-
type="checkbox" checked={scDocIds}
179-
onChange={(e) => setScDocIds(e.target.checked)} />
180-
doc_ids
181-
</label>
182-
<label style={{ display: "flex", gap: 3, alignItems: "center" }}>
183-
<input data-testid="train-side-channel-token_ids"
184-
type="checkbox" checked={scTokenIds}
185-
onChange={(e) => setScTokenIds(e.target.checked)} />
186-
token_ids
187-
</label>
188-
</div>
189168
<button data-testid="run-pipeline-train"
190169
onClick={() => { setOpen(false);
191-
const sc: string[] = [];
192-
if (scDocIds) sc.push("doc_ids");
193-
if (scTokenIds) sc.push("token_ids");
194170
p.onRunPipeline("train",
195-
{ num_steps: trainNumSteps,
196-
side_channels: sc }); }}
171+
{ num_steps: trainNumSteps }); }}
197172
disabled={!!p.trainDisabled}
198173
title={p.trainDisabled?.reason ?? ""}
199174
style={{ ...menuItem,

vbgui/src/components/sidebar/SideChannelsTab.tsx

Lines changed: 46 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -11,8 +11,10 @@ import type {
1111
export interface SideChannelsTabProps {
1212
sideChannels: SideChannelState;
1313
availableChannels: string[];
14+
selectedTrainChannels: string[];
1415
gotchas: GotchaState[];
1516
onApply: (next: SideChannelState) => void;
17+
onTrainChannelsChange: (next: string[]) => void;
1618
}
1719

1820
const MODES: SideChannelMode[] = ["off", "auto", "require", "if_available"];
@@ -28,7 +30,8 @@ const FAIL_POLICIES: InferenceFailPolicy[] = [
2830
const ADAPTERS = ["none", "cpp", "rust", "go", "python"] as const;
2931

3032
export function SideChannelsTab({
31-
sideChannels, availableChannels, gotchas, onApply,
33+
sideChannels, availableChannels, selectedTrainChannels, gotchas, onApply,
34+
onTrainChannelsChange,
3235
}: SideChannelsTabProps): JSX.Element {
3336
const [draft, setDraft] = useState<SideChannelState>(sideChannels);
3437
const [platform, setPlatform] = useState({
@@ -44,6 +47,10 @@ export function SideChannelsTab({
4447
useEffect(() => setDraft(sideChannels), [sideChannels]);
4548

4649
const available = useMemo(() => new Set(availableChannels), [availableChannels]);
50+
const selectedTrain = useMemo(
51+
() => new Set(selectedTrainChannels),
52+
[selectedTrainChannels],
53+
);
4754
const requiredErrors = gotchas.filter((g) =>
4855
g.id.startsWith("side_channel_required_"));
4956
const platformPreview = renderPlatform(platform);
@@ -69,6 +76,29 @@ export function SideChannelsTab({
6976
</div>
7077
</section>
7178

79+
<section data-testid="side-channel-train-selection" style={section}>
80+
<h4 style={heading}>Train Inputs</h4>
81+
{availableChannels.length === 0 ? (
82+
<div data-testid="side-channel-train-empty" style={muted}>
83+
no available token channels
84+
</div>
85+
) : (
86+
<div style={{ display: "flex", flexWrap: "wrap", gap: 8 }}>
87+
{availableChannels.map((name) => (
88+
<label key={name}
89+
style={{ display: "flex", gap: 4, alignItems: "center" }}>
90+
<input data-testid={`side-channel-train-${name}`}
91+
type="checkbox"
92+
checked={selectedTrain.has(name)}
93+
onChange={(e) =>
94+
setTrainChannel(name, e.target.checked)} />
95+
{name}
96+
</label>
97+
))}
98+
</div>
99+
)}
100+
</section>
101+
72102
<section style={section}>
73103
{Object.entries(draft.families).map(([name, family]) => {
74104
const present = family.columns.filter((c) => available.has(c));
@@ -258,6 +288,21 @@ families=${enabledFamilies.join(",") || "none"}`}
258288
},
259289
});
260290
}
291+
292+
function setTrainChannel(name: string, checked: boolean) {
293+
const selected = new Set(selectedTrainChannels);
294+
if (checked) {
295+
selected.add(name);
296+
} else {
297+
selected.delete(name);
298+
}
299+
const ordered = [
300+
...availableChannels.filter((c) => selected.has(c)),
301+
...selectedTrainChannels.filter((c) =>
302+
selected.has(c) && !available.has(c)),
303+
];
304+
onTrainChannelsChange(ordered);
305+
}
261306
}
262307

263308
function renderPlatform(platform: Record<string, string>): string {

vbgui/tests/Sidebar.test.tsx

Lines changed: 27 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -210,8 +210,10 @@ describe("SideChannelsTab", () => {
210210
const onApply = vi.fn();
211211
render(<SideChannelsTab sideChannels={INITIAL_SPEC.side_channels}
212212
availableChannels={["doc_ids", "token_ids"]}
213+
selectedTrainChannels={[]}
213214
gotchas={[]}
214-
onApply={onApply} />);
215+
onApply={onApply}
216+
onTrainChannelsChange={() => {}} />);
215217
fireEvent.change(screen.getByTestId("side-channel-family-platform-mode"),
216218
{ target: { value: "require" } });
217219
fireEvent.change(screen.getByTestId("side-channel-family-platform-dropout"),
@@ -228,8 +230,10 @@ describe("SideChannelsTab", () => {
228230
it("renders inference and platform preview controls", () => {
229231
render(<SideChannelsTab sideChannels={INITIAL_SPEC.side_channels}
230232
availableChannels={["platform_ids"]}
233+
selectedTrainChannels={[]}
231234
gotchas={[]}
232-
onApply={() => {}} />);
235+
onApply={() => {}}
236+
onTrainChannelsChange={() => {}} />);
233237
fireEvent.change(screen.getByTestId("side-channel-inference-source"),
234238
{ target: { value: "parse_if_possible" } });
235239
fireEvent.change(screen.getByTestId("side-channel-platform-os"),
@@ -243,16 +247,34 @@ describe("SideChannelsTab", () => {
243247
it("surfaces required-family contract probe errors", () => {
244248
render(<SideChannelsTab sideChannels={INITIAL_SPEC.side_channels}
245249
availableChannels={[]}
250+
selectedTrainChannels={[]}
246251
gotchas={[{
247252
id: "side_channel_required_platform",
248253
severity: "error",
249254
message: "required side-channel family 'platform'",
250255
}]}
251-
onApply={() => {}} />);
256+
onApply={() => {}}
257+
onTrainChannelsChange={() => {}} />);
252258
expect(screen.getByTestId(
253259
"side-channel-probe-error-side_channel_required_platform",
254260
).textContent).toContain("platform");
255261
});
262+
263+
it("selects train side-channel inputs inside the side-channel tab", () => {
264+
const onTrainChannelsChange = vi.fn();
265+
render(<SideChannelsTab sideChannels={INITIAL_SPEC.side_channels}
266+
availableChannels={["doc_ids", "token_ids"]}
267+
selectedTrainChannels={["doc_ids"]}
268+
gotchas={[]}
269+
onApply={() => {}}
270+
onTrainChannelsChange={onTrainChannelsChange} />);
271+
expect(screen.getByTestId("side-channel-train-doc_ids"))
272+
.toHaveProperty("checked", true);
273+
fireEvent.click(screen.getByTestId("side-channel-train-token_ids"));
274+
expect(onTrainChannelsChange).toHaveBeenCalledWith([
275+
"doc_ids", "token_ids",
276+
]);
277+
});
256278
});
257279

258280
describe("Sidebar", () => {
@@ -261,12 +283,14 @@ describe("Sidebar", () => {
261283
rewriters: INITIAL_SPEC.rewriters,
262284
sideChannels: INITIAL_SPEC.side_channels,
263285
availableSideChannels: ["doc_ids", "token_ids"],
286+
selectedTrainSideChannels: [],
264287
sharding: INITIAL_SPEC.sharding,
265288
gotchas: INITIAL_SPEC.gotchas, proposals: [],
266289
onLossApply: () => {}, onOptimApply: () => {},
267290
onRewriterAdd: () => {}, onRewriterRemove: () => {},
268291
onRewriterReorder: () => {},
269292
onSideChannelsApply: () => {},
293+
onTrainSideChannelsChange: () => {},
270294
onShardingChange: () => {}, onShardingAccept: () => {},
271295
};
272296

vbgui/tests/TopBar.test.tsx

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

134+
it("does not render legacy train side-channel checkboxes", () => {
135+
render(<TopBar {...defaultTopProps()} />);
136+
fireEvent.click(screen.getByTestId("run-pipeline-toggle"));
137+
expect(screen.queryByTestId("train-side-channel-doc_ids")).toBeNull();
138+
expect(screen.queryByTestId("train-side-channel-token_ids")).toBeNull();
139+
});
140+
134141
it("topology selector emits onTopologyChange", () => {
135142
const onTopologyChange = vi.fn();
136143
render(<TopBar {...defaultTopProps({ onTopologyChange })} />);

0 commit comments

Comments
 (0)