Diversity-Weighted Leave-One-Out Policy Baseline / diversity_baseline.py
Mechanism confirmed, baseline not beaten
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()