| import torch |
| from torch import nn |
|
|
|
|
| class ReuseUpdater: |
| def __call__(self, previous_memory, *args, **kwargs): |
| return previous_memory |
|
|
|
|
| class GatedResidualMemoryUpdater(nn.Module): |
| def __init__( |
| self, |
| channels, |
| ego_dim=0, |
| feature_format='spatial', |
| residual_scale_init=0.0, |
| ): |
| super().__init__() |
| if feature_format not in ('spatial', 'tokens'): |
| raise ValueError("feature_format must be 'spatial' or 'tokens'") |
| self.channels = channels |
| self.ego_dim = ego_dim |
| self.feature_format = feature_format |
| in_channels = channels * 2 + ego_dim |
| if feature_format == 'spatial': |
| self.gate = nn.Conv2d(in_channels, channels, kernel_size=1) |
| self.residual = nn.Conv2d(in_channels, channels, kernel_size=1) |
| else: |
| self.gate = nn.Linear(in_channels, channels) |
| self.residual = nn.Linear(in_channels, channels) |
| self.residual_scale = nn.Parameter(torch.tensor(float(residual_scale_init))) |
|
|
| def forward(self, previous_memory, partial_current=None, ego_motion_embedding=None): |
| if isinstance(previous_memory, list): |
| return [ |
| self.forward(prev, part, ego_motion_embedding) |
| for prev, part in zip(previous_memory, partial_current) |
| ] |
| if isinstance(previous_memory, tuple): |
| return tuple( |
| self.forward(prev, part, ego_motion_embedding) |
| for prev, part in zip(previous_memory, partial_current) |
| ) |
| if partial_current is None: |
| raise ValueError('partial_current is required for gated residual updates') |
| if previous_memory.shape != partial_current.shape: |
| raise ValueError( |
| f'partial_current shape {tuple(partial_current.shape)} must match ' |
| f'previous_memory shape {tuple(previous_memory.shape)}') |
|
|
| x = self._concat_inputs(previous_memory, partial_current, ego_motion_embedding) |
| gate = torch.sigmoid(self.gate(x)) |
| residual = self.residual(x) |
| if self.feature_format == 'spatial': |
| gate = self._from_spatial(gate, previous_memory) |
| residual = self._from_spatial(residual, previous_memory) |
| return previous_memory + self.residual_scale * gate * residual |
|
|
| def _concat_inputs(self, previous_memory, partial_current, ego_motion_embedding): |
| if self.feature_format == 'spatial': |
| if previous_memory.dim() != 5: |
| raise ValueError('spatial memory must have shape B,N_cam,C,H,W') |
| pieces = [previous_memory, partial_current] |
| if ego_motion_embedding is not None: |
| pieces.append(self._expand_spatial_ego(ego_motion_embedding, previous_memory)) |
| return self._to_spatial(torch.cat(pieces, dim=2)) |
|
|
| if previous_memory.dim() != 4: |
| raise ValueError('token memory must have shape B,N_cam,N_token,C') |
| pieces = [previous_memory, partial_current] |
| if ego_motion_embedding is not None: |
| pieces.append(self._expand_token_ego(ego_motion_embedding, previous_memory)) |
| return torch.cat(pieces, dim=-1) |
|
|
| @staticmethod |
| def _to_spatial(x): |
| b, n_cam, channels, h, w = x.shape |
| return x.reshape(b * n_cam, channels, h, w) |
|
|
| @staticmethod |
| def _from_spatial(x, reference): |
| b, n_cam, _, h, w = reference.shape |
| return x.reshape(b, n_cam, -1, h, w) |
|
|
| @staticmethod |
| def _expand_spatial_ego(ego_motion_embedding, reference): |
| b, n_cam, _, h, w = reference.shape |
| return ego_motion_embedding[:, None, :, None, None].expand(b, n_cam, -1, h, w) |
|
|
| @staticmethod |
| def _expand_token_ego(ego_motion_embedding, reference): |
| b, n_cam, n_token, _ = reference.shape |
| return ego_motion_embedding[:, None, None, :].expand(b, n_cam, n_token, -1) |
|
|
|
|
| def build_memory_updater(spec, **kwargs): |
| if callable(spec) and not isinstance(spec, str): |
| return spec |
| if isinstance(spec, dict): |
| spec = spec.copy() |
| name = spec.pop('type', spec.pop('name', 'reuse')) |
| spec.update(kwargs) |
| else: |
| name = spec |
| spec = kwargs |
|
|
| if name == 'reuse': |
| return ReuseUpdater() |
| if name == 'gated_residual': |
| return GatedResidualMemoryUpdater(**spec) |
| raise KeyError(f'Unsupported temporal memory updater: {name}') |
|
|