Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
camera_encoder.py204 linesDownload Raw Back to models
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