From e861941adf730934901d9d34ee4976848896dcf0 Mon Sep 17 00:00:00 2001 From: Yaoyiran Li Date: Fri, 2 Oct 2026 17:51:22 -0700 Subject: [PATCH] Internal change PiperOrigin-RevId: 992633618 --- recml/core/ops/binary_cross_entropy_ops.py | 160 ++++++++++++++++-- .../core/ops/binary_cross_entropy_ops_test.py | 64 +++++++ .../ops/binary_focal_cross_entropy_ops.py | 130 ++++++++++++-- .../binary_focal_cross_entropy_ops_test.py | 74 ++++++++ 4 files changed, 404 insertions(+), 24 deletions(-) diff --git a/recml/core/ops/binary_cross_entropy_ops.py b/recml/core/ops/binary_cross_entropy_ops.py index 10d46ff..d92c888 100644 --- a/recml/core/ops/binary_cross_entropy_ops.py +++ b/recml/core/ops/binary_cross_entropy_ops.py @@ -247,6 +247,8 @@ 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 @@ -254,6 +256,7 @@ class BCEConfig: 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 @@ -261,6 +264,49 @@ class BCEConfig: 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"], @@ -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: @@ -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, :] @@ -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)), @@ -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)) @@ -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"]] @@ -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. @@ -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( @@ -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) @@ -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 @@ -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"], @@ -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 @@ -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( @@ -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 @@ -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 @@ -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 ) @@ -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 diff --git a/recml/core/ops/binary_cross_entropy_ops_test.py b/recml/core/ops/binary_cross_entropy_ops_test.py index 474b136..122bad4 100644 --- a/recml/core/ops/binary_cross_entropy_ops_test.py +++ b/recml/core/ops/binary_cross_entropy_ops_test.py @@ -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() diff --git a/recml/core/ops/binary_focal_cross_entropy_ops.py b/recml/core/ops/binary_focal_cross_entropy_ops.py index ef714b4..f9d32f0 100644 --- a/recml/core/ops/binary_focal_cross_entropy_ops.py +++ b/recml/core/ops/binary_focal_cross_entropy_ops.py @@ -54,6 +54,12 @@ _replicate_hidden_dim = ( bce_ops._replicate_hidden_dim # pylint: disable=protected-access ) +_apply_bwd_pos_correction = ( + bce_ops._apply_bwd_pos_correction # pylint: disable=protected-access +) +_fwd_pos_correction = ( + bce_ops._fwd_pos_correction # pylint: disable=protected-access +) _BLOCK_N = bce_ops._BLOCK_N # pylint: disable=protected-access @@ -79,6 +85,20 @@ class FocalBCEConfig(bce_ops.BCEConfig): apply_class_balancing: bool = False +def _focal_bce_loss_diff( + config: FocalBCEConfig, logits: jt.Float[jt.Array, "N L"] +) -> jt.Float[jt.Array, "N L"]: + """Computes FocalBCE(logits, y=1) - FocalBCE(logits, y=0).""" + probs = jax.nn.sigmoid(logits) + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p(jnp.exp(-jnp.abs(logits))) + loss_pos = jnp.power(1.0 - probs, config.gamma) * (loss_zero - logits) + loss_neg = jnp.power(probs, config.gamma) * loss_zero + if config.apply_class_balancing: + loss_pos = config.alpha * loss_pos + loss_neg = (1.0 - config.alpha) * loss_neg + return loss_pos - loss_neg + + def _focal_bce_fwd_chunk( config: FocalBCEConfig, activations_2d: jt.Float[jt.Array, "N D"], @@ -166,6 +186,7 @@ def _focal_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: @@ -175,7 +196,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 = _focal_bce_fwd_chunk( - config, activations_2d, embeddings, targets_2d, j, block_v, vocab + config, activations_2d, embeddings, scan_targets, j, block_v, vocab ) predictions_chunk = (logits > 0.0) & valid_mask[None, :] @@ -209,6 +230,17 @@ 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, + functools.partial(_focal_bce_loss_diff, config), + ) + 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)), @@ -223,12 +255,22 @@ def v_body_no_metrics( j: jt.Int[jt.Array, ""], ) -> tuple[jt.Float[jt.Array, "N"], None]: loss_sum, _, _, _ = _focal_bce_fwd_chunk( - config, activations_2d, embeddings, targets_2d, j, block_v, vocab + config, 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, + functools.partial(_focal_bce_loss_diff, config), + )[0] + ) return jnp.reshape(loss_final, (batch, seq_len)) @@ -376,6 +418,7 @@ def cut_binary_focal_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"]] @@ -406,8 +449,8 @@ def cut_binary_focal_cross_entropy( D]`` with a single extra leading axis. embeddings: Output embedding / unembedding weights of shape ``[V, D]``. targets: Target token ids of shape ``[B, N, L]``, i.e. the ``L`` positive - labels of each sequence position. For 4D activations, either a matching - 4D ``[X, B, N, L]`` array or a 3D ``[B, N, L]`` array shared across the + labels of each sequence position. For 4D activations, either a matching 4D + ``[X, B, N, L]`` array or a 3D ``[B, N, L]`` array shared across the leading axis. weights: Per-sequence-position loss weights, typically a padding/validity mask. Must have exactly the shape of the per-position losses, i.e. @@ -427,6 +470,8 @@ def cut_binary_focal_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. @@ -491,6 +536,7 @@ def cut_binary_focal_cross_entropy( emb_spec=emb_spec, use_pallas=use_pallas, global_vocab=vocab_size, + optimize_large_l=optimize_large_l, ) res = _cut_binary_focal_cross_entropy( @@ -585,8 +631,6 @@ def _focal_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) @@ -600,14 +644,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) + y_true_float = y_true_chunk.astype(probs.dtype) + else: + y_true_float = 0.0 # 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 - y_true_float = y_true_chunk.astype(probs.dtype) g_bce = probs - y_true_float abs_g = jnp.abs(g_bce) p_t = 1.0 - abs_g @@ -673,6 +721,36 @@ def loop_body(n_idx, _): d_emb_ref[...] = d_emb_scratch[...] +def _focal_bce_grad_diff( + config: FocalBCEConfig, logits: jt.Float[jt.Array, "N L"] +) -> jt.Float[jt.Array, "N L"]: + """Computes g_focal(logits, y=1) - g_focal(logits, y=0).""" + probs = jax.nn.sigmoid(logits) + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p(jnp.exp(-jnp.abs(logits))) + grads = [] + for y in (1.0, 0.0): + g_bce = probs - y + abs_g = jnp.abs(g_bce) + focal_factor = jnp.power(abs_g, config.gamma) + bce_loss = loss_zero - y * logits + if config.gamma >= 1.0: + g = g_bce * ( + focal_factor + + config.gamma + * jnp.power(abs_g, config.gamma - 1.0) + * (1.0 - abs_g) + * bce_loss + ) + elif config.gamma == 0.0: + g = g_bce + else: + g = focal_factor * (g_bce + config.gamma * (1.0 - y - probs) * bce_loss) + if config.apply_class_balancing: + g = (y * (2.0 * config.alpha - 1.0) + (1.0 - config.alpha)) * g + grads.append(g) + return grads[0] - grads[1] + + def _focal_bce_bwd_pallas_chunked_n( config: FocalBCEConfig, d_loss: jt.Float[jt.Array, "B N"], @@ -765,7 +843,7 @@ def _focal_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 @@ -782,14 +860,14 @@ def _focal_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( @@ -886,6 +964,19 @@ def _focal_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, + functools.partial(_focal_bce_grad_diff, config), + ) + d_activations = jnp.reshape(d_act_2d, (batch, seq_len, hidden)) return d_activations, d_emb @@ -927,6 +1018,7 @@ def _focal_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 @@ -946,9 +1038,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 ) @@ -1010,6 +1102,18 @@ 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, + functools.partial(_focal_bce_grad_diff, config), + ) + d_activations = jnp.reshape(d_act_final, (batch, seq_len, hidden)) return d_activations, d_embeddings diff --git a/recml/core/ops/binary_focal_cross_entropy_ops_test.py b/recml/core/ops/binary_focal_cross_entropy_ops_test.py index 733d28c..6f1c2d0 100644 --- a/recml/core/ops/binary_focal_cross_entropy_ops_test.py +++ b/recml/core/ops/binary_focal_cross_entropy_ops_test.py @@ -962,6 +962,80 @@ def run_cut(act, emb, use_pallas): np.testing.assert_allclose(g_act_pure, g_act_pallas, atol=1e-5, rtol=1e-5) np.testing.assert_allclose(g_emb_pure, g_emb_pallas, atol=1e-5, rtol=1e-5) + @parameterized.named_parameters( + ('l8_pallas_gamma2', 8, True, 2.0, True), + ('l16_pallas_gamma2', 16, True, 2.0, True), + ('l32_pallas_gamma2', 32, True, 2.0, True), + ('l32_pure_jax_gamma2', 32, False, 2.0, True), + ('l32_pallas_gamma05', 32, True, 0.5, False), + ('l64_pallas_gamma2', 64, True, 2.0, True), + ('l64_pure_jax_gamma2', 64, False, 2.0, True), + ) + def test_optimize_large_l_forward_backward_and_metrics( + self, num_labels, use_pallas, gamma, apply_class_balancing + ): + 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_focal_cross_entropy_ops.cut_binary_focal_cross_entropy( + activations, + embeddings, + targets, + block_v=block_v, + gamma=gamma, + alpha=0.25, + apply_class_balancing=apply_class_balancing, + return_metrics=True, + use_pallas=use_pallas, + optimize_large_l=False, + ) + res_opt = binary_focal_cross_entropy_ops.cut_binary_focal_cross_entropy( + activations, + embeddings, + targets, + block_v=block_v, + gamma=gamma, + alpha=0.25, + apply_class_balancing=apply_class_balancing, + 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_focal_cross_entropy_ops.cut_binary_focal_cross_entropy( + act, + emb, + targets, + block_v=block_v, + gamma=gamma, + alpha=0.25, + apply_class_balancing=apply_class_balancing, + 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-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_opt, g_emb_ref, atol=1e-4, rtol=1e-4) + if __name__ == '__main__': absltest.main()