最近进度有点慢了,主要在给老师赶项目,身心俱疲,但还是想要坚持学习,祝愿每个苦命的牛马都能身体健康,哎~
另外VAD的代码实在是过于复杂,只能一点点看,不过本篇博客应该算是覆盖了VAD的整体框架了,如有不对的地方,还请各位大佬在评论区指出,多谢多谢~

之前都没看明白,原来!!在基于 ​MMDetection/MMLab 框架​(如 MMDetection3D、MMYOLO 等)的项目中,VAD_base_e2e.py 这类配置文件的核心作用是为训练脚本 train.py 和测试脚本 test.py提供统一的参数定义

MMCV相关知识: 官方文档:欢迎来到 MMCV 的中文文档! — mmcv 1.4.0 文档
知乎:
(3 封私信 / 11 条消息) MMCV 核心组件分析(四): Config - 知乎
注册机制:(3 封私信 / 11 条消息) MMCV 核心组件分析(五): Registry - 知乎
hook中有不少是自动实现的,没有显式调用,这点有点烦~(3 封私信 / 11 条消息) MMCV 核心组件分析(六): Hook - 知乎
训练的核心:负责 OpenMMLab 中所有框架的训练过程调度(2 封私信) MMCV 核心组件分析(七): Runner - 知乎

mmdet3d: MMDetection3D is an open source object detection toolbox based on PyTorch, towards the next-generation platform for general 3D detection. It is a part of the OpenMMLab project.

下面我详细解释给出的train.py和VAD_base_e2e.py(即配置文件)之间的关系和交互流程。

0. 大体结构:

个人感觉VAD整体的代码量确实不小,这里留一个脉络如下:

VAD: (VAD.py)
    - ResNet50 + FPN
    - VAD_Head (VAD_head.py)
    - train_cfg

    VADHead:
        - ego_agent_decoder
        - ego_map_decoder
        - motion_decoder
        - motion_map_decoder # CustomTransformerDecoder 输出query
        - transformer # VADPerceptionTransformer
            encoder # BEVFormerEncoder
            decoder # DetectionTransformerDecoder
            map_decoder # MapDetectionTransformerDecoder
        - bbox_coder # CustomNMSFreeCoder
        - map_bbox_coder # MapNMSFreeCoder
        - positional_encoding # LearnedPositionalEncoding
        - loss_cls 
        - loss ... 

1. 命令行的作用

你的命令如下:
argparse 中,像 parser.add_argument('config', ...) 这样没有加 – 前缀的参数叫做位置参数(positional argument),它是必须要提供的,并且会按照顺序自动赋值

python -m torch.distributed.run --nproc_per_node=8 --master_port=2333 tools/train.py projects/configs/VAD/VAD_base.py --launcher pytorch --deterministic --work-dir path/to/save/outputs
  • train.py:主训练脚本(你贴的就是它的内容)。
  • projects/configs/VAD/VAD_base.py:模型和训练的配置文件(你给的VAD_base_e2e.py)。
  • 其它参数如--launcher pytorch等用于分布式训练和运行环境设置。

2. 交互流程详解

步骤一:train.py 解析命令行参数

  • train.py 首先通过 argparse 解析命令行参数,获取配置文件路径、工作目录、分布式参数等。

步骤二:加载配置文件

  • train.py 通过 Config.fromfile(args.config) 加载你指定的配置文件(即VAD_base_e2e.py)。
  • 这个配置文件是一个Python脚本,定义了模型结构、数据集、优化器、训练流程等所有细节。

步骤三:合并命令行参数与配置

  • 如果你在命令行用 --work-dir--resume-from 等参数指定了内容,会覆盖配置文件中的同名参数。

步骤四:插件与自定义模块导入

  • 如果配置文件里有 plugin=True,train.py 会自动导入你指定的插件目录(如mmdet3d_plugin),以便支持自定义的模型、头、损失等。

步骤五:构建模型和数据集

  • train.py 读取配置文件中的 model 字典,调用 build_model 构建模型。
  • 读取 data 字典,调用 build_dataset 构建训练/验证数据集。

步骤六:训练主循环

  • train.py 调用 custom_train_model(或标准的train_model),传入模型、数据集、配置等,开始训练。
  • 训练过程中会自动使用配置文件中定义的优化器、学习率策略、损失函数、评估方式等。

3. 代码部分:

参数. VAD_base_e2e.py

./configs/VAD/VAD_base_e2e.py 参数文件:
其中的关键参数就是:model
这定义了整个模型的框架和组成其中的大框架有两部分:VADVADHead

mmcv/mmdet/mmdet3d 等 MM 系列框架通过 type=‘VADHead’ 这种配置,自动实例化并调用了你自定义的 VADHead 类。

  • 1、配置文件到模型实例化:- MM 系列框架会读取配置文件,利用注册机制(如 @HEADS.register_module())找到 VADHead 类,并实例化它(pts_bbox_head)
pts_bbox_head=dict(
    type='VADHead',
    ...
)
  • 2、注册机制:在 VAD_head.py 里写了:这会把 VADHead 注册到 HEADS 这个注册表里,type='VADHead’就能找到它。
@HEADS.register_module()
class VADHead(DETRHead):
    ...
    1. 构建模型:mmcv框架会根据配置自动构建整个模型(如 VAD,其中 pts_bbox_head 就是 VADHead 的实例,pts_bbox_head作为model的一个属性。
    1. 前向推理时的调用:在训练或推理时,VAD 模型的 forward 方法会被调用。VAD 的 forward 方法内部会调用 self.pts_bbox_head.forward(…),即 VADHead 的 forward。你自定义的 forward 方法就会被自动执行,返回定义的 outs 字典。
model:class VAD

比较简单~ 定义了train和test的主入口 并返回各分支的损失
注:pts_bbox_head 就是 VADHead

# 定义了train和test的主入口 并返回各分支的损失
    def forward(self, return_loss=True, **kwargs):
        if return_loss:
            return self.forward_train(**kwargs)
        else:
            return self.forward_test(**kwargs)

losses = self.pts_bbox_head.loss(*loss_inputs, img_metas=img_metas) # VADHead
        return losses
model:class VADHead

重点看看VADHead.py 的 forward 函数:
inputs:

mlvl_feats:#上游数据 (batch_size, num_cam, channels, height, width)
prev_bev: #previous bev featues   [batch_size, bev_h * bev_w, embed_dims]

outputs : Dictionary

outs = {
            'bev_embed': bev_embed,
            'all_cls_scores': outputs_classes,# 目标检测分支的预测
            'all_bbox_preds': outputs_coords,
            'all_traj_preds': outputs_trajs.repeat(outputs_coords.shape[0], 1, 1, 1, 1), 
            'all_traj_cls_scores': outputs_trajs_classes.repeat(outputs_coords.shape[0], 1, 1, 1),
            'map_all_cls_scores': map_outputs_classes, # 地图分支的预测
            'map_all_bbox_preds': map_outputs_coords,
            'map_all_pts_preds': map_outputs_pts_coords,
            'enc_cls_scores': None, # 这部分是encoder阶段的输出 只在双阶段下才有
            'enc_bbox_preds': None,
            'map_enc_cls_scores': None,
            'map_enc_bbox_preds': None,
            'map_enc_pts_preds': None,
            'ego_fut_preds': outputs_ego_trajs,# ego轨迹预测 [batch, ego_fut_mode, fut_ts, 2]
        }

首先值得说下这个 outputs,这里执行了执行了 BEVFormerEncoder DetectionTransformerDecoder MapDetectionTransformerDecoder 具体去看 class VADPerceptionTransformer

else: # true 输入多模态特征、查询、分支等,输出多层解码结果(目标和地图)。
            outputs = self.transformer( # 执行了DetectionTransformerDecoder and MapDetectionTransformerDecoder
                mlvl_feats,             # map_init_reference     and      map_inter_references
                bev_queries,
                object_query_embeds,
                map_query_embeds,
                self.bev_h,
                self.bev_w,
                grid_length=(self.real_h / self.bev_h,
                             self.real_w / self.bev_w),
                bev_pos=bev_pos,
                reg_branches=self.reg_branches if self.with_box_refine else None,  # noqa:E501
                cls_branches=self.cls_branches if self.as_two_stage else None,
                map_reg_branches=self.map_reg_branches if self.with_box_refine else None,  # noqa:E501
                map_cls_branches=self.map_cls_branches if self.as_two_stage else None,
                img_metas=img_metas,
                prev_bev=prev_bev
        )

outputs: (bev_embed, inter_states, init_reference_out, inter_references_out, map_inter_states, map_init_reference_out, map_inter_references_out)
  1. bev_embed:BEV特征(鸟瞰图空间特征)。
  2. inter_states:目标检测分支 transformer 解码器每一层的输出特征(目标 query 的特征)。每层所有目标 query 的高维特征,后续用于分类、回归、轨迹预测等分支。[num_layers, num_query, batch_size, embed_dims]
  3. init_reference_out:目标检测分支的初始参考点(ref point)。
    • shape[batch_size, num_query, 3]
    • 作用:每个目标 query 的初始空间位置(归一化坐标),用于解码器的增量回归。
  4. inter_references_out:目标检测分支每一层的参考点(ref point)。
    • shape[num_layers, batch_size, num_query, 3]
    • 作用:每层 refined 后的参考点,供回归分支做增量回归。
  5. map_inter_states 地图元素分支 transformer 解码器每一层的输出特征(map query 的特征)。
    • shape[num_layers, num_map_query, batch_size, embed_dims]
    • 作用:每层所有地图 query 的高维特征,后续用于地图元素的分类、回归、点集预测等。
  6. map_init_reference_out 地图元素分支的初始参考点。
    • shape[batch_size, num_map_query, 2]
    • 作用:每个地图 query 的初始空间位置(归一化坐标)。
  7. map_inter_references_out 地图元素分支每一层的参考点。
    • shape[num_layers, batch_size, num_map_query, 2]
    • 作用:每层 refined 后的地图 query 参考点,供回归分支做增量回归。

后续的hs 对应 inter_states 目标检测的输出特征 map_hs 同理

接下来用hs map_hs 做分类和位置预测

		for lvl in range(hs.shape[0]): # 对每一层num_decoder_layers 共6层
...
# 针对地图元素分支(如车道线、路界等)做类似处理
        for lvl in range(map_hs.shape[0]):
        ...

然后调用motion_decodermotion_map_decoder 对应论文中的 Vectorized Motion Transformer 输出 motion_hs 这里很关键,后续构成了 agent_query

# motion_decoder  最后输出query ./VAD/VAD_transformer.py
motion_hs = self.motion_decoder(
	query=motion_query,
	key=motion_query,
	value=motion_query,
	query_pos=motion_pos,
	key_pos=motion_pos,
	key_padding_mask=invalid_motion_idx)

ca_motion_query = self.motion_map_decoder(
	query=ca_motion_query,
	key=map_query,
	value=map_query,
	query_pos=motion_pos,
	key_pos=map_pos,
	key_padding_mask=key_padding_mask)
motion_hs = torch.cat([motion_hs, ca_motion_query], dim=-1)  # [B, A, fut_mode, 2D]

agent_query = motion_hs.reshape(batch, num_agent, -1)

motion_hs:关键中间输出:对应论文中的 updated agent query

好了,到这里我们已经得到了agent query 这个比较关键的参数了~
不过细心的读者会发现,不对啊,首先在论文当中agent query 的生成用到了map query,其次这里的ca_motion_query啥玩意,为啥要加这一部分,事实上,这部分就是来自map query
motion_map_decoder将map query中的特征融入到了agent query里
![[VADsub1.png|300]]
在这里插入图片描述

那么回头说说map query,这里就简单带过一下:map query来自于地图元素分支(如车道线、路界等)transformer解码器的输出特征map_hs
,其shape为[num_layers, batch_size, num_map_query, embed_dim] 这里取最后一层解码器输出,这里view先理解为reshape -> [batch_size, map_num_vec, map_num_pts_per_vec, embed_dim]
select_and_pad_pred_map()对地图元素(如车道线、路界等)的特征进行置信度筛选、距离过滤和 batch 统一 padding,并生成 mask,方便后续与 agent 进行交互。
输出:map_query:筛选、pad、过滤后的地图特征 [B*A, max_P, D]
map_pos:对应的地图元素坐标 [B*A, max_P, 2]
key_padding_mask:mask,True 表示该位置无效 [B*A, max_P]

map_query = map_hs[-1].view(batch_size, self.map_num_vec, self.map_num_pts_per_vec, -1)

map_query = self.lane_encoder(map_query)  
map_query, map_pos, key_padding_mask = self.select_and_pad_pred_map(
    motion_coords, map_query, map_score, map_pos, map_thresh=self.map_thresh, dis_thresh=self.dis_thresh,
pe_normalization=self.pe_normalization, use_fix_pad=True)

在沿着VAD结构进行梳理之前,先介绍一下 ego query 是怎么来的
Embedding(1, embed_dims=256) :- 创建了一个可学习的嵌入向量(embedding),用于表示“自车(ego vehicle)”的查询特征。shape为[1, 256]
接下来,可以关注一下ego_query的shape的变化,最后这里第‘1’维为1,因为自己只有一个嘛,即ego

self.ego_query = nn.Embedding(1, self.embed_dims) # [1, 256]

ego_his_feats = self.ego_query.weight.unsqueeze(0).repeat(batch, 1, 1)# [1,256] -> [1, 1, embed_dims] -> [batch, 1, 256]
ego_query = ego_his_feats# [batch, 1, 256]

下面自车和agent和map进行交互,对应论文中的planning transformer
ego <-> agent interaction 对应论文中planning transformer的中间步骤,得到ego_agent_query!!

ego_agent_query = self.ego_agent_decoder(
            query=ego_query.permute(1, 0, 2),
            key=agent_query.permute(1, 0, 2),
            value=agent_query.permute(1, 0, 2),
            query_pos=ego_pos_emb.permute(1, 0, 2),
            key_pos=agent_pos_emb.permute(1, 0, 2),
            key_padding_mask=agent_mask)

类似的,ego <-> map interaction,当然这里的query变为上面得到的ego_agent_query,这里输出得到 ego_map_query~

ego_map_query = self.ego_map_decoder(
	query=ego_agent_query,
	key=map_query.permute(1, 0, 2),
	value=map_query.permute(1, 0, 2),
	query_pos=ego_pos_emb.permute(1, 0, 2),
	key_pos=map_pos_emb.permute(1, 0, 2),
	key_padding_mask=map_mask)

上述过程对应论文中的这里:
![[VADsub2.png|200]]
在这里插入图片描述

emm?这里为什么呢又将得到的ego_map_query和前面的ego_agent_query拼起来?嘶~~
自车未来轨迹预测(Ego prediction)需要同时融合 自车与其他agent的交互特征 和 自车与地图的交互特征,但这里的ego_map_query同样也嵌入了agent的信息

elif self.ego_his_encoder is None and self.ego_lcf_feat_idx is None:  # true           
	ego_feats = torch.cat(
		[ego_agent_query.permute(1, 0, 2),
		 ego_map_query.permute(1, 0, 2)],
		dim=-1
	)  # [1, B, D] -> [B, 1, 2D]  

轨迹预测:
![[VADsub3.png|300]]
在这里插入图片描述

# Ego prediction
	outputs_ego_trajs = self.ego_fut_decoder(ego_feats)
	outputs_ego_trajs = outputs_ego_trajs.reshape(outputs_ego_trajs.shape[0], self.ego_fut_mode, self.fut_ts, 2)

最终输出: 不过这里的outs还有个bug 还没有融合 driving command!!!!

outs = {
            'bev_embed': bev_embed,
            'all_cls_scores': outputs_classes,# [num_layers, batch, num_query, num_classes(+1)]每个query的类别概率
            'all_bbox_preds': outputs_coords,# [num_layers, batch, num_query, box_code_size]
            'all_traj_preds': outputs_trajs.repeat(outputs_coords.shape[0], 1, 1, 1, 1), # [num_layers, batch, num_query, fut_mode, fut_ts, 2]
            'all_traj_cls_scores': outputs_trajs_classes.repeat(outputs_coords.shape[0], 1, 1, 1),
            'map_all_cls_scores': map_outputs_classes,
            'map_all_bbox_preds': map_outputs_coords,
            'map_all_pts_preds': map_outputs_pts_coords,
            'enc_cls_scores': None, # 这部分是encoder阶段的输出 只在双阶段下才有
            'enc_bbox_preds': None,
            'map_enc_cls_scores': None,
            'map_enc_bbox_preds': None,
            'map_enc_pts_preds': None,
            'ego_fut_preds': outputs_ego_trajs,# ego轨迹预测 [batch, ego_fut_mode, fut_ts, 2]
        }

TODO:

  • 1、transformer=dict(type='VADPerceptionTransformer',...)
  • 2、4个decodertype='CustomTransformerDecoder'
    class CustomTransformerDecoder:注意其中layers的定义
def forward(self,
                query,
                key=None,
                value=None,
                query_pos=None,
                key_pos=None,
                attn_masks=None,
                key_padding_mask=None,
                *args,
                **kwargs):
        """Forward function for `Detr3DTransformerDecoder`.
        Args:
            query (Tensor): Input query with shape
                `(num_query, bs, embed_dims)`.
        Returns:
            Tensor: Results with shape [1, num_query, bs, embed_dims] when
                return_intermediate is `False`, otherwise it has shape
                [num_layers, num_query, bs, embed_dims].
        """
        intermediate = []
        for lid, layer in enumerate(self.layers): # layers -> transformerlayers定义
            query = layer(
                query=query,
                key=key,
                value=value,
                query_pos=query_pos,
                key_pos=key_pos,
                attn_masks=attn_masks,
                key_padding_mask=key_padding_mask,
                *args,
                **kwargs)

            if self.return_intermediate: # false
                intermediate.append(query)

        if self.return_intermediate:
            return torch.stack(intermediate)

        return query # [1, num_query, bs, embed_dims] 一层的 MultiheadAttention

训练. ./mmdet3d_plugin/VAD/apis/mmdet_train.py

data_loaders = [
        build_dataloader(
            ds,
            cfg.data.samples_per_gpu,
            cfg.data.workers_per_gpu,
            # cfg.gpus will be ignored if distributed
            len(cfg.gpu_ids),
            dist=distributed,
            seed=cfg.seed,
            shuffler_sampler=cfg.data.shuffler_sampler,  # dict(type='DistributedGroupSampler'),
            nonshuffler_sampler=cfg.data.nonshuffler_sampler,  # dict(type='DistributedSampler'),
        ) for ds in dataset
    ]

# put model on gpus
...

# build runner
optimizer = build_optimizer(model, cfg.optimizer)
cfg.runner = {
            'type': 'EpochBasedRunner',
            'max_epochs': cfg.total_epochs
        }
...

# register hooks
...

runner.run(data_loaders, cfg.workflow)

Related Information

__init__.py:在Python中,__init__.py文件用于将一个目录标记为Python的包。这个机制允许Python进行模块导入和组织代码的分层结构。 -> 当然这也会有包的 嵌套 python中__init__.py的主要作用和用途_python中init.py的作用-CSDN博客

Checkpoint:在深度学习中,指在模型训练过程中保存的模型状态。
包括:

  1. 模型权重:模型的所有参数,包括权重和偏置。
  2. 优化器状态:优化器的状态,包括动量、学习率等。
  3. 训练状态:当前的训练轮数(epoch)、批次(batch)编号等。
  4. 其他元数据:如学习率调度器的状态、自定义指标等。

Python的argparse模块:用于处理命令行参数和选项;很有意思,相当于构建了一套关于命令行的函数~[每天一个python技巧]argparse.ArgumentParser()用法解析+action的使用分析-CSDN博客

get()函数:利用键来获取值;get("xx",0)拟定初值语句,相当于给xx=0;当让如果xx本来在字典中有对应的值的话,就直接获取,而不对其进行赋值python字典中get()函数的用法个人小结-CSDN博客

hasattr() 函数用于判断对象是否包含对应的属性。

isinstance(object,classtype):isinstance()用来判断一个对象是否是一个已知的类型

  • object – 实例对象。
  • classtype – 可以是直接或间接类名、基本类型或者由它们组成的元组。

Hook:简单来说就是能够改变程序执行流程的一种技术统称。Hook 技术无处不在,各大编程语言、软件、深度学习框架都有大量应用 Hook 技术。在我们熟知的 pytorch 中某个 tensor 或者 module 都有 register_hook(hook_fn) 函数,通过注册 hook,可以拦截和修改某些中间变量的值。

PyTorch 的广播机制是:如果两个张量在某一维的大小不同,但其中一个为1,则自动扩展为另一个的大小。 - 广播机制是“扩展”,不是“相加”。

补充mmcv运行机制

@HEADS.register_module() 是 ​​MMDetection/MMCV 框架中实现模块化设计的核心机制​​,其作用是 将自定义的 VADHead 类注册到名为 HEADS 的全局注册表中,从而允许在配置文件中通过字符串名称调用该类。
实现细节:

    1. 注册:注册到全局字典
      HEADS 是 mmdet.models 中预定义的 Registry 类实例(一个全局字典)
      @HEADS.register_module() 作为​​类装饰器​​,将 VADHead 类以键值对形式存入 HEADS._module_dict 字典中:
# 注册后,HEADS 内部存储变为:
_module_dict = {
    'VADHead': <class '__main__.VADHead'>,  # 类对象的引用
    # 其他已注册的 Head 类...
}
    1. 配置:在其他代码中调用
      通过配置文件动态构建​​,在配置文件中只需指定 type 字段为注册名,框架自动解析并实例化:
model = dict(
    bbox_head=dict(
        type='VADHead',  # 对应注册名称
        in_channels=256,
        num_classes=80,
        # 其他参数...
    )
)

測試记录:

$ CUDA_VISIBLE_DEVICES=0 python tools/test.py projects/configs/VAD/VAD_base_e2e.py ckpts/VAD_base.pth --launcher none --eval bbox --tmpdir tmp

projects.mmdet3d_plugin
生成的結果在VAD/test/VAD_tiny_e2e/Wed_Jun_11_16_06_10_2025/pts_bbox
文件:results_nusc.pkl
ps: 遇到一个小问题之前的路径输入是采用"./xx/xx.py";但其实原本就是在相对路径下面,所以,采用"xx/xx.py"即可

可視化

$ python tools/analysis_tools/visualization.py --result-path ./test/VAD_tiny_e2e/Wed_Jun_11_16_06_10_2025/pts_bbox/results_nusc.pkl --save-path ./visualization/results

–result-path 後面跟測試生成的結果
–save-path 後面是可視化結果的保存位置,可以自己指定 會生成六張圖片,且最前方的截圖會帶有生成的預測軌跡
如果提前创建好保存的文件夹的话,能够生成.mp4的视频

注:看一个端到端的代码,先看他的config里都用了哪些模型,比如VAD里的VAD_base_e2e.py

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐