File size: 19,501 Bytes
251713e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
"""PyTorch Lightning Module for MAVT training.

Supports 3-stage curriculum via `training_stage` parameter:
  1 — image only,  SigLIP2 fully frozen,  LR = 1e-4
  2 — +video,      SigLIP2 last 4 unfrozen, LR = 5e-5
  3 — +3D,         SigLIP2 fully unfrozen, LR = 2e-5

To move between stages: start training with the next stage config and pass
`--ckpt_path <prev_stage_checkpoint>` to LightningCLI.
"""

from __future__ import annotations
from typing import Any, Dict, Optional

import torch
import torch.nn as nn
import torch.nn.functional as F
import lightning as L
from lightning.pytorch.utilities import grad_norm

from mavt.model.mavt import MAVT
from mavt.losses.losses import MAVTLoss


_STAGE_LR = {1: 1e-4, 2: 5e-5, 3: 2e-5}
_STAGE_SIGLIP2_FROZEN_BLOCKS = {1: 10, 2: 6, 3: 0}  # number of transformer blocks frozen
_STAGE_W_SEM = {1: 0.5, 2: 0.3, 3: 0.2}             # default cosine-distill weight per stage


class MAVTLightningModule(L.LightningModule):
    """Lightning module for end-to-end MAVT training."""

    def __init__(
        self,
        # Model
        embed_dim: int = 1152,
        num_heads: int = 16,
        num_blocks: int = 12,
        patch_size: int = 16,
        t_patch: int = 2,
        latent_dim: int = 32,
        kl_weight: float = 1e-4,
        semantic_dim: int = 768,
        dec_dim: int = 768,
        num_dec_attn_blocks: int = 4,
        r_s: int = 2,
        r_t: int = 1,
        use_gradient_checkpointing: bool = True,
        mlp_ratio: float = 4.0,
        dropout: float = 0.0,
        # C-D split (cd-split commit 6368dfb)
        local_detail_window_size: int = 1,
        local_detail_temporal_window_size: int = 1,
        # Loss
        w_l1: float = 1.0,
        w_lpips: float = 0.1,
        w_kl: float = 1.0,      # passthrough — KL pre-scaled by VAEHead.kl_weight
        w_clip: float = 0.0,
        w_sem: float = 0.0,
        w_aux: float = 0.01,
        w_temp: float = 0.0,    # temporal consistency weight (video only)
        use_lpips: bool = True,
        use_clip: bool = False,
        active_modalities: list = None,  # e.g. ['image', 'video']; None → all three
        # Curriculum
        training_stage: int = 1,
        siglip2_model_name: str = "google/siglip2-base-patch16-224",
        init_siglip2: bool = True,
        use_semantic_distill: bool = False,
        # Cross-stage weight transfer (loads model weights only, NOT optimizer
        # / scheduler / step state — use --ckpt_path for true resume instead).
        init_from_ckpt: Optional[str] = None,
        # Optimiser
        weight_decay: float = 0.01,
        grad_clip: float = 1.0,
        warmup_steps: int = 1000,
        total_steps: int = 200_000,
    ):
        super().__init__()
        self.save_hyperparameters()

        self.model = MAVT(
            embed_dim=embed_dim, num_heads=num_heads, num_blocks=num_blocks,
            patch_size=patch_size, t_patch=t_patch,
            latent_dim=latent_dim, kl_weight=kl_weight,
            semantic_dim=semantic_dim, dec_dim=dec_dim,
            num_dec_attn_blocks=num_dec_attn_blocks, r_s=r_s, r_t=r_t,
            use_gradient_checkpointing=use_gradient_checkpointing,
            mlp_ratio=mlp_ratio, dropout=dropout,
            local_detail_window_size=local_detail_window_size,
            local_detail_temporal_window_size=local_detail_temporal_window_size,
        )

        _active_mods = tuple(active_modalities) if active_modalities else ('image', 'video', 'threed')
        self.loss_fn = MAVTLoss(
            w_l1=w_l1, w_lpips=w_lpips, w_kl=w_kl,
            w_clip=w_clip, w_sem=w_sem, w_aux=w_aux,
            w_temp=w_temp,
            use_lpips=use_lpips, use_clip=use_clip,
            active_modalities=_active_mods,
        )

        # Frozen vision teacher (loaded lazily in setup() to keep __init__ light)
        self.semantic_teacher: Optional[nn.Module] = None
        self._teacher_image_size: int = 224

    # ------------------------------------------------------------------ #
    #  Setup                                                               #
    # ------------------------------------------------------------------ #

    def setup(self, stage: str) -> None:
        hp = self.hparams
        if stage != 'fit':
            return
        # Eager slot pooler creation — must run BEFORE configure_optimizers
        # so the new params land in the optimizer's param_groups.
        self._prepare_cd_split_poolers()
        # Sync EMA modalities from DataModule — single source of truth
        self._sync_ema_modalities()
        if hp.init_siglip2:
            frozen = _STAGE_SIGLIP2_FROZEN_BLOCKS[hp.training_stage]
            self.model.load_siglip2_weights(hp.siglip2_model_name, frozen)
        if hp.use_semantic_distill and self.semantic_teacher is None:
            self._load_semantic_teacher(hp.siglip2_model_name)
        # Cross-stage weight transfer (after siglip2 / teacher are in place so
        # they get overwritten by ckpt values when present).
        if hp.init_from_ckpt:
            self._load_weights_from_ckpt(hp.init_from_ckpt)

    def _load_weights_from_ckpt(self, path: str) -> None:
        # weights_only=False is required because Lightning ckpts contain a
        # full pickle (state_dict + hparams + callbacks). Source is our own
        # filesystem so untrusted-pickle risk is N/A.
        ckpt = torch.load(path, map_location='cpu', weights_only=False)
        sd = ckpt.get('state_dict', ckpt)
        missing, unexpected = self.load_state_dict(sd, strict=False)
        kept_missing = [
            k for k in missing
            if not k.startswith('semantic_teacher.')
            and not k.startswith('model.cd_split._content_poolers.')
            and not k.startswith('model.cd_split._detail_poolers.')
        ]
        print(f"[init_from_ckpt] loaded {path}")
        print(f"[init_from_ckpt] missing (kept random init): {len(kept_missing)} keys "
              f"+ {len(missing) - len(kept_missing)} expected (teacher/new poolers)")
        if unexpected:
            print(f"[init_from_ckpt] unexpected (dropped): {len(unexpected)} keys")
        print("[init_from_ckpt] NOTE: optimizer state / LR scheduler / global_step are NOT "
              "restored — this is a soft restart. For exact resume of the SAME stage, "
              "use Lightning --ckpt_path instead.")

    def _prepare_cd_split_poolers(self) -> None:
        """Read active modality + resolution from the attached DataModule and
        eagerly create every SlotPooler the trainer will need."""
        dm = getattr(self.trainer, 'datamodule', None)
        if dm is None or not hasattr(dm, 'hparams'):
            return
        dm_hp = dm.hparams
        active = getattr(dm_hp, 'active_modalities', None) or []
        specs = []
        for modality in active:
            if modality == 'image':
                specs.append({
                    'modality': 'image',
                    'resolution': dm_hp.image_resolution,
                })
            elif modality == 'video':
                specs.append({
                    'modality': 'video',
                    'resolution': dm_hp.video_resolution,
                    'frames':     dm_hp.video_frames,
                    't_patch':    self.hparams.t_patch,
                })
            elif modality == 'threed':
                specs.append({
                    'modality': 'threed',
                    'resolution': dm_hp.triplane_res,
                })
        if specs:
            self.model.prepare_for_modalities(specs)

    def _sync_ema_modalities(self) -> None:
        """Use DataModule as single source of truth for active_modalities.

        Overrides whatever was set in model config so data.active_modalities
        and the EMA weighter never drift apart.
        """
        dm = getattr(self.trainer, 'datamodule', None)
        if dm is not None and hasattr(dm, 'hparams'):
            active = getattr(dm.hparams, 'active_modalities', None)
            if active:
                self.loss_fn.ema_weighter.active_modalities = tuple(active)
                return
        # Fallback: use model config param if DataModule not available
        fallback = getattr(self.hparams, 'active_modalities', None)
        if fallback:
            self.loss_fn.ema_weighter.active_modalities = tuple(fallback)

    def _load_semantic_teacher(self, model_name: str) -> None:
        """Load frozen SigLIP2 vision tower as teacher for cosine distillation."""
        try:
            from transformers import AutoModel
            siglip = AutoModel.from_pretrained(model_name)
            teacher = siglip.vision_model
            for p in teacher.parameters():
                p.requires_grad_(False)
            teacher.eval()
            self.semantic_teacher = teacher
            try:
                self._teacher_image_size = int(siglip.config.vision_config.image_size)
            except AttributeError:
                self._teacher_image_size = 224
        except Exception as exc:  # noqa: BLE001
            print(f"[lightning_module] semantic teacher load failed ({exc}); "
                  f"distillation disabled this run")
            self.semantic_teacher = None

    def train(self, mode: bool = True):  # type: ignore[override]
        """Keep frozen teacher in eval mode regardless of train()/eval() calls."""
        super().train(mode)
        if self.semantic_teacher is not None:
            self.semantic_teacher.eval()
        return self

    def on_save_checkpoint(self, checkpoint: Dict[str, Any]) -> None:
        """Strip frozen teacher weights from checkpoints to keep them small."""
        state = checkpoint.get('state_dict', {})
        for k in list(state.keys()):
            if k.startswith('semantic_teacher.'):
                del state[k]

    def on_load_checkpoint(self, checkpoint: Dict[str, Any]) -> None:
        """Inject the freshly-loaded teacher's params back into the checkpoint
        state_dict so Lightning's strict load does not fail with missing keys.

        Order at training-resume time:
          1. setup('fit')               → teacher loaded from HF (deterministic)
          2. on_load_checkpoint(ckpt)   → we inject teacher keys here
          3. self.load_state_dict(ckpt) → strict=True now sees a complete dict

        Teacher weights are deterministic from the HF model name, so injecting
        the current values is equivalent to whatever the checkpoint would have
        contained — no behavior change, only avoids the strict-load error.
        """
        if self.semantic_teacher is None:
            return
        sd = checkpoint.setdefault('state_dict', {})
        for k, v in self.semantic_teacher.state_dict().items():
            sd.setdefault(f'semantic_teacher.{k}', v)

    # ------------------------------------------------------------------ #
    #  Training step                                                       #
    # ------------------------------------------------------------------ #

    def _step(self, batch: Dict, log_prefix: str) -> torch.Tensor:
        x       = batch['data']
        modality = batch['modality']

        out = self.model(x, modality, decode=True)

        # For video: decoder reconstructs in patch-grid temporal space (Tp = T//t_patch).
        # Target must match. Use first frame of each temporal patch group.
        if modality == 'video':
            t_patch = self.hparams.t_patch
            target = x[:, :, ::t_patch]  # (B, 3, Tp, H, W)
        else:
            target = x  # image and threed pass through as-is

        # Vision-vision distillation: forward x_proxy through frozen teacher.
        teacher_embed: Optional[torch.Tensor] = None
        if self.semantic_teacher is not None:
            with torch.no_grad():
                proxy = self._make_teacher_input(x, modality, self._teacher_image_size)
                teacher_embed = self.semantic_teacher(pixel_values=proxy).pooler_output

        losses = self.loss_fn(
            pred=out.reconstruction,
            target=target,
            loss_kl=out.loss_kl,
            slot_diversity=out.cd_metrics['slot_diversity'],
            modality=modality,
            semantic_embed=out.semantic,
            teacher_embed=teacher_embed,
        )

        # Logging — aggregate (all modalities combined)
        for k, v in losses.items():
            self.log(f'{log_prefix}/{k}', v, on_step=True, on_epoch=True,
                     prog_bar=(k == 'loss'), sync_dist=True)
        # Per-modality breakdown (diagnose image vs video separately)
        for k in ('loss', 'loss_recon', 'loss_l1', 'loss_kl'):
            if k in losses:
                self.log(f'{log_prefix}/{k}_{modality}', losses[k],
                         on_step=True, on_epoch=True, sync_dist=True)
        for k, v in out.cd_metrics.items():
            self.log(f'{log_prefix}/cd_{k}', v, on_step=False, on_epoch=True,
                     sync_dist=True)
        self.log(f'{log_prefix}/modality_{modality}', 1.0,
                 on_step=False, on_epoch=True, sync_dist=False)

        return losses['loss']

    def training_step(self, batch: Dict, batch_idx: int) -> torch.Tensor:
        return self._step(batch, 'train')

    def validation_step(self, batch: Dict, batch_idx: int) -> None:
        with torch.no_grad():
            self._step(batch, 'val')
            # Log sample reconstructions to wandb/tensorboard every N steps
            if batch_idx == 0:
                self._log_images(batch)

    # ------------------------------------------------------------------ #
    #  Teacher input proxy                                                 #
    # ------------------------------------------------------------------ #

    @staticmethod
    def _make_teacher_input(x: torch.Tensor, modality: str,
                            target_size: int) -> torch.Tensor:
        """Project the multi-modal input down to a single (B, 3, S, S) image
        the SigLIP2 vision teacher can consume.

        image  : x as-is.
        video  : middle frame.
        threed : XY plane (front view, closest to natural-image distribution).
        """
        if modality == 'image':
            proxy = x
        elif modality == 'video':
            T = x.shape[2]
            proxy = x[:, :, T // 2]                # (B, 3, H, W)
        elif modality == 'threed':
            proxy = x[:, 0]                         # (B, 3, H, W) plane XY
        else:
            raise ValueError(f"Unknown modality: {modality}")

        if proxy.shape[-1] != target_size or proxy.shape[-2] != target_size:
            proxy = F.interpolate(
                proxy, size=(target_size, target_size),
                mode='bilinear', align_corners=False,
            )
        return proxy

    # ------------------------------------------------------------------ #
    #  Visualisation                                                       #
    # ------------------------------------------------------------------ #

    def _log_images(self, batch: Dict, n: int = 4) -> None:
        try:
            x       = batch['data'][:n]
            modality = batch['modality']
            out = self.model(x, modality, decode=True)

            if modality == 'image':
                grid_in  = _to_grid(x)
                grid_out = _to_grid(out.reconstruction)
            elif modality == 'video':
                # Log a temporal strip: 4 evenly-spaced frames per clip stacked
                # horizontally so reviewers can spot temporal coherence (vs.
                # only a single first frame).
                B, C, T, H, W = x.shape
                Tp = out.reconstruction.shape[2]
                t_in = torch.linspace(0, T - 1, 4).long()
                t_out = torch.linspace(0, Tp - 1, 4).long()
                in_strip  = x[:, :, t_in].permute(0, 2, 1, 3, 4).reshape(B * 4, C, H, W)
                out_strip = out.reconstruction[:, :, t_out].permute(0, 2, 1, 3, 4).reshape(B * 4, C, H, W)
                grid_in  = _to_grid(in_strip,  nrow=4)
                grid_out = _to_grid(out_strip, nrow=4)
            elif modality == 'threed':
                # Log XY plane (plane index 0)
                grid_in  = _to_grid(x[:, 0])
                grid_out = _to_grid(out.reconstruction[:, 0])
            else:
                return

            loggers = self.loggers if isinstance(self.loggers, (list, tuple)) else [self.loggers]
            for logger in loggers:
                if hasattr(logger, 'log_image'):
                    logger.log_image(key=f'val/{modality}_input',  images=[grid_in])
                    logger.log_image(key=f'val/{modality}_recon',  images=[grid_out])
        except Exception:  # noqa: BLE001
            pass

    # ------------------------------------------------------------------ #
    #  Optimiser                                                           #
    # ------------------------------------------------------------------ #

    def configure_optimizers(self):
        hp = self.hparams
        lr = _STAGE_LR[hp.training_stage]

        # Separate RGAT params for potential different LR (currently same LR)
        rgat_params, other_params = [], []
        for name, p in self.model.named_parameters():
            if not p.requires_grad:
                continue
            if 'rgat' in name.lower() or 'rgat4d' in name.lower():
                rgat_params.append(p)
            else:
                other_params.append(p)

        param_groups = [{'params': other_params, 'lr': lr}]
        if rgat_params:
            param_groups.append({'params': rgat_params, 'lr': lr})

        optimizer = torch.optim.AdamW(param_groups, weight_decay=hp.weight_decay)

        # Linear warmup + cosine decay
        def lr_lambda(step: int) -> float:
            if step < hp.warmup_steps:
                return step / max(1, hp.warmup_steps)
            progress = (step - hp.warmup_steps) / max(1, hp.total_steps - hp.warmup_steps)
            return max(0.0, 0.5 * (1.0 + torch.cos(torch.tensor(torch.pi * progress)).item()))

        scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
        return {
            'optimizer': optimizer,
            'lr_scheduler': {
                'scheduler': scheduler,
                'interval': 'step',
                'frequency': 1,
            },
        }

    def configure_gradient_clipping(self, optimizer, gradient_clip_val=None, gradient_clip_algorithm=None):
        self.clip_gradients(optimizer, gradient_clip_val=self.hparams.grad_clip,
                            gradient_clip_algorithm='norm')


# --------------------------------------------------------------------------- #
#  Visualisation helper                                                         #
# --------------------------------------------------------------------------- #

def _to_grid(x: torch.Tensor, nrow: int = 4) -> Any:
    """Convert (B, 3, H, W) tensor to a PIL Image grid for logging."""
    try:
        from torchvision.utils import make_grid
        from PIL import Image
        import numpy as np
        x = x.detach().cpu().float().clamp(-1, 1)
        x = (x + 1) / 2                         # [0, 1]
        grid = make_grid(x, nrow=nrow, normalize=False)
        arr = (grid.permute(1, 2, 0).numpy() * 255).astype('uint8')
        return Image.fromarray(arr)
    except Exception:  # noqa: BLE001
        return None