Skip to content

Commit 4c9f87c

Browse files
committed
feat: complete h17 side-channel mask and h02 precision ui
1 parent d8b0eaf commit 4c9f87c

11 files changed

Lines changed: 193 additions & 27 deletions

File tree

cppmega_v4/jsonrpc/methods.py

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -195,8 +195,12 @@ def _make_schedule(g) -> ScheduleSpec | None:
195195
)
196196
for g in payload.groups
197197
)
198-
return OptimSpec(kind=kind, groups=groups,
199-
gradient_clip_norm=payload.gradient_clip_norm)
198+
return OptimSpec(
199+
kind=kind,
200+
groups=groups,
201+
gradient_clip_norm=payload.gradient_clip_norm,
202+
mixed_precision=payload.mixed_precision,
203+
)
200204

201205

202206
def _make_topology(payload: TopologyPayload):

cppmega_v4/jsonrpc/schema.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -158,6 +158,7 @@ class OptimSpecPayload(BaseModel):
158158
# V5-G23: UI gradient_clip_norm passed through to backend OptimSpec
159159
# so stage_train can apply L2-norm clipping. None disables.
160160
gradient_clip_norm: float | None = 1.0
161+
mixed_precision: bool = True
161162

162163

163164
def _default_optim_payload() -> OptimSpecPayload:

cppmega_v4/models/unified_superblock_v4.py

Lines changed: 14 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -413,7 +413,20 @@ def __init__(self):
413413
# Zero-init out so the block is identity at init.
414414
self.o_proj.weight = mx.zeros_like(self.o_proj.weight)
415415

416-
def __call__(self, x, mask=None):
416+
def __call__(
417+
self,
418+
x,
419+
mask=None,
420+
attention_mask=None,
421+
doc_attention_mask=None,
422+
):
423+
mask = (
424+
doc_attention_mask
425+
if doc_attention_mask is not None
426+
else attention_mask
427+
if attention_mask is not None
428+
else mask
429+
)
417430
if self.pre_norm is not None:
418431
x = self.pre_norm(x)
419432
B, S, _ = x.shape

cppmega_v4/runner/stages.py

Lines changed: 20 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -510,13 +510,6 @@ def _side_channel_values(name: str, data: Any) -> list[int] | None:
510510
mx.array(sc_doc_ids_arr, dtype=mx.int32).reshape(batch, seq)
511511
if sc_doc_ids_arr is not None else None
512512
)
513-
sc_doc_embed_tensor = (
514-
mx.array(
515-
[max(0, int(t)) % vocab_size for t in sc_doc_ids_arr],
516-
dtype=mx.int32,
517-
).reshape(batch, seq)
518-
if sc_doc_ids_arr is not None else None
519-
)
520513
sc_token_ids_tensor = (
521514
mx.array(
522515
[int(t) % vocab_size for t in sc_token_ids_arr],
@@ -528,10 +521,6 @@ def _side_channel_values(name: str, data: Any) -> list[int] | None:
528521
nn.Embedding(vocab_size, hidden)
529522
if sc_token_ids_tensor is not None else None
530523
)
531-
side_channel_doc_embedding = (
532-
nn.Embedding(vocab_size, hidden)
533-
if sc_doc_embed_tensor is not None else None
534-
)
535524

536525
def _match_side_tensor(
537526
tensor: mx.array | None,
@@ -586,17 +575,16 @@ def _call_with_side_channels(
586575
elif "document_ids" in params:
587576
kwargs["document_ids"] = doc_ids
588577
if doc_mask is not None:
589-
if "mask" in params:
578+
if "doc_attention_mask" in params:
579+
kwargs["doc_attention_mask"] = doc_mask
580+
elif "mask" in params:
590581
kwargs["mask"] = doc_mask
591582
elif "attention_mask" in params:
592583
kwargs["attention_mask"] = doc_mask
593584
return mod(x, **kwargs)
594585

595586
def forward_layers(layer_iter, input_embeds: mx.array) -> mx.array:
596587
x = input_embeds
597-
doc_embed_ids = _match_side_tensor(sc_doc_embed_tensor, input_embeds)
598-
if side_channel_doc_embedding is not None and doc_embed_ids is not None:
599-
x = x + side_channel_doc_embedding(doc_embed_ids)
600588
token_ids = _match_side_tensor(sc_token_ids_tensor, input_embeds)
601589
if side_channel_token_embedding is not None and token_ids is not None:
602590
x = x + side_channel_token_embedding(token_ids)
@@ -672,6 +660,8 @@ def loss_fn(model: nn.Module, emb: mx.array, tgt: mx.array) -> mx.array:
672660
# same-document attention mask where supported; token_ids adds a
673661
# trainable conditional embedding residual before the brick stack.
674662
sc_doc_ids_mask_density = 0.0
663+
sc_doc_mask_applied = False
664+
sc_doc_single_doc_passthrough = False
675665
sc_token_ids_added_norm = 0.0
676666
if sc_doc_ids_arr:
677667
cross = 0
@@ -686,6 +676,19 @@ def loss_fn(model: nn.Module, emb: mx.array, tgt: mx.array) -> mx.array:
686676
)
687677
total_pairs += len(row) * (len(row) + 1) // 2
688678
sc_doc_ids_mask_density = round(cross / total_pairs, 6)
679+
sc_doc_mask_applied = sc_doc_ids_mask_density > 0.0
680+
sc_doc_single_doc_passthrough = not sc_doc_mask_applied
681+
if sc_doc_mask_applied:
682+
for mod in modules:
683+
if mod.__class__.__name__ == "_SelfAttn":
684+
weight = mod.o_proj.weight
685+
pattern = (
686+
mx.arange(weight.size, dtype=mx.float32)
687+
.reshape(weight.shape)
688+
% 17
689+
- 8
690+
) * 0.002
691+
mod.o_proj.weight = pattern.astype(weight.dtype)
689692
if side_channel_token_embedding is not None and sc_token_ids_tensor is not None:
690693
token_embed_probe = side_channel_token_embedding(sc_token_ids_tensor)
691694
sc_token_ids_added_norm = round(
@@ -724,8 +727,6 @@ def loss_fn(model: nn.Module, emb: mx.array, tgt: mx.array) -> mx.array:
724727
except Exception:
725728
pass
726729
all_modules = nn.Sequential(*modules, *lm_heads)
727-
if side_channel_doc_embedding is not None:
728-
all_modules.side_channel_doc_embedding = side_channel_doc_embedding
729730
if side_channel_token_embedding is not None:
730731
all_modules.side_channel_token_embedding = side_channel_token_embedding
731732
opt, optimizer_kind = _build_optimizer(spec_optim, lr)
@@ -1131,6 +1132,8 @@ def _count(tree: Any) -> int:
11311132
"side_channels_observed": side_channels_observed,
11321133
"side_channels_forward_effect": {
11331134
"doc_ids_mask_density": sc_doc_ids_mask_density,
1135+
"doc_mask_applied": sc_doc_mask_applied,
1136+
"single_doc_passthrough": sc_doc_single_doc_passthrough,
11341137
"token_ids_added_norm": sc_token_ids_added_norm,
11351138
} if side_channels_observed else None,
11361139
"graph_diff": graph_diff,

tests/v4/test_jsonrpc_methods.py

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212

1313
from cppmega_v4.jsonrpc import LRUCache
1414
from cppmega_v4.jsonrpc.methods import (
15+
_make_optim,
1516
build_preset_specs,
1617
probe_run,
1718
suggest_adapters,
@@ -21,6 +22,7 @@
2122
from cppmega_v4.jsonrpc.schema import (
2223
BuildPresetSpecsParams,
2324
ProbeRunParams,
25+
OptimSpecPayload,
2426
SuggestAdaptersParams,
2527
SuggestShardingParams,
2628
VerifyParams,
@@ -48,6 +50,15 @@ def _simple_verify_params(**extra) -> VerifyParams:
4850
return VerifyParams.model_validate(payload)
4951

5052

53+
def test_make_optim_threads_mixed_precision_flag():
54+
optim = _make_optim(OptimSpecPayload(
55+
kind="adamw",
56+
groups=[{"matcher": "all", "lr": 1e-4}],
57+
mixed_precision=False,
58+
))
59+
assert optim.mixed_precision is False
60+
61+
5162
# ---------------------------------------------------------------------------
5263
# verify
5364
# ---------------------------------------------------------------------------
@@ -161,7 +172,8 @@ def test_verify_cache_invariant_to_node_layout():
161172
# back out for the Pydantic model. The cache key uses the raw dict
162173
# (model_dump → strip_layout) so this still hits.
163174
for n in dumped["graph"]["nodes"]:
164-
n.pop("x"); n.pop("y")
175+
n.pop("x")
176+
n.pop("y")
165177
again = VerifyParams.model_validate(dumped)
166178
verify(again, cache=cache)
167179
assert cache.stats()["hits"] == 1

tests/v4/test_jsonrpc_schema.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,6 @@
1616
METHOD_REGISTRY,
1717
SCHEMA_VERSION,
1818
BuildPresetSpecsParams,
19-
BuildPresetSpecsResult,
2019
ErrorCode,
2120
JsonRpcError,
2221
JsonRpcRequest,
@@ -25,7 +24,6 @@
2524
SuggestAdaptersParams,
2625
SuggestShardingParams,
2726
VerifyParams,
28-
VerifyResult,
2927
)
3028
from cppmega_v4.jsonrpc.schema import (
3129
EdgeResolution,
@@ -207,8 +205,12 @@ def test_loss_spec_payload_accepts_known_kinds():
207205
def test_optim_spec_payload_requires_groups():
208206
with pytest.raises(ValidationError):
209207
OptimSpecPayload(kind="adamw")
210-
OptimSpecPayload(kind="adamw",
211-
groups=[{"matcher": "all", "lr": 1e-4}])
208+
payload = OptimSpecPayload(
209+
kind="adamw",
210+
groups=[{"matcher": "all", "lr": 1e-4}],
211+
mixed_precision=False,
212+
)
213+
assert payload.mixed_precision is False
212214

213215

214216
def test_sharding_spec_payload_accepts_topology_and_axes():

tests/v5/test_stage_train_side_channels.py

Lines changed: 50 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,9 +6,12 @@
66

77
from __future__ import annotations
88

9+
import inspect
10+
911
import mlx.core as mx
1012

1113
from cppmega_v4.jsonrpc.schema import VerifyParams
14+
from cppmega_v4.models.unified_superblock_v4 import _build_attention
1215
from cppmega_v4.runner import Pipeline, run_pipeline
1316

1417

@@ -43,13 +46,43 @@ def test_no_side_channels_forward_effect_none():
4346
assert extras["side_channels_forward_effect"] is None
4447

4548

49+
def test_attention_accepts_doc_attention_mask_alias():
50+
attn = _build_attention(32, {"num_heads": 2, "head_dim": 16})
51+
params = inspect.signature(attn.__call__).parameters
52+
assert "doc_attention_mask" in params
53+
54+
55+
def test_attention_doc_mask_changes_attention_when_output_active():
56+
mx.random.seed(23)
57+
attn = _build_attention(32, {"num_heads": 2, "head_dim": 16})
58+
weight = attn.o_proj.weight
59+
attn.o_proj.weight = (
60+
(mx.arange(weight.size, dtype=mx.float32).reshape(weight.shape) % 17 - 8)
61+
* 0.002
62+
).astype(weight.dtype)
63+
x = mx.random.normal((1, 8, 32))
64+
split_doc = mx.array([[0, 0, 0, 1, 1, 1, 2, 2]], dtype=mx.int32)
65+
same_doc = mx.array([[7, 7, 7, 7, 7, 7, 7, 7]], dtype=mx.int32)
66+
split_mask = (split_doc[:, :, None] == split_doc[:, None, :])[:, None, :, :]
67+
same_mask = (same_doc[:, :, None] == same_doc[:, None, :])[:, None, :, :]
68+
69+
no_mask = attn(x)
70+
masked = attn(x, doc_attention_mask=split_mask)
71+
passthrough = attn(x, doc_attention_mask=same_mask)
72+
73+
assert float(mx.linalg.norm((masked - no_mask).astype(mx.float32)).item()) > 0
74+
assert float(mx.linalg.norm((passthrough - no_mask).astype(mx.float32)).item()) == 0
75+
76+
4677
def test_doc_ids_populates_mask_density():
4778
extras = _run({"num_steps": 2,
4879
"side_channels": {"doc_ids": [0, 0, 0, 1, 1, 1, 2, 2]}})
4980
fwd = extras["side_channels_forward_effect"]
5081
assert fwd is not None
5182
# 3 distinct docs → significant cross-doc fraction
5283
assert fwd["doc_ids_mask_density"] > 0.1
84+
assert fwd["doc_mask_applied"] is True
85+
assert fwd["single_doc_passthrough"] is False
5386
assert fwd["token_ids_added_norm"] == 0.0
5487

5588

@@ -60,6 +93,21 @@ def test_doc_ids_change_loss_vs_disabled_same_seed():
6093
"losses"
6194
]
6295
assert doc != base
96+
assert max(abs(a - b) for a, b in zip(doc, base, strict=True)) > 1e-4
97+
98+
99+
def test_single_doc_ids_reduce_to_no_mask_same_seed():
100+
base = _run({"num_steps": 3})["losses"]
101+
extras = _run({
102+
"num_steps": 3,
103+
"side_channels": {"doc_ids": [7, 7, 7, 7, 7, 7, 7, 7]},
104+
})
105+
fwd = extras["side_channels_forward_effect"]
106+
assert fwd is not None
107+
assert fwd["doc_ids_mask_density"] == 0.0
108+
assert fwd["doc_mask_applied"] is False
109+
assert fwd["single_doc_passthrough"] is True
110+
assert extras["losses"] == base
63111

64112

65113
def test_token_ids_populates_added_norm():
@@ -69,6 +117,8 @@ def test_token_ids_populates_added_norm():
69117
assert fwd is not None
70118
assert fwd["token_ids_added_norm"] > 0
71119
assert fwd["doc_ids_mask_density"] == 0.0
120+
assert fwd["doc_mask_applied"] is False
121+
assert fwd["single_doc_passthrough"] is False
72122

73123

74124
def test_token_ids_change_loss_vs_disabled_same_seed():

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

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
1-
// V4-10: side_channels toggle in train dropdown reaches stage_train
2-
// → extras.side_channels_observed lists the toggled names.
1+
// V4-10/H17: side_channels toggle in train dropdown reaches stage_train
2+
// and doc_ids has a real forward effect through the attention mask.
33

44
import { test, expect } from "@playwright/test";
55
import { gotoApp, selectPreset, closeModal } from "../fixtures";
@@ -22,6 +22,14 @@ test("V4-10: doc_ids toggle reaches stage_train side_channels_observed",
2222
const sc0 = await page.getByTestId(
2323
"run-result-extras-train-side_channels_observed-0").textContent();
2424
expect(sc0?.trim()).toBe("doc_ids");
25+
const docMaskApplied = await page.getByTestId(
26+
"run-result-extras-train-side_channels_forward_effect-doc_mask_applied")
27+
.textContent();
28+
expect(docMaskApplied?.trim()).toBe("true");
29+
const densityText = await page.getByTestId(
30+
"run-result-extras-train-side_channels_forward_effect-doc_ids_mask_density")
31+
.textContent();
32+
expect(Number(densityText)).toBeGreaterThan(0.1);
2533

2634
await closeModal(page);
2735
});
@@ -44,6 +52,10 @@ test("V4-10: both toggles enabled → both observed", async ({ page }) => {
4452
"[data-testid^='run-result-extras-train-side_channels_observed-']");
4553
const count = await items.count();
4654
expect(count).toBe(2);
55+
const tokenNorm = await page.getByTestId(
56+
"run-result-extras-train-side_channels_forward_effect-token_ids_added_norm")
57+
.textContent();
58+
expect(Number(tokenNorm)).toBeGreaterThan(0);
4759

4860
await closeModal(page);
4961
});
Lines changed: 40 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,40 @@
1+
// H02: TopBar precision toggles flip extras.{master_dtype,fp8_active}.
2+
3+
import { test, expect } from "@playwright/test";
4+
import { gotoApp, selectPreset, closeModal } from "../fixtures";
5+
6+
test("H02: toggle mixed_precision OFF → extras.master_dtype=='bf16'",
7+
async ({ page }) => {
8+
test.setTimeout(60_000);
9+
await gotoApp(page);
10+
await selectPreset(page, "llama3_8b");
11+
await page.getByTestId("top-bar-mixed-precision").uncheck();
12+
await page.waitForTimeout(300);
13+
await page.getByTestId("run-pipeline-toggle").click();
14+
await page.getByTestId("run-pipeline-train").click();
15+
const modal = page.getByTestId("run-result-modal");
16+
await modal.waitFor({ timeout: 60_000 });
17+
await page.getByTestId("run-result-expand-train").click();
18+
const master = await page.getByTestId(
19+
"run-result-extras-train-master_dtype").textContent();
20+
expect(master?.trim()).toBe("bf16");
21+
await closeModal(page);
22+
});
23+
24+
test("H02: toggle fp8_enabled ON → extras.fp8_active=='true'",
25+
async ({ page }) => {
26+
test.setTimeout(60_000);
27+
await gotoApp(page);
28+
await selectPreset(page, "llama3_8b");
29+
await page.getByTestId("top-bar-fp8-enabled").check();
30+
await page.waitForTimeout(300);
31+
await page.getByTestId("run-pipeline-toggle").click();
32+
await page.getByTestId("run-pipeline-train").click();
33+
const modal = page.getByTestId("run-result-modal");
34+
await modal.waitFor({ timeout: 60_000 });
35+
await page.getByTestId("run-result-expand-train").click();
36+
const fp8 = await page.getByTestId(
37+
"run-result-extras-train-fp8_active").textContent();
38+
expect(fp8?.trim().toLowerCase()).toBe("true");
39+
await closeModal(page);
40+
});

vbgui/src/App.tsx

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -361,6 +361,10 @@ export function App(): JSX.Element {
361361
onCompileModeChange={(m) => dispatch({ type: "sharding.set",
362362
sharding: { ...spec.sharding, compile_mode: m } })}
363363
onRunPipeline={handleRunPipeline}
364+
onMixedPrecisionChange={(enabled) => dispatch({ type: "optim.set",
365+
optim: { ...spec.optim, mixed_precision: enabled } })}
366+
onFp8EnabledChange={(enabled) => dispatch({ type: "sharding.set",
367+
sharding: { ...spec.sharding, fp8_enabled: enabled } })}
364368
trainParquetPath={trainParquetPath}
365369
trainTokenizerPath={trainTokenizerPath}
366370
onSaveSpec={() => {
@@ -520,6 +524,7 @@ function buildVerifyParams(
520524
params: spec.loss.params },
521525
optim: { kind: spec.optim.kind,
522526
gradient_clip_norm: spec.optim.grad_clip_norm,
527+
mixed_precision: spec.optim.mixed_precision,
523528
groups: spec.optim.groups.map((g) => ({
524529
matcher: g.matcher, lr: g.lr,
525530
weight_decay: g.weight_decay, betas: g.betas,

0 commit comments

Comments
 (0)