Skip to content

Commit dd50d39

Browse files
author
Codex
committed
Path C: sync sparse MLA static ABI inputs into banks
1 parent 1bb242a commit dd50d39

2 files changed

Lines changed: 130 additions & 1 deletion

File tree

cppmega_mlx/models/hybrid_lm.py

Lines changed: 76 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1272,6 +1272,58 @@ def path_c_fused_in_region_parameter_bank_aliases(
12721272
break
12731273
return out
12741274

1275+
def _path_c_static_real_abi_input_values(
1276+
self,
1277+
abi_map: Mapping[str, Any],
1278+
) -> dict[str, mx.array]:
1279+
"""Return small non-trainable real-ABI inputs required by Path C.
1280+
1281+
The fused train-block ABI contains a few real inputs that are neither
1282+
trainable parameters nor per-batch runtime tensors. For Sparse-MLA
1283+
these are the softmax scale and optional attention sinks. They must be
1284+
written into model-owned banks along with parameter values; leaving the
1285+
zero-initialized bank slots in place makes the forward/backward kernel
1286+
run with ``sm_scale=0`` and kills q/kv gradients.
1287+
"""
1288+
1289+
attention = self.config.attention_config("dsa")
1290+
sm_scale = float(attention.q_head_dim) ** -0.5
1291+
values: dict[str, mx.array] = {}
1292+
for logical_name, info in abi_map.items():
1293+
if not isinstance(info, Mapping):
1294+
continue
1295+
shape = tuple(
1296+
int(dim)
1297+
for dim in tuple(info.get("logical_shape", info.get("shape", ())))
1298+
)
1299+
if str(logical_name).endswith("sparse_mla_sm_scale"):
1300+
values[str(logical_name)] = mx.array(
1301+
[sm_scale], dtype=mx.float32
1302+
).reshape(shape or (1,))
1303+
elif str(logical_name).endswith("sparse_mla_sinks"):
1304+
values[str(logical_name)] = mx.zeros(shape, dtype=mx.float32)
1305+
elif str(logical_name).endswith("sparse_mla_has_sinks"):
1306+
values[str(logical_name)] = mx.zeros(shape, dtype=mx.int32)
1307+
return values
1308+
1309+
def _sync_path_c_static_real_abi_inputs_into_bank(
1310+
self,
1311+
*,
1312+
abi_map: Mapping[str, Any],
1313+
buffers: Mapping[str, Any],
1314+
) -> tuple[list[str], list[tuple[str, str]]]:
1315+
synced: list[str] = []
1316+
skipped: list[tuple[str, str]] = []
1317+
for logical_name, value in sorted(
1318+
self._path_c_static_real_abi_input_values(abi_map).items()
1319+
):
1320+
try:
1321+
write_into_bank_slot(abi_map, buffers, logical_name, value)
1322+
synced.append(logical_name)
1323+
except Exception as exc:
1324+
skipped.append((logical_name, f"{type(exc).__name__}: {exc}"))
1325+
return synced, skipped
1326+
12751327
def _path_c_lookup_parameter_holder(
12761328
self,
12771329
parameter_name: str,
@@ -1378,6 +1430,13 @@ def sync_path_c_in_region_parameters_into_bank(
13781430
)
13791431
synced: list[str] = []
13801432
skipped: list[tuple[str, str]] = []
1433+
static_synced, static_skipped = (
1434+
self._sync_path_c_static_real_abi_inputs_into_bank(
1435+
abi_map=abi_map,
1436+
buffers=buffers,
1437+
)
1438+
)
1439+
skipped.extend(static_skipped)
13811440
for parameter_name, info in sorted(aliases.items()):
13821441
tensor = self._path_c_get_parameter_tensor(parameter_name)
13831442
if tensor is None:
@@ -1414,6 +1473,11 @@ def sync_path_c_in_region_parameters_into_bank(
14141473
{"parameter_name": name, "reason": reason}
14151474
for name, reason in skipped
14161475
],
1476+
"static_real_abi_inputs_synced": static_synced,
1477+
"static_real_abi_inputs_skipped": [
1478+
{"logical_name": name, "reason": reason}
1479+
for name, reason in static_skipped
1480+
],
14171481
}
14181482

14191483
def bind_path_c_in_region_parameter_views_into_bank(
@@ -1486,6 +1550,13 @@ def bind_path_c_in_region_parameter_views_into_bank(
14861550
)
14871551
bound: list[str] = []
14881552
skipped: list[tuple[str, str]] = []
1553+
static_synced, static_skipped = (
1554+
self._sync_path_c_static_real_abi_inputs_into_bank(
1555+
abi_map=abi_map,
1556+
buffers=buffers,
1557+
)
1558+
)
1559+
skipped.extend(static_skipped)
14891560
for parameter_name, info in sorted(aliases.items()):
14901561
tensor = self._path_c_get_parameter_tensor(parameter_name)
14911562
if tensor is None:
@@ -1524,6 +1595,11 @@ def bind_path_c_in_region_parameter_views_into_bank(
15241595
"in_region_parameter_count": len(bound),
15251596
"in_region_parameter_names": tuple(sorted(bound)),
15261597
"in_region_parameter_bank_aliases": dict(aliases),
1598+
"static_real_abi_inputs_synced": static_synced,
1599+
"static_real_abi_inputs_skipped": [
1600+
{"logical_name": name, "reason": reason}
1601+
for name, reason in static_skipped
1602+
],
15271603
}
15281604

15291605
def path_c_fused_first_in_region_layer_index(

tests/test_hybrid_lm_path_c_physical_abi_bank_owner.py

Lines changed: 54 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,8 @@
99

1010
from __future__ import annotations
1111

12+
from typing import Any
13+
1214
import mlx.core as mx
1315
import pytest
1416

@@ -188,7 +190,7 @@ def test_bind_path_c_in_region_parameter_views_into_bank_replaces_attributes(
188190
bank = bank_owner.buffers[sample_info["bank"]]
189191
slot = bank[sample_info["offset"] : sample_info["offset"] + sample_info["size"]]
190192
parts = sample_small.split(".")
191-
holder: object = model
193+
holder: Any = model
192194
for part in parts[:-1]:
193195
holder = holder[int(part)] if part.isdigit() else getattr(holder, part)
194196
param_value = getattr(holder, parts[-1])
@@ -226,3 +228,54 @@ def test_sync_path_c_in_region_parameters_into_bank_updates_bank_slots(
226228
bank = bank_owner.buffers[bank_name]
227229
slot = bank[info["offset"] : info["offset"] + info["size"]]
228230
assert mx.allclose(slot, sentinel).item() is True
231+
232+
233+
def test_bind_path_c_syncs_sparse_mla_static_real_abi_inputs() -> None:
234+
model = build_local_gb10_quarter_tiny_smoke_model()
235+
sequence_length = 127
236+
bank_owner = model.make_path_c_physical_abi_bank_owner(
237+
sequence_length=sequence_length,
238+
)
239+
assert bank_owner is not None
240+
report = model.bind_path_c_in_region_parameter_views_into_bank(
241+
bank_owner,
242+
sequence_length=sequence_length,
243+
)
244+
prim_func = model.path_c_fused_train_block_prim_func(
245+
sequence_length=sequence_length,
246+
)
247+
assert prim_func is not None
248+
abi_map = dict(
249+
getattr(prim_func, "_cppmega_path_c_physical_buffer_abi_map", {})
250+
or {}
251+
)
252+
253+
def bank_slot(logical_name: str) -> mx.array:
254+
info = abi_map[logical_name]
255+
bank = bank_owner.buffers[str(info["bank"])]
256+
offset = int(info["offset"])
257+
size = int(info["size"])
258+
return bank[offset : offset + size]
259+
260+
sm_scale_name = next(
261+
name for name in abi_map if name.endswith("sparse_mla_sm_scale")
262+
)
263+
sinks_name = next(
264+
name for name in abi_map if name.endswith("sparse_mla_sinks")
265+
)
266+
has_sinks_name = next(
267+
name for name in abi_map if name.endswith("sparse_mla_has_sinks")
268+
)
269+
sm_scale = bank_slot(sm_scale_name)
270+
sinks = bank_slot(sinks_name)
271+
has_sinks = bank_slot(has_sinks_name)
272+
mx.eval(sm_scale, sinks, has_sinks)
273+
274+
assert report["static_real_abi_inputs_synced"] == [
275+
has_sinks_name,
276+
sinks_name,
277+
sm_scale_name,
278+
]
279+
assert sm_scale.tolist() == [0.5]
280+
assert sinks.tolist() == [0.0, 0.0, 0.0, 0.0]
281+
assert has_sinks.tolist() == [0]

0 commit comments

Comments
 (0)