Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 20 additions & 1 deletion projects/mmdet3d_plugin/uniad/modules/temporal_self_attention.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,24 @@
'_ext', ['ms_deform_attn_backward', 'ms_deform_attn_forward'])


def _select_previous_bev(value, batch_size, num_bev_queue):
"""Select each sample's history BEV from the interleaved queue.

The encoder packs temporal values as ``[prev_0, curr_0, prev_1,
curr_1, ...]``. Flattened indexing with ``value[:batch_size]`` only
selects the correct history for the first sample when ``batch_size > 1``.
Reshaping the queue restores the sample/queue axes before selecting the
history slot.
"""
if value.shape[0] != batch_size * num_bev_queue:
raise ValueError(
'Expected value batch dimension to equal '
f'batch_size * num_bev_queue ({batch_size} * {num_bev_queue}), '
f'but got {value.shape[0]}'
)
return value.reshape(batch_size, num_bev_queue, *value.shape[1:])[:, 0]


@ATTENTION.register_module()
class TemporalSelfAttention(BaseModule):
"""An attention module used in BEVFormer based on Deformable-Detr.
Expand Down Expand Up @@ -191,7 +209,8 @@ def forward(self,
assert (spatial_shapes[:, 0] * spatial_shapes[:, 1]).sum() == num_value
assert self.num_bev_queue == 2

query = torch.cat([value[:bs], query], -1)
query = torch.cat(
[_select_previous_bev(value, bs, self.num_bev_queue), query], -1)
value = self.value_proj(value)

if key_padding_mask is not None:
Expand Down
168 changes: 168 additions & 0 deletions tests/test_temporal_self_attention.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,168 @@
import importlib.util
import sys
import types
from pathlib import Path

import pytest
import torch
from torch import nn


def _stub(monkeypatch, name, package=False, **attributes):
module = types.ModuleType(name)
if package:
module.__path__ = []
for key, value in attributes.items():
setattr(module, key, value)
monkeypatch.setitem(sys.modules, name, module)
return module


def _load_temporal_attention(monkeypatch):
"""Load the module while replacing its optional simulator dependencies."""
package_name = "_uniad_temporal_attention_test"
source_dir = Path(__file__).resolve().parents[1] / "projects/mmdet3d_plugin/uniad/modules"
package = _stub(monkeypatch, package_name, package=True)
package.__path__ = [str(source_dir)]
_stub(
monkeypatch,
f"{package_name}.multi_scale_deformable_attn_function",
MultiScaleDeformableAttnFunction_fp32=object,
)

_stub(monkeypatch, "mmcv", package=True)
_stub(monkeypatch, "mmcv.ops", package=True)
_stub(
monkeypatch,
"mmcv.ops.multi_scale_deform_attn",
multi_scale_deformable_attn_pytorch=lambda *args, **kwargs: None,
)
_stub(
monkeypatch,
"mmcv.cnn",
xavier_init=lambda *args, **kwargs: None,
constant_init=lambda *args, **kwargs: None,
)
_stub(monkeypatch, "mmcv.cnn.bricks", package=True)

class Registry:
def register_module(self):
return lambda cls: cls

_stub(monkeypatch, "mmcv.cnn.bricks.registry", ATTENTION=Registry())
_stub(monkeypatch, "mmcv.runner", package=True)

class BaseModule(nn.Module):
def __init__(self, init_cfg=None):
super().__init__()

_stub(
monkeypatch,
"mmcv.runner.base_module",
BaseModule=BaseModule,
ModuleList=nn.ModuleList,
Sequential=nn.Sequential,
)
_stub(monkeypatch, "mmcv.utils", package=True)
ext_loader = _stub(
monkeypatch,
"mmcv.utils.ext_loader",
load_ext=lambda *args, **kwargs: types.SimpleNamespace(),
)
utils = sys.modules["mmcv.utils"]
utils.ConfigDict = dict
utils.build_from_cfg = lambda *args, **kwargs: None
utils.deprecated_api_warning = lambda *args, **kwargs: None
utils.to_2tuple = lambda value: (value, value)
utils.ext_loader = ext_loader

module_name = f"{package_name}.temporal_self_attention"
spec = importlib.util.spec_from_file_location(
module_name, source_dir / "temporal_self_attention.py"
)
module = importlib.util.module_from_spec(spec)
monkeypatch.setitem(sys.modules, module_name, module)
spec.loader.exec_module(module)
return module


def test_previous_bev_selection_preserves_sample_alignment(monkeypatch):
temporal_attention = _load_temporal_attention(monkeypatch)
value = torch.tensor(
[
[[10.0, 10.1]], # sample 0, previous frame
[[20.0, 20.1]], # sample 0, current frame
[[30.0, 30.1]], # sample 1, previous frame
[[40.0, 40.1]], # sample 1, current frame
]
)

previous = temporal_attention._select_previous_bev(value, 2, 2)

torch.testing.assert_close(previous, value[[0, 2]])
assert not torch.equal(previous, value[:2])


def test_temporal_attention_uses_each_sample_history(monkeypatch):
temporal_attention = _load_temporal_attention(monkeypatch)
attention = temporal_attention.TemporalSelfAttention(
embed_dims=2, num_heads=1, num_levels=1, num_points=1, dropout=0.0
)

class CaptureProjection(nn.Module):
def __init__(self, output_size, capture=False):
super().__init__()
self.output_size = output_size
self.capture = capture

def forward(self, input_tensor):
if self.capture:
self.input = input_tensor.detach().clone()
return input_tensor.new_zeros(
input_tensor.shape[0], input_tensor.shape[1], self.output_size
)

offsets = CaptureProjection(4, capture=True)
attention.sampling_offsets = offsets
attention.attention_weights = CaptureProjection(2)
attention.value_proj = nn.Identity()
attention.output_proj = nn.Identity()
attention.dropout = nn.Identity()
monkeypatch.setattr(
temporal_attention,
"multi_scale_deformable_attn_pytorch",
lambda value, spatial_shapes, sampling_locations, attention_weights: value.new_zeros(
value.shape[0], sampling_locations.shape[1], value.shape[-2] * value.shape[-1]
),
)

attention(
query=torch.tensor([[[100.0, 101.0]], [[200.0, 201.0]]]),
value=torch.tensor(
[
[[10.0, 11.0]], # sample 0, previous frame
[[20.0, 21.0]], # sample 0, current frame
[[30.0, 31.0]], # sample 1, previous frame
[[40.0, 41.0]], # sample 1, current frame
]
),
reference_points=torch.zeros(4, 1, 1, 2),
spatial_shapes=torch.tensor([[1, 1]]),
level_start_index=torch.tensor([0]),
)

torch.testing.assert_close(
offsets.input,
torch.tensor(
[
[[10.0, 11.0, 100.0, 101.0]],
[[30.0, 31.0, 200.0, 201.0]],
]
),
)


def test_previous_bev_selection_rejects_incomplete_queue(monkeypatch):
temporal_attention = _load_temporal_attention(monkeypatch)
with pytest.raises(ValueError, match=r"batch_size \* num_bev_queue"):
temporal_attention._select_previous_bev(torch.zeros(3, 1, 4), 2, 2)