Team Ai
Modelpublic

Codeseys/composer-replication-framework

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
kl_logging.py111 linesDownload Raw Back to altered_minds
1"""kl_logging.py — dual_kl_logger (ADR-013, framework-side, generic).2 3The washout/amplification instrument. Given per-token logprobs from three4forward passes on the SAME answer+reasoning tokens:5 6  - policy:         the model currently being RL-trained7  - altered_init:   the altered SFT checkpoint the run STARTED from (the locus8                    of the cognitive-distortion signature)9  - unaltered_base: the original base model BEFORE personality SFT10 11returns ``{'kl_to_altered_init': float, 'kl_to_base': float}``.12 13NEITHER KL is optimized by default — both are diagnostics:14  - ``kl_to_altered_init`` rising means the policy is moving AWAY from the15    altered checkpoint (task-RL is *changing* the alteration).16  - ``kl_to_base`` measures distance to the unaltered base. If17    ``kl_to_base`` SHRINKS while ``kl_to_altered_init`` grows, the alteration18    is WASHING OUT (the policy drifts back toward base). If ``kl_to_base``19    GROWS faster than ``kl_to_altered_init``, the alteration is being AMPLIFIED20    (the policy moves further from base than the altered init already was) —21    the ADR-013 amplification hypothesis, most likely on the SDPO channel.22 23Token-mean KL is used (mean over the masked answer+reasoning tokens), the24standard diagnostic convention. The math is the discrete KL between the two25softmax distributions implied by the logprob tensors:26 27    KL(p || q) = sum_v p_v (log p_v - log q_v)28 29where ``p`` is the policy's per-token distribution. This is unit-testable on30toy tensors: KL(p || p) == 0, and KL grows monotonically as the policy moves.31"""32from __future__ import annotations33 34from typing import Any35 36import torch37 38__all__ = ["dual_kl_logger", "token_mean_kl"]39 40 41def _as_log_probs(logprobs: torch.Tensor) -> torch.Tensor:42    """Normalize an input that may be raw logits OR already-log-probs to valid43    log-probabilities along the last (vocab) dim.44 45    We re-apply ``log_softmax`` defensively: it is idempotent on a genuine46    log-prob tensor up to floating point (log_softmax of log-probs == log-probs47    since they already sum-exp to 1), and converts raw logits correctly. This48    makes the logger robust to either calling convention.49    """50    return torch.log_softmax(logprobs.to(torch.float64), dim=-1)51 52 53def token_mean_kl(54    policy_logprobs: torch.Tensor,55    ref_logprobs: torch.Tensor,56    mask: torch.Tensor | None = None,57) -> float:58    """Token-mean KL(policy || ref) over distributions on the last dim.59 60    Args:61        policy_logprobs: (..., V) logits or log-probs for the policy.62        ref_logprobs:    (..., V) logits or log-probs for the reference.63        mask: optional (...,) mask of tokens to include (1/True = include). If64            None, all tokens count.65 66    Returns:67        scalar token-mean KL as a python float (>= 0 up to float error).68    """69    log_p = _as_log_probs(policy_logprobs)70    log_q = _as_log_probs(ref_logprobs)71    p = log_p.exp()72    # per-token KL: sum over vocab of p * (log p - log q)73    per_token = (p * (log_p - log_q)).sum(dim=-1)  # (...,)74 75    if mask is not None:76        m = mask.to(per_token.dtype)77        denom = m.sum()78        if float(denom) == 0.0:79            return 0.080        return float((per_token * m).sum() / denom)81    return float(per_token.mean())82 83 84def dual_kl_logger(85    policy_logprobs: torch.Tensor,86    altered_init_logprobs: torch.Tensor,87    unaltered_base_logprobs: torch.Tensor,88    mask: torch.Tensor | None = None,89    **_: Any,90) -> dict[str, float]:91    """Compute the two diagnostic KLs for a step.92 93    Args:94        policy_logprobs:        (..., V) policy logits/log-probs on the95            answer+reasoning tokens.96        altered_init_logprobs:  (..., V) for the altered SFT init.97        unaltered_base_logprobs:(..., V) for the unaltered base.98        mask: optional (...,) token mask (answer+reasoning tokens to score).99 100    Returns:101        ``{'kl_to_altered_init': float, 'kl_to_base': float}``.102    """103    return {104        "kl_to_altered_init": token_mean_kl(105            policy_logprobs, altered_init_logprobs, mask106        ),107        "kl_to_base": token_mean_kl(108            policy_logprobs, unaltered_base_logprobs, mask109        ),110    }111