hugging-apps/echo-memory
0
1"""2Camera Encoder for RT (Rotation-Translation) matrix injection.3CAM paper [2506.03141]: Maps camera pose to DiT hidden dimension for spatial attention conditioning.4 5Design notes:6- Primary purpose: dimension alignment (action/RT [12] -> DiT hidden D). One-layer MLP is sufficient; no need for deeper encoder.7- Per-frame MLP: each frame's RT [12] -> hidden_size independently (no temporal context).8- Optional zero-init scale: conditioning starts weak (scale=0) and grows with training for stability.9- shallow: single Linear(12, D) to match CAM "single-layer MLP" wording (ablation). When shallow=True10 and separate_t_r=False, the encoder is exactly one layer (merged MLP): RT [12] -> D.11- separate_t_r: encode translation (t) and rotation (R) with separate MLPs then add, for scale balance.12 Use shallow=True and separate_t_r=False for merged single-layer MLP (RT not split).13- explicit_yaw: add a signed yaw scalar branch (Z-only) so CW/CCW are explicitly encoded; helps when14 the model is insensitive to rotation direction (Zhou et al. CVPR 2019, sign continuity).15- sincos_yaw: add [cos(yaw), sin(yaw)] branch (2D) for direction; sin carries sign explicitly.16 17Design limitations / caveats:18- No input normalization: t (translation) and R (rotation) have different scales; one Linear(12,D) may be19 sensitive to units. Caller should use consistent RT scale or relative RT; optional input LayerNorm not implemented.20- Per-frame only: no temporal context (each frame encoded independently). Fine for CAM ablation; temporal modeling not supported.21- 16-dim input: layout for flattened 4x4 is unspecified; yaw branches are disabled. Prefer 12-dim in practice.22- explicit_yaw and sincos_yaw can both be True (redundant encoding of yaw); usually use one.23"""24 25import torch26import torch.nn as nn27from typing import Optional28 29# For Z-only rotation: yaw = atan2(R_21, R_11); R is row-major [R_11,R_12,R_13, R_21,...]30def _yaw_from_rt_12(rt: torch.Tensor) -> torch.Tensor:31 """rt [..., 12] -> yaw in [-1, 1] (normalized by pi)."""32 R11 = rt[..., 3]33 R21 = rt[..., 6]34 yaw_rad = torch.atan2(R21, R11)35 return yaw_rad / 3.14159265358979336 37 38def _sincos_yaw_from_rt_12(rt: torch.Tensor) -> torch.Tensor:39 """rt [..., 12] -> [..., 2] (cos(yaw), sin(yaw))."""40 R11 = rt[..., 3]41 R21 = rt[..., 6]42 yaw_rad = torch.atan2(R21, R11)43 return torch.stack([torch.cos(yaw_rad), torch.sin(yaw_rad)], dim=-1)44 45 46class CameraEncoder(nn.Module):47 """48 Encode RT matrices (camera pose) to DiT hidden dimension.49 50 Input: rt_matrices [B, F, 12] or [B, F, 16]51 - 12: [t_x, t_y, t_z, R_11..R_33] (3 translation + 9 rotation), R row-major. No input normalization:52 t and R often differ in scale (e.g. t in meters, R in [-1,1]); single Linear(12,D) may be sensitive to units.53 - 16: 4x4 matrix flattened (layout/order unspecified; yaw branches disabled when rt_dim=16).54 Output: camera_emb [B, F, D] where D = hidden_size (scaled by learnable scale, default 0-init).55 """56 57 def __init__(58 self,59 rt_dim: int = 12,60 hidden_size: int = 5120,61 mlp_hidden_mult: int = 4,62 eps: float = 1e-6,63 zero_init_scale: bool = False,64 full_zero_init: bool = False,65 shallow: bool = False,66 separate_t_r: bool = False,67 explicit_yaw: bool = False,68 sincos_yaw: bool = False,69 conditioning_scale: float = 1.0,70 r_mlp_no_layernorm: bool = False,71 ):72 super().__init__()73 self.rt_dim = rt_dim74 self.hidden_size = hidden_size75 self.zero_init_scale = zero_init_scale76 self.full_zero_init = full_zero_init77 self.shallow = shallow78 self.separate_t_r = separate_t_r79 self.explicit_yaw = explicit_yaw and rt_dim == 1280 self.sincos_yaw = sincos_yaw and rt_dim == 1281 self.conditioning_scale = float(conditioning_scale)82 self.r_mlp_no_layernorm = r_mlp_no_layernorm and separate_t_r83 dtype = torch.get_default_dtype()84 85 if separate_t_r:86 # Plan B: separate t (3) and R (9) encoders for scale balance; only for rt_dim=12.87 assert rt_dim == 12, "separate_t_r only supported for rt_dim=12"88 mid = max(hidden_size // 2, 256)89 self.t_mlp = nn.Sequential(90 nn.Linear(3, mid),91 nn.LayerNorm(mid, eps=eps),92 nn.GELU(),93 nn.Linear(mid, hidden_size),94 nn.LayerNorm(hidden_size, eps=eps),95 )96 if r_mlp_no_layernorm:97 # No LayerNorm on R so sign of R_12/R_21 (yaw direction) is not normalized away.98 self.r_mlp = nn.Sequential(99 nn.Linear(9, mid),100 nn.GELU(),101 nn.Linear(mid, hidden_size),102 )103 else:104 self.r_mlp = nn.Sequential(105 nn.Linear(9, mid),106 nn.LayerNorm(mid, eps=eps),107 nn.GELU(),108 nn.Linear(mid, hidden_size),109 nn.LayerNorm(hidden_size, eps=eps),110 )111 self.mlp = None112 elif shallow:113 # Merged single-layer MLP: one Linear(rt_dim, hidden_size), no separate t/R.114 self.mlp = nn.Linear(rt_dim, hidden_size)115 assert isinstance(self.mlp, nn.Linear), "shallow path must be exactly one Linear layer"116 if full_zero_init:117 nn.init.zeros_(self.mlp.weight)118 nn.init.zeros_(self.mlp.bias)119 else:120 mid_dim = hidden_size * mlp_hidden_mult121 self.mlp = nn.Sequential(122 nn.Linear(rt_dim, mid_dim),123 nn.LayerNorm(mid_dim, eps=eps),124 nn.GELU(),125 nn.Linear(mid_dim, mid_dim),126 nn.LayerNorm(mid_dim, eps=eps),127 nn.GELU(),128 nn.Linear(mid_dim, hidden_size),129 nn.LayerNorm(hidden_size, eps=eps),130 )131 132 if self.explicit_yaw:133 self.yaw_embed = nn.Linear(1, hidden_size)134 else:135 self.yaw_embed = None136 if self.sincos_yaw:137 self.sincos_embed = nn.Linear(2, hidden_size)138 else:139 self.sincos_embed = None140 141 # Learnable scale: when zero_init_scale=True, init to 0 so conditioning grows with training (stable).142 # When full_zero_init=True, skip scale (GF-ICL style: Linear output directly, no extra scale).143 if full_zero_init:144 self.scale = None # no scale, use 1.0 in forward145 else:146 self.scale = nn.Parameter(torch.zeros(1) if zero_init_scale else torch.ones(1))147 148 def is_single_layer_merged(self) -> bool:149 """True if encoder is exactly one Linear(12, D) with no separate t/R (merged MLP)."""150 return self.shallow and not self.separate_t_r and self.mlp is not None and isinstance(self.mlp, nn.Linear)151 152 def forward(self, rt_matrices: torch.Tensor) -> torch.Tensor:153 """154 Args:155 rt_matrices: [B, F, 12] or [B, F, 16]156 Returns:157 camera_emb: [B, F, hidden_size], scaled by self.scale.158 """159 d = rt_matrices.dtype160 if self.separate_t_r:161 t = rt_matrices[..., :3].to(d)162 r = rt_matrices[..., 3:12].to(d)163 out = self.t_mlp(t) + self.r_mlp(r)164 else:165 out = self.mlp(rt_matrices.to(d))166 if self.yaw_embed is not None and rt_matrices.shape[-1] >= 12:167 yaw_norm = _yaw_from_rt_12(rt_matrices[..., :12]).unsqueeze(-1).to(d)168 out = out + self.yaw_embed(yaw_norm)169 if self.sincos_embed is not None and rt_matrices.shape[-1] >= 12:170 sincos = _sincos_yaw_from_rt_12(rt_matrices[..., :12]).to(d)171 out = out + self.sincos_embed(sincos)172 scale = self.scale.to(d) if self.scale is not None else torch.ones(1, device=out.device, dtype=out.dtype)173 return out * scale * self.conditioning_scale174 175 176def expand_camera_emb_to_tokens(177 camera_emb: torch.Tensor,178 num_frames: int,179 h: int,180 w: int,181) -> torch.Tensor:182 """183 Expand per-frame camera_emb [B, F, D] to per-token [B, N, D]184 where N = F * h * w (tokens ordered as frame0_all_patches, frame1_all_patches, ...).185 186 Dimension alignment (与 DiT patchify 一致):187 - Encoder 输出: 每帧一个向量 [B, F, D],即相当于 [B, F, 1, D](F 帧每帧 1 个 embedding)。188 - 对齐方式: 在空间维上把该 1 重复 H×W 次,得到 [B, F, h*w, D],再展平为 [B, F*h*w, D]。189 - Token 顺序: frame0 的 h*w 个 token 共用 frame0 的 camera_emb,frame1 的 h*w 个 token 共用 frame1 的 camera_emb,与190 wan_video_dit patchify 的 rearrange(..., 'b c f h w -> b (f h w) c') 顺序一致(帧优先,再空间)。191 192 Args:193 camera_emb: [B, F, D]194 num_frames: F (must equal camera_emb.shape[1]; used for assertion only).195 h, w: spatial grid (patches per frame)196 Returns:197 [B, F*h*w, D]198 """199 B, F, D = camera_emb.shape200 if F != num_frames:201 raise ValueError(f"expand_camera_emb_to_tokens: camera_emb has F={F}, num_frames={num_frames}")202 # [B, F, D] -> [B, F, 1, D] (每帧 1 个) -> expand 到 [B, F, h*w, D] (每帧重复 H×W 次) -> [B, F*h*w, D]203 return camera_emb.unsqueeze(2).expand(B, F, h * w, D).reshape(B, F * h * w, D)204 