MCPcopy Create free account
hub / github.com/RightNow-AI/autokernel / cross_entropy_kernel

Function cross_entropy_kernel

kernels/cross_entropy.py:21–67  ·  view source on GitHub ↗

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,
)

Source from the content-addressed store, hash-verified

19
20@triton.jit
21def 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
70def kernel_fn(logits: torch.Tensor, targets: torch.Tensor) -> torch.Tensor:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected