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