Fused cross-entropy: log_softmax + nll_loss per row. One program per row (batch element).
(
logits_ptr,
targets_ptr,
losses_ptr,
n_cols,
stride_logits_row,
BLOCK_SIZE: tl.constexpr,
)
| 19 | |
| 20 | @triton.jit |
| 21 | def cross_entropy_kernel( |
| 22 | logits_ptr, |
| 23 | targets_ptr, |
| 24 | losses_ptr, |
| 25 | n_cols, |
| 26 | stride_logits_row, |
| 27 | BLOCK_SIZE: tl.constexpr, |
| 28 | ): |
| 29 | """ |
| 30 | Fused cross-entropy: log_softmax + nll_loss per row. |
| 31 | One program per row (batch element). |
| 32 | """ |
| 33 | row_idx = tl.program_id(0) |
| 34 | |
| 35 | row_start = logits_ptr + row_idx * stride_logits_row |
| 36 | col_offsets = tl.arange(0, BLOCK_SIZE) |
| 37 | mask = col_offsets < n_cols |
| 38 | |
| 39 | # Load logits row in float32 |
| 40 | logits = tl.load(row_start + col_offsets, mask=mask, other=float("-inf")).to(tl.float32) |
| 41 | |
| 42 | # Numerically stable log-softmax |
| 43 | # Step 1: find max |
| 44 | row_max = tl.max(logits, axis=0) |
| 45 | |
| 46 | # Step 2: subtract max, exp, sum |
| 47 | logits_shifted = logits - row_max |
| 48 | exp_logits = tl.exp(logits_shifted) |
| 49 | sum_exp = tl.sum(exp_logits, axis=0) |
| 50 | log_sum_exp = tl.log(sum_exp) |
| 51 | |
| 52 | # log_softmax = logits_shifted - log_sum_exp |
| 53 | # We only need the value at the target index |
| 54 | |
| 55 | # Load target for this row |
| 56 | target = tl.load(targets_ptr + row_idx) |
| 57 | |
| 58 | # Get the logit at the target position |
| 59 | # log_softmax[target] = logits[target] - max - log(sum(exp(logits - max))) |
| 60 | target_logit = tl.load(row_start + target).to(tl.float32) |
| 61 | log_softmax_target = (target_logit - row_max) - log_sum_exp |
| 62 | |
| 63 | # NLL loss = -log_softmax[target] |
| 64 | loss = -log_softmax_target |
| 65 | |
| 66 | # Store per-sample loss |
| 67 | tl.store(losses_ptr + row_idx, loss) |
| 68 | |
| 69 | |
| 70 | def kernel_fn(logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor: |
nothing calls this directly
no outgoing calls
no test coverage detected