Leplanner / code /losses.py
nottygian's picture
Push code package
872cf4d verified
Raw
History Blame Contribute Delete
6.41 kB
"""Controller objective: latent goal reaching plus a thresholded support term.
The controller is never asked to reproduce dataset actions. Gradients come
from what the frozen world model predicts the plan will *cause*; the dataset
only supplies goals and a notion of which actions are in-distribution.
"""
import torch
import torch.nn.functional as F
from torch import nn
def path_weights(horizon, device=None, dtype=None):
"""Late-weighted path coefficients ``w_j ~ (j/H)^2`` over ``j=1..H-1``."""
j = torch.arange(1, horizon, device=device, dtype=dtype or torch.float32)
w = (j / horizon).pow(2)
return w / w.sum()
def goal_loss(distances, alpha, weights):
"""``d_H + alpha * sum_j w_j d_j`` for one refinement.
Args:
distances: ``(B, H)`` per-step latent goal distances.
alpha: Path-loss coefficient.
weights: ``(H-1,)`` path weights.
"""
terminal = distances[:, -1]
if alpha == 0 or distances.size(1) < 2:
return terminal
return terminal + alpha * (distances[:, :-1] * weights).sum(dim=1)
def arrival_hold_loss(distances, goal_offset, hold_weight):
"""``d_q + hold_weight * mean(d_{q+1..H})`` for one refinement.
The fixed-terminal objective ``d_H`` means "be at the goal exactly H
blocks from now". Under receding-horizon execution the deadline resets
to H after every replan, so the controller keeps deferring arrival and
approaches the goal asymptotically without landing on it. Indexing the
arrival term by the offset the goal was actually relabeled from ties the
deadline to the state instead of to the plan, and the hold term stops the
controller from touching the goal and leaving.
Args:
distances: ``(B, H)`` per-step latent goal distances.
goal_offset: ``(B,)`` long, in ``1..H`` — how many transitions ahead
this sample's goal was taken from.
hold_weight: Coefficient on staying near the goal after arrival.
"""
B, H = distances.shape
q = goal_offset.clamp(1, H)
arrival = distances.gather(1, (q - 1).unsqueeze(1)).squeeze(1)
# mean over j > q, skipping samples where the deadline is the last block
steps = torch.arange(H, device=distances.device).unsqueeze(0)
after = (steps >= q.unsqueeze(1)).float()
count = after.sum(dim=1)
hold = (distances * after).sum(dim=1) / count.clamp(min=1)
return arrival + hold_weight * torch.where(
count > 0, hold, torch.zeros_like(hold)
)
def refinement_loss(
distances_per_iter, alpha=0.05, goal_offset=None, hold_weight=None
):
"""``2^k``-weighted average of the goal loss across refinements.
Later refinements matter more, but every iteration gets a direct signal so
early plans stay usable if computation is stopped short.
Passing ``goal_offset`` selects the horizon-matched arrival-and-hold
objective; otherwise this is the fixed-terminal loss ``d_H + alpha*path``.
"""
ref = distances_per_iter[0]
horizon = ref.size(1)
weights = path_weights(horizon, ref.device, ref.dtype)
if goal_offset is None:
def per_sample(d):
return goal_loss(d, alpha, weights)
else:
def per_sample(d):
return arrival_hold_loss(d, goal_offset, hold_weight)
rho = torch.tensor(
[2.0**k for k in range(len(distances_per_iter))],
device=ref.device,
dtype=ref.dtype,
)
per_iter = torch.stack([per_sample(d).mean() for d in distances_per_iter])
return (rho * per_iter).sum() / rho.sum()
class BehaviorDensity(nn.Module):
"""Conditional Gaussian mixture ``beta(b | C)`` over real action blocks.
Trained separately on real latent histories and real five-action blocks.
It is a support model, not a policy: the controller is only penalized for
leaving the region the dataset actually covers.
"""
def __init__(
self,
latent_dim=192,
block_dim=10,
num_context=3,
components=16,
width=256,
min_log_std=-5.0,
max_log_std=2.0,
):
super().__init__()
self.block_dim = block_dim
self.components = components
self.min_log_std = min_log_std
self.max_log_std = max_log_std
self.net = nn.Sequential(
nn.Linear(num_context * latent_dim, width),
nn.GELU(),
nn.Linear(width, width),
nn.GELU(),
)
self.logits = nn.Linear(width, components)
self.means = nn.Linear(width, components * block_dim)
self.log_stds = nn.Linear(width, components * block_dim)
def log_prob(self, ctx_emb, block):
"""Log density of ``block`` ``(B, A)`` given context ``(B, N, D)``."""
h = self.net(ctx_emb.flatten(1))
B = h.size(0)
logits = self.logits(h)
means = self.means(h).view(B, self.components, self.block_dim)
log_stds = self.log_stds(h).view(B, self.components, self.block_dim)
log_stds = log_stds.clamp(self.min_log_std, self.max_log_std)
x = block.unsqueeze(1) # (B, 1, A)
z = (x - means) / log_stds.exp()
comp = -0.5 * (z.pow(2) + 1.8378770664093453) - log_stds
return torch.logsumexp(
F.log_softmax(logits, dim=-1) + comp.sum(-1), dim=-1
)
def nll_per_dim(self, ctx_emb, block):
"""``r(C, b) = -log beta(b | C) / A`` — the support score."""
return -self.log_prob(ctx_emb, block) / self.block_dim
def support_loss(density, contexts, blocks, threshold):
"""Squared hinge on plans that fall outside the dataset's action support.
Args:
density: Trained :class:`BehaviorDensity` (frozen during controller
training).
contexts: ``(M, N, D)`` latent histories along the imagined rollouts.
blocks: ``(M, A)`` the action blocks proposed at those histories.
threshold: ``c_95``, the 95th-percentile score on held-out real data.
Returns:
Scalar loss, and the fraction of blocks that violated the threshold.
"""
score = density.nll_per_dim(contexts, blocks)
violation = (score - threshold).clamp(min=0)
return violation.pow(2).mean(), (score > threshold).float().mean()