MCPcopy Create free account
hub / github.com/RL-Align/RL-Kernel / compute_reference_kl

Function compute_reference_kl

rl_engine/testing/reference_ops.py:95–106  ·  view source on GitHub ↗

Compute the common GRPO/PPO reference KL approximation.

(
    current_logps: torch.Tensor,
    ref_logps: torch.Tensor,
    mask: torch.Tensor | None = None,
)

Source from the content-addressed store, hash-verified

93
94
95def compute_reference_kl(
96 current_logps: torch.Tensor,
97 ref_logps: torch.Tensor,
98 mask: torch.Tensor | None = None,
99) -> torch.Tensor:
100 """Compute the common GRPO/PPO reference KL approximation."""
101
102 diff = ref_logps.float() - current_logps.float()
103 kl = torch.exp(diff) - diff - 1.0
104 if mask is not None:
105 kl = kl.masked_fill(~_bool_mask(mask, device=kl.device), 0.0)
106 return kl
107
108
109def summarize_kernel_drift(

Callers 6

trainMethod · 0.90
_minimal_rl_lossFunction · 0.90
_reference_lossFunction · 0.90
grpo_lossFunction · 0.90
_selected_logprob_rowFunction · 0.90

Calls 1

_bool_maskFunction · 0.70

Tested by 3

_minimal_rl_lossFunction · 0.72
_reference_lossFunction · 0.72