diff --git a/recml/core/ops/binary_cross_entropy_ops.py b/recml/core/ops/binary_cross_entropy_ops.py new file mode 100644 index 0000000..93f75c8 --- /dev/null +++ b/recml/core/ops/binary_cross_entropy_ops.py @@ -0,0 +1,1091 @@ +# Copyright 2024 RecML authors . +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Binary cross-entropy loss implementation with optimized memory footprint. + +This implementation computes BCE loss without materializing the [B, N, V] logits +matrix in memory, by chunking the vocabulary dimension. +""" + +import dataclasses +import functools +import jax +from jax.experimental import pallas as pl +from jax.experimental.pallas import tpu as pltpu +import jax.numpy as jnp +import jaxtyping as jt +import numpy as np + +EPS = 1e-8 + + +def _auto_block_v(n: int, vocab_size: int) -> int: + """Automatically picks block_v to target ~32MB intermediate logits.""" + # Estimate local n per device to handle sharded runs + try: + num_devices = jax.device_count() + except: # pylint: disable=bare-except + num_devices = 1 + local_n = max(n // num_devices, 1) + + # Target 32MB for intermediate logits tensor [local_n, block_v] + # in float32 (4 bytes). + target_elements = (32 * 1024 * 1024) // 4 + block_v = target_elements // local_n + + # Align to MXU size (128) + block_v = max(128, (block_v // 128) * 128) + + # Clamp to reasonable range [256, 8192] + block_v = max(256, min(8192, block_v)) + # Don't exceed vocab_size + block_v = min(vocab_size, block_v) + return block_v + + +def _get_sharding(x) -> jax.sharding.Sharding | None: + if hasattr(x, "sharding"): + return x.sharding + if hasattr(x, "aval") and hasattr(x.aval, "sharding"): + return x.aval.sharding + return None + + +def _replicate_hidden_dim(x): + """Replicates the hidden dimension of the input tensor if sharded.""" + sharding = _get_sharding(x) + if isinstance(sharding, jax.sharding.NamedSharding): + mesh = sharding.mesh + if not mesh.empty: + spec = sharding.spec + new_spec_list = list(spec) + if new_spec_list: + new_spec_list[-1] = None + new_spec = jax.sharding.PartitionSpec(*new_spec_list) + return jax.lax.with_sharding_constraint( + x, jax.sharding.NamedSharding(mesh, new_spec) + ) + return x + + +@dataclasses.dataclass +class BCEConfig: + """Configuration for the binary cross-entropy loss.""" + + block_v: int + block_n: int = 256 + compute_metrics: 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 + bias_spec: jax.sharding.PartitionSpec | None = None + + +def _bce_fwd_local( + config: BCEConfig, + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], +) -> tuple[ + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], +]: + """Computes the sum of Loss(x_v, target_v) over all V block-wise, and metrics.""" + block_v = config.block_v + batch, seq_len, hidden = activations.shape + vocab = embeddings.shape[0] + + n = batch * seq_len + # NOMUTANTS -- v_blocks is calculated from block_v and vocab. + v_blocks = int(np.ceil(vocab / block_v)) + + activations_2d = jnp.reshape(activations, (n, hidden)) + targets_2d = jnp.reshape(targets, (n, -1)) + + if config.compute_metrics: + + def v_body( + carry: tuple[jt.Float[jt.Array, "N"], ...], + j: jt.Int[jt.Array, ""], + ) -> tuple[tuple[jt.Float[jt.Array, "N"], ...], None]: + loss_acc, tp_acc, fp_acc, fn_acc, tn_acc = carry + actual_start = jnp.maximum(0, jnp.minimum(j * block_v, vocab - block_v)) + emb_chunk = jax.lax.dynamic_slice_in_dim( + embeddings, actual_start, block_v + ) + bias_chunk = jax.lax.dynamic_slice_in_dim(bias, actual_start, block_v) + logits = ( + jnp.matmul( + activations_2d, + jnp.transpose(emb_chunk), + preferred_element_type=jnp.float32, + ) + + bias_chunk + ) + + chunk_indices = actual_start + jnp.arange(block_v) + valid_mask = (chunk_indices >= j * block_v) & (chunk_indices < vocab) + + # Fused BCE Loss: BCE(x, y) = BCE(x, 0) - y * x + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p( + jnp.exp(-jnp.abs(logits)) + ) + + targets_match = targets_2d[:, :, None] == chunk_indices[None, None, :] + targets_chunk = jnp.any(targets_match, axis=1) + + loss_chunk = loss_zero - targets_chunk * logits + loss_chunk = loss_chunk * valid_mask[None, :] + loss_sum = jnp.sum(loss_chunk, axis=-1) + + predictions_chunk = logits > 0.0 + + tp_chunk = targets_chunk & predictions_chunk + fp_chunk = (~targets_chunk) & predictions_chunk + fn_chunk = targets_chunk & (~predictions_chunk) + tn_chunk = (~targets_chunk) & (~predictions_chunk) + + tp_chunk = tp_chunk & valid_mask[None, :] + fp_chunk = fp_chunk & valid_mask[None, :] + fn_chunk = fn_chunk & valid_mask[None, :] + tn_chunk = tn_chunk & valid_mask[None, :] + + tp_sum = jnp.sum(tp_chunk, axis=-1).astype(jnp.float32) + fp_sum = jnp.sum(fp_chunk, axis=-1).astype(jnp.float32) + fn_sum = jnp.sum(fn_chunk, axis=-1).astype(jnp.float32) + tn_sum = jnp.sum(tn_chunk, axis=-1).astype(jnp.float32) + + return ( + loss_acc + loss_sum, + tp_acc + tp_sum, + fp_acc + fp_sum, + fn_acc + fn_sum, + tn_acc + tn_sum, + ), None + + init = ( + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + ) + (loss_final, tp_final, fp_final, fn_final, tn_final), _ = jax.lax.scan( + jax.checkpoint(v_body), init, jnp.arange(v_blocks) + ) + return ( + jnp.reshape(loss_final, (batch, seq_len)), + jnp.reshape(tp_final, (batch, seq_len)), + jnp.reshape(fp_final, (batch, seq_len)), + jnp.reshape(fn_final, (batch, seq_len)), + jnp.reshape(tn_final, (batch, seq_len)), + ) + else: + + def v_body_no_metrics( + loss_acc: jt.Float[jt.Array, "N"], + j: jt.Int[jt.Array, ""], + ) -> tuple[jt.Float[jt.Array, "N"], None]: + actual_start = jnp.maximum(0, jnp.minimum(j * block_v, vocab - block_v)) + emb_chunk = jax.lax.dynamic_slice_in_dim( + embeddings, actual_start, block_v + ) + bias_chunk = jax.lax.dynamic_slice_in_dim(bias, actual_start, block_v) + logits = ( + jnp.matmul( + activations_2d, + jnp.transpose(emb_chunk), + preferred_element_type=jnp.float32, + ) + + bias_chunk + ) + + chunk_indices = actual_start + jnp.arange(block_v) + valid_mask = (chunk_indices >= j * block_v) & (chunk_indices < vocab) + + # Fused BCE Loss: BCE(x, y) = BCE(x, 0) - y * x + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p( + jnp.exp(-jnp.abs(logits)) + ) + + targets_match = targets_2d[:, :, None] == chunk_indices[None, None, :] + targets_chunk = jnp.any(targets_match, axis=1) + + loss_chunk = loss_zero - targets_chunk * logits + loss_chunk = loss_chunk * valid_mask[None, :] + loss_sum = jnp.sum(loss_chunk, axis=-1) + return loss_acc + loss_sum, None + + init = jnp.zeros((n,), dtype=jnp.float32) + loss_final, _ = jax.lax.scan( + jax.checkpoint(v_body_no_metrics), init, jnp.arange(v_blocks) + ) + dummy = jnp.zeros((batch, seq_len), dtype=jnp.float32) + return ( + jnp.reshape(loss_final, (batch, seq_len)), + dummy, + dummy, + dummy, + dummy, + ) + + +@functools.partial(jax.custom_vjp, nondiff_argnums=(0,)) +def _cut_binary_cross_entropy( + config: BCEConfig, + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "... B N L"], +) -> tuple[ + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], +]: + """Computes the non-differentiable path of cut BCE loss and metrics.""" + outputs, _ = _cut_binary_cross_entropy_fwd( + config, activations, embeddings, bias, targets + ) + return outputs + + +def _cut_binary_cross_entropy_fwd( + config: BCEConfig, + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], +) -> tuple[ + tuple[ + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + ], + tuple[ + jt.Float[jt.Array, "B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], + jt.Int[jt.Array, "B N L"], + ], +]: + """Computes forward mode of cut BCE loss.""" + replicated_activations = _replicate_hidden_dim(activations) + if activations.ndim == 4: + if targets.ndim == 3: + targets_in_axis = None + else: + targets_in_axis = 0 + fwd_vmap = jax.vmap( + functools.partial(_bce_fwd_local, config), + in_axes=(0, None, None, targets_in_axis), + ) + loss_y0, tp, fp, fn, tn = fwd_vmap( + replicated_activations, embeddings, bias, targets + ) + else: + loss_y0, tp, fp, fn, tn = _bce_fwd_local( + config, replicated_activations, embeddings, bias, targets + ) + vocab_size = embeddings.shape[0] + losses = loss_y0 / vocab_size + + return (losses, tp, fp, fn, tn), ( + activations, + embeddings, + bias, + targets, + ) + + +def _cut_binary_cross_entropy_bwd( + config: BCEConfig, + res: tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], + jt.Int[jt.Array, "... B N L"], + ], + d_outputs: tuple[ + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + ], +) -> tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], + None, +]: + """Computes the backward mode of cut BCE loss.""" + d_losses, _, _, _, _ = d_outputs + activations, embeddings, bias, targets = res + d_activations, d_embeddings, d_bias = _bce_bwd_sharded( + config, d_losses, activations, embeddings, bias, targets + ) + return d_activations, d_embeddings, d_bias, None + + +_cut_binary_cross_entropy.defvjp( + _cut_binary_cross_entropy_fwd, _cut_binary_cross_entropy_bwd +) + + +def cut_binary_cross_entropy( + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + targets: jt.Int[jt.Array, "... B N L"], + bias: jt.Float[jt.Array, "V"] | None = None, + weights: jt.Float[jt.Array, "... B N"] | None = None, + *, + return_per_target_losses: bool = False, + return_metrics: bool = False, + block_v: int | None = None, + mesh: jax.sharding.Mesh | None = None, + act_spec: jax.sharding.PartitionSpec | None = None, + emb_spec: jax.sharding.PartitionSpec | None = None, + bias_spec: jax.sharding.PartitionSpec | None = None, +) -> ( + jt.Float[jt.Array, ""] + | tuple[jt.Float[jt.Array, ""], jt.Float[jt.Array, "... B N"]] + | tuple[ + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + ] + | tuple[ + jt.Float[jt.Array, ""], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + ] +): + """Computes binary cross entropy loss over unmaterialized logits. + + Args: + activations: Hidden-state outputs of shape ``[B, N, D]``. + embeddings: Output embedding / unembedding weights of shape ``[V, D]``. + targets: Target token ids of shape ``[B, N, L]``. + bias: Optional bias of shape ``[V]``. + weights: Per-token loss weights of shape ``[B, N]``. + return_per_target_losses: If True, also return the per-target loss tensor. + return_metrics: If True, also return TP, FP, FN, TN metric counts. + block_v: Vocab-axis block size. Auto-picked if omitted. + mesh: Optional mesh to use for sharding. + act_spec: Optional partition spec for activations. + emb_spec: Optional partition spec for embeddings. + bias_spec: Optional partition spec for bias. + + Returns: + Scalar loss, optionally paired with per-target losses and/or metrics. + """ + vocab_size = embeddings.shape[0] + if bias is None: + bias = jnp.zeros(vocab_size, dtype=embeddings.dtype) + + # Prevent collective communications inside the loop by forcing replication of + # weights. + sharding = _get_sharding(embeddings) + if ( + isinstance(sharding, jax.sharding.NamedSharding) + and not sharding.mesh.empty + ): + replicated_sharding = jax.sharding.NamedSharding( + sharding.mesh, jax.sharding.PartitionSpec() + ) + embeddings = jax.lax.with_sharding_constraint( + embeddings, replicated_sharding + ) + bias = jax.lax.with_sharding_constraint(bias, replicated_sharding) + + if block_v is None: + n = activations.shape[-3] * activations.shape[-2] + block_v = _auto_block_v(n, vocab_size) + else: + block_v = min(block_v, vocab_size) + + losses, tp, fp, fn, tn = _cut_binary_cross_entropy( + BCEConfig( + block_v=block_v, + compute_metrics=return_metrics, + mesh=mesh, + act_spec=act_spec, + emb_spec=emb_spec, + bias_spec=bias_spec, + ), + activations, + embeddings, + bias, + targets, + ) + + if weights is not None: + losses = losses * weights + weight_sum = jnp.sum(weights) + tp_sum = jnp.sum(tp * weights) + fp_sum = jnp.sum(fp * weights) + fn_sum = jnp.sum(fn * weights) + tn_sum = jnp.sum(tn * weights) + else: + weight_sum = np.prod(targets.shape[:-1]) + tp_sum = jnp.sum(tp) + fp_sum = jnp.sum(fp) + fn_sum = jnp.sum(fn) + tn_sum = jnp.sum(tn) + + loss = jnp.sum(losses) / (weight_sum + EPS) + + if return_metrics: + if return_per_target_losses: + return loss, losses, tp_sum, fp_sum, fn_sum, tn_sum + return loss, tp_sum, fp_sum, fn_sum, tn_sum + + if return_per_target_losses: + return loss, losses + + return loss + + +def _pallas_lane() -> int: + """TPU MXU lane width - VMEM minor-axis tile constraint.""" + if any(d.platform == "tpu" for d in jax.devices()): + return pltpu.get_tpu_info().num_lanes + return 128 + + +def _pallas_vmem_budget() -> int: + """Per-scoped Pallas allocation VMEM budget for the live TPU.""" + if any(d.platform == "tpu" for d in jax.devices()): + cap = pltpu.get_tpu_info().vmem_capacity_bytes - 8 * 1024 * 1024 + return max(16 * 1024 * 1024, cap) + return 16 * 1024 * 1024 + + +def _max_safe_block_v(vmem_budget: int, padded_d: int) -> int: + """Returns maximum safe block_v to prevent VMEM OOM.""" + # On 16MB VMEM (TPU v3/v4), max block_v * padded_d is 131072. + # On 128MB VMEM (TPU v5e/v6e), max block_v * padded_d is 1048576. + if vmem_budget <= 32 * 1024 * 1024: + max_product = 131072 + max_absolute = 1024 + else: + max_product = 16777216 + max_absolute = 32768 + + # NOMUTANTS -- max_safe is calculated based on VMEM budget. + max_safe = max_product // padded_d + # Ensure it is a multiple of 128 + max_safe = (max_safe // 128) * 128 + return max(256, min(max_safe, max_absolute)) + + +def _pallas_interpret() -> bool: + """Run Pallas in interpret mode whenever no TPU is present (dev / CI).""" + return not any(d.platform == "tpu" for d in jax.devices()) + + +def _check_vocab_replicated_in_d(emb_spec: jax.sharding.PartitionSpec) -> None: + if len(emb_spec) > 1 and emb_spec[1] is not None: + raise NotImplementedError( + "Embeddings sharded along the hidden dimension D are not supported; " + f"got embedding sharding {emb_spec}." + ) + + +def _bce_bwd_kernel( + emb_ref, # [block_v, 128] VMEM + act_ref, # [padded_n, 256] HBM + bias_ref, # [8, block_v] VMEM + tgt_ref, # [l_padded, padded_n] HBM + dloss_ref, # [l_padded, padded_n] HBM + d_emb_ref, # [block_v, 128] HBM output + d_bias_ref, # [block_v, 128] HBM output + d_act_partials_ref, # [1, padded_n, 256] HBM output + d_emb_scratch, # [block_v, 128] VMEM scratch + d_bias_scratch, # [block_v] VMEM scratch + *, + block_v: int, + block_n: int, + n_blocks: int, + vocab: int, + vocab_offset: int, + labels: int, + n_real: int, +): + """Pallas TPU per-shard chunked backward for BCE.""" + v_idx = pl.program_id(0) + + # Initialize accumulators in VMEM to 0 + d_emb_scratch[...] = jnp.zeros(d_emb_scratch.shape, jnp.float32) + d_bias_scratch[...] = jnp.zeros(d_bias_scratch.shape, jnp.float32) + + v_start_local = v_idx * block_v + v_start_global = v_start_local + vocab_offset + + # Valid vocab mask + chunk_indices_local = v_start_local + jnp.arange(block_v)[None, :] + valid_vocab_mask = chunk_indices_local < vocab + + # Load inputs from VMEM + emb = emb_ref[...] + emb = jnp.where(valid_vocab_mask.T, emb, 0.0) + + bias_val = bias_ref[...] + bias = bias_val[0, :] + bias = jnp.where(valid_vocab_mask[0], bias, 0.0) + + # Broadcasted iota for target matching + col_ids = jax.lax.broadcasted_iota(jnp.int32, (block_n, block_v), 1) + col_ids_global = col_ids + v_start_global + + def loop_body(n_idx, _): + # Load act, tgt, dloss slices for this n_idx from HBM to VMEM + act = act_ref[pl.ds(n_idx * block_n, block_n), :] + tgt_val = tgt_ref[:, pl.ds(n_idx * block_n, block_n)] + dloss_val = dloss_ref[:, pl.ds(n_idx * block_n, block_n)] + + # Slice real dloss + dloss = dloss_val[0, :] + dloss = dloss[:, None] + + # Force loading of padded columns/rows by computing a dummy sum + dummy_sum = ( + jnp.sum(tgt_val).astype(jnp.float32) + + jnp.sum(dloss_val) + + jnp.sum(bias_val) + ) + + # Compute logits + logits = ( + jax.lax.dot_general( + act, + emb, + (((1,), (1,)), ((), ())), + precision=jax.lax.Precision.DEFAULT, + preferred_element_type=jnp.float32, + ) + + bias[None, :] + + dummy_sum * 0.0 + ) + probs = jax.nn.sigmoid(logits) + + # Target matching + y_true_chunk = jnp.zeros((block_n, block_v), dtype=jnp.bool_) + for l in range(labels): + target_l = tgt_val[l, :] + target_l = target_l[:, None] + is_valid_target = target_l >= 0 + match = (target_l == col_ids_global) & is_valid_target + y_true_chunk = y_true_chunk | match + + # Valid batch mask for this loop step + batch_indices_local = n_idx * block_n + jnp.arange(block_n)[:, None] + valid_batch_mask = batch_indices_local < n_real + + # Gradients w.r.t logits, masked + g = probs - y_true_chunk.astype(probs.dtype) + g = jnp.where(valid_vocab_mask & valid_batch_mask, g, 0.0) + + # Scale + deriv = (dloss / vocab) * g + + # Accumulate d_emb + d_emb_contrib = jax.lax.dot_general( + deriv.astype(act.dtype), + act, + (((0,), (0,)), ((), ())), + precision=jax.lax.Precision.DEFAULT, + preferred_element_type=jnp.float32, + ) + d_emb_scratch[...] = d_emb_scratch[...] + d_emb_contrib + + # Accumulate d_bias + d_bias_contrib = jnp.sum(deriv, axis=0) + d_bias_scratch[...] = d_bias_scratch[...] + d_bias_contrib + + # Compute d_act and write directly to HBM + d_act_contrib = jax.lax.dot_general( + deriv.astype(emb.dtype), + emb, + (((1,), (0,)), ((), ())), + precision=jax.lax.Precision.DEFAULT, + preferred_element_type=jnp.float32, + ) + d_act_partials_ref[0, pl.ds(n_idx * block_n, block_n), :] = d_act_contrib + + return None + + # Run reduction loop over n_blocks + jax.lax.fori_loop(0, n_blocks, loop_body, None) + + # Store accumulated results to HBM + d_emb_ref[...] = d_emb_scratch[...] + d_bias_ref[...] = d_bias_scratch[...][..., None] * ( + jnp.arange(d_bias_ref.shape[1]) == 0 + ) + + +def _bce_bwd_pallas( + config: BCEConfig, + d_loss: jt.Float[jt.Array, "B N"], + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], + vocab_offset: int = 0, +) -> tuple[ + jt.Float[jt.Array, "B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Pallas TPU per-shard chunked backward for BCE.""" + block_n = config.block_n + batch, seq_len, hidden = activations.shape + padded_d = ((hidden + 127) // 128) * 128 + vocab = embeddings.shape[0] + vmem_budget = _pallas_vmem_budget() + max_safe = _max_safe_block_v(vmem_budget, padded_d) + vocab_padded_128 = ((vocab + 127) // 128) * 128 + block_v = min(config.block_v, vocab, max_safe) + block_v = ((block_v + 127) // 128) * 128 + block_v = min(block_v, vocab_padded_128, max_safe) + n = batch * seq_len + 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] + + # Pad activations, targets, d_loss to multiples of block_n if necessary + padded_n = n_blocks * block_n + if padded_n > n: + pad_len = padded_n - n + activations_2d = jnp.pad( + jnp.reshape(activations, (n, hidden)), ((0, pad_len), (0, 0)) + ) + dloss_2d = jnp.pad(jnp.reshape(d_loss, (n, 1)), ((0, pad_len), (0, 0))) + targets_2d = jnp.pad( + jnp.reshape(targets, (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)) + # Pad activations and embeddings to padded_d columns + if padded_d > hidden: + activations_padded = jnp.pad( + activations_2d, ((0, 0), (0, padded_d - hidden)) + ) + else: + activations_padded = activations_2d + + if vocab_padded > vocab or padded_d > hidden: + embeddings_padded = jnp.pad( + embeddings, ((0, vocab_padded - vocab), (0, padded_d - hidden)) + ) + else: + embeddings_padded = embeddings + + # Transpose and pad targets: (n, labels) -> (labels, n) -> + # (l_padded, padded_n) + l_padded = max(8, ((labels + 7) // 8) * 8) + targets_t = jnp.transpose(targets_2d, (1, 0)) + if l_padded > labels or padded_n > n: + targets_padded = jnp.pad( + targets_t, + ((0, l_padded - labels), (0, padded_n - n)), + constant_values=-1, + ).astype(jnp.int32) + else: + targets_padded = targets_t.astype(jnp.int32) + + # Transpose and pad dloss: (n, 1) -> (1, n) -> (l_padded, padded_n) + dloss_t = jnp.transpose(dloss_2d, (1, 0)) + dloss_padded = jnp.pad( + dloss_t, ((0, l_padded - 1), (0, padded_n - n)) + ).astype(embeddings.dtype) + + # Reshape, transpose and pad bias: (vocab,) -> (1, vocab) -> (8, vocab_padded) + bias_t = jnp.reshape(bias, (1, vocab)) + if vocab_padded > vocab: + bias_padded = jnp.pad(bias_t, ((0, 7), (0, vocab_padded - vocab))).astype( + bias.dtype + ) + else: + bias_padded = jnp.pad(bias_t, ((0, 7), (0, 0))).astype(bias.dtype) + + d_emb_padded, d_bias_padded, d_act_partials = pl.pallas_call( + functools.partial( + _bce_bwd_kernel, + block_v=block_v, + block_n=block_n, + n_blocks=n_blocks, + vocab=vocab, + vocab_offset=vocab_offset, + labels=labels, + n_real=n, + ), + out_shape=[ + jax.ShapeDtypeStruct((vocab_padded, padded_d), embeddings.dtype), + jax.ShapeDtypeStruct((vocab_padded, 8), bias.dtype), + jax.ShapeDtypeStruct( + (v_blocks, padded_n, padded_d), activations.dtype + ), + ], + grid=(v_blocks,), + in_specs=[ + pl.BlockSpec((block_v, padded_d), lambda v: (v, 0)), # emb + pl.BlockSpec((padded_n, padded_d), lambda v: (0, 0)), # act + pl.BlockSpec((8, block_v), lambda v: (0, v)), # bias + pl.BlockSpec((l_padded, padded_n), lambda v: (0, 0)), # tgt + pl.BlockSpec((l_padded, padded_n), lambda v: (0, 0)), # dloss + ], + out_specs=[ + pl.BlockSpec((block_v, padded_d), lambda v: (v, 0)), # d_emb + pl.BlockSpec((block_v, 8), lambda v: (v, 0)), # d_bias + pl.BlockSpec((1, padded_n, padded_d), lambda v: (v, 0, 0)), # d_act + ], + scratch_shapes=[ + pltpu.VMEM((block_v, padded_d), jnp.float32), # d_emb_scratch + pltpu.VMEM((block_v,), jnp.float32), # d_bias_scratch + ], + compiler_params=pltpu.CompilerParams( + dimension_semantics=("parallel",), + vmem_limit_bytes=_pallas_vmem_budget(), + ), + interpret=_pallas_interpret(), + )( + embeddings_padded, + activations_padded, + bias_padded, + targets_padded, + dloss_padded, + ) + + # d_emb_padded = jax.lax.optimization_barrier(d_emb_padded) + # d_bias_padded = jax.lax.optimization_barrier(d_bias_padded) + if vocab < vocab_padded or hidden < padded_d: + d_emb = d_emb_padded[:vocab, :hidden] + else: + d_emb = d_emb_padded + d_bias = d_bias_padded[:vocab, 0] + + if hidden < padded_d: + d_act_2d = jnp.sum(d_act_partials[:, :, :hidden], axis=0) + else: + d_act_2d = jnp.sum(d_act_partials, axis=0) + + if padded_n > n: + d_act_2d = d_act_2d[:n, :] + d_activations = jnp.reshape(d_act_2d, (batch, seq_len, hidden)) + return d_activations, d_emb, d_bias + + +def _max_safe_chunk_n(vmem_budget: int, padded_d: int) -> int: + """Returns maximum safe chunk_n to prevent VMEM OOM.""" + # We want act (single buffered) and d_act_partials (double buffered) + # to take at most ~30% of VMEM budget on low-VMEM devices (v4) to leave + # room for compilation/gradient/scan state. + # VMEM usage approx: chunk_n * padded_d * 4 * 3 (1 load + 2 store). + if vmem_budget <= 32 * 1024 * 1024: + target_bytes = int(vmem_budget * 0.3) + else: + target_bytes = int(vmem_budget * 0.6) + max_safe = target_bytes // (padded_d * 4 * 3) + # Round down to multiple of 128 for MXU alignment + max_safe = (max_safe // 128) * 128 + return max(512, min(max_safe, 32768)) + + +def _bce_bwd_pallas_chunked_n( + config: BCEConfig, + d_loss: jt.Float[jt.Array, "B N"], + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], + vocab_offset: int = 0, +) -> tuple[ + jt.Float[jt.Array, "B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Wraps _bce_bwd_pallas by chunking the N (sequence) dimension in JAX.""" + batch, seq_len, hidden = activations.shape + n = batch * seq_len + + padded_d = ((hidden + 127) // 128) * 128 + vmem_budget = _pallas_vmem_budget() + chunk_n = _max_safe_chunk_n(vmem_budget, padded_d) + print(f"JETS_DEBUG: chunk_n={chunk_n}, n={n}", flush=True) + if n <= chunk_n: + return _bce_bwd_pallas( + config, d_loss, activations, embeddings, bias, targets, vocab_offset + ) + # Reshape inputs to 2D + activations_2d = jnp.reshape(activations, (n, hidden)) + dloss_2d = jnp.reshape(d_loss, (n, 1)) + targets_2d = jnp.reshape(targets, (n, -1)) + + n_chunks = (n + chunk_n - 1) // chunk_n + padded_n = n_chunks * chunk_n + + if padded_n > n: + pad_len = padded_n - n + activations_padded = jnp.pad(activations_2d, ((0, pad_len), (0, 0))) + dloss_padded = jnp.pad(dloss_2d, ((0, pad_len), (0, 0))) + targets_padded = jnp.pad( + targets_2d, ((0, pad_len), (0, 0)), constant_values=-1 + ) + else: + activations_padded = activations_2d + dloss_padded = dloss_2d + targets_padded = targets_2d + + act_chunks = jnp.reshape(activations_padded, (n_chunks, chunk_n, hidden)) + dloss_chunks = jnp.reshape(dloss_padded, (n_chunks, chunk_n)) + tgt_chunks = jnp.reshape(targets_padded, (n_chunks, chunk_n, -1)) + + act_chunks_3d = jnp.reshape(act_chunks, (n_chunks, 1, chunk_n, hidden)) + dloss_chunks_3d = jnp.reshape(dloss_chunks, (n_chunks, 1, chunk_n)) + tgt_chunks_3d = jnp.reshape(tgt_chunks, (n_chunks, 1, chunk_n, -1)) + + # Kernel block_n is fixed to 256 + kernel_config = BCEConfig( + block_v=config.block_v, + block_n=256, + compute_metrics=config.compute_metrics, + ) + + def loop_body(carry, x): + d_emb_acc, d_bias_acc = carry + act_chunk, dloss_chunk, tgt_chunk = x + + d_act_chunk, d_emb_contrib, d_bias_contrib = _bce_bwd_pallas( + kernel_config, + dloss_chunk, + act_chunk, + embeddings, + bias, + tgt_chunk, + vocab_offset=vocab_offset, + ) + + return ( + d_emb_acc + d_emb_contrib, + d_bias_acc + d_bias_contrib, + ), d_act_chunk + + init = (jnp.zeros_like(embeddings), jnp.zeros_like(bias)) + (d_embeddings, d_bias), d_act_chunks = jax.lax.scan( + loop_body, + init, + (act_chunks_3d, dloss_chunks_3d, tgt_chunks_3d), + ) + + d_act_padded = jnp.reshape(d_act_chunks, (padded_n, hidden)) + if padded_n > n: + d_act_2d = d_act_padded[:n, :] + else: + d_act_2d = d_act_padded + + d_activations = jnp.reshape(d_act_2d, (batch, seq_len, hidden)) + return d_activations, d_embeddings, d_bias + + +def _bce_bwd_sharded( + config: BCEConfig, + d_loss: jt.Float[jt.Array, "... B N"], + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "... B N L"], +) -> tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Sharding-aware dispatcher for BCE backward.""" + mesh = None + act_spec = None + emb_spec = None + bias_spec = None + + if config.mesh is not None and config.act_spec is not None: + mesh = config.mesh + act_spec = config.act_spec + emb_spec = config.emb_spec + bias_spec = config.bias_spec + act_sharding = jax.sharding.NamedSharding(mesh, act_spec) + is_sharded = True + else: + act_sharding = _get_sharding(activations) + emb_sharding = _get_sharding(embeddings) + bias_sharding = _get_sharding(bias) + is_sharded = ( + isinstance(act_sharding, jax.sharding.NamedSharding) + and not act_sharding.mesh.empty + and isinstance(emb_sharding, jax.sharding.NamedSharding) + and isinstance(bias_sharding, jax.sharding.NamedSharding) + ) + if is_sharded: + act_spec = act_sharding.spec # pyrefly: ignore[missing-attribute] + emb_spec = emb_sharding.spec + bias_spec = bias_sharding.spec + mesh = act_sharding.mesh # pyrefly: ignore[missing-attribute] + + if not is_sharded: + if activations.ndim == 4: + return _bce_bwd_loop_fallback( + config, d_loss, activations, embeddings, bias, targets + ) + return _bce_bwd_pallas_chunked_n( + config, d_loss, activations, embeddings, bias, targets + ) + + _check_vocab_replicated_in_d(emb_spec) # pyrefly: ignore[bad-argument-type] + + hidden_axis_name = act_spec[-1] # pyrefly: ignore[unsupported-operation] + if hidden_axis_name is not None: + replicated_activations = _replicate_hidden_dim(activations) + new_spec_list = list(act_spec) # pyrefly: ignore[bad-argument-type] + new_spec_list[-1] = None + local_act_spec = jax.sharding.PartitionSpec(*new_spec_list) + else: + replicated_activations = activations + local_act_spec = act_spec + + vocab_axis_name = emb_spec[0] # pyrefly: ignore[unsupported-operation] + + dp_axes = [] + for axis in local_act_spec: # pyrefly: ignore[not-iterable] + if axis is not None and axis != vocab_axis_name: + dp_axes.append(axis) + + def _bwd_local_with_reduction(d_loss_, act_, emb_, bias_, tgt_): + if vocab_axis_name is not None: + vocab_offset = jax.lax.axis_index(vocab_axis_name) * emb_.shape[0] + else: + vocab_offset = 0 + if act_.ndim == 4: + d_act, d_emb, d_bias = _bce_bwd_loop_fallback( + config, d_loss_, act_, emb_, bias_, tgt_, vocab_offset=vocab_offset + ) + else: + d_act, d_emb, d_bias = _bce_bwd_pallas_chunked_n( + config, d_loss_, act_, emb_, bias_, tgt_, vocab_offset=vocab_offset + ) + if dp_axes: + d_emb = jax.lax.psum(d_emb, axis_name=dp_axes) + d_bias = jax.lax.psum(d_bias, axis_name=dp_axes) + if vocab_axis_name is not None: + d_act = jax.lax.psum(d_act, axis_name=vocab_axis_name) + return d_act, d_emb, d_bias + + d_act_replicated, d_emb, d_bias = jax.shard_map( + _bwd_local_with_reduction, + mesh=mesh, + in_specs=( + jax.sharding.PartitionSpec(*local_act_spec[:-1]), # pyrefly: ignore[unsupported-operation] + local_act_spec, + emb_spec, + bias_spec, + local_act_spec, + ), + out_specs=( + local_act_spec, + emb_spec, + bias_spec, + ), + check_vma=False, + )( + d_loss, + replicated_activations, + embeddings, + bias, + targets, + ) + + if hidden_axis_name is not None: + d_activations = jax.lax.with_sharding_constraint( + d_act_replicated, act_sharding + ) + else: + d_activations = d_act_replicated + + return d_activations, d_emb, d_bias + + +def _bce_bwd_loop_fallback( + config: BCEConfig, + d_loss: jt.Float[jt.Array, "... B N"], + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "... B N L"], + vocab_offset: int = 0, +) -> tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Sequential loop fallback over group dimension. + + Applicable to activations with an extra dimension reshaped before batch axis. + + Args: + config: Focal BCE config. + d_loss: Gradient of the loss with respect to the logits. + activations: Hidden-state outputs of shape ``[B, N, D]``. + embeddings: Output embedding / unembedding weights of shape ``[V, D]``. + bias: Optional bias of shape ``[V]``. + targets: Target token ids of shape ``[B, N, L]``. + vocab_offset: Vocab offset for the current chunk. + + Returns: + Gradient of the loss with respect to the activations, embeddings, and bias. + """ + groups = activations.shape[0] + d_acts = [] + d_emb = jnp.zeros_like(embeddings) + d_bias = jnp.zeros_like(bias) + for g in range(groups): + tgt_g = targets if targets.ndim == 3 else targets[g] + d_act_g, d_emb_g, d_bias_g = _bce_bwd_pallas_chunked_n( + config, + d_loss[g], + activations[g], + embeddings, + bias, + tgt_g, + vocab_offset, + ) + d_acts.append(d_act_g) + d_emb += d_emb_g + d_bias += d_bias_g + return jnp.stack(d_acts, axis=0), d_emb, d_bias diff --git a/recml/core/ops/binary_cross_entropy_ops_test.py b/recml/core/ops/binary_cross_entropy_ops_test.py new file mode 100644 index 0000000..ef0fa88 --- /dev/null +++ b/recml/core/ops/binary_cross_entropy_ops_test.py @@ -0,0 +1,634 @@ +# Copyright 2024 RecML authors . +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tests for binary_cross_entropy_ops.""" + +import time + +from unittest import mock + +from absl import logging +from absl.testing import absltest +from absl.testing import parameterized +import jax +from jax.experimental.pallas import tpu as pltpu +import jax.numpy as jnp +import keras +import numpy as np +from recml.core.ops import binary_cross_entropy_ops + + +def _naive_bce(activations, embeddings, bias, targets, weights=None): + """Naive implementation that materializes the full logits matrix.""" + vocab_size = embeddings.shape[0] + logits = jnp.matmul(activations, embeddings.T) + bias # (B, N, V) + + # targets: (B, N, L) -> multi_hot: (B, N, V) + one_hot = jax.nn.one_hot(targets, vocab_size, axis=-1) # (B, N, L, V) + multi_hot = jnp.max(one_hot, axis=-2) # (B, N, V) + + # Compute stable BCE loss per class + # Loss = max(x, 0) - x * y + log(1 + exp(-|x|)) + losses = ( + jnp.maximum(logits, 0.0) + - logits * multi_hot + + jnp.log1p(jnp.exp(-jnp.abs(logits))) + ) + loss_per_target = jnp.mean(losses, axis=-1) # (B, N) + + if weights is not None: + loss_per_target = loss_per_target * weights + weight_sum = jnp.sum(weights) + else: + weight_sum = np.prod(targets.shape[:-1]) + + loss = jnp.sum(loss_per_target) / (weight_sum + 1e-8) + return loss, loss_per_target + + +class BinaryCrossEntropyOpsTest(parameterized.TestCase): + + def setUp(self): + super().setUp() + if jax.devices()[0].platform == 'tpu': + + vmem = pltpu.get_tpu_info().vmem_capacity_bytes + logging.info( + 'JETS_DEBUG: VMEM capacity: %d bytes (%.2f MB)', + vmem, + vmem / 1024 / 1024, + ) + + def test_get_sharding(self): + class ObjWithSharding: + sharding = 'dummy_sharding_1' + + class ObjWithAvalSharding: + + class Aval: + sharding = 'dummy_sharding_2' + + aval = Aval() + + class ObjWithAvalWithoutSharding: + + class Aval: + pass + + aval = Aval() + + class ObjWithNoSharding: + pass + + self.assertEqual( + binary_cross_entropy_ops._get_sharding(ObjWithSharding()), + 'dummy_sharding_1', + ) + self.assertEqual( + binary_cross_entropy_ops._get_sharding(ObjWithAvalSharding()), + 'dummy_sharding_2', + ) + self.assertIsNone( + binary_cross_entropy_ops._get_sharding(ObjWithAvalWithoutSharding()) + ) + self.assertIsNone( + binary_cross_entropy_ops._get_sharding(ObjWithNoSharding()) + ) + + @parameterized.named_parameters( + ('standard', 2, 256, 128, 1024, 4, 256), + ('unaligned_seq_len', 2, 130, 128, 1024, 4, 256), + ('unaligned_vocab', 2, 256, 128, 1000, 4, 256), + ('single_label', 2, 256, 128, 1024, 1, 256), + ('small_block_v', 2, 256, 128, 1024, 4, 128), + ) + def test_cut_bce_correctness( + self, batch, seq_len, hidden_dim, vocab_size, num_labels, block_v + ): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + key = jax.random.PRNGKey(0) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + # naive BCE + def run_naive(act, emb, b_val): + loss, _ = _naive_bce(act, emb, b_val, targets) + return loss + + grad_naive_fn = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + loss_naive = run_naive(activations, embeddings, bias) + g_act_naive, g_emb_naive, g_bias_naive = grad_naive_fn( + activations, embeddings, bias + ) + + # cut BCE + def run_cut(act, emb, b_val): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + # Compare + np.testing.assert_allclose(loss_cut, loss_naive, atol=1e-5, rtol=1e-5) + np.testing.assert_allclose(g_act_cut, g_act_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_cut, g_emb_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_bias_cut, g_bias_naive, atol=1e-4, rtol=1e-4) + + @parameterized.named_parameters( + ('4d_act_4d_tgt', (2, 2, 128, 64), (2, 2, 128, 4)), + ('4d_act_3d_tgt', (3, 2, 128, 64), (2, 128, 4)), + ) + def test_cut_bce_4d_correctness(self, act_shape, tgt_shape): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + vocab_size, block_v = 512, 256 + hidden_dim = act_shape[-1] + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, act_shape) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint(key_tgt, tgt_shape, 0, vocab_size) + + def run_naive(act, emb, b_val): + loss, _ = _naive_bce(act, emb, b_val, targets) + return loss + + def run_cut(act, emb, b_val): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + loss_naive = run_naive(activations, embeddings, bias) + loss_cut = run_cut(activations, embeddings, bias) + np.testing.assert_allclose(loss_cut, loss_naive, rtol=1e-3, atol=1e-3) + + # Test backward grad in 4D (exercises _bce_bwd_loop_fallback) + grad_naive_fn = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + + g_act_naive, g_emb_naive, g_bias_naive = grad_naive_fn( + activations, embeddings, bias + ) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + np.testing.assert_allclose(g_act_cut, g_act_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_cut, g_emb_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_bias_cut, g_bias_naive, atol=1e-4, rtol=1e-4) + + def test_cut_bce_with_sharded_embeddings(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + batch, seq_len, hidden_dim, vocab_size, num_labels = 2, 128, 128, 512, 4 + key = jax.random.PRNGKey(0) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + devices = jax.devices() + mesh = jax.sharding.Mesh(np.array(devices), ('devices',)) + act_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec('devices', None, None) + ) + emb_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec('devices', None) + ) + bias_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec('devices',) + ) + activations_sharded = jax.device_put(activations, act_sharding) + embeddings_sharded = jax.device_put(embeddings, emb_sharding) + bias_sharded = jax.device_put(bias, bias_sharding) + + with mock.patch.object( + jax.lax, + 'with_sharding_constraint', + wraps=jax.lax.with_sharding_constraint, + ) as mock_fwd_constraint: + binary_cross_entropy_ops.cut_binary_cross_entropy( + activations_sharded, + embeddings_sharded, + targets, + bias=bias_sharded, + block_v=256, + ) + self.assertEqual(mock_fwd_constraint.call_count, 3) + + def run_cut(act, emb, b_val): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=256 + ) + + with mock.patch.object( + jax.lax, + 'with_sharding_constraint', + wraps=jax.lax.with_sharding_constraint, + ) as mock_sharding_constraint: + grad_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + g_act, g_emb, g_bias = grad_fn( + activations_sharded, embeddings_sharded, bias_sharded + ) + self.assertTrue(mock_sharding_constraint.called) + self.assertIsNotNone(g_act) + self.assertIsNotNone(g_emb) + self.assertIsNotNone(g_bias) + + def test_cut_bce_correctness_large_sequence(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + # Test with sequence length larger than chunk_n to trigger scan loop + batch, seq_len, hidden_dim, vocab_size, num_labels = 2, 2048, 128, 512, 4 + block_v = 256 + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + # naive BCE + def run_naive(act, emb, b_val): + loss, _ = _naive_bce(act, emb, b_val, targets) + return loss + + grad_naive_fn = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + loss_naive = run_naive(activations, embeddings, bias) + g_act_naive, g_emb_naive, g_bias_naive = grad_naive_fn( + activations, embeddings, bias + ) + + # cut BCE + def run_cut(act, emb, b_val): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + # Compare + np.testing.assert_allclose(loss_cut, loss_naive, atol=1e-5, rtol=1e-5) + np.testing.assert_allclose(g_act_cut, g_act_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_cut, g_emb_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_bias_cut, g_bias_naive, atol=1e-4, rtol=1e-4) + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + def test_cut_bce_vs_keras(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + batch, seq_len, hidden_dim, vocab_size, num_labels = 2, 128, 64, 512, 4 + block_v = 256 + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + # Convert targets to multi-hot for Keras + one_hot = jax.nn.one_hot(targets, vocab_size, axis=-1) # (B, N, L, V) + multi_hot = jnp.max(one_hot, axis=-2) # (B, N, V) + + # Keras BCE version + def run_keras(act, emb, b_val): + logits = jnp.matmul(act, emb.T) + b_val # (B, N, V) + loss_per_token = keras.losses.binary_crossentropy( + multi_hot, logits, from_logits=True + ) + return jnp.mean(loss_per_token) + + grad_keras_fn = jax.jit(jax.grad(run_keras, argnums=(0, 1, 2))) + loss_keras = run_keras(activations, embeddings, bias) + g_act_keras, g_emb_keras, g_bias_keras = grad_keras_fn( + activations, embeddings, bias + ) + + # Cut BCE version + def run_cut(act, emb, b_val): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + # Compare + np.testing.assert_allclose(loss_cut, loss_keras, atol=1e-5, rtol=1e-5) + np.testing.assert_allclose(g_act_cut, g_act_keras, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_cut, g_emb_keras, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_bias_cut, g_bias_keras, atol=1e-4, rtol=1e-4) + + def test_benchmark(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + # Large scale benchmark (N = 65536) + b, m, d, v, l = 512, 128, 128, 100000, 4 + block_v = 2048 + + print(f'\n--- Running Benchmark (B={b}, M={m}, D={d}, V={v}, L={l}) ---') + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (b, m, d)) + embeddings = jax.random.normal(key_emb, (v, d)) + bias = jax.random.normal(key_bias, (v,)) + targets = jax.random.randint(key_tgt, (b, m, l), 0, v) + + # --- Naive version (Skip to avoid OOM at large scale) --- + + # --- Pallas TPU custom VJP version (with JAX scan over N) --- + def run_cut(act, emb, b_val, bv=block_v): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=bv + ) + + grad_cut = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + logging.info('Compiling cut (block_v=%d)...', block_v) + t0 = time.time() + grad_cut(activations, embeddings, bias)[0].block_until_ready() + logging.info( + 'Cut (block_v=%d) compiled in %.2f s', + block_v, + time.time() - t0, + ) + + num_steps = 20 + logging.info( + 'Benchmarking cut (block_v=%d) (%d steps)...', + block_v, + num_steps, + ) + t0 = time.time() + for _ in range(num_steps): + g_act_cut, _, _ = grad_cut(activations, embeddings, bias) + g_act_cut.block_until_ready() + t_cut = (time.time() - t0) / num_steps + logging.info( + 'Cut (block_v=%d) step time: %.2f ms', + block_v, + t_cut * 1000, + ) + + def test_benchmark_small_scale(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + b, m, d, v, l = 16, 16, 128, 100000, 4 + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (b, m, d)) + embeddings = jax.random.normal(key_emb, (v, d)) + bias = jax.random.normal(key_bias, (v,)) + targets = jax.random.randint(key_tgt, (b, m, l), 0, v) + + # --- Naive version --- + def run_naive(act, emb, b_val): + loss, _ = _naive_bce(act, emb, b_val, targets) + return loss + + grad_naive = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + logging.info('Compiling naive...') + t0 = time.time() + grad_naive(activations, embeddings, bias)[0].block_until_ready() + logging.info('Naive compiled in %.2f s', time.time() - t0) + + num_steps = 20 + logging.info('Benchmarking naive (%d steps)...', num_steps) + t0 = time.time() + for _ in range(num_steps): + g_act_naive, _, _ = grad_naive(activations, embeddings, bias) + g_act_naive.block_until_ready() + t_naive = (time.time() - t0) / num_steps + logging.info('Naive step time: %.2f ms', t_naive * 1000) + + # --- Cut versions --- + try: + vmem = pltpu.get_tpu_info().vmem_capacity_bytes + except Exception: # pylint: disable=broad-except,broad-exception-caught + vmem = 16 * 1024 * 1024 + + if vmem <= 16 * 1024 * 1024: + block_sizes = [128, 256, 512, 1024, 2048] + else: + block_sizes = [128, 256, 512, 1024, 2048, 4096, 8192] + + for bv in block_sizes: + def run_cut(act, emb, b_val, block_v=bv): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_cut = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + logging.info('Compiling cut (block_v=%d)...', bv) + t0 = time.time() + grad_cut(activations, embeddings, bias)[0].block_until_ready() + logging.info('Cut (block_v=%d) compiled in %.2f s', bv, time.time() - t0) + + logging.info('Benchmarking cut (block_v=%d) (%d steps)...', bv, num_steps) + t0 = time.time() + for _ in range(num_steps): + g_act_cut, _, _ = grad_cut(activations, embeddings, bias) + g_act_cut.block_until_ready() + t_cut = (time.time() - t0) / num_steps + logging.info( + 'Cut (block_v=%d) step time: %.2f ms (Speedup: %.2fx)', + bv, + t_cut * 1000, + t_naive / t_cut, + ) + + def test_benchmark_medium_scale(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + b, m, d, v, l = 192, 128, 128, 20000, 4 + print( + f'\n--- Running Medium Scale Benchmark (B={b}, M={m}, D={d}, V={v},' + f' L={l}) ---' + ) + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (b, m, d)) + embeddings = jax.random.normal(key_emb, (v, d)) + bias = jax.random.normal(key_bias, (v,)) + targets = jax.random.randint(key_tgt, (b, m, l), 0, v) + + # --- Naive version --- + def run_naive(act, emb, b_val): + loss, _ = _naive_bce(act, emb, b_val, targets) + return loss + + grad_naive = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + logging.info('Compiling naive...') + t0 = time.time() + grad_naive(activations, embeddings, bias)[0].block_until_ready() + logging.info('Naive compiled in %.2f s', time.time() - t0) + + num_steps = 20 + logging.info('Benchmarking naive (%d steps)...', num_steps) + t0 = time.time() + for _ in range(num_steps): + g_act_naive, _, _ = grad_naive(activations, embeddings, bias) + g_act_naive.block_until_ready() + t_naive = (time.time() - t0) / num_steps + logging.info('Naive step time: %.2f ms', t_naive * 1000) + + # --- Cut versions --- + block_sizes = [128, 256, 512, 1024, 2048, 4096] + + for bv in block_sizes: + def run_cut(act, emb, b_val, block_v=bv): + return binary_cross_entropy_ops.cut_binary_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_cut = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + logging.info('Compiling cut (block_v=%d)...', bv) + t0 = time.time() + grad_cut(activations, embeddings, bias)[0].block_until_ready() + logging.info('Cut (block_v=%d) compiled in %.2f s', bv, time.time() - t0) + + logging.info('Benchmarking cut (block_v=%d) (%d steps)...', bv, num_steps) + t0 = time.time() + for _ in range(num_steps): + g_act_cut, _, _ = grad_cut(activations, embeddings, bias) + g_act_cut.block_until_ready() + t_cut = (time.time() - t0) / num_steps + logging.info('Cut (block_v=%d) step time: %.2f ms (Speedup: %.2fx)', + bv, t_cut * 1000, t_naive / t_cut) + + def test_cut_bce_block_v_capped_at_vocab(self): + activations = jnp.ones((2, 64, 32)) + embeddings = jnp.ones((200, 32)) + bias = jnp.zeros((200,)) + targets = jnp.zeros((2, 64, 2), dtype=jnp.int32) + + with mock.patch.object( + binary_cross_entropy_ops, + '_cut_binary_cross_entropy', + wraps=binary_cross_entropy_ops._cut_binary_cross_entropy, + ) as mock_fn: + binary_cross_entropy_ops.cut_binary_cross_entropy( + activations, embeddings, targets, bias=bias, block_v=1000 + ) + config = mock_fn.call_args[0][0] + self.assertEqual(config.block_v, 200) + + def test_pallas_vmem_budget(self): + budget = binary_cross_entropy_ops._pallas_vmem_budget() + if any(d.platform == 'tpu' for d in jax.devices()): + expected = max( + 16 * 1024 * 1024, + pltpu.get_tpu_info().vmem_capacity_bytes - 8 * 1024 * 1024, + ) + self.assertEqual(budget, expected) + else: + self.assertEqual(budget, 16 * 1024 * 1024) + + def test_max_safe_block_v(self): + val = binary_cross_entropy_ops._max_safe_block_v(16 * 1024 * 1024, 256) + self.assertEqual(val, 512) + + def test_pallas_interpret(self): + is_interpret = binary_cross_entropy_ops._pallas_interpret() + has_tpu = any(d.platform == 'tpu' for d in jax.devices()) + self.assertEqual(is_interpret, not has_tpu) + + def test_pallas_lane(self): + lane = binary_cross_entropy_ops._pallas_lane() + if any(d.platform == 'tpu' for d in jax.devices()): + self.assertEqual(lane, pltpu.get_tpu_info().num_lanes) + else: + self.assertEqual(lane, 128) + + def test_check_vocab_replicated_in_d(self): + with self.assertRaises(NotImplementedError): + binary_cross_entropy_ops._check_vocab_replicated_in_d( + jax.sharding.PartitionSpec('devices', 'devices') + ) + binary_cross_entropy_ops._check_vocab_replicated_in_d( + jax.sharding.PartitionSpec('devices') + ) + + def test_auto_block_v(self): + num_devices = jax.device_count() + self.assertEqual( + binary_cross_entropy_ops._auto_block_v(3277 * num_devices, 10000), 2432 + ) + self.assertEqual( + binary_cross_entropy_ops._auto_block_v(4096 * num_devices, 10000), 2048 + ) + self.assertEqual( + binary_cross_entropy_ops._auto_block_v(1000000 * num_devices, 10000), + 256, + ) + self.assertEqual( + binary_cross_entropy_ops._auto_block_v(100 * num_devices, 10000), 8192 + ) + self.assertEqual( + binary_cross_entropy_ops._auto_block_v(1024 * num_devices, 500), 500 + ) + + +if __name__ == '__main__': + absltest.main() + +if __name__ == '__main__': + absltest.main() + diff --git a/recml/core/ops/binary_focal_cross_entropy.py b/recml/core/ops/binary_focal_cross_entropy.py new file mode 100644 index 0000000..45e8e42 --- /dev/null +++ b/recml/core/ops/binary_focal_cross_entropy.py @@ -0,0 +1,1018 @@ +# Copyright 2024 RecML authors . +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Memory-efficient JAX operations for binary focal cross-entropy loss. + +Computes exact binary focal cross-entropy loss over large vocabulary +without materializing full [batch, seq_len, vocab_size] logit tensors in HBM. +""" + +import dataclasses +import functools +import jax +from jax.experimental import pallas as pl +from jax.experimental.pallas import tpu as pltpu +import jax.numpy as jnp +import jaxtyping as jt +import numpy as np + +from recml.core.ops import binary_cross_entropy_ops as bce_ops + +EPS = bce_ops.EPS +_auto_block_v = bce_ops._auto_block_v # pylint: disable=protected-access +_check_vocab_replicated_in_d = ( + bce_ops._check_vocab_replicated_in_d # pylint: disable=protected-access +) +_get_sharding = bce_ops._get_sharding # pylint: disable=protected-access +_max_safe_block_v = bce_ops._max_safe_block_v # pylint: disable=protected-access +_max_safe_chunk_n = bce_ops._max_safe_chunk_n # pylint: disable=protected-access +_pallas_interpret = bce_ops._pallas_interpret # pylint: disable=protected-access +_pallas_lane = bce_ops._pallas_lane # pylint: disable=protected-access +_pallas_vmem_budget = ( + bce_ops._pallas_vmem_budget # pylint: disable=protected-access +) +_replicate_hidden_dim = ( + bce_ops._replicate_hidden_dim # pylint: disable=protected-access +) + + +@dataclasses.dataclass +class FocalBCEConfig: + """Configuration for the binary focal cross-entropy loss.""" + + block_v: int + block_n: int = 256 + compute_metrics: bool = False + gamma: float = 2.0 + alpha: float = 0.25 + apply_class_balancing: 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 + bias_spec: jax.sharding.PartitionSpec | None = None + + +def _focal_bce_fwd_local( + config: FocalBCEConfig, + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], +) -> tuple[ + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], +]: + """Computes the sum of FocalLoss(x_v, target_v) over all V block-wise, and metrics.""" + block_v = config.block_v + batch, seq_len, hidden = activations.shape + vocab = embeddings.shape[0] + + n = batch * seq_len + v_blocks = int(np.ceil(vocab / block_v)) + + activations_2d = jnp.reshape(activations, (n, hidden)) + targets_2d = jnp.reshape(targets, (n, -1)) + + if config.compute_metrics: + + def v_body( + carry: tuple[jt.Float[jt.Array, "N"], ...], + j: jt.Int[jt.Array, ""], + ) -> tuple[tuple[jt.Float[jt.Array, "N"], ...], None]: + loss_acc, tp_acc, fp_acc, fn_acc, tn_acc = carry + actual_start = jnp.maximum(0, jnp.minimum(j * block_v, vocab - block_v)) + emb_chunk = jax.lax.dynamic_slice_in_dim( + embeddings, actual_start, block_v + ) + bias_chunk = jax.lax.dynamic_slice_in_dim(bias, actual_start, block_v) + logits = ( + jnp.matmul( + activations_2d, + jnp.transpose(emb_chunk), + preferred_element_type=jnp.float32, + ) + + bias_chunk + ) + + chunk_indices = actual_start + jnp.arange(block_v) + valid_mask = (chunk_indices >= j * block_v) & (chunk_indices < vocab) + + targets_match = targets_2d[:, :, None] == chunk_indices[None, None, :] + targets_chunk = jnp.any(targets_match, axis=1) + targets_float = targets_chunk.astype(logits.dtype) + + probs = jax.nn.sigmoid(logits) + p_t = targets_float * probs + (1.0 - targets_float) * (1.0 - probs) + focal_factor = jnp.power(1.0 - p_t, config.gamma) + + # Fused BCE Loss: BCE(x, y) = BCE(x, 0) - y * x + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p( + jnp.exp(-jnp.abs(logits)) + ) + bce_loss_chunk = loss_zero - targets_float * logits + loss_chunk = focal_factor * bce_loss_chunk + if config.apply_class_balancing: + weight = targets_float * config.alpha + (1.0 - targets_float) * ( + 1.0 - config.alpha + ) + loss_chunk = weight * loss_chunk + + loss_chunk = loss_chunk * valid_mask[None, :] + loss_sum = jnp.sum(loss_chunk, axis=-1) + + predictions_chunk = logits > 0.0 + + tp_chunk = targets_chunk & predictions_chunk + fp_chunk = (~targets_chunk) & predictions_chunk + fn_chunk = targets_chunk & (~predictions_chunk) + tn_chunk = (~targets_chunk) & (~predictions_chunk) + + tp_chunk = tp_chunk & valid_mask[None, :] + fp_chunk = fp_chunk & valid_mask[None, :] + fn_chunk = fn_chunk & valid_mask[None, :] + tn_chunk = tn_chunk & valid_mask[None, :] + + tp_sum = jnp.sum(tp_chunk, axis=-1).astype(jnp.float32) + fp_sum = jnp.sum(fp_chunk, axis=-1).astype(jnp.float32) + fn_sum = jnp.sum(fn_chunk, axis=-1).astype(jnp.float32) + tn_sum = jnp.sum(tn_chunk, axis=-1).astype(jnp.float32) + + return ( + loss_acc + loss_sum, + tp_acc + tp_sum, + fp_acc + fp_sum, + fn_acc + fn_sum, + tn_acc + tn_sum, + ), None + + init = ( + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + jnp.zeros((n,), dtype=jnp.float32), + ) + (loss_final, tp_final, fp_final, fn_final, tn_final), _ = jax.lax.scan( + jax.checkpoint(v_body), init, jnp.arange(v_blocks) + ) + return ( + jnp.reshape(loss_final, (batch, seq_len)), + jnp.reshape(tp_final, (batch, seq_len)), + jnp.reshape(fp_final, (batch, seq_len)), + jnp.reshape(fn_final, (batch, seq_len)), + jnp.reshape(tn_final, (batch, seq_len)), + ) + else: + + def v_body_no_metrics( + loss_acc: jt.Float[jt.Array, "N"], + j: jt.Int[jt.Array, ""], + ) -> tuple[jt.Float[jt.Array, "N"], None]: + actual_start = jnp.maximum(0, jnp.minimum(j * block_v, vocab - block_v)) + emb_chunk = jax.lax.dynamic_slice_in_dim( + embeddings, actual_start, block_v + ) + bias_chunk = jax.lax.dynamic_slice_in_dim(bias, actual_start, block_v) + logits = ( + jnp.matmul( + activations_2d, + jnp.transpose(emb_chunk), + preferred_element_type=jnp.float32, + ) + + bias_chunk + ) + + chunk_indices = actual_start + jnp.arange(block_v) + valid_mask = (chunk_indices >= j * block_v) & (chunk_indices < vocab) + + targets_match = targets_2d[:, :, None] == chunk_indices[None, None, :] + targets_chunk = jnp.any(targets_match, axis=1) + targets_float = targets_chunk.astype(logits.dtype) + + probs = jax.nn.sigmoid(logits) + p_t = targets_float * probs + (1.0 - targets_float) * (1.0 - probs) + focal_factor = jnp.power(1.0 - p_t, config.gamma) + + # Fused BCE Loss: BCE(x, y) = BCE(x, 0) - y * x + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p( + jnp.exp(-jnp.abs(logits)) + ) + bce_loss_chunk = loss_zero - targets_float * logits + loss_chunk = focal_factor * bce_loss_chunk + if config.apply_class_balancing: + weight = targets_float * config.alpha + (1.0 - targets_float) * ( + 1.0 - config.alpha + ) + loss_chunk = weight * loss_chunk + + loss_chunk = loss_chunk * valid_mask[None, :] + loss_sum = jnp.sum(loss_chunk, axis=-1) + return loss_acc + loss_sum, None + + init = jnp.zeros((n,), dtype=jnp.float32) + loss_final, _ = jax.lax.scan( + jax.checkpoint(v_body_no_metrics), init, jnp.arange(v_blocks) + ) + dummy = jnp.zeros((batch, seq_len), dtype=jnp.float32) + return ( + jnp.reshape(loss_final, (batch, seq_len)), + dummy, + dummy, + dummy, + dummy, + ) + + +@functools.partial(jax.custom_vjp, nondiff_argnums=(0,)) +def _cut_binary_focal_cross_entropy( + config: FocalBCEConfig, + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "... B N L"], +) -> tuple[ + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], +]: + """Computes the non-differentiable path of cut Focal BCE loss and metrics.""" + outputs, _ = _cut_binary_focal_cross_entropy_fwd( + config, activations, embeddings, bias, targets + ) + return outputs + + +def _cut_binary_focal_cross_entropy_fwd( + config: FocalBCEConfig, + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], +) -> tuple[ + tuple[ + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + jt.Float[jt.Array, "B N"], + ], + tuple[ + jt.Float[jt.Array, "B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], + jt.Int[jt.Array, "B N L"], + ], +]: + """Computes forward mode of cut Focal BCE loss.""" + replicated_activations = _replicate_hidden_dim(activations) + if activations.ndim == 4: + if targets.ndim == 3: + targets_in_axis = None + else: + targets_in_axis = 0 + fwd_vmap = jax.vmap( + functools.partial(_focal_bce_fwd_local, config), + in_axes=(0, None, None, targets_in_axis), + ) + loss_y0, tp, fp, fn, tn = fwd_vmap( + replicated_activations, embeddings, bias, targets + ) + else: + loss_y0, tp, fp, fn, tn = _focal_bce_fwd_local( + config, replicated_activations, embeddings, bias, targets + ) + vocab_size = embeddings.shape[0] + losses = loss_y0 / vocab_size + + return (losses, tp, fp, fn, tn), ( + activations, + embeddings, + bias, + targets, + ) + + +def _cut_binary_focal_cross_entropy_bwd( + config: FocalBCEConfig, + res: tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], + jt.Int[jt.Array, "... B N L"], + ], + d_outputs: tuple[ + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, "... B N"], + ], +) -> tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], + None, +]: + """Computes the backward mode of cut Focal BCE loss.""" + d_losses, _, _, _, _ = d_outputs + activations, embeddings, bias, targets = res + d_activations, d_embeddings, d_bias = _focal_bce_bwd_sharded( + config, d_losses, activations, embeddings, bias, targets + ) + return d_activations, d_embeddings, d_bias, None + + +_cut_binary_focal_cross_entropy.defvjp( + _cut_binary_focal_cross_entropy_fwd, _cut_binary_focal_cross_entropy_bwd +) + + +def cut_binary_focal_cross_entropy( + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + targets: jt.Int[jt.Array, "... B N L"], + bias: jt.Float[jt.Array, "V"] | None = None, + weights: jt.Float[jt.Array, "... B N"] | None = None, + *, + gamma: float = 2.0, + alpha: float = 0.25, + apply_class_balancing: bool = False, + return_per_target_losses: bool = False, + return_metrics: bool = False, + block_v: int | None = None, + mesh: jax.sharding.Mesh | None = None, + act_spec: jax.sharding.PartitionSpec | None = None, + emb_spec: jax.sharding.PartitionSpec | None = None, + bias_spec: jax.sharding.PartitionSpec | None = None, +) -> ( + jt.Float[jt.Array, ""] + | tuple[jt.Float[jt.Array, ""], jt.Float[jt.Array, "... B N"]] + | tuple[ + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + ] + | tuple[ + jt.Float[jt.Array, ""], + jt.Float[jt.Array, "... B N"], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + jt.Float[jt.Array, ""], + ] +): + """Computes binary focal cross entropy loss over unmaterialized logits.""" + vocab_size = embeddings.shape[0] + if bias is None: + bias = jnp.zeros(vocab_size, dtype=embeddings.dtype) + + sharding = _get_sharding(embeddings) + if ( + isinstance(sharding, jax.sharding.NamedSharding) + and not sharding.mesh.empty + ): + replicated_sharding = jax.sharding.NamedSharding( + sharding.mesh, jax.sharding.PartitionSpec() + ) + embeddings = jax.lax.with_sharding_constraint( + embeddings, replicated_sharding + ) + bias = jax.lax.with_sharding_constraint(bias, replicated_sharding) + + if block_v is None: + n = activations.shape[-3] * activations.shape[-2] + block_v = _auto_block_v(n, vocab_size) + else: + block_v = min(block_v, vocab_size) + + losses, tp, fp, fn, tn = _cut_binary_focal_cross_entropy( + FocalBCEConfig( + block_v=block_v, + compute_metrics=return_metrics, + gamma=gamma, + alpha=alpha, + apply_class_balancing=apply_class_balancing, + mesh=mesh, + act_spec=act_spec, + emb_spec=emb_spec, + bias_spec=bias_spec, + ), + activations, + embeddings, + bias, + targets, + ) + + if weights is not None: + losses = losses * weights + weight_sum = jnp.sum(weights) + tp_sum = jnp.sum(tp * weights) + fp_sum = jnp.sum(fp * weights) + fn_sum = jnp.sum(fn * weights) + tn_sum = jnp.sum(tn * weights) + else: + weight_sum = np.prod(targets.shape[:-1]) + tp_sum = jnp.sum(tp) + fp_sum = jnp.sum(fp) + fn_sum = jnp.sum(fn) + tn_sum = jnp.sum(tn) + + loss = jnp.sum(losses) / (weight_sum + EPS) + + if return_metrics: + if return_per_target_losses: + return loss, losses, tp_sum, fp_sum, fn_sum, tn_sum + return loss, tp_sum, fp_sum, fn_sum, tn_sum + + if return_per_target_losses: + return loss, losses + + return loss + + +cut_binary_cross_entropy = cut_binary_focal_cross_entropy + + +def _focal_bce_bwd_kernel( + emb_ref, # [block_v, 128] VMEM + act_ref, # [padded_n, 256] HBM + bias_ref, # [8, block_v] VMEM + tgt_ref, # [l_padded, padded_n] HBM + dloss_ref, # [l_padded, padded_n] HBM + d_emb_ref, # [block_v, 128] HBM output + d_bias_ref, # [block_v, 128] HBM output + d_act_partials_ref, # [1, padded_n, 256] HBM output + d_emb_scratch, # [block_v, 128] VMEM scratch + d_bias_scratch, # [block_v] VMEM scratch + *, + block_v: int, + block_n: int, + n_blocks: int, + vocab: int, + vocab_offset: int, + labels: int, + n_real: int, + gamma: float, + alpha: float, + apply_class_balancing: bool, +): + """Pallas TPU per-shard chunked backward for Focal BCE.""" + v_idx = pl.program_id(0) + + # Initialize accumulators in VMEM to 0 + d_emb_scratch[...] = jnp.zeros(d_emb_scratch.shape, jnp.float32) + d_bias_scratch[...] = jnp.zeros(d_bias_scratch.shape, jnp.float32) + + v_start_local = v_idx * block_v + v_start_global = v_start_local + vocab_offset + + # Valid vocab mask + chunk_indices_local = v_start_local + jnp.arange(block_v)[None, :] + valid_vocab_mask = chunk_indices_local < vocab + + # Load inputs from VMEM + emb = emb_ref[...] + emb = jnp.where(valid_vocab_mask.T, emb, 0.0) + + bias_val = bias_ref[...] + bias = bias_val[0, :] + bias = jnp.where(valid_vocab_mask[0], bias, 0.0) + + # Broadcasted iota for target matching + col_ids = jax.lax.broadcasted_iota(jnp.int32, (block_n, block_v), 1) + col_ids_global = col_ids + v_start_global + + def loop_body(n_idx, _): + # Load act, tgt, dloss slices for this n_idx from HBM to VMEM + act = act_ref[pl.ds(n_idx * block_n, block_n), :] + tgt_val = tgt_ref[:, pl.ds(n_idx * block_n, block_n)] + dloss_val = dloss_ref[:, pl.ds(n_idx * block_n, block_n)] + + # Slice real dloss + dloss = dloss_val[0, :] + dloss = dloss[:, None] + + # Force loading of padded columns/rows by computing a dummy sum + dummy_sum = ( + jnp.sum(tgt_val).astype(jnp.float32) + + jnp.sum(dloss_val) + + jnp.sum(bias_val) + ) + + # Compute logits + logits = ( + jax.lax.dot_general( + act, + emb, + (((1,), (1,)), ((), ())), + precision=jax.lax.Precision.DEFAULT, + preferred_element_type=jnp.float32, + ) + + bias[None, :] + + dummy_sum * 0.0 + ) + probs = jax.nn.sigmoid(logits) + + # Target matching + y_true_chunk = jnp.zeros((block_n, block_v), dtype=jnp.bool_) + for l in range(labels): + target_l = tgt_val[l, :] + target_l = target_l[:, None] + is_valid_target = target_l >= 0 + match = (target_l == col_ids_global) & is_valid_target + y_true_chunk = y_true_chunk | match + + # Valid batch mask for this loop step + batch_indices_local = n_idx * block_n + jnp.arange(block_n)[:, None] + 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 + p_t = y_true_float * probs + (1.0 - y_true_float) * (1.0 - probs) + focal_factor = jnp.power(1.0 - p_t, gamma) + focal_factor_m1 = jnp.where( + gamma == 0.0, + 0.0, + jnp.power(1.0 - p_t, jnp.maximum(0.0, gamma - 1.0)), + ) + loss_zero = jnp.maximum(logits, 0.0) + jnp.log1p( + jnp.exp(-jnp.abs(logits)) + ) + bce_loss_chunk = loss_zero - y_true_float * logits + + g_focal = g_bce * ( + focal_factor + gamma * focal_factor_m1 * p_t * bce_loss_chunk + ) + if apply_class_balancing: + weight = y_true_float * alpha + (1.0 - y_true_float) * (1.0 - alpha) + g_focal = weight * g_focal + + g = jnp.where(valid_vocab_mask & valid_batch_mask, g_focal, 0.0) + + # Scale + deriv = (dloss / vocab) * g + + # Accumulate d_emb + d_emb_contrib = jax.lax.dot_general( + deriv.astype(act.dtype), + act, + (((0,), (0,)), ((), ())), + precision=jax.lax.Precision.DEFAULT, + preferred_element_type=jnp.float32, + ) + d_emb_scratch[...] = d_emb_scratch[...] + d_emb_contrib + + # Accumulate d_bias + d_bias_contrib = jnp.sum(deriv, axis=0) + d_bias_scratch[...] = d_bias_scratch[...] + d_bias_contrib + + # Compute d_act and write directly to HBM + d_act_contrib = jax.lax.dot_general( + deriv.astype(emb.dtype), + emb, + (((1,), (0,)), ((), ())), + precision=jax.lax.Precision.DEFAULT, + preferred_element_type=jnp.float32, + ) + d_act_partials_ref[0, pl.ds(n_idx * block_n, block_n), :] = d_act_contrib + + return None + + # Run reduction loop over n_blocks + jax.lax.fori_loop(0, n_blocks, loop_body, None) + + # Store accumulated results to HBM + d_emb_ref[...] = d_emb_scratch[...] + d_bias_ref[...] = d_bias_scratch[...][..., None] * ( + jnp.arange(d_bias_ref.shape[1]) == 0 + ) + + +def _focal_bce_bwd_pallas_chunked_n( + config: FocalBCEConfig, + d_loss: jt.Float[jt.Array, "B N"], + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], + vocab_offset: int = 0, +) -> tuple[ + jt.Float[jt.Array, "B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Wraps _focal_bce_bwd_pallas by chunking the N (sequence) dimension in JAX.""" + batch, seq_len, hidden = activations.shape + n = batch * seq_len + + padded_d = ((hidden + 127) // 128) * 128 + vmem_budget = _pallas_vmem_budget() + chunk_n = _max_safe_chunk_n(vmem_budget, padded_d) + if n <= chunk_n: + return _focal_bce_bwd_pallas( + config, d_loss, activations, embeddings, bias, targets, vocab_offset + ) + activations_2d = jnp.reshape(activations, (n, hidden)) + dloss_2d = jnp.reshape(d_loss, (n, 1)) + targets_2d = jnp.reshape(targets, (n, -1)) + + n_chunks = (n + chunk_n - 1) // chunk_n + padded_n = n_chunks * chunk_n + + if padded_n > n: + activations_2d = jnp.pad(activations_2d, ((0, padded_n - n), (0, 0))) + dloss_2d = jnp.pad(dloss_2d, ((0, padded_n - n), (0, 0))) + targets_2d = jnp.pad( + targets_2d, ((0, padded_n - n), (0, 0)), constant_values=-1 + ) + + act_chunks_3d = jnp.reshape(activations_2d, (n_chunks, 1, chunk_n, hidden)) + dloss_chunks_3d = jnp.reshape(dloss_2d, (n_chunks, 1, chunk_n)) + tgt_chunks_3d = jnp.reshape(targets_2d, (n_chunks, 1, chunk_n, -1)) + + kernel_config = FocalBCEConfig( + block_v=config.block_v, + block_n=256, + compute_metrics=config.compute_metrics, + gamma=config.gamma, + alpha=config.alpha, + apply_class_balancing=config.apply_class_balancing, + ) + + def loop_body(carry, x): + d_emb_acc, d_bias_acc = carry + act_chunk, dloss_chunk, tgt_chunk = x + + d_act_chunk, d_emb_contrib, d_bias_contrib = _focal_bce_bwd_pallas( + kernel_config, + dloss_chunk, + act_chunk, + embeddings, + bias, + tgt_chunk, + vocab_offset=vocab_offset, + ) + + return ( + d_emb_acc + d_emb_contrib, + d_bias_acc + d_bias_contrib, + ), d_act_chunk + + init = (jnp.zeros_like(embeddings), jnp.zeros_like(bias)) + (d_embeddings, d_bias), d_act_chunks = jax.lax.scan( + loop_body, init, (act_chunks_3d, dloss_chunks_3d, tgt_chunks_3d) + ) + + d_act_flat = jnp.reshape(d_act_chunks, (padded_n, hidden))[:n] + d_act = jnp.reshape(d_act_flat, (batch, seq_len, hidden)) + return d_act, d_embeddings, d_bias + + +def _focal_bce_bwd_pallas( + config: FocalBCEConfig, + d_loss: jt.Float[jt.Array, "B N"], + activations: jt.Float[jt.Array, "B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "B N L"], + vocab_offset: int = 0, +) -> tuple[ + jt.Float[jt.Array, "B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Pallas TPU per-shard chunked backward for Focal BCE.""" + block_n = config.block_n + batch, seq_len, hidden = activations.shape + padded_d = ((hidden + 127) // 128) * 128 + vocab = embeddings.shape[0] + vmem_budget = _pallas_vmem_budget() + max_safe = _max_safe_block_v(vmem_budget, padded_d) + vocab_padded_128 = ((vocab + 127) // 128) * 128 + block_v = min(config.block_v, vocab, max_safe) + block_v = ((block_v + 127) // 128) * 128 + block_v = min(block_v, vocab_padded_128, max_safe) + n = batch * seq_len + 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] + + # Pad activations, targets, d_loss to multiples of block_n if necessary + padded_n = n_blocks * block_n + if padded_n > n: + pad_len = padded_n - n + activations_2d = jnp.pad( + jnp.reshape(activations, (n, hidden)), ((0, pad_len), (0, 0)) + ) + dloss_2d = jnp.pad(jnp.reshape(d_loss, (n, 1)), ((0, pad_len), (0, 0))) + targets_2d = jnp.pad( + jnp.reshape(targets, (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)) + # Pad activations and embeddings to padded_d columns + if padded_d > hidden: + activations_padded = jnp.pad( + activations_2d, ((0, 0), (0, padded_d - hidden)) + ) + else: + activations_padded = activations_2d + + if vocab_padded > vocab or padded_d > hidden: + embeddings_padded = jnp.pad( + embeddings, ((0, vocab_padded - vocab), (0, padded_d - hidden)) + ) + else: + embeddings_padded = embeddings + + # Transpose and pad targets: (n, labels) -> (labels, n) -> + # (l_padded, padded_n) + l_padded = max(8, ((labels + 7) // 8) * 8) + targets_t = jnp.transpose(targets_2d, (1, 0)) + if l_padded > labels or padded_n > n: + targets_padded = jnp.pad( + targets_t, + ((0, l_padded - labels), (0, padded_n - n)), + constant_values=-1, + ).astype(jnp.int32) + else: + targets_padded = targets_t.astype(jnp.int32) + + # Transpose and pad dloss: (n, 1) -> (1, n) -> (l_padded, padded_n) + dloss_t = jnp.transpose(dloss_2d, (1, 0)) + dloss_padded = jnp.pad( + dloss_t, ((0, l_padded - 1), (0, padded_n - n)) + ).astype(embeddings.dtype) + + # Reshape, transpose and pad bias: (vocab,) -> (1, vocab) -> (8, vocab_padded) + bias_t = jnp.reshape(bias, (1, vocab)) + if vocab_padded > vocab: + bias_padded = jnp.pad(bias_t, ((0, 7), (0, vocab_padded - vocab))).astype( + bias.dtype + ) + else: + bias_padded = jnp.pad(bias_t, ((0, 7), (0, 0))).astype(bias.dtype) + + d_emb_padded, d_bias_padded, d_act_partials = pl.pallas_call( + functools.partial( + _focal_bce_bwd_kernel, + block_v=block_v, + block_n=block_n, + n_blocks=n_blocks, + vocab=vocab, + vocab_offset=vocab_offset, + labels=labels, + n_real=n, + gamma=config.gamma, + alpha=config.alpha, + apply_class_balancing=config.apply_class_balancing, + ), + out_shape=[ + jax.ShapeDtypeStruct((vocab_padded, padded_d), embeddings.dtype), + jax.ShapeDtypeStruct((vocab_padded, 8), bias.dtype), + jax.ShapeDtypeStruct( + (v_blocks, padded_n, padded_d), activations.dtype + ), + ], + grid=(v_blocks,), + in_specs=[ + pl.BlockSpec((block_v, padded_d), lambda v: (v, 0)), # emb + pl.BlockSpec((padded_n, padded_d), lambda v: (0, 0)), # act + pl.BlockSpec((8, block_v), lambda v: (0, v)), # bias + pl.BlockSpec((l_padded, padded_n), lambda v: (0, 0)), # tgt + pl.BlockSpec((l_padded, padded_n), lambda v: (0, 0)), # dloss + ], + out_specs=[ + pl.BlockSpec((block_v, padded_d), lambda v: (v, 0)), # d_emb + pl.BlockSpec((block_v, 8), lambda v: (v, 0)), # d_bias + pl.BlockSpec((1, padded_n, padded_d), lambda v: (v, 0, 0)), # d_act + ], + scratch_shapes=[ + pltpu.VMEM((block_v, padded_d), jnp.float32), # d_emb_scratch + pltpu.VMEM((block_v,), jnp.float32), # d_bias_scratch + ], + compiler_params=pltpu.CompilerParams( + dimension_semantics=("parallel",), + vmem_limit_bytes=_pallas_vmem_budget(), + ), + interpret=_pallas_interpret(), + )( + embeddings_padded, + activations_padded, + bias_padded, + targets_padded, + dloss_padded, + ) + + if vocab < vocab_padded or hidden < padded_d: + d_emb = d_emb_padded[:vocab, :hidden] + else: + d_emb = d_emb_padded + d_bias = d_bias_padded[:vocab, 0] + + if hidden < padded_d: + d_act_2d = jnp.sum(d_act_partials[:, :, :hidden], axis=0) + else: + d_act_2d = jnp.sum(d_act_partials, axis=0) + + if padded_n > n: + d_act_2d = d_act_2d[:n, :] + d_activations = jnp.reshape(d_act_2d, (batch, seq_len, hidden)) + return d_activations, d_emb, d_bias + + +def _focal_bce_bwd_sharded( + config: FocalBCEConfig, + d_loss: jt.Float[jt.Array, "... B N"], + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "... B N L"], +) -> tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Sharding-aware dispatcher for Focal BCE backward.""" + mesh = None + act_spec = None + emb_spec = None + bias_spec = None + + if config.mesh is not None and config.act_spec is not None: + mesh = config.mesh + act_spec = config.act_spec + emb_spec = config.emb_spec + bias_spec = config.bias_spec + act_sharding = jax.sharding.NamedSharding(mesh, act_spec) + is_sharded = True + else: + act_sharding = _get_sharding(activations) + emb_sharding = _get_sharding(embeddings) + bias_sharding = _get_sharding(bias) + is_sharded = ( + isinstance(act_sharding, jax.sharding.NamedSharding) + and not act_sharding.mesh.empty + and isinstance(emb_sharding, jax.sharding.NamedSharding) + and isinstance(bias_sharding, jax.sharding.NamedSharding) + ) + if is_sharded: + act_spec = act_sharding.spec # pyrefly: ignore[missing-attribute] + emb_spec = emb_sharding.spec + bias_spec = bias_sharding.spec + mesh = act_sharding.mesh # pyrefly: ignore[missing-attribute] + + if not is_sharded: + if activations.ndim == 4: + return _focal_bce_bwd_loop_fallback( + config, d_loss, activations, embeddings, bias, targets + ) + return _focal_bce_bwd_pallas_chunked_n( + config, d_loss, activations, embeddings, bias, targets + ) + + _check_vocab_replicated_in_d(emb_spec) # pyrefly: ignore[bad-argument-type] + + hidden_axis_name = act_spec[-1] # pyrefly: ignore[unsupported-operation] + if hidden_axis_name is not None: + replicated_activations = _replicate_hidden_dim(activations) + new_spec_list = list(act_spec) # pyrefly: ignore[bad-argument-type] + new_spec_list[-1] = None + local_act_spec = jax.sharding.PartitionSpec(*new_spec_list) + else: + replicated_activations = activations + local_act_spec = act_spec + + vocab_axis_name = emb_spec[0] # pyrefly: ignore[unsupported-operation] + + dp_axes = [] + for axis in local_act_spec: # pyrefly: ignore[not-iterable] + if axis is not None and axis != vocab_axis_name: + dp_axes.append(axis) + + def _bwd_local_with_reduction(d_loss_, act_, emb_, bias_, tgt_): + if vocab_axis_name is not None: + vocab_offset = jax.lax.axis_index(vocab_axis_name) * emb_.shape[0] + else: + vocab_offset = 0 + if act_.ndim == 4: + d_act, d_emb, d_bias = _focal_bce_bwd_loop_fallback( + config, d_loss_, act_, emb_, bias_, tgt_, vocab_offset=vocab_offset + ) + else: + d_act, d_emb, d_bias = _focal_bce_bwd_pallas_chunked_n( + config, d_loss_, act_, emb_, bias_, tgt_, vocab_offset=vocab_offset + ) + if dp_axes: + d_emb = jax.lax.psum(d_emb, axis_name=dp_axes) + d_bias = jax.lax.psum(d_bias, axis_name=dp_axes) + if vocab_axis_name is not None: + d_act = jax.lax.psum(d_act, axis_name=vocab_axis_name) + return d_act, d_emb, d_bias + + d_act_replicated, d_emb, d_bias = jax.shard_map( + _bwd_local_with_reduction, + mesh=mesh, + in_specs=( + jax.sharding.PartitionSpec(*local_act_spec[:-1]), # pyrefly: ignore[unsupported-operation] + local_act_spec, + emb_spec, + bias_spec, + local_act_spec, + ), + out_specs=( + local_act_spec, + emb_spec, + bias_spec, + ), + check_vma=False, + )( + d_loss, + replicated_activations, + embeddings, + bias, + targets, + ) + + if hidden_axis_name is not None: + d_activations = jax.lax.with_sharding_constraint( + d_act_replicated, act_sharding + ) + else: + d_activations = d_act_replicated + + return d_activations, d_emb, d_bias + + +def _focal_bce_bwd_loop_fallback( + config: FocalBCEConfig, + d_loss: jt.Float[jt.Array, "... B N"], + activations: jt.Float[jt.Array, "... B N D"], + embeddings: jt.Float[jt.Array, "V D"], + bias: jt.Float[jt.Array, "V"], + targets: jt.Int[jt.Array, "... B N L"], + vocab_offset: int = 0, +) -> tuple[ + jt.Float[jt.Array, "... B N D"], + jt.Float[jt.Array, "V D"], + jt.Float[jt.Array, "V"], +]: + """Sequential loop fallback over group dimension. + + Applicable to activations with an extra dimension reshaped before batch axis. + + Args: + config: Focal BCE config. + d_loss: Gradient of the loss with respect to the logits. + activations: Hidden-state outputs of shape ``[B, N, D]``. + embeddings: Output embedding / unembedding weights of shape ``[V, D]``. + bias: Optional bias of shape ``[V]``. + targets: Target token ids of shape ``[B, N, L]``. + vocab_offset: Vocab offset for the current chunk. + + Returns: + Gradient of the loss with respect to the activations, embeddings, and bias. + """ + groups = activations.shape[0] + d_acts = [] + d_emb = jnp.zeros_like(embeddings) + d_bias = jnp.zeros_like(bias) + for g in range(groups): + tgt_g = targets if targets.ndim == 3 else targets[g] + d_act_g, d_emb_g, d_bias_g = _focal_bce_bwd_pallas_chunked_n( + config, + d_loss[g], + activations[g], + embeddings, + bias, + tgt_g, + vocab_offset, + ) + d_acts.append(d_act_g) + d_emb += d_emb_g + d_bias += d_bias_g + return jnp.stack(d_acts, axis=0), d_emb, d_bias diff --git a/recml/core/ops/binary_focal_cross_entropy_test.py b/recml/core/ops/binary_focal_cross_entropy_test.py new file mode 100644 index 0000000..8610360 --- /dev/null +++ b/recml/core/ops/binary_focal_cross_entropy_test.py @@ -0,0 +1,440 @@ +# Copyright 2024 RecML authors . +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Tests for binary_focal_cross_entropy.""" + +from unittest import mock + +from absl import logging +from absl.testing import absltest +from absl.testing import parameterized +import jax +from jax.experimental.pallas import tpu as pltpu +import jax.numpy as jnp +import keras +import numpy as np +from recml.core.ops import binary_focal_cross_entropy + + +def _naive_focal_bce( + activations, + embeddings, + bias, + targets, + weights=None, + gamma=2.0, + alpha=0.25, + apply_class_balancing=False, +): + """Naive implementation that materializes the full logits matrix for focal loss.""" + vocab_size = embeddings.shape[0] + logits = jnp.matmul(activations, embeddings.T) + bias # (B, N, V) + + # targets: (B, N, L) -> multi_hot: (B, N, V) + one_hot = jax.nn.one_hot(targets, vocab_size, axis=-1) # (B, N, L, V) + multi_hot = jnp.max(one_hot, axis=-2) # (B, N, V) + + probs = jax.nn.sigmoid(logits) + p_t = multi_hot * probs + (1.0 - multi_hot) * (1.0 - probs) + focal_factor = jnp.power(1.0 - p_t, gamma) + + # Compute stable BCE loss per class + # Loss = max(x, 0) - x * y + log(1 + exp(-|x|)) + bce_losses = ( + jnp.maximum(logits, 0.0) + - logits * multi_hot + + jnp.log1p(jnp.exp(-jnp.abs(logits))) + ) + if gamma == 0.0: + losses = bce_losses + else: + losses = focal_factor * bce_losses + if apply_class_balancing: + weight = multi_hot * alpha + (1.0 - multi_hot) * (1.0 - alpha) + losses = weight * losses + + loss_per_target = jnp.mean(losses, axis=-1) + + if weights is not None: + loss_per_target = loss_per_target * weights + weight_sum = jnp.sum(weights) + else: + weight_sum = np.prod(targets.shape[:-1]) + + loss = jnp.sum(loss_per_target) / (weight_sum + 1e-7) + return loss, loss_per_target + + +class BinaryFocalCrossEntropyTest(parameterized.TestCase): + + def setUp(self): + super().setUp() + if jax.devices()[0].platform == 'tpu': + vmem = pltpu.get_tpu_info().vmem_capacity_bytes + logging.info( + 'JETS_DEBUG: VMEM capacity: %d bytes (%.2f MB)', + vmem, + vmem / 1024 / 1024, + ) + + @parameterized.named_parameters( + ('standard_gamma2', 2, 256, 128, 1024, 4, 256, 2.0, 0.25, True), + ('no_balancing', 2, 256, 128, 1024, 4, 256, 2.0, 0.25, False), + ('gamma0', 2, 256, 128, 1024, 4, 256, 0.0, 0.25, True), + ('unaligned_vocab', 2, 256, 128, 1000, 4, 256, 2.0, 0.25, True), + ('unaligned_seq_len', 2, 130, 128, 1024, 4, 256, 1.5, 0.25, True), + ('single_label', 2, 256, 128, 1024, 1, 256, 2.0, 0.25, True), + ) + def test_cut_focal_bce_correctness( + self, + batch, + seq_len, + hidden_dim, + vocab_size, + num_labels, + block_v, + gamma, + alpha, + apply_class_balancing, + ): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + key = jax.random.PRNGKey(0) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + # naive Focal BCE + def run_naive(act, emb, b_val): + loss, _ = _naive_focal_bce( + act, + emb, + b_val, + targets, + gamma=gamma, + alpha=alpha, + apply_class_balancing=apply_class_balancing, + ) + return loss + + grad_naive_fn = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + loss_naive = run_naive(activations, embeddings, bias) + g_act_naive, g_emb_naive, g_bias_naive = grad_naive_fn( + activations, embeddings, bias + ) + + # cut Focal BCE + def run_cut(act, emb, b_val): + return binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + act, + emb, + targets, + bias=b_val, + block_v=block_v, + gamma=gamma, + alpha=alpha, + apply_class_balancing=apply_class_balancing, + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + # Compare + np.testing.assert_allclose(loss_cut, loss_naive, atol=1e-5, rtol=1e-5) + np.testing.assert_allclose(g_act_cut, g_act_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_cut, g_emb_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_bias_cut, g_bias_naive, atol=1e-4, rtol=1e-4) + + @parameterized.named_parameters( + ('4d_act_4d_tgt', (2, 2, 64, 32), (2, 2, 64, 2)), + ('4d_act_3d_tgt', (2, 2, 64, 32), (2, 64, 2)), + ) + def test_cut_focal_bce_4d(self, act_shape, tgt_shape): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + vocab_size, block_v = 256, 128 + hidden_dim = act_shape[-1] + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, act_shape) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint(key_tgt, tgt_shape, 0, vocab_size) + + def run_naive(act, emb, b_val): + loss, _ = _naive_focal_bce(act, emb, b_val, targets) + return loss + + def run_cut(act, emb, b_val): + return binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_naive_fn = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + loss_naive = run_naive(activations, embeddings, bias) + g_act_naive, g_emb_naive, g_bias_naive = grad_naive_fn( + activations, embeddings, bias + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + np.testing.assert_allclose(loss_cut, loss_naive, rtol=1e-3, atol=1e-3) + np.testing.assert_allclose(g_act_cut, g_act_naive, atol=1e-3, rtol=1e-3) + np.testing.assert_allclose(g_emb_cut, g_emb_naive, atol=1e-3, rtol=1e-3) + np.testing.assert_allclose(g_bias_cut, g_bias_naive, atol=1e-3, rtol=1e-3) + + def test_focal_bce_v_blocks_exact(self): + activations = jnp.ones((2, 64, 32)) + embeddings = jnp.ones((250, 32)) + bias = jnp.zeros((250,)) + targets = jnp.zeros((2, 64, 2), dtype=jnp.int32) + block_v = 100 + + # vocab = 250, block_v = 100 -> v_blocks = 3 (ceil(250/100) = 3) + with mock.patch.object(jax.lax, 'scan', wraps=jax.lax.scan) as mock_scan: + binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + activations, embeddings, targets, bias=bias, block_v=block_v + ) + self.assertEqual(mock_scan.call_args[0][2].shape[0], 3) + + def test_cut_focal_bce_sharded(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + batch, seq_len, hidden_dim, vocab_size, num_labels = 2, 64, 32, 256, 2 + key = jax.random.PRNGKey(0) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + devices = jax.devices() + mesh = jax.sharding.Mesh(np.array(devices), ('devices',)) + act_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec('devices', None, None) + ) + emb_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec('devices', None) + ) + bias_sharding = jax.sharding.NamedSharding( + mesh, jax.sharding.PartitionSpec('devices',) + ) + activations_sharded = jax.device_put(activations, act_sharding) + embeddings_sharded = jax.device_put(embeddings, emb_sharding) + bias_sharded = jax.device_put(bias, bias_sharding) + + with mock.patch.object( + jax.lax, + 'with_sharding_constraint', + wraps=jax.lax.with_sharding_constraint, + ) as mock_fwd_constraint: + binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + activations_sharded, + embeddings_sharded, + targets, + bias=bias_sharded, + block_v=128, + ) + self.assertEqual(mock_fwd_constraint.call_count, 3) + + def run_cut(act, emb, b_val): + return binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + act, emb, targets, bias=b_val, block_v=128 + ) + + with mock.patch.object( + jax.lax, + 'with_sharding_constraint', + wraps=jax.lax.with_sharding_constraint, + ) as mock_sharding_constraint: + grad_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + g_act, g_emb, g_bias = grad_fn( + activations_sharded, embeddings_sharded, bias_sharded + ) + self.assertTrue(mock_sharding_constraint.called) + self.assertIsNotNone(g_act) + self.assertIsNotNone(g_emb) + self.assertIsNotNone(g_bias) + + def test_cut_focal_bce_correctness_large_sequence(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + batch, seq_len, hidden_dim, vocab_size, num_labels = 2, 2048, 128, 512, 4 + block_v = 256 + + key = jax.random.PRNGKey(42) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + def run_naive(act, emb, b_val): + loss, _ = _naive_focal_bce(act, emb, b_val, targets) + return loss + + grad_naive_fn = jax.jit(jax.grad(run_naive, argnums=(0, 1, 2))) + loss_naive = run_naive(activations, embeddings, bias) + g_act_naive, g_emb_naive, g_bias_naive = grad_naive_fn( + activations, embeddings, bias + ) + + def run_cut(act, emb, b_val): + return binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + act, emb, targets, bias=b_val, block_v=block_v + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + np.testing.assert_allclose(loss_cut, loss_naive, atol=1e-5, rtol=1e-5) + np.testing.assert_allclose(g_act_cut, g_act_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_emb_cut, g_emb_naive, atol=1e-4, rtol=1e-4) + np.testing.assert_allclose(g_bias_cut, g_bias_naive, atol=1e-4, rtol=1e-4) + + def test_cut_focal_bce_metrics(self): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + batch, seq_len, hidden_dim, vocab_size, num_labels = 2, 64, 32, 128, 2 + key = jax.random.PRNGKey(1) + activations = jax.random.normal(key, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key, (vocab_size, hidden_dim)) + bias = jax.random.normal(key, (vocab_size,)) + targets = jax.random.randint( + key, (batch, seq_len, num_labels), 0, vocab_size + ) + + loss, tp, fp, fn, tn = ( + binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + activations, + embeddings, + targets, + bias=bias, + return_metrics=True, + gamma=2.0, + alpha=0.25, + apply_class_balancing=True, + ) + ) + self.assertIsNotNone(loss) + self.assertIsNotNone(tp) + self.assertIsNotNone(fp) + self.assertIsNotNone(fn) + self.assertIsNotNone(tn) + + @parameterized.named_parameters( + ('standard_gamma2', 2, 64, 32, 128, 2, 64, 2.0, 0.25, True), + ('no_balancing', 2, 64, 32, 128, 2, 64, 2.0, 0.25, False), + ) + def test_cut_focal_bce_vs_keras( + self, + batch, + seq_len, + hidden_dim, + vocab_size, + num_labels, + block_v, + gamma, + alpha, + apply_class_balancing, + ): + if jax.devices()[0].platform != 'tpu': + self.skipTest('Skipping TPU test.') + + key = jax.random.PRNGKey(200) + key_act, key_emb, key_bias, key_tgt = jax.random.split(key, 4) + + activations = jax.random.normal(key_act, (batch, seq_len, hidden_dim)) + embeddings = jax.random.normal(key_emb, (vocab_size, hidden_dim)) + bias = jax.random.normal(key_bias, (vocab_size,)) + targets = jax.random.randint( + key_tgt, (batch, seq_len, num_labels), 0, vocab_size + ) + + def run_keras(act, emb, b_val): + logits = jnp.matmul(act, emb.T) + b_val + one_hot = jax.nn.one_hot(targets, vocab_size, axis=-1) + multi_hot = jnp.max(one_hot, axis=-2) + loss_fn = keras.losses.BinaryFocalCrossentropy( + from_logits=True, + gamma=gamma, + alpha=alpha, + apply_class_balancing=apply_class_balancing, + ) + return jnp.mean(loss_fn(multi_hot, logits)) + + grad_keras_fn = jax.jit(jax.grad(run_keras, argnums=(0, 1, 2))) + loss_keras = run_keras(activations, embeddings, bias) + g_act_keras, g_emb_keras, g_bias_keras = grad_keras_fn( + activations, embeddings, bias + ) + + def run_cut(act, emb, b_val): + return binary_focal_cross_entropy.cut_binary_focal_cross_entropy( + act, + emb, + targets, + bias=b_val, + block_v=block_v, + gamma=gamma, + alpha=alpha, + apply_class_balancing=apply_class_balancing, + ) + + grad_cut_fn = jax.jit(jax.grad(run_cut, argnums=(0, 1, 2))) + loss_cut = run_cut(activations, embeddings, bias) + g_act_cut, g_emb_cut, g_bias_cut = grad_cut_fn( + activations, embeddings, bias + ) + + np.testing.assert_allclose(loss_cut, loss_keras, atol=1e-2, rtol=1e-2) + np.testing.assert_allclose(g_act_cut, g_act_keras, atol=1e-2, rtol=1e-2) + np.testing.assert_allclose(g_emb_cut, g_emb_keras, atol=1e-2, rtol=1e-2) + np.testing.assert_allclose(g_bias_cut, g_bias_keras, atol=1e-2, rtol=1e-2) + + def test_check_vocab_replicated_in_d(self): + with self.assertRaises(NotImplementedError): + binary_focal_cross_entropy._check_vocab_replicated_in_d( + jax.sharding.PartitionSpec('devices', 'devices') + ) + + +if __name__ == '__main__': + absltest.main()