import torch
def policy_loss(logits, sampler_log_probs, actions, advantages,
mode="sc", cap=2.0):
"""logits, sampler_log_probs: [N, V].
actions: [N] long; advantages: [N]; N > 0.
Full, finite, normalized sampler log-probabilities.
Same device; all rows valid (no padding).
Modes: raw, sc, is, tis_sc. This is a local surrogate.
"""
if mode not in ("raw", "sc", "is", "tis_sc"):
raise ValueError("Unknown mode")
if mode == "tis_sc" and cap <= 0:
raise ValueError("cap must be positive")
logp = torch.log_softmax(logits.float(), dim=-1)
logq = sampler_log_probs.detach().float()
q = logq.exp()
if mode in ("is", "tis_sc"):
u = (logp.detach() - logq).exp()
if mode == "tis_sc":
u = u.clamp_max(cap)
else:
u = torch.ones_like(logp)
index = actions[:, None]
chosen_logp = logp.gather(-1, index).squeeze(-1)
chosen_u = u.gather(-1, index).squeeze(-1)
surrogate = chosen_u.detach() * chosen_logp
if mode in ("sc", "tis_sc"):
correction = ((q * u).detach() * logp).sum(dim=-1)
surrogate = surrogate - correction
return -(advantages.detach() * surrogate).mean()