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