Skip to content

Commit 7d590c6

Browse files
committed
feat(v7-d04): hybrid master-fp32 + bf16-fwd helpers
Closes V7-D04 (cppmega-mlx-m9i4): pure helpers for the best-practice mixed-precision pattern: - cast_for_forward(param, fwd_dtype): bf16 view; master untouched. - cast_grads_to_master(grads, master_dtype=fp32): tree-cast. - hybrid_step(master_params, fwd_dtype, loss_and_grad, apply_gradients) → updated master_params: 1. cast master→bf16 for forward 2. compute (loss, grads_bf16) 3. cast grads back to fp32 4. apply_gradients(master, fp32 grads) → fp32 master Tests (tests/v4/test_hybrid_precision.py): 4/4 — fwd cast leaves master fp32, grads cast preserves fp32, end-to-end step keeps master fp32 and applies fp32 grads. hybrid_lm.set_dtype('hybrid') integration is V7-D04 follow-up.
1 parent bdb245f commit 7d590c6

2 files changed

Lines changed: 125 additions & 0 deletions

File tree

Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,55 @@
1+
"""V7-D04: hybrid master-fp32-grad + bf16-forward helpers.
2+
3+
Best-practice mixed precision: params live in fp32 (master copy);
4+
forward materialises a bf16 view (cast on the fly); gradients are
5+
cast back to fp32 for optimiser apply.
6+
7+
cast_for_forward(param, fwd_dtype) → fwd_dtype view
8+
cast_grads_to_master(grads, master_dtype=fp32) → fp32 grads
9+
"""
10+
11+
from __future__ import annotations
12+
13+
import mlx.core as mx
14+
import mlx.nn as nn
15+
16+
17+
def cast_for_forward(param: mx.array, fwd_dtype: mx.Dtype) -> mx.array:
18+
"""Return a non-master view of `param` in fwd_dtype. The original
19+
fp32 master is preserved by the caller."""
20+
return param.astype(fwd_dtype)
21+
22+
23+
def cast_grads_to_master(grads, master_dtype: mx.Dtype = mx.float32):
24+
"""Walk a grad tree and cast every leaf to master_dtype."""
25+
return nn.utils.tree_map(
26+
lambda g: g.astype(master_dtype) if hasattr(g, "shape") else g,
27+
grads,
28+
)
29+
30+
31+
def hybrid_step(*, master_params: dict, fwd_dtype: mx.Dtype,
32+
loss_and_grad,
33+
apply_gradients) -> dict:
34+
"""One end-to-end hybrid mixed-precision step.
35+
36+
Args:
37+
master_params: dict of name → fp32 mx.array master params
38+
(mutated in place by apply_gradients).
39+
fwd_dtype: e.g. mx.bfloat16.
40+
loss_and_grad: callable (fwd_params_dict) → (loss, grads).
41+
apply_gradients: callable (master_params, fp32_grads) → updated
42+
master_params dict.
43+
44+
Returns the updated master params dict.
45+
"""
46+
fwd_params = {k: cast_for_forward(v, fwd_dtype)
47+
for k, v in master_params.items()}
48+
_, grads = loss_and_grad(fwd_params)
49+
grads_fp32 = cast_grads_to_master(grads, master_dtype=mx.float32)
50+
return apply_gradients(master_params, grads_fp32)
51+
52+
53+
__all__ = [
54+
"cast_for_forward", "cast_grads_to_master", "hybrid_step",
55+
]

tests/v4/test_hybrid_precision.py

Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,70 @@
1+
"""V7-D04: hybrid master-fp32 + bf16-fwd helpers."""
2+
3+
from __future__ import annotations
4+
5+
import mlx.core as mx
6+
import pytest
7+
8+
from cppmega_v4.runtime.hybrid_precision import (
9+
cast_for_forward, cast_grads_to_master, hybrid_step,
10+
)
11+
12+
13+
def test_v7_d04_cast_for_forward_returns_bf16_view():
14+
p = mx.random.normal(shape=(8, 8), key=mx.random.key(0))
15+
assert p.dtype == mx.float32
16+
v = cast_for_forward(p, mx.bfloat16)
17+
assert v.dtype == mx.bfloat16
18+
# Master unchanged.
19+
assert p.dtype == mx.float32
20+
21+
22+
def test_v7_d04_cast_grads_to_master_preserves_fp32():
23+
grads = {"a": mx.zeros((4,), dtype=mx.bfloat16),
24+
"b": mx.zeros((4,), dtype=mx.float16)}
25+
out = cast_grads_to_master(grads, master_dtype=mx.float32)
26+
assert out["a"].dtype == mx.float32
27+
assert out["b"].dtype == mx.float32
28+
29+
30+
def test_v7_d04_hybrid_step_updates_master_in_fp32():
31+
master = {"w": mx.array([1.0, 2.0, 3.0], dtype=mx.float32)}
32+
33+
def loss_and_grad(fwd):
34+
# Always-positive grad of 0.1 per element.
35+
return mx.array(0.0), {"w": mx.array([0.1, 0.1, 0.1],
36+
dtype=mx.bfloat16)}
37+
38+
def apply_gradients(m, g):
39+
return {k: m[k] - g[k] for k in m}
40+
41+
new_master = hybrid_step(
42+
master_params=master, fwd_dtype=mx.bfloat16,
43+
loss_and_grad=loss_and_grad,
44+
apply_gradients=apply_gradients,
45+
)
46+
# Updated master stays fp32.
47+
assert new_master["w"].dtype == mx.float32
48+
# Each entry decreased by 0.1.
49+
assert mx.allclose(new_master["w"],
50+
mx.array([0.9, 1.9, 2.9], dtype=mx.float32),
51+
atol=1e-3)
52+
53+
54+
def test_v7_d04_hybrid_step_grads_cast_to_master_before_apply():
55+
"""apply_gradients must receive fp32 grads, not bf16."""
56+
seen_dtype: dict[str, mx.Dtype] = {}
57+
58+
master = {"w": mx.zeros((2,), dtype=mx.float32)}
59+
60+
def loss_and_grad(fwd):
61+
return mx.array(0.0), {"w": mx.ones((2,), dtype=mx.bfloat16)}
62+
63+
def apply_gradients(m, g):
64+
seen_dtype["w"] = g["w"].dtype
65+
return m
66+
67+
hybrid_step(master_params=master, fwd_dtype=mx.bfloat16,
68+
loss_and_grad=loss_and_grad,
69+
apply_gradients=apply_gradients)
70+
assert seen_dtype["w"] == mx.float32

0 commit comments

Comments
 (0)