Diversity-Weighted Leave-One-Out Policy Baseline / diversity_baseline.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1"""Diversity-weighted leave-one-out baseline for grouped policy samples."""
 2import numpy as np
 3
 4
 5def diversity_baseline(costs, embeddings, p=1.0, eps=1e-8):
 6    """Return baseline, minimization advantage, normalized weights, distances."""
 7    costs = np.asarray(costs, dtype=float)
 8    z = np.asarray(embeddings, dtype=float)
 9    if costs.ndim != 1 or z.ndim != 2 or len(costs) != len(z):
10        raise ValueError("costs must be [B], embeddings must be [B,d]")
11    if len(costs) < 2 or p < 0 or eps <= 0:
12        raise ValueError("need B>=2, p>=0, eps>0")
13    z = z - z.mean(axis=0, keepdims=True)
14    d = ((z[:, None, :] - z[None, :, :]) ** 2).sum(axis=-1)
15    w = (d + eps) ** p
16    np.fill_diagonal(w, 0.0)
17    denom = w.sum(axis=1)
18    baseline = (w @ costs) / denom
19    advantage = baseline - costs
20    return baseline, advantage, w / denom[:, None], d
21
22
23def uniform_loo(costs):
24    costs = np.asarray(costs, dtype=float)
25    return (costs.sum() - costs) / (len(costs) - 1)
26
27
28def effective_count(normalized_weights):
29    q = np.asarray(normalized_weights)
30    return 1.0 / (q * q).sum(axis=-1)
31
32
33def torch_diversity_loss(log_probs, costs, embeddings, p=1.0, eps=1e-8,
34                         normalize_advantage=False):
35    """SSPO loss; log_probs are [B] summed trajectory log probabilities.
36
37    Baseline quantities are detached as required. Costs may be tensors or arrays;
38    embeddings may retain gradients, but no gradient is allowed through them here.
39    """
40    import torch
41    if log_probs.ndim != 1 or embeddings.ndim != 2 or costs.ndim != 1:
42        raise ValueError("log_probs/costs must be [B], embeddings must be [B,d]")
43    if log_probs.shape[0] != costs.shape[0] or costs.shape[0] != embeddings.shape[0]:
44        raise ValueError("batch dimensions must agree")
45    z = embeddings.detach() - embeddings.detach().mean(dim=0, keepdim=True)
46    d = ((z[:, None, :] - z[None, :, :]) ** 2).sum(dim=-1)
47    w = (d + eps).pow(p)
48    eye = torch.eye(w.shape[0], dtype=torch.bool, device=w.device)
49    w = w.masked_fill(eye, 0.0)
50    denom = w.sum(dim=1)
51    b = (w @ costs.detach()) / denom
52    adv = b - costs.detach()
53    if normalize_advantage:
54        adv = adv / costs.detach().std().clamp_min(eps)
55    return -(adv.detach() * log_probs).mean()