BEVdet模型解析
BEVdet模åè§£æ
å°å¹³çº¿å¼åè 2026-08-14 0 é 读6åéBEVdet模åè§£æ
BEVDet模å代ç è§£æä¸å®ç°ç»è
ä¸ãæ¨¡åæ¶ææ¦è¿°BEVDetéç¨åé¶æ®µå¤çæµç¨å®æ3Dç®æ æ£æµä»»å¡ï¼
Image-view Encoderï¼å¯¹ç¯è§ç¸æºå¾åè¿è¡ç¹å¾æå View Transformerï¼å°å¾åè§è§ç¹å¾è½¬æ¢ä¸ºé¸ç°å¾(BEV)ç¹å¾ BEV Encoderï¼å¯¹BEVç¹å¾è¿è¡ç¼ç å¢å¼º Headï¼å®ææç»çç®æ æ£æµ
è¯¥æ¶æéè¿view transformerçææ¾å¼çBEVç¹å¾è¡¨ç¤ºï¼å ¶å½¢å¼ç±»ä¼¼äºç¹äºç¹å¾ãä¸ºäºæåæ§è½ï¼æ¨¡åéç¨CUDAå évoxel poolingæä½ï¼å¹¶å¨æ£æµå¤´ä¸ä½¿ç¨ä¼åçNMSç®æ³ã
äºãæ ¸å¿ä»£ç æµç¨è§£æï¼
1.tools/test.py æµè¯å ¥å£
outputs = single_gpu_test(...)
# -> mmdet3d/apis/test.py
else:
...
2.mmdet3d/apis/test.py æ¨çè°åº¦
return self.forward_train(**kwargs)
else:
return self.forward_test(**kwargs)
# -> mmdet3d/models/detectors/base.py
3.mmdet3d/models/detectors/bevdet.py BEVDet主ä½å®ç°
def __init__(...):
...
def forward_test(...):
if not isinstance(img_inputs[0][0], list):
return self.simple_test(...)
def simple_test(...):
img_feats, _, _ = self.extract_feat(...)
# åècenterpoint
bbox_pts = self.simple_test_pts(img_feats, img_metas, rescale=rescale)
def extract_feat(...):
img_feats, depth = self.extract_img_feat(...)
pts_feats = None
return (img_feats, pts_feats, depth)
def extract_img_feat(...):
# æåç¯è§å¾ççç¹å¾
x = self.image_encoder(img[0])
# BEV ç¹å¾
x, depth = self.img_view_transformer([x] + img[1:7])
# -> mmdet3d/models/necks/view_transformer.py
x = self.bev_encoder(x)
return [x],depth
ä¸ãè§è§è½¬æ¢æ ¸å¿å®ç°
3.1.mmdet3d/models/necks/view_transformer.py Voxel poolingçå ³é®æ¥éª¤ä¸ºvoxel_pooling_prepare_v2ï¼ä¸ºäºæ´å¥½ççè§£ï¼å¨ä»£ç 䏿¹åå¤äºå¾ä¾æ¥è¿è¡çè§£ã
def create_frustum(...):
...
def forward(self, input)ï¼
""" Transform image-view feature into bird-eye-view feature.
Args:
input: [image-view feature,rots,trans,intrins,post_rots,post_trans]
image-view feature:ç¯è§å¾çç¹å¾
rots:ç±ç¸æºåæ ç³»->è½¦èº«åæ ç³»çæè½¬ç©éµ
trans:ç¸æºåæ ç³»->è½¦èº«åæ ç³»ç平移ç©éµ
intrinsic:ç¸æºå
å
post_rots:ç±å¾åå¢å¼ºå¼èµ·çæè½¬ç©éµ
post_trans:ç±å¾åå¢å¼ºå¼èµ·ç平移ç©éµ
"""
# LIFT, x:[6, 139, 16, 44]
# åself.Dä¸ºé¢æµç离æ£è·ç¦»ï¼åself.out_channels为深度ç¹å¾
x = self.depth_net(x)
# 深度
depth_digit = x[:, :self.D, ...]
# ç¹å¾
tran_feat = x[:, self.D:self.D + self.out_channels, ...]
# 深度æ¦çåå¸
depth = depth_digit.softmax(dim=1)
# 转åå°bev空é´
return self.view_transform(input, depth, tran_feat)
def view_transform(...):
return self.view_transform_core(input, depth, tran_feat)
def view_transform_core(...):
'''
Args:
input:[1, 6, 512, 16, 44],ç¯è§ç¸æºç¹å¾
depth:[6, 59, 16, 44]ï¼# 深度æ¦çåå¸
tran_feat: [6, 80, 16, 44]ï¼æ·±åº¦ç¹å¾
'''
if ...:
...
else:
# è·å¾ç¹äº
coor = self.get_lidar_coor(*input[1:7])
# å°ç¹äºæå½±å°BEV空é´
# è®²è§£é¾æ¥å¯åè https://zhuanlan.zhihu.com/p/586637783
bev_feat = self.voxel_pooling_v2(...)
# bev_feat:[1, 80, 128, 128] depth:[6, 59, 16, 44]
return bev_feat,depth
def get_lidar_coor(...):
# self.frustum è§é¥
# å廿°æ®å¢å¼ºç平移ç©éµ
points = self.frustum.to(rots) - post_trans.view(B, N, 1, 1, 1, 3)
# ä¹ä»¥å¾åé¢å¤ççæè½¬ç©éµçéç©éµ
points = torch.inverse(post_rots).view(B, N, 1, 1, 1, 3, 3).matmul(points.unsqueeze(-1))
# å¾ååæ ç³» -> å½ä¸åç¸æºåæ ç³» -> ç¸æºåæ ç³» -> è½¦èº«åæ ç³»
# lamda * [xs, ys, 1 ] -> lamda * xs ,lamda * ys , lamdaï¼å¨å¤ä¸ªé¡¹ç®ä¸é½æä½ç°ï¼åç´ åæ ç³»è½¬ç¸æºåæ ç³»
points = torch.cat((points[..., :2, :] * points[..., 2:3, :], points[..., 2:3, :]), 5)
# ç¸æºå
å
combine = rots.matmul(torch.inverse(cam2imgs))
# ç¸æºåæ ç³»è½¬è½¦èº«åæ ç³»
points = combine.view(B, N, 1, 1, 1, 3, 3).matmul(points).squeeze(-1)
points += trans.view(B, N, 1, 1, 1, 3)
# bad 为BEV ç¹å¾ä¸çå¢å¼ºç©éµï¼è¿é为åä½ç©éµ
# è§£éæ¥æºä¸º https://github.com/Megvii-BaseDetection/BEVDepth/issues/44
points = bda.view(B, 1, 1, 1, 1, 3,3).matmul(points.unsqueeze(-1)).squeeze(-1)
return points
def voxel_pooling_v2(self, coor, depth, feat):
"""
Args:
coor:è½¦èº«åæ ç³»ä¸çè§é¥ç¹åæ
depth:ç¦»æ£æ·±åº¦æ¦çåå¸
feat:深度ç¹å¾
"""
ranks_bev, ranks_depth, ranks_feat, interval_starts, interval_lengths = self.voxel_pooling_prepare_v2(coor)
def voxel_pooling_prepare_v2(...):
"""Data preparation for voxel pooling
"""
B, N, D, H, W, _ = coor.shape
num_points = B * N * D * H * W # æ»è§é¥ç¹ä¸ªæ°
ranks_depth = torch.range(0, num_points - 1, dtype=torch.int, device=coor.device) # 0~249215
# æ¯ä¸å±featçä½ç½®ç´¢å¼ [0,1,2,3..4223,0,1,2...,4223,...,0,1,2...,4223]
ranks_feat = ...
# å°åç¹ç§»å¨å°å·¦ä¸è§å¹¶ä¸å°åæ 系转å°BEV空é´ç尺度
# [-51.2,51.2] -> [0,102.4] -> [0,128]
coor = ((coor - self.grid_lower_bound.to(coor)) / self.grid_interval.to(coor))
coor = coor.long().view(num_points, 3
# è®°å½å½åè§é¥ç¹å¨åªä¸ªbatch
batch_idx = torch.range(0, B - 1).reshape(B, 1). expand(B, num_points // B).reshape(num_points, 1).to(coor)
coor = torch.cat((coor, batch_idx), 1)
# è¿æ»¤æä¸å¨bev空é´ä¸çè§é¥ç¹
kept = (coor[:, 0] >= 0) & (coor[:, 0] < self.grid_size[0]) & \
(coor[:, 1] >= 0) & (coor[:, 1] < self.grid_size[1]) & \
(coor[:, 2] >= 0) & (coor[:, 2] < self.grid_size[2])
if len(kept) == 0:
return None, None, None, None, None
# æéBEV空é´ä¸çè§é¥ç¹
coor, ranks_depth, ranks_feat = coor[kept], ranks_depth[kept], ranks_feat[kept]
# å©ç¨è§é¥ ç¹çbatch,x,y 计ç®åº è§é¥ç¹å¨BEVç¹å¾ä¸çå
¨å±ç´¢å¼(128*128)
ranks_bev = coor[:, 3] * (self.grid_size[2] * self.grid_size[1] * self.grid_size[0])
ranks_bev += coor[:, 2] * (self.grid_size[1] * self.grid_size[0])
ranks_bev += coor[:, 1] * self.grid_size[0] + coor[:, 0]
# æåº,å°BEV空é´ä¸ï¼å
¨å±ç´¢å¼ä¸ºç¸åç弿åå¨ä¸èµ·
order = ranks_bev.argsort()
ranks_bev, ranks_depth, ranks_feat = ranks_bev[order], ranks_depth[order], ranks_feat[order]
kept = torch.ones(ranks_bev.shape[0], device=ranks_bev.device, dtype=torch.bool)
# é使¯è¾ï¼å¯ä»¥ä½¿å¾ç´¢å¼ä½ç½®ç¸åçï¼æ¶ä¸ªä½ç½®ä¸ºTrueï¼å¦å¾æç¤ºã
kept[1:] = ranks_bev[1:] != ranks_bev[:-1]
interval_starts = torch.where(kept)[0].int()
if len(interval_starts) == 0:
return None, None, None, None, None
interval_lengths = torch.zeros_like(interval_starts)
# æ¯ä¸ªä¸ºTrueçç´¢å¼ä½ç½®ï¼ååç´¯å çé¿åº¦
interval_lengths[:-1] = interval_starts[1:] - interval_starts[:-1]
interval_lengths[-1] = ranks_bev.shape[0] - interval_starts[-1]
return ranks_bev.int().contiguous(), ranks_depth.int().contiguous(
), ranks_feat.int().contiguous(), interval_starts.int().contiguous(
), interval_lengths.int().contiguous()
Voxel Pooling å¾ä¾
3.2.mmdet3d/ops/bev_pool_v2/src/bev_pool_cuda.cu
"""
Args:
c:80ï¼bevç¹å¾channel维度
n_intervals:Nd,ä½ç½®ä¸ºtrueçç´¢å¼çéå
å
¶ä»åæ°è§ä¸æ¹ç voxel_pooling_prepare_v2彿°
"""
# ç´¢å¼ä½ç½®ä¸ºTrueçè§é¥ç¹,æ¯ä¸ªè§é¥ç¹çç¹å¾æ·±åº¦ä¸º80 ä¸å
±å¼è¾ è§é¥ç¹ä¸ªæ°*80个thread
# å
±æ(int)ceil(((double)n_intervals * c / 256)) 个block ï¼æ¯ä¸ªblockæ 256ä¸ªçº¿ç¨ ,为æ¯ä¸ªæ·±åº¦ç¹å¾çæ¯ä¸å±(80å±)å建ä¸ä¸ªthread
bev_pool_v2_kernel<<<(int)ceil(((double)n_intervals * c / 256)), 256>>>(...);
}
__global__ void bev_pool_v2_kernel(...) {
// out:è¾åºçbevç¹å¾ [1,1,128,128,80]
int idx = blockIdx.x * blockDim.x + threadIdx.x; //å½åthreadçå
¨å±ç´¢å¼
int index = idx / c; // å½åå¤çåªä¸ä¸ªè§é¥ç¹
int cur_c = idx % c; // å½åå¤çåªä¸ä¸ªè§é¥ç¹ç第 cur_c å±çæ°æ® (å
±80å±)
if (index >= n_intervals) return;
int interval_start = interval_starts[index]; // 为Trueçç´¢å¼
int interval_length = interval_lengths[index]; // ååç´¯å å¤å°ä¸ªé¿åº¦
float psum = 0; //æå±æ·±åº¦ç¹å¾çç´¯å å
const float* cur_depth;
const float* cur_feat;
// ç´¯å
for(int i = 0; i < interval_length; i++){
cur_depth = depth + ranks_depth[interval_start+i]; # è§é¥ç¹ç颿µæ·±åº¦
cur_feat = feat + ranks_feat[interval_start+i] * c + cur_c; # è§é¥ç¹æ·±åº¦ç¹å¾
psum += *cur_feat * *cur_depth; # ç¸ä¹
}
const int* cur_rank = ranks_bev + interval_start; // ranks_bev + interval_start å¨bevç¹å¾çä½ç½®ç´¢å¼(128*128)ä¸çä½ç½®ç´¢å¼
float* cur_out = out + *cur_rank * c + cur_c; // å¨BEVç¹å¾ä¸çä½ç½®ç´¢å¼(128*128*80)ä¸çä½ç½®ç´¢å¼
*cur_out = psum;
}
æ»ç»ï¼
éè¿å¯¹BEVDetçäºè§£ï¼è¿ä¸æ¥çè§£äºLSSçææ³ï¼åæ¶ä¹çè§£äºVoxel Poolingä¸çåç§æä½ï¼åç»å¸æè½å¤åæ·±å ¥äºè§£å ç¯èªé¡¶èä¸ç论æï¼å¦PETR\PETRv2çã
Aitishiku.com