Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
160 changes: 149 additions & 11 deletions recml/core/ops/binary_cross_entropy_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -247,20 +247,66 @@ class BCEConfig:
slice of the embedding table, but the forward pass normalizes the loss by
the full vocabulary, so the backward must use this value rather than
``embeddings.shape[0]``.
optimize_large_l: Whether to use the target-free vocabulary pass + sparse
positive-target correction. Defaults to False.
"""

block_v: int
block_n: int = _BLOCK_N
compute_metrics: bool = False
use_pallas: bool = True
global_vocab: int | None = None
optimize_large_l: bool = False

# Sharding specs for VJP backward pass optimization
mesh: jax.sharding.Mesh | None = None
act_spec: jax.sharding.PartitionSpec | None = None
emb_spec: jax.sharding.PartitionSpec | None = None


def _unique_valid_targets(
targets_2d: jt.Int[jt.Array, "N L"],
vocab: int,
) -> tuple[jt.Int[jt.Array, "N L"], jt.Bool[jt.Array, "N L"]]:
"""Returns sorted in-bounds target indices in [0, vocab) and a unique mask."""
sorted_t = jnp.sort(targets_2d, axis=-1)
valid_pos = (sorted_t >= 0) & (sorted_t < vocab)
if targets_2d.shape[-1] > 1:
valid_pos = valid_pos & jnp.concatenate(
[jnp.ones_like(valid_pos[:, :1]), sorted_t[:, 1:] != sorted_t[:, :-1]],
axis=-1,
)
return jnp.where(valid_pos, sorted_t, 0), valid_pos


def _fwd_pos_correction(
activations_2d: jt.Float[jt.Array, "N D"],
embeddings: jt.Float[jt.Array, "V D"],
targets_2d: jt.Int[jt.Array, "N L"],
loss_diff_fn: (
Callable[[jt.Float[jt.Array, "N L"]], jt.Float[jt.Array, "N L"]] | None
) = None,
) -> tuple[
jt.Float[jt.Array, "N"],
jt.Float[jt.Array, "N"],
jt.Float[jt.Array, "N"],
]:
"""Returns (delta_loss_sum, tp_pos, num_pos) for positive targets."""
safe_t, valid_pos = _unique_valid_targets(targets_2d, embeddings.shape[0])
pos_logits = jnp.einsum(
"nd,nld->nl",
activations_2d,
embeddings[safe_t],
preferred_element_type=jnp.float32,
precision=jax.lax.Precision.DEFAULT,
)
delta = -pos_logits if loss_diff_fn is None else loss_diff_fn(pos_logits)
delta_loss_sum = jnp.sum(jnp.where(valid_pos, delta, 0.0), axis=-1)
tp_pos = jnp.sum((pos_logits > 0.0) & valid_pos, axis=-1).astype(jnp.float32)
num_pos = jnp.sum(valid_pos, axis=-1).astype(jnp.float32)
return delta_loss_sum, tp_pos, num_pos


def _bce_fwd_chunk(
activations_2d: jt.Float[jt.Array, "N D"],
embeddings: jt.Float[jt.Array, "V D"],
Expand Down Expand Up @@ -336,6 +382,7 @@ def _bce_fwd_local(

activations_2d = jnp.reshape(activations, (n, hidden))
targets_2d = jnp.reshape(targets, (n, -1))
scan_targets = targets_2d[:, :0] if config.optimize_large_l else targets_2d

if config.compute_metrics:

Expand All @@ -345,7 +392,7 @@ def v_body(
) -> tuple[tuple[jt.Float[jt.Array, "N"], ...], None]:
loss_acc, tp_acc, fp_acc, fn_acc, tn_acc = carry
loss_sum, logits, targets_chunk, valid_mask = _bce_fwd_chunk(
activations_2d, embeddings, targets_2d, j, block_v, vocab
activations_2d, embeddings, scan_targets, j, block_v, vocab
)

predictions_chunk = (logits > 0.0) & valid_mask[None, :]
Expand Down Expand Up @@ -379,6 +426,14 @@ def v_body(
(loss_final, tp_final, fp_final, fn_final, tn_final), _ = jax.lax.scan(
v_body, init, jnp.arange(v_blocks)
)
if config.optimize_large_l:
delta_loss, tp_final, num_pos = _fwd_pos_correction(
activations_2d, embeddings, targets_2d
)
loss_final = loss_final + delta_loss
fp_final = fp_final - tp_final
fn_final = num_pos - tp_final
tn_final = tn_final - fn_final
return (
jnp.reshape(loss_final, (batch, seq_len)),
jnp.reshape(tp_final, (batch, seq_len)),
Expand All @@ -393,12 +448,17 @@ def v_body_no_metrics(
j: jt.Int[jt.Array, ""],
) -> tuple[jt.Float[jt.Array, "N"], None]:
loss_sum, _, _, _ = _bce_fwd_chunk(
activations_2d, embeddings, targets_2d, j, block_v, vocab
activations_2d, embeddings, scan_targets, j, block_v, vocab
)
return loss_acc + loss_sum, None

init = jnp.zeros((n,), dtype=jnp.float32)
loss_final, _ = jax.lax.scan(v_body_no_metrics, init, jnp.arange(v_blocks))
if config.optimize_large_l:
loss_final = (
loss_final
+ _fwd_pos_correction(activations_2d, embeddings, targets_2d)[0]
)
return jnp.reshape(loss_final, (batch, seq_len))


Expand Down Expand Up @@ -541,6 +601,7 @@ def cut_binary_cross_entropy(
act_spec: jax.sharding.PartitionSpec | None = None,
emb_spec: jax.sharding.PartitionSpec | None = None,
use_pallas: bool = True,
optimize_large_l: bool = False,
) -> (
jt.Float[jt.Array, ""]
| tuple[jt.Float[jt.Array, ""], jt.Float[jt.Array, "... B N"]]
Expand Down Expand Up @@ -589,6 +650,8 @@ def cut_binary_cross_entropy(
emb_spec: Optional partition spec for embeddings.
use_pallas: If True, use Pallas kernels; if False, use pure JAX
implementation.
optimize_large_l: If True, use the target-free vocabulary pass + sparse
positive-target correction. Defaults to False.

Returns:
Scalar loss, optionally paired with per-target losses and/or metrics.
Expand Down Expand Up @@ -651,6 +714,7 @@ def cut_binary_cross_entropy(
emb_spec=emb_spec,
use_pallas=use_pallas,
global_vocab=vocab_size,
optimize_large_l=optimize_large_l,
)

res = _cut_binary_cross_entropy(
Expand Down Expand Up @@ -865,8 +929,6 @@ def _bce_bwd_kernel(

def loop_body(n_idx, _):
act = act_ref[pl.ds(n_idx * block_n, block_n), :]
tgt_val = tgt_ref[:labels, pl.ds(n_idx * block_n, block_n)]
# (labels, block_n)
dloss_val = dloss_ref[0, pl.ds(n_idx * block_n, block_n)] # (block_n,)

# Compute logits: (block_n, block_v)
Expand All @@ -880,14 +942,18 @@ def loop_body(n_idx, _):
probs = jax.nn.sigmoid(logits)

# Target matching across labels: (block_n, block_v)
y_true_chunk = jnp.any(tgt_val[:, :, None] == chunk_indices_local, axis=0)
if labels > 0:
tgt_val = tgt_ref[:labels, pl.ds(n_idx * block_n, block_n)]
y_true_chunk = jnp.any(tgt_val[:, :, None] == chunk_indices_local, axis=0)
g = probs - y_true_chunk.astype(probs.dtype)
else:
g = probs

# Valid batch mask for this loop step: (block_n, 1)
batch_indices_local = n_idx * block_n + local_n_iota
valid_batch_mask = batch_indices_local < n_real

# Gradients w.r.t logits, masked: (block_n, block_v)
g = probs - y_true_chunk.astype(probs.dtype)
scale = jnp.where(valid_batch_mask, (dloss_val * inv_vocab)[:, None], 0.0)
deriv = scale * g

Expand Down Expand Up @@ -920,6 +986,54 @@ def loop_body(n_idx, _):
d_emb_ref[...] = d_emb_scratch[...]


def _apply_bwd_pos_correction(
d_act_2d: jt.Float[jt.Array, "N D"],
d_emb: jt.Float[jt.Array, "V D"],
activations_2d: jt.Float[jt.Array, "N D"],
embeddings: jt.Float[jt.Array, "V D"],
targets_2d: jt.Int[jt.Array, "N L"],
dloss_2d: jt.Float[jt.Array, "N 1"],
inv_vocab: float,
grad_diff_fn: (
Callable[[jt.Float[jt.Array, "N L"]], jt.Float[jt.Array, "N L"]] | None
) = None,
) -> tuple[jt.Float[jt.Array, "N D"], jt.Float[jt.Array, "V D"]]:
"""Applies the sparse positive-target gradient correction."""
safe_t, valid_pos = _unique_valid_targets(targets_2d, embeddings.shape[0])
pos_emb = embeddings[safe_t]
raw_diff = (
-1.0
if grad_diff_fn is None
else grad_diff_fn(
jnp.einsum(
"nd,nld->nl",
activations_2d,
pos_emb,
preferred_element_type=jnp.float32,
precision=jax.lax.Precision.DEFAULT,
)
)
)
scale_2d = (dloss_2d * inv_vocab).astype(jnp.float32)
delta_deriv = jnp.where(valid_pos, raw_diff * scale_2d, 0.0).astype(
embeddings.dtype
)
d_act_pos = jnp.einsum(
"nl,nld->nd",
delta_deriv,
pos_emb,
preferred_element_type=jnp.float32,
precision=jax.lax.Precision.DEFAULT,
).astype(d_act_2d.dtype)
unique_t = jnp.where(valid_pos, safe_t, -1)
d_emb_updates = delta_deriv[:, :, None] * activations_2d[:, None, :].astype(
embeddings.dtype
)
return d_act_2d + d_act_pos, d_emb.at[unique_t].add(
d_emb_updates, mode="drop"
)


def _bce_bwd_pallas(
config: BCEConfig,
d_loss: jt.Float[jt.Array, "B N"],
Expand All @@ -941,7 +1055,7 @@ def _bce_bwd_pallas(
n_blocks = (n + block_n - 1) // block_n
v_blocks = (vocab + block_v - 1) // block_v
vocab_padded = v_blocks * block_v
labels = targets.shape[-1]
labels = 0 if config.optimize_large_l else targets.shape[-1]

# Rebase global target ids onto this shard's local rows here rather than
# inside the kernel: under vocab sharding `vocab_offset` is a tracer, and a
Expand All @@ -958,14 +1072,14 @@ def _bce_bwd_pallas(
)
dloss_2d = jnp.pad(jnp.reshape(d_loss, (n, 1)), ((0, pad_len), (0, 0)))
targets_2d = jnp.pad(
jnp.reshape(targets, (n, labels)),
jnp.reshape(targets[..., :labels], (n, labels)),
((0, pad_len), (0, 0)),
constant_values=-1,
)
else:
activations_2d = jnp.reshape(activations, (n, hidden))
dloss_2d = jnp.reshape(d_loss, (n, 1))
targets_2d = jnp.reshape(targets, (n, labels))
targets_2d = jnp.reshape(targets[..., :labels], (n, labels))
# Pad activations and embeddings to padded_d columns
if padded_d > hidden:
activations_padded = jnp.pad(
Expand Down Expand Up @@ -1059,6 +1173,18 @@ def _bce_bwd_pallas(

if padded_n > n:
d_act_2d = d_act_2d[:n, :]

if config.optimize_large_l:
d_act_2d, d_emb = _apply_bwd_pos_correction(
d_act_2d,
d_emb,
activations_2d[:n],
embeddings,
jnp.reshape(targets, (n, -1)),
dloss_2d[:n],
1.0 / global_vocab,
)

d_activations = jnp.reshape(d_act_2d, (batch, seq_len, hidden))
return d_activations, d_emb

Expand Down Expand Up @@ -1340,6 +1466,7 @@ def _bce_bwd_pure_jax(

# Chunked scan backward pass without tensor padding
v_blocks = int(np.ceil(vocab / block_v))
scan_targets = targets_2d[:, :0] if config.optimize_large_l else targets_2d

def scan_body(carry, j):
d_act_acc, d_emb_acc = carry
Expand All @@ -1359,9 +1486,9 @@ def scan_body(carry, j):
valid_mask = chunk_indices >= j * block_v

targets_chunk = jnp.zeros((n, block_v), dtype=jnp.bool_)
rel_targets = targets_2d - (actual_start + vocab_offset)
rel_targets = scan_targets - (actual_start + vocab_offset)
chunk_cols = jnp.arange(block_v)[None, :]
for l_idx in range(targets_2d.shape[-1]):
for l_idx in range(scan_targets.shape[-1]):
targets_chunk = targets_chunk | (
rel_targets[:, l_idx : l_idx + 1] == chunk_cols
)
Expand Down Expand Up @@ -1401,6 +1528,17 @@ def scan_body(carry, j):
scan_body, init, jnp.arange(v_blocks)
)

if config.optimize_large_l:
d_act_final, d_embeddings = _apply_bwd_pos_correction(
d_act_final,
d_embeddings,
activations_2d,
embeddings,
targets_2d - vocab_offset,
dloss_2d,
inv_vocab,
)

d_activations = jnp.reshape(d_act_final, (batch, seq_len, hidden))
return d_activations, d_embeddings

Expand Down
64 changes: 64 additions & 0 deletions recml/core/ops/binary_cross_entropy_ops_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -1024,6 +1024,70 @@ def test_check_vocab_replicated_in_d(self):
)
binary_cross_entropy_ops._check_vocab_replicated_in_d(P('devices'))

@parameterized.named_parameters(
('l8_pallas', 8, True),
('l16_pallas', 16, True),
('l32_pallas', 32, True),
('l32_pure_jax', 32, False),
('l64_pallas', 64, True),
('l64_pure_jax', 64, False),
)
def test_optimize_large_l_forward_backward_and_metrics(
self, num_labels, use_pallas
):
batch, seq_len, hidden_dim, vocab_size = 2, 32, 64, 512
block_v = 128
key_act, key_emb, key_tgt, key_pad = jax.random.split(
jax.random.PRNGKey(42 + num_labels), 4
)
activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim))
embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim))
raw_targets = jax.random.randint(
key_tgt, (batch, seq_len, num_labels), 0, 64
)
pad_mask = jax.random.bernoulli(key_pad, 0.2, raw_targets.shape)
targets = jnp.where(pad_mask, -1, raw_targets)

res_ref = binary_cross_entropy_ops.cut_binary_cross_entropy(
activations,
embeddings,
targets,
block_v=block_v,
return_metrics=True,
use_pallas=use_pallas,
optimize_large_l=False,
)
res_opt = binary_cross_entropy_ops.cut_binary_cross_entropy(
activations,
embeddings,
targets,
block_v=block_v,
return_metrics=True,
use_pallas=use_pallas,
optimize_large_l=True,
)
for val_opt, val_ref in zip(res_opt, res_ref):
np.testing.assert_allclose(val_opt, val_ref, atol=1e-5, rtol=1e-5)

def loss_fn(act, emb, opt):
return binary_cross_entropy_ops.cut_binary_cross_entropy(
act,
emb,
targets,
block_v=block_v,
use_pallas=use_pallas,
optimize_large_l=opt,
)

g_act_ref, g_emb_ref = jax.jit(
jax.grad(lambda a, e: loss_fn(a, e, False), argnums=(0, 1))
)(activations, embeddings)
g_act_opt, g_emb_opt = jax.jit(
jax.grad(lambda a, e: loss_fn(a, e, True), argnums=(0, 1))
)(activations, embeddings)
np.testing.assert_allclose(g_act_opt, g_act_ref, atol=1e-5, rtol=1e-5)
np.testing.assert_allclose(g_emb_opt, g_emb_ref, atol=1e-5, rtol=1e-5)


if __name__ == '__main__':
absltest.main()
Loading
Loading