echo / code /diffsynth /models /memory /geometry_spatial_memory.py
amonshano's picture
Add Echo-Memory codebase used for this run (CC BY 4.0, JD Echo Team)
00c7b31 verified
Raw
History Blame Contribute Delete
3.42 kB
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)))