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)))