Fused LogP¶
Fused LogP computes selected token log probabilities from model logits. It targets RL
post-training workloads where repeated log_softmax + gather operations create memory
pressure at large group sizes.
Entry Point¶
from rl_engine.kernels.registry import kernel_registry
logp_op = kernel_registry.get_op("logp")
output = logp_op(logits, token_ids)
The PyTorch native reference also exposes the Issue #108 interface:
from rl_engine.kernels.ops.pytorch.loss.logp import NativeLogpOp
logp_ref = NativeLogpOp()
output = logp_ref.forward(logits, token_ids)
reference = logp_ref.forward_fp32(logits, token_ids)
apply(...) and apply_fp32(...) remain available as backward-compatible aliases.
Backends¶
| Backend | Wrapper | Native symbol | Notes |
|---|---|---|---|
| CUDA SM90 | FusedLogpSM90Op |
_C.fused_logp_sm90 |
Experimental TMA-oriented path for 2D contiguous bf16 logits on Hopper-class GPUs. It is disabled by default and requires RL_KERNEL_ENABLE_EXPERIMENTAL_SM90_LOGP=1; otherwise the wrapper delegates to the CUDA generic fallback. |
| CUDA generic | FusedLogpGenericOp |
_C.fused_logp |
Generic compiled extension fallback. |
| PyTorch native | NativeLogpOp |
None | PyTorch baseline/reference path. |
Tensor Contract¶
| Argument | Shape | Dtype | Requirements |
|---|---|---|---|
logits |
[N, V] |
bfloat16 for the experimental SM90 fast path; fp16/fp32 use generic fallback |
Contiguous, on the target device for the experimental SM90 fast path. |
token_ids / labels |
[N] |
Converted to int32 |
Same logical device as logits. |
| Output | [N] |
Backend-defined tensor dtype | One selected log probability per row. |
Reference Semantics¶
ref = torch.log_softmax(logits.float(), dim=-1)
ref = torch.gather(ref, dim=-1, index=token_ids.unsqueeze(-1).long()).squeeze(-1)
Tests¶
tests/test_logp.py covers the PyTorch reference contract, dtype behavior,
backward-compatible aliases, batch invariance, and registry dispatch. The existing
operator accuracy tests continue to validate native/CUDA fused API compatibility.
Implementation Files¶
rl_engine/kernels/registry.pyrl_engine/kernels/ops/pytorch/loss/logp.pyrl_engine/kernels/ops/cuda/loss/logp.pycsrc/ops.cppcsrc/fused_logp_kernel.cucsrc/cuda/fused_logp_sm90.cutests/test_logp.py