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