TL;DR
- 문제: GT wrist를 넣고 있고,
torch.no_grad() 안에서 VAE decode → gradient가 transformer까지 안 흐름
- 해결: z_0 prediction (
z_0_pred = z_noisy - sigma * pred)으로 gradient 살리고, 3단 loss 구조 설계
- Loss A: Ego 1장 + Wrist 13장 → 14-view scene reconstruction
- Loss B: Per-timestep [Ego_t, Wrist_t] → same-timestep 3D consistency
- Loss C: Ego camera pose consistency (고정 카메라 → 동일 pose 강제)
1 현재 문제
| # | 문제 | 영향 |
| 1 | 49 frame 중 5개만 사용 | 정보 손실 |
| 2 | GT wrist를 3D loss에 넣음 | 생성 품질과 무관한 loss |
| 3 | torch.no_grad() 안에서 VAE decode | gradient 단절 |
| 4 | Transformer 생성 결과 시각화 없음 | 디버깅 불가 |
2 z_0 Prediction
Flow matching에서 1-step denoising으로 z_0를 예측:
z_sigma = (1 - sigma) * z_0 + sigma * noise
pred ≈ noise - z_0
z_0_pred = z_noisy - sigma * pred ← gradient 살아있음
Split:
ego_z0: z_0_pred[:, :, :, :, :32] → (B, 16, 13, 32, 32)
wrist_z0: z_0_pred[:, :, :, :, 32:] → (B, 16, 13, 32, 32)
핵심: z_0_pred가 requires_grad=True이면 stitching_layer → AnySplat → render까지 backward 가능. frozen 모듈이라도 weight gradient만 없고 input gradient는 흐름.
3 3-Loss 구조
Loss A: 14-view Scene Reconstruction
ego_z0[:, :, 0:1, :, :] → 1장 (고정 카메라 anchor)
wrist_z0[:, :, :, :, :] → 13장 (이동 카메라, 다양한 viewpoint)
concat → (B, 16, 14, 32, 32)
→ stitching_layer → AnySplat → 3D Gaussians
→ gsplat render from each viewpoint → GT와 L1
Loss B: Per-timestep Consistency
for t in [0, 3, 6, 9]:
[ego_z0[:,:,t], wrist_z0[:,:,t]] = 2-view latent
→ stitching → AnySplat → 3D Gaussians
→ render ego cam → ego_gt_t 비교
→ render wrist cam → wrist_gt_t 비교
같은 시점의 ego와 wrist가 동일한 3D scene을 봐야 한다는 제약.
Loss C: Ego Pose Consistency
AnySplat이 예측한 ego view의 camera pose가 전 timestep에서 동일해야 함 (고정 카메라).
4 Gradient Flow
L1 loss (rendered vs GT)
↑ gsplat render (differentiable)
↑ AnySplat (frozen weights, input grad flows)
↑ stitching_layer Conv3D (frozen weights, input grad flows)
↑ z_0_pred = z_noisy - sigma * pred
↑ pred = transformer(z_noisy) ← LoRA trainable
주의: torch.no_grad() 대신 requires_grad_(False)로 weight만 freeze. no_grad()는 input gradient도 끊어버림.
5 구현 주의사항
Conv3D temporal kernel
stitching_layer kernel (5,3,3) → temporal dim 최소 5 필요. Loss B의 2-view (T=2)는 너무 짧아서 padding 또는 최소 5 views 구성 필요.
feedforward_image 해상도
AnySplat은 448×448 입력 기대. z_0_pred decode 결과는 256×256이므로 resize 필요.
메모리
14 views × 256 tokens = 3584 tokens. Aggregator global attention O(n²) → ~50MB. 문제없음.