File size: 10,791 Bytes
343ac31 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 | """The assembled objective for the recursive planner.
Six terms, per ``recursive_planner_design.pdf`` section 8 with the halting head
(Change 10, OPTIONAL) dropped:
1. arrival-and-hold at the relabel offset ``q`` (Change 4)
2. late-rising path weighting, coefficient ``alpha`` (Change 5)
3. deep supervision over cycles, ``lambda_cycle`` (Change 3)
4. support hinge against the fitted GMM,
``lambda_support`` (Change 8)
5. manifold anchor on the answer, ``lambda_anchor`` (Change 7)
6. pre-tanh saturation barrier, ``lambda_sat`` (Change 6)
No cross-entropy against dataset actions, no entropy term, no progress term.
Weight decay plays the role of ``lambda_reg``.
Terms 1, 2 and 4 reuse the phase-1 implementations unchanged — they take
per-step distances and do not care how the actions were produced.
"""
import torch
from torch import nn
from lejepa_control.losses import arrival_hold_loss, path_weights, support_loss
__all__ = [
'DistanceScale',
'anchor_loss',
'cycle_weights',
'deep_supervision_loss',
'planner_loss',
'saturation_loss',
]
class DistanceScale(nn.Module):
"""EMA of the median context-to-goal distance; ``d~ = d / (m + eps)``.
Every loss coefficient in section 10 assumes this normalization is in
place. Without it the goal term's magnitude depends on the encoder and the
goal-offset curriculum, and the six weights stop meaning what the document
says they mean.
The median, not the mean, because relabeled pairs at the far end of the
offset curriculum produce a long right tail that would drag a mean around
every time the curriculum advances.
Kept as a module so the running scale rides along in the checkpoint and a
resumed run does not restart with a cold normalizer.
"""
def __init__(self, momentum=0.99, eps=1e-6):
super().__init__()
self.momentum = momentum
self.eps = eps
self.register_buffer('scale', torch.zeros(()))
self.register_buffer('initialized', torch.zeros((), dtype=torch.bool))
@torch.no_grad()
def update(self, ctx_emb, goal_emb):
"""Fold this batch's median context-to-goal distance into the EMA."""
median = (ctx_emb[:, -1] - goal_emb).pow(2).mean(dim=-1).median()
if bool(self.initialized):
self.scale.mul_(self.momentum).add_(median, alpha=1 - self.momentum)
else:
self.scale.copy_(median)
self.initialized.fill_(True)
return self.scale
def normalize(self, distances):
if not bool(self.initialized):
return distances
return distances / (self.scale + self.eps)
def cycle_weights(cycles, device=None, dtype=None):
"""``rho_j = 2^j / sum_i 2^i`` over ``j = 1..T``.
For ``T = 3`` this is ``[0.143, 0.286, 0.571]``: the last cycle — the one
that runs at deployment — carries 57% of the weight, while cycle 1 still
receives 14%, enough that a budget-cut deployment at ``T = 1`` emits a
trained action rather than an untrained one.
"""
j = torch.arange(1, cycles + 1, device=device, dtype=dtype or torch.float32)
rho = torch.pow(2.0, j)
return rho / rho.sum()
def deep_supervision_loss(cycle_distances, weights=None):
"""``sum_j rho_j * d_hat^(j)``, averaged over batch and horizon steps.
Args:
cycle_distances: ``(B, H, T)`` normalized per-cycle lookahead
distances from the planner.
weights: ``(T,)`` ``rho``; computed from the tensor's own ``T`` if
omitted.
Note that under the Change-1 gradient policy only the final cycle's term
carries a gradient — cycles ``1..T-1`` are produced inside the ``no_grad``
region and enter as constants. They still shape the logged value and are
the acceptance-test diagnostic. Moving the boundary to gradient-supervise
the earlier cycles is the explicit trade-off the document names; see
``--cycle-grad-boundary`` in the training script.
"""
if weights is None:
weights = cycle_weights(
cycle_distances.size(-1),
cycle_distances.device,
cycle_distances.dtype,
)
return (cycle_distances * weights).sum(dim=-1).mean()
def anchor_loss(action_embed, answers):
"""``(1/W) * ||phi(psi(y)) - y||^2`` over every supervised answer.
Change 7. Without it ``y`` is whatever ``g`` emits, under no constraint to
resemble ``phi(b)`` for any real block ``b``. Over ``T`` cycles and ``H``
steps it drifts into a region of answer space ``phi`` never maps into;
``psi`` still decodes it to a valid block, so nothing crashes and the
failure is silent — the encoder becomes dead weight and ``g``'s input
degenerates into an unconstrained hidden state.
Returns:
The loss, and the relative anchor ratio ``||.||^2 / ||y||^2`` used as
the health metric (healthy: under 0.05 and stable).
"""
y_hat = action_embed.round_trip(answers)
err = (y_hat - answers).pow(2)
loss = err.mean()
ratio = (
err.sum(dim=-1) / answers.pow(2).sum(dim=-1).clamp(min=1e-8)
).mean()
return loss, ratio.detach()
def saturation_loss(raw, limit=2.0):
"""``mean(max(0, |r| - limit)^2)`` on ``psi``'s pre-tanh output.
Change 6. Entropy is moot for a continuous answer; the real degeneracy is
that the goal loss rewards extreme actions, driving ``r`` into the flat
region where ``tanh'(r) -> 0`` and the dimension freezes at +-1 with
nothing able to pull it back. At ``|r| = 4`` the tanh gradient is already
attenuated ~43x. The barrier is zero inside ``|r| <= 2``, which covers
~96% of the action range, and grows quadratically outside it.
Returns:
The loss, and the fraction of activations past the limit (healthy:
under 10%).
"""
excess = (raw.abs() - limit).clamp(min=0)
return excess.pow(2).mean(), (raw.abs() > limit).float().mean().detach()
def planner_loss(
out,
goal_offset,
scale,
action_embed,
density=None,
c95=None,
hold_weight=0.5,
alpha=0.05,
lambda_cycle=0.3,
lambda_support=0.01,
lambda_anchor=0.05,
lambda_sat=1e-3,
sat_limit=2.0,
terminal_only=False,
path_weighting='late',
gamma=0.9,
):
"""Assemble the six-term objective from one planner rollout.
Args:
out: The dict returned by :class:`~lejepa_control_2.planner.RecursivePlanner`.
goal_offset: ``(B,)`` long, the ``q`` each sample's goal was relabeled
from — the deadline the arrival term is indexed by.
scale: :class:`DistanceScale`, already updated for this batch.
action_embed: The planner's ``phi``/``psi`` pair.
density / c95: Fitted :class:`BehaviorDensity` and its held-out 95th
percentile threshold. Both ``None`` disables the support term.
hold_weight, alpha, lambda_*: Section 10 coefficients.
Returns:
``(loss, metrics)`` where ``metrics`` holds detached scalars and the
per-step / per-cycle distance curves for the diagnostics log.
"""
d = scale.normalize(out['distances']) # (B, H)
H = d.size(1)
# --- 1. arrival-and-hold at the relabel offset (Change 4) -------------
# terminal_only is ablation 5's control: score at a fixed d_H instead. The
# deadline then resets to H at every replan, so deferring arrival forever
# is optimal under the loss.
if terminal_only:
loss_arrival = d[:, -1].mean()
else:
loss_arrival = arrival_hold_loss(d, goal_offset, hold_weight).mean()
# --- 2. late-rising path weighting (Change 5) -------------------------
if H > 1 and alpha != 0:
if path_weighting == 'discount':
# ablation 6's control: a discount puts the most weight on step 1,
# which in PushT prefers jamming against the block over the detour
# the task actually requires
k = torch.arange(1, H, device=d.device, dtype=d.dtype)
w = torch.pow(gamma, k)
w = w / w.sum()
else:
w = path_weights(H, d.device, d.dtype)
loss_path = (d[:, :-1] * w).sum(dim=1).mean()
else:
loss_path = d.new_zeros(())
loss = loss_arrival + alpha * loss_path
metrics = {
'arrival': loss_arrival.detach(),
'path': loss_path.detach(),
}
# --- 3. deep supervision over cycles (Change 3) -----------------------
if lambda_cycle != 0 and out['cycle_distances'] is not None:
cycle_d = scale.normalize(out['cycle_distances']) # (B, H, T)
loss_cycle = deep_supervision_loss(cycle_d)
loss = loss + lambda_cycle * loss_cycle
metrics['cycle'] = loss_cycle.detach()
# --- 4. support hinge (Change 8) --------------------------------------
# scored whenever a density model is available, penalized only when
# lambda_support is on — a violation fraction that rises while the goal
# loss falls is the signature of model exploitation, and it shows up
# 1-2k steps before real-env success starts dropping
if density is not None and c95 is not None:
ctx = out['contexts'].flatten(0, 1) # (B*H, N, D)
blocks = out['blocks'].flatten(0, 1) # (B*H, A)
loss_support, violation = support_loss(density, ctx, blocks, c95)
metrics['support'] = loss_support.detach()
metrics['violation'] = violation.detach()
if lambda_support != 0:
loss = loss + lambda_support * loss_support
# --- 5. manifold anchor (Change 7) ------------------------------------
# computed even when disabled: the anchor ratio is the drift diagnostic,
# and a stage that is not yet paying for the anchor is exactly the stage
# where you want to watch it
loss_anchor, anchor_ratio = anchor_loss(action_embed, out['answers'])
metrics['anchor'] = loss_anchor.detach()
metrics['anchor_ratio'] = anchor_ratio
if lambda_anchor != 0:
loss = loss + lambda_anchor * loss_anchor
# --- 6. pre-tanh saturation barrier (Change 6) ------------------------
loss_sat, sat_fraction = saturation_loss(out['raw'], sat_limit)
metrics['sat'] = loss_sat.detach()
metrics['sat_fraction'] = sat_fraction
if lambda_sat != 0:
loss = loss + lambda_sat * loss_sat
metrics['loss'] = loss.detach()
# curves for the diagnostics log: healthy per-step is decreasing with its
# minimum near q, healthy per-cycle is strictly decreasing in j
metrics['per_step'] = d.mean(dim=0).detach()
if out['cycle_distances'] is not None:
metrics['per_cycle'] = (
scale.normalize(out['cycle_distances']).mean(dim=(0, 1)).detach()
)
return loss, metrics
|