Skip to content
Merged
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
56 changes: 56 additions & 0 deletions test/optimizations/test_bevfusion_implementations.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,62 @@
from turbo_physai.optimizations.models.bevfusion import transfusion


def test_ddp_forward_supports_new_pytorch():
class DDPStub(torch.nn.Module):
def __init__(self):
super().__init__()
self.module = torch.nn.Linear(2, 1)

def _run_ddp_forward(self, *inputs, **kwargs):
assert self._use_replicated_tensor_module is False
return self.module(*inputs, **kwargs)

model = DDPStub()
wrapped = training.ddp_forward_compat_wrapper(
DDPStub._run_ddp_forward, {}
)

assert wrapped(model, torch.ones(1, 2)).shape == (1, 1)
assert model._use_replicated_tensor_module is False


@pytest.mark.parametrize("compiled", [False, True])
def test_bev_pool_fp16_casts_factorized_namedtuples(monkeypatch, compiled):
def pool(self, geometry, factors):
return geometry, factors

monkeypatch.setattr(depth, "base_transform_bev_pool", pool)
wrapped = depth.base_transform_bev_pool_wrapper(None, {})
if compiled:
wrapped = torch.compile(wrapped, backend="eager")

model = torch.nn.Module()
model.fp16_enabled = True
geometry = depth.PreparedGeometry(
torch.ones(1, dtype=torch.float16),
torch.ones(1, dtype=torch.long),
torch.ones(1, dtype=torch.bool),
1,
)
factors = depth.DepthFeatureFactorization(
torch.ones(1, dtype=torch.float16),
torch.ones(1, dtype=torch.float16),
)
actual_geometry, actual_factors = wrapped(model, geometry, factors)
assert isinstance(actual_geometry, depth.PreparedGeometry)
assert isinstance(actual_factors, depth.DepthFeatureFactorization)
assert actual_geometry.coords.dtype == torch.float32
assert actual_geometry.ranks.dtype == torch.long
assert actual_geometry.kept.dtype == torch.bool
assert actual_factors.depth.dtype == torch.float32
assert actual_factors.features.dtype == torch.float32

model.fp16_enabled = False
unchanged_geometry, unchanged_factors = wrapped(model, geometry, factors)
assert unchanged_geometry is geometry
assert unchanged_factors is factors


def test_extract_camera_features_preserves_camera_contract(monkeypatch):
captured = {}

Expand Down
7 changes: 7 additions & 0 deletions turbo_physai/optimizations/models/bevfusion/catalog.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,13 @@
"turbo_physai.optimizations.models.bevfusion.training.training_wrapper"
),
),
wrap(
target="mmcv.parallel.distributed.MMDistributedDataParallel._run_ddp_forward",
replacement=(
"turbo_physai.optimizations.models.bevfusion.training."
"ddp_forward_compat_wrapper"
),
),
)

GAUSSIAN = group(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
"model_commit": "326653dc06e0938edf1aae7d01efcd158ba83de5",
"config": {
"path": "optimization.yaml",
"sha256": "eadc1cf9155cf8279fad9fdae9188d24b9e136edecf6abb920175fd71ba4aed1"
"sha256": "ec41d78a5c8c5b8c053f47bb1996ddab72ee683e647902f71d07b951d13baf2d"
},
"inputs": [
{
Expand Down Expand Up @@ -35,7 +35,7 @@
"role": "catalog",
"base": "package",
"path": "optimizations/models/bevfusion/catalog.py",
"sha256": "309abcd7ea5c9f7cd64a89fab30ded23b62922397e711018fa155f5af9f88e21"
"sha256": "8fdfa085391af2e3bec60deaab92d1f66f0195d517c5f4877b55156eb184441c"
}
]
}
Original file line number Diff line number Diff line change
Expand Up @@ -124,9 +124,13 @@ optimization_groups:
source_hashes:
mmdet3d.apis.train.train_model:
- source-v1:6215bc796d762e6fea35f7289bf6edc46f9f406e377520e00bc7a500144080b1
mmcv.parallel.distributed.MMDistributedDataParallel._run_ddp_forward:
- source-v1:f4e196a85655850dbae1bf752ea1519a86083ba518b95a34cc5afbcf76bb5274
ast_hashes:
mmdet3d.apis.train.train_model:
- ast-v1:b1e7df6eca71976ac570963e0ee3419fa294faa6501a9c9321613779442bf2a7
mmcv.parallel.distributed.MMDistributedDataParallel._run_ddp_forward:
- ast-v1:7d601b5fefa14e6288d8b597d1b68043025af1567dc704226bdd24c17faab5d9
- id: bevfusion.gaussian
enabled: true
options: {}
Expand Down
34 changes: 31 additions & 3 deletions turbo_physai/optimizations/models/bevfusion/depth.py
Original file line number Diff line number Diff line change
Expand Up @@ -641,12 +641,40 @@ def base_transform_get_geometry(


def base_transform_bev_pool_wrapper(original, options):
"""Retain BaseTransform's MMCV FP32 boundary around optimized pooling."""
"""Cast pooling inputs to FP32 without MMCV's broken NamedTuple recursion."""

del original, options
from mmcv.runner import force_fp32
import functools
import torch

return force_fp32()(base_transform_bev_pool)
def fp32(value):
if isinstance(value, torch.Tensor) and value.dtype == torch.float16:
return value.float()
return value

@functools.wraps(base_transform_bev_pool)
def wrapped(self, geom_feats, x):
if not getattr(self, "fp16_enabled", False):
return base_transform_bev_pool(self, geom_feats, x)

if isinstance(geom_feats, PreparedGeometry):
geom_feats = PreparedGeometry(
fp32(geom_feats.coords),
geom_feats.ranks,
geom_feats.kept,
geom_feats.batch_size,
)
else:
geom_feats = fp32(geom_feats)
if isinstance(x, DepthFeatureFactorization):
x = DepthFeatureFactorization(fp32(x.depth), fp32(x.features))
else:
x = fp32(x)

with torch.amp.autocast("cuda", enabled=False):
return base_transform_bev_pool(self, geom_feats, x)

return wrapped


def base_transform_bev_pool_prepared(self, x, coords, ranks, kept):
Expand Down
14 changes: 14 additions & 0 deletions turbo_physai/optimizations/models/bevfusion/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,20 @@ def _option_flag(options, name, env_name, default=False):
return _env_flag(env_name, default)


def ddp_forward_compat_wrapper(original, options):
"""Support MMCV's DDP forward on newer PyTorch versions."""

del options

@functools.wraps(original)
def wrapped(self, *args, **kwargs):
if not hasattr(self, "_use_replicated_tensor_module"):
self._use_replicated_tensor_module = False
return original(self, *args, **kwargs)

return wrapped


def parse_losses(self, losses):
"""Reduce all scalar losses with one collective and one host transfer."""

Expand Down