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