| |
| """Quick demo inference: 1 image + 1 video → save GT vs recon visualizations, |
| plus cosine similarity to SigLIP2 teacher (understanding). |
| |
| Usage: |
| PYTHONPATH=src .venv/bin/python infer_demo.py \\ |
| --ckpt checkpoints/stage1_3/balanced/.../loss=0.1962.ckpt \\ |
| --image_shards_dir dataset/image10k/train \\ |
| --video_shards_dir dataset/dataset_10m \\ |
| --outdir results/infer_demo |
| """ |
| from __future__ import annotations |
| import argparse, inspect, json |
| from pathlib import Path |
|
|
| import torch |
| import torch.nn.functional as F |
| from PIL import Image |
| import numpy as np |
| from torchvision.utils import make_grid |
|
|
| from mavt.training.lightning_module import MAVTLightningModule |
| from mavt.data.datasets import WDSImageDataset, ShardVideoDataset |
|
|
|
|
| def to_pil(t: torch.Tensor) -> Image.Image: |
| """t: (3,H,W) in [-1,1] → PIL (H,W,3).""" |
| t = ((t.clamp(-1, 1) + 1) * 0.5 * 255).byte().permute(1, 2, 0).cpu().numpy() |
| return Image.fromarray(t) |
|
|
|
|
| def video_to_strip(clip: torch.Tensor, n_frames: int = 8) -> Image.Image: |
| """clip: (3,T,H,W) in [-1,1] → strip of n_frames horizontally.""" |
| T = clip.shape[1] |
| idx = torch.linspace(0, T - 1, n_frames).long() |
| frames = clip[:, idx].permute(1, 0, 2, 3) |
| grid = make_grid(((frames.clamp(-1, 1) + 1) * 0.5), |
| nrow=n_frames, padding=2, pad_value=1.0) |
| arr = (grid.clamp(0, 1) * 255).byte().permute(1, 2, 0).cpu().numpy() |
| return Image.fromarray(arr) |
|
|
|
|
| def load_module(ckpt_path: str, device): |
| print(f'[infer] loading {ckpt_path}') |
| ckpt = torch.load(ckpt_path, map_location='cpu', weights_only=False) |
| raw_hp = dict(ckpt.get('hyper_parameters', {})) |
| state = ckpt.get('state_dict', {}) |
|
|
| valid = set(inspect.signature(MAVTLightningModule.__init__).parameters) |
| hparams = {k: v for k, v in raw_hp.items() if k in valid} |
| module = MAVTLightningModule(**hparams) |
|
|
| |
| combos = set() |
| for k in state.keys(): |
| if k.startswith('model.cd_split._content_poolers.'): |
| shape = k.split('.')[3] |
| if '_' in shape and all(s.isdigit() for s in shape.split('_')): |
| a, b = shape.split('_') |
| combos.add((int(a), int(b))) |
| for n_c, n_d in sorted(combos): |
| module.model.cd_split.prepare_poolers(n_c, n_d) |
| print(f'[infer] pre-created poolers: {sorted(combos)}') |
|
|
| missing, unexpected = module.load_state_dict(state, strict=False) |
| real_missing = [k for k in missing if not k.startswith('semantic_teacher.')] |
| print(f'[infer] load: {len(real_missing)} missing (excl teacher), {len(unexpected)} unexpected') |
|
|
| module.eval().to(device) |
| return module, hparams |
|
|
|
|
| @torch.no_grad() |
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument('--ckpt', required=True) |
| ap.add_argument('--image_shards_dir', required=True) |
| ap.add_argument('--video_shards_dir', required=True) |
| ap.add_argument('--outdir', default='results/infer_demo') |
| ap.add_argument('--image_idx', type=int, default=0) |
| ap.add_argument('--video_idx', type=int, default=0) |
| ap.add_argument('--video_max_shards', type=int, default=2) |
| ap.add_argument('--device', default='cuda' if torch.cuda.is_available() else 'cpu') |
| ap.add_argument('--threed_dir', default=None, |
| help='Optional path to 3d_objects/renders/<id>/{oxoy,oxoz,oyoz}.png. ' |
| 'If given, run threed inference in addition to image+video.') |
| ap.add_argument('--threed_idx', type=int, default=0) |
| ap.add_argument('--threed_resolution', type=int, default=256) |
| args = ap.parse_args() |
|
|
| device = torch.device(args.device) |
| outdir = Path(args.outdir) |
| outdir.mkdir(parents=True, exist_ok=True) |
|
|
| module, hp = load_module(args.ckpt, device) |
| autocast = torch.amp.autocast(device_type=device.type, dtype=torch.bfloat16, |
| enabled=device.type == 'cuda') |
|
|
| |
| teacher_name = hp.get('siglip2_model_name', 'google/siglip2-base-patch16-224') |
| print(f'[infer] loading teacher: {teacher_name}') |
| from transformers import AutoModel |
| siglip = AutoModel.from_pretrained(teacher_name) |
| teacher = siglip.vision_model.to(device).eval() |
| for p in teacher.parameters(): |
| p.requires_grad_(False) |
| teacher_size = int(siglip.config.vision_config.image_size) |
|
|
| results = {'ckpt': args.ckpt} |
|
|
| |
| print('[infer] ===== image =====') |
| ds_img = WDSImageDataset(args.image_shards_dir, 256) |
| sample = ds_img[args.image_idx] |
| x = sample['data'].unsqueeze(0).to(device) |
| print(f' caption: {sample.get("caption", "")[:80]}') |
|
|
| with autocast: |
| out = module.model(x, 'image', decode=True) |
| recon = out.reconstruction.float().clamp(-1, 1) |
|
|
| to_pil(x[0]).save(outdir / 'image_input.png') |
| to_pil(recon[0]).save(outdir / 'image_recon.png') |
|
|
| |
| pair = torch.cat([x[0], recon[0]], dim=2) |
| to_pil(pair).save(outdir / 'image_side_by_side.png') |
|
|
| |
| teacher_in = F.interpolate(x, size=teacher_size, mode='bilinear', align_corners=False) |
| with autocast: |
| t_emb = teacher(pixel_values=teacher_in).pooler_output.float() |
| cos_img = F.cosine_similarity(out.semantic.float(), t_emb, dim=-1).item() |
|
|
| |
| rec01 = (recon.clamp(-1, 1) + 1) * 0.5 |
| tgt01 = (x.clamp(-1, 1) + 1) * 0.5 |
| mse = F.mse_loss(rec01, tgt01).item() |
| psnr = -10 * np.log10(mse + 1e-12) |
|
|
| results['image'] = { |
| 'shape': list(x.shape), |
| 'caption': sample.get('caption', ''), |
| 'cos_sim_teacher': cos_img, |
| 'recon_psnr_single': psnr, |
| 'recon_l1_single': F.l1_loss(rec01, tgt01).item(), |
| 'files': { |
| 'input': str(outdir / 'image_input.png'), |
| 'recon': str(outdir / 'image_recon.png'), |
| 'side_by_side': str(outdir / 'image_side_by_side.png'), |
| }, |
| } |
| print(f' cos_sim={cos_img:.4f}, single PSNR={psnr:.2f}, L1={results["image"]["recon_l1_single"]:.4f}') |
|
|
| |
| print('[infer] ===== video (caveat: random video poolers if stage1 ckpt) =====') |
| ds_vid = ShardVideoDataset(args.video_shards_dir, n_frames=16, resolution=256, |
| max_shards=args.video_max_shards) |
| sample = ds_vid[args.video_idx] |
| x = sample['data'].unsqueeze(0).to(device) |
| print(f' caption: {sample.get("caption", "")[:80]}, shape: {tuple(x.shape)}') |
|
|
| with autocast: |
| out_v = module.model(x, 'video', decode=True) |
| recon_v = out_v.reconstruction.float().clamp(-1, 1) |
| print(f' recon shape: {tuple(recon_v.shape)}') |
|
|
| |
| t_patch = int(hp.get('t_patch', 2)) |
| tgt_v = x[:, :, ::t_patch] |
|
|
| video_to_strip(x[0], n_frames=8).save(outdir / 'video_input_strip.png') |
| video_to_strip(recon_v[0], n_frames=min(8, recon_v.shape[2])).save(outdir / 'video_recon_strip.png') |
| video_to_strip(tgt_v[0], n_frames=min(8, tgt_v.shape[2])).save(outdir / 'video_gt_subsampled_strip.png') |
|
|
| |
| mid_frame = x[:, :, x.shape[2] // 2] |
| teacher_in = F.interpolate(mid_frame, size=teacher_size, mode='bilinear', align_corners=False) |
| with autocast: |
| t_emb = teacher(pixel_values=teacher_in).pooler_output.float() |
| cos_vid = F.cosine_similarity(out_v.semantic.float(), t_emb, dim=-1).item() |
|
|
| rec01 = (recon_v.clamp(-1, 1) + 1) * 0.5 |
| tgt01 = (tgt_v.clamp(-1, 1) + 1) * 0.5 |
| mse = F.mse_loss(rec01, tgt01).item() |
| psnr_v = -10 * np.log10(mse + 1e-12) |
|
|
| results['video'] = { |
| 'input_shape': list(x.shape), |
| 'recon_shape': list(recon_v.shape), |
| 'caption': sample.get('caption', ''), |
| 'cos_sim_teacher_midframe': cos_vid, |
| 'recon_psnr_single': psnr_v, |
| 'recon_l1_single': F.l1_loss(rec01, tgt01).item(), |
| 'caveat': 'stage1 ckpt has no trained video pooler — recon is roughly random', |
| 'files': { |
| 'input_strip': str(outdir / 'video_input_strip.png'), |
| 'recon_strip': str(outdir / 'video_recon_strip.png'), |
| 'gt_subsampled_strip': str(outdir / 'video_gt_subsampled_strip.png'), |
| }, |
| } |
| print(f' cos_sim={cos_vid:.4f}, single PSNR={psnr_v:.2f} (caveat: random video pooler)') |
|
|
| |
| json_path = outdir / 'summary.json' |
| json_path.write_text(json.dumps(results, indent=2)) |
| print(f'[infer] wrote {json_path}') |
|
|
| |
| if args.threed_dir is not None: |
| print('[infer] ===== threed =====') |
| from mavt.data.datasets import UniversalThreeDDataset |
| ds_3d = UniversalThreeDDataset(args.threed_dir, resolution=args.threed_resolution) |
| if len(ds_3d) == 0: |
| print(f'[infer] no threed objects found in {args.threed_dir}') |
| else: |
| idx = min(args.threed_idx, len(ds_3d) - 1) |
| sample_3d = ds_3d[idx] |
| x_3d = sample_3d['data'].unsqueeze(0).to(device) |
| print(f' caption: {sample_3d.get("caption", "")[:80]}, shape: {tuple(x_3d.shape)}') |
|
|
| with autocast: |
| out_3d = module.model(x_3d, 'threed', decode=True) |
| recon_3d = out_3d.reconstruction.float().clamp(-1, 1) |
| print(f' recon shape: {tuple(recon_3d.shape)}') |
|
|
| |
| xy = x_3d[:, 0] |
| teacher_in = F.interpolate(xy, size=teacher_size, mode='bilinear', |
| align_corners=False) |
| with autocast: |
| t_emb_3d = teacher(pixel_values=teacher_in).pooler_output.float() |
| cos_3d = F.cosine_similarity(out_3d.semantic.float(), t_emb_3d, dim=-1).item() |
|
|
| |
| rec01_3d = (recon_3d.clamp(-1, 1) + 1) * 0.5 |
| tgt01_3d = (x_3d.clamp(-1, 1) + 1) * 0.5 |
| per_plane = {} |
| for i, p_name in enumerate(('oxoy', 'oxoz', 'oyoz')): |
| mse = F.mse_loss(rec01_3d[0, i], tgt01_3d[0, i]).item() |
| per_plane[p_name] = { |
| 'psnr': -10 * np.log10(mse + 1e-12), |
| 'l1': F.l1_loss(rec01_3d[0, i], tgt01_3d[0, i]).item(), |
| } |
|
|
| |
| pair = torch.cat([tgt01_3d[0], rec01_3d[0]], dim=2) |
| pair_flat = pair.reshape(3, 3 * 2, args.threed_resolution, args.threed_resolution) |
| grid = make_grid(pair_flat, nrow=3, padding=4, pad_value=1.0) |
| arr = (grid.clamp(0, 1).permute(1, 2, 0).numpy() * 255).astype('uint8') |
| Image.fromarray(arr).save(outdir / 'threed_plane_grid.png') |
|
|
| results['threed'] = { |
| 'input_shape': list(x_3d.shape), |
| 'recon_shape': list(recon_3d.shape), |
| 'caption': sample_3d.get('caption', ''), |
| 'cos_sim_teacher_xy': cos_3d, |
| 'per_plane': per_plane, |
| 'caveat': ('stage1/2 ckpts have no trained threed pooler — ' |
| 'recon is roughly random unless training_stage=3'), |
| 'files': { |
| 'plane_grid': str(outdir / 'threed_plane_grid.png'), |
| }, |
| } |
| psnr_str = ' '.join(f'{k}={v["psnr"]:.2f}' for k, v in per_plane.items()) |
| print(f' cos_sim={cos_3d:.4f}, per-plane PSNR: {psnr_str}') |
|
|
| json_path.write_text(json.dumps(results, indent=2)) |
| print(f'[infer] updated {json_path}') |
|
|
|
|
| if __name__ == '__main__': |
| main() |
|
|