File size: 3,416 Bytes
00c7b31 | 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 | import math
import torch
import torch.nn as nn
import torch.nn.functional as F
class GeometrySpatialMemory(nn.Module):
"""Encode a rendered static 3D scene into compact conditioning tokens.
The input must be VAE latents of a geometry-rendered video produced from a
persistent 3D representation (for example a TSDF-fused point cloud). This
module intentionally does not infer geometry from ordinary context tokens.
"""
def __init__(
self,
dim: int,
latent_channels: int = 16,
patch_size=(1, 2, 2),
grid_size: int = 8,
temporal_bins: int = 4,
num_tokens: int = 64,
):
super().__init__()
self.dim = int(dim)
self.latent_channels = int(latent_channels)
self.patch_size = tuple(int(x) for x in patch_size)
self.grid_size = int(grid_size)
self.temporal_bins = int(temporal_bins)
self.num_tokens = int(num_tokens)
self.geometry_patch_embedding = nn.Conv3d(
self.latent_channels,
self.dim,
kernel_size=self.patch_size,
stride=self.patch_size,
)
source_tokens = self.temporal_bins * self.grid_size * self.grid_size
self.register_buffer(
"geometry_config",
torch.tensor(
[1, self.grid_size, self.temporal_bins],
dtype=torch.int64,
),
persistent=True,
)
self.geometry_to_tokens = nn.Parameter(torch.empty(source_tokens, self.num_tokens))
self.norm = nn.LayerNorm(self.dim)
nn.init.normal_(self.geometry_to_tokens, std=0.02)
def initialize_from_dit_patch_embedding(self, patch_embedding: nn.Conv3d) -> None:
"""Initialize the geometry encoder from the pretrained video tokenizer."""
if not isinstance(patch_embedding, nn.Conv3d):
return
if patch_embedding.weight.shape != self.geometry_patch_embedding.weight.shape:
return
with torch.no_grad():
self.geometry_patch_embedding.weight.copy_(patch_embedding.weight)
if (
self.geometry_patch_embedding.bias is not None
and patch_embedding.bias is not None
and self.geometry_patch_embedding.bias.shape == patch_embedding.bias.shape
):
self.geometry_patch_embedding.bias.copy_(patch_embedding.bias)
def forward(self, geometry_latents: torch.Tensor) -> torch.Tensor:
if geometry_latents is None or geometry_latents.ndim != 5:
raise ValueError(
"GeometrySpatialMemory expects geometry VAE latents shaped (B, C, F, H, W)."
)
if int(geometry_latents.shape[1]) != self.latent_channels:
raise ValueError(
f"GeometrySpatialMemory channel mismatch: input={geometry_latents.shape[1]} "
f"expected={self.latent_channels}"
)
features = self.geometry_patch_embedding(geometry_latents)
features = F.adaptive_avg_pool3d(
features,
(self.temporal_bins, self.grid_size, self.grid_size),
)
features = features.flatten(2).transpose(1, 2)
mix = torch.softmax(self.geometry_to_tokens, dim=0)
memory = torch.einsum("bsd,sm->bmd", features, mix)
return self.norm(memory / math.sqrt(max(self.temporal_bins, 1)))
|