Index
2026-07-08 — Analysis

VGGT-Ω Cosmos VAE latent 채널 정규화 — 진단 및 적용

VGGRPO | Phase 4 reward backbone — Cosmos-Predict2 VAE 타깃 학습 트랙

TL;DR

60.7×
채널간 std 비율
1.79M
stats samples/채널
3
적용 스크립트

1 배경 / 목적

07/04부터 진행 중인 VGGT-Ω LGM Cosmos VAE aggregator vs DINOv2 stitch ablation(job 36540 등, aggregator step 35.9K/50K)에서 사용하는 Cosmos-Predict2 tokenizer는 Wan 2.1 VAE 아키텍처를 그대로 채용했지만 NVIDIA가 별도로 fine-tune한 가중치다. 이 fine-tune 과정에서 latent의 채널별 통계 분포가 Wan과 완전히 달라졌을 가능성을 점검했다.

기존 한계: LGM Conv3d connector(16→1024 채널)는 지금까지 Wan VAE 기준으로 설계됐고, Cosmos VAE로 전환하면서 입력 latent 분포가 동일하다는 암묵적 가정이 검증되지 않은 상태였다.

2 작업 내용

Cosmos VAE로 96 clip × 3 cam × 8 frame을 인코딩해 채널별 mean/std를 Welford single-pass 알고리즘(수치적으로 안정, 1-pass)으로 계산했다. 결과는 극단적으로 불균일했다:

Per-channel mean (16채널): [2.67, -1.73, -3.36, 0.60, 5.43, 0.66, 1.06, 4.00, 1.00, 0.72, 3.98, -1.19, 4.69, -1.48, 4.38, -3.51] → 범위 [-3.51, 5.43], 채널마다 큰 bias Per-channel std (16채널): [8.09, 1.98, 1.66, 2.13, 2.30, 5.30, 4.67, 1.87, 12.57, 60.74, 4.57, 9.58, 2.41, 5.24, 4.27, 3.67] ↑ 채널 10 = 60.74 (다른 채널의 30~60배) Wan VAE 대조군: 전 채널 std ~1.7로 균등

구현 (파일: scripts/compute_cosmos_latent_stats.py, 실행 ~3분 login node 1 GPU):

python scripts/compute_cosmos_latent_stats.py --num_clips 100 --num_frames 8 → experiments/cosmos_latent_stats.pt 저장 {mean[16], std[16], n_samples, num_clips, num_frames}

load_cosmos_vae()(scripts/train_lgm_gt_v0.py) 마지막에 stats를 VAE 객체에 _latent_mean/_latent_std로 attach. load_wan_vae()에는 이 코드가 없어 Wan은 attach 자체가 안 됨.

def normalize_latent(z, vae): if not hasattr(vae, "_latent_mean"): return z # Wan: no-op mean = vae._latent_mean.to(z.device, z.dtype) # [16] std = vae._latent_std.to(z.device, z.dtype) # [16] shape = [1] * z.dim(); shape[-4] = 16 return (z - mean.view(*shape)) / std.view(*shape).clamp(min=1e-4)

적용 지점: vae.encode() 직후, Conv3d connector 직전, 총 3개 스크립트의 정확히 같은 위치에 삽입.

스크립트용도적용 지점
train_lgm_robocasa_vggt_v0.py학습vae_encode_views()
find_stitch_layer_vggt.pyStitch search (aggregator)Pass 1(train) + Pass 3(eval)
find_stitch_layer_vggt_dinov2.pyStitch search (DINOv2)Pass 1 + Pass 3, 동일 패턴

대안으로 Conv3d weight init만 조정하는 방법도 검토했으나, 입력 분포 자체를 표준화하는 편이 stitch search closed-form 해(ℓ̂)와도 일관되게 맞물려 latent 정규화를 채택했다.

3 결과

정규화 적용 전후 데이터 흐름:

[video] (T=8/32, 3, 336, 592) [0,1] ↓ Wan/Cosmos VAE encode (frozen, bf16) [z_raw] Cosmos: std [1.7~60.7] 불균일 | Wan: std ~1.7 균등 ↓ F.interpolate 시간축 → T=32 [z_interp] ↓ normalize_latent (Cosmos만 (z-mean)/std, Wan은 no-op) [z_norm] Cosmos: 채널별 unit variance로 통일 | Wan: 변경 없음 ↓ LGMConnectorVGGT.Conv3d(16→1024, k=5, stride=(1,2,2)) [patches] → Aggregator/DINOv2 → [depth, depth_conf, pose_enc]
핵심 발견: 정규화 적용은 코드 수준에서 Wan 파이프라인을 완전히 무변경으로 유지하면서(hasattr 분기로 no-op 보장) Cosmos 경로에만 격리 적용됐다 — 3개 스크립트 동일 위치 삽입으로 학습/stitch search 간 불일치 리스크 제거.

미완료: 정규화 반영된 stitch search(job 36679/36680)는 이 시점 재실행 중 — Conv3d init을 정규화된 입력 분포에 맞춰 재산출해야 비로소 학습 첫 step부터 안정적 forward가 보장된다. 정규화 적용 자체의 loss/수렴 개선 수치는 아직 미측정 (다음 섹션 참조).

4 Takeaway

의미

Cosmos VAE 채널10의 std가 60.74로 나머지 대비 30~60배 큰 것은 단순 스케일 문제가 아니라 gradient가 사실상 단일 채널로 쏠려 나머지 15채널의 geometry 정보가 학습에 거의 기여하지 못하는 구조적 위험이었다. 07/04부터 진행 중인 no-norm 축 학습(job 36540, aggregator step 35.9K/50K)은 이 문제를 안은 채 진행된 것으로, 향후 norm 적용 축과의 비교가 이 가설의 실질적 검증이 된다. Wan VAE 경로는 코드 변경 없이 격리됐으므로 기존 Wan 학습 결과의 재현성에는 영향 없음.

5 Next Steps

미해결 한계

정규화가 실제로 downstream depth/pose 예측 품질을 개선하는지는 아직 정량 확인 전. stitch search 재실행(job 36679/36680) 완료 후 새 Cosmos-norm ablation을 시작해야 loss curve로 비교 가능.

다음 실험

Stitch search(norm 반영) 완료 → 새 Conv3d init으로 Cosmos-norm 학습 시작. 최종적으로 {Wan, Cosmos-nonorm, Cosmos-norm} × {aggregator, DINOv2} 6-way 비교 진행, no-norm 대비 norm 축의 초기 loss 안정성 및 AbsRel 개선폭을 확인할 것.