TAOCenterPose训练Skill tao-train-centerpose

该技能用于训练、评估、导出和推理CenterPose模型,实现6-DoF物体姿态估计与关键点检测。支持AutoML、多GPU训练、TensorRT部署,适用于机器人、智能汽车等场景的物体姿态回归。关键词:CenterPose、姿态估计、关键点检测、6-DoF、物体姿态、TAO训练、TensorRT。

视觉模型训练 0 次安装 0 次浏览 更新于 9/6/2026
名称 tao-train-centerpose
描述 CenterPose用于关键点/姿态估计。检测物体中心并为6-DoF物体姿态估计回归关键点位置。当需要为TAO CenterPose模型进行训练、评估、导出或运行推理时使用。触发短语包括“训练CenterPose”、“6-DoF物体姿态”、“关键点估计”、“物体姿态回归”。
开源协议 Apache-2.0 compatibility: 需要docker + nvidia-container-toolkit。 metadata:
版本 “0.1.0”
作者 NVIDIA Corporation allowed-tools: Read Bash tags: - pose - estimation

CenterPose

独立安装? 如果本会话不是由TAO技能库插件初始化的,请先运行tao-setup技能(主机预检、凭据、跨技能发现)。

CenterPose用于关键点/姿态估计。检测物体中心并回归关键点位置。用于6-DoF物体姿态估计。

设置 model.backbone.pretrained_backbone_path

对于TAO部署TensorRT操作(gen_trt_engine、TensorRT evaluate和TensorRT inference),请使用本技能references/文件夹中的部署规范模板,文件名前缀为spec_template_deploy_*.yaml

数据类模式

生成的TAO Core模式位于 schemas/<action>.schema.json,其中schemas/manifest.json列出了可用的操作。每个生成的模式还会从模式顶层default字段发出references/spec_template_<action>.yaml。AutoML的启用状态在references/skill_info.yaml中的模型层通过automl_enabled声明。某个操作的可运行AutoML要求存在schemas/<action>.schema.jsonreferences/spec_template_<action>.yaml并能正常解析。使用打包的选定操作模式获取automl_default_parametersautoml_disabled_parameters、默认值、最小/最大边界、枚举、选项权重、数学条件、依赖关系和常用参数。不要在运行时期望~/tao-core;维护者在打包技能库之前会重新生成模式/模板。

训练操作策略

该模型在模型层启用了AutoML。在处理任何训练阶段请求之前,请阅读references/skill_info.yaml,并根据显式的automl_policy值或用户的工作流请求来解析运行覆盖。默认使用automl_policy: on,仅在新的启动提示中暴露on/off。将“关闭AutoML”、“禁用AutoML”、“无HPO”或“普通训练”等短语视为本次运行的automl_policy: off。当automl_policy: onautoml_enabled: true,并且schemas/train.schema.jsonreferences/spec_template_train.yaml都已打包时,默认使用此模型的skill_dir通过tao-skill-bank:tao-run-automl路由训练操作。保留数据集、规范、输出目录、GPU/平台设置、父检查点和automl_policy的工作流/应用程序覆盖。仅当automl_policy: off或打包的训练架构/模板缺失时,才使用直接模型训练;在缺少架构的情况下,报告AutoML已启用但在此模型上不可运行,直到生成架构。

evaluateinferenceexport和部署流程等非训练操作保留在此模型技能中。每次运行的automl_policy覆盖不会改变模型元数据。

训练要求

  • 数据集类型: centerpose
  • **格式:**默认
  • 监控指标: val_3DIoU

每个操作的数据集要求

操作 规范键 来源 文件 列表?
evaluate dataset.test_data eval_dataset test.tar.gz
gen_trt_engine gen_trt_engine.tensorrt.calibration.cal_image_dir calibration_dataset train.tar.gz
inference dataset.inference_data inference_dataset val.tar.gz
train dataset.train_data train_datasets train.tar.gz
train dataset.val_data eval_dataset val.tar.gz

典型规范覆盖

数据源覆盖对于每个操作都是强制性的——代理必须根据上面的“每个操作的数据集要求”表构造数据源路径,并将其包含在spec_overrides中。

TRAIN_DIR = "/path/to/extracted/train"
VAL_DIR = "/path/to/extracted/val"
TEST_DIR = "/path/to/extracted/test"
INFER_DIR = VAL_DIR
CAL_IMAGE_DIRS = ["/path/to/extracted/train/<sequence_or_image_dir>"]

训练(强制数据源):

{
    "train.num_epochs": 30,
    "train.checkpoint_interval": 10,
    "train.validation_interval": 10,
    "train.num_gpus": 1,
    "dataset.category": "bike",
    "dataset.batch_size": 4,
    "dataset.train_data": TRAIN_DIR,
    "dataset.val_data": VAL_DIR,
}

评估(强制数据源):

{
    "dataset.category": "bike",
    "dataset.test_data": TEST_DIR,
}

推理(强制数据源):

{
    "dataset.category": "bike",
    "dataset.inference_data": INFER_DIR,
}

生成TRT引擎(强制数据源):

{
    "gen_trt_engine.tensorrt.calibration.cal_image_dir": CAL_IMAGE_DIRS,
}

评估数据集

可选。验证集和测试集作为单独的压缩包提供。

重要参数

  • dataset.num_classes:物体类别数。默认1。
  • dataset.num_joints:每个物体的关键点数。固定为8(边界框关键点)。有效范围:恰好8。
  • dataset.input_res:输入分辨率。固定为512。输出分辨率固定为128。
  • dataset.category:物体类别名称。默认为“cereal_box”。
  • model.backbone.model_type:默认为fan_small。架构中可用的骨干选项有限。
  • train.optim.lr:学习率。默认6e-5。MultiStep调度器,lr_steps=[90, 120],lr_decay=0.1。
  • train.loss_config:丰富的损失配置,包含开关:mse_loss, obj_scale, obj_scale_uncertainty, hps_uncertainty, reg_bbox, hm_hp。权重:wh_weight=0.1, off_weight=1, hp_weight=1。
  • inference.use_pnp:使用PnP进行6-DoF姿态估计。默认True。需要相机内参(focal_length_x/y, principle_point_x/y)。
  • export.input_width:导出输入大小。固定为512x512。opset_version=16。

多GPU/多节点

启动方式: Lightning管理(单个python进程,Lightning生成工作进程)。

规范键 描述 默认值
train.num_gpus GPU数量 1
train.gpu_ids GPU设备索引 [0]
  • 策略:auto(Lightning自动选择最佳策略)
  • 没有显式的num_nodesdistributed_strategy配置——仅单节点
  • sync_batchnorm

导出/TRT默认值

  • 导出输入:512x512(固定),opset 16
  • TRT数据类型:FP32、FP16、INT8
  • TRT opt_batch_size:4,max_batch_size:8

硬件

最低1个GPU,推荐2个GPU。每个GPU 16GB以上显存。CenterPose在输入分辨率和关键点数量方面中等程度占用内存。

错误模式

num_joints不匹配:确保dataset.num_joints与标注中的关键点数量匹配。

为本地Docker解压S3压缩包:starter-kit S3数据打包为train.tar.gzval.tar.gztest.tar.gz,但CenterPose TAO操作使用解压后的文件夹。解压每个压缩包,并将dataset.train_datadataset.val_datadataset.test_datadataset.inference_data设置为解压后的分割目录。

检查点交接:CenterPose训练会写入具体的检查点,如model_epoch_000_step_00008.pth和一个centerpose_model_latest.pth符号链接。对评估、推理、导出和恢复使用SDK/模型检查点解析器或确切的epoch/step检查点。仅当用户明确要求最新的检查点时使用符号链接。

TAO部署后处理器兼容性:使用从技能固定的部署镜像或所选平台解析的部署镜像。成功的gen_trt_engine运行不能证明部署evaluateinference有效;请分别检查这些操作的退出代码和日志,特别是对于CenterPose后处理器错误,例如TypeError: only 0-dimensional arrays can be converted to Python scalars

规范参数/父模型推理

模型特定的推理映射属于此MD文件,而不属于config.json。生成的运行器应在create_job()之前读取本部分并使用SDK辅助函数应用映射。这反映了旧的微服务infer_params.py流程。

来自TAO Core centerpose.config.json的推理映射:

操作 规范字段 推理函数 含义
evaluate encryption_key key 加密密钥
evaluate evaluate.checkpoint parent_model 从父作业结果文件夹推断的模型文件
evaluate evaluate.trt_engine parent_model 从父作业结果文件夹推断的模型文件
evaluate results_dir output_dir 当前作业结果目录
export encryption_key key 加密密钥
export export.checkpoint parent_model 从父作业结果文件夹推断的模型文件
export export.onnx_file create_onnx_file 输出的ONNX路径
export results_dir output_dir 当前作业结果目录
gen_trt_engine encryption_key key 加密密钥
gen_trt_engine gen_trt_engine.onnx_file parent_model 从父作业结果文件夹推断的模型文件
gen_trt_engine gen_trt_engine.tensorrt.calibration.cal_cache_file create_cal_cache 校准缓存路径
gen_trt_engine gen_trt_engine.trt_engine create_engine_file 输出的TensorRT引擎路径
gen_trt_engine results_dir output_dir 当前作业结果目录
inference encryption_key key 加密密钥
inference inference.checkpoint parent_model 从父作业结果文件夹推断的模型文件
inference inference.trt_engine parent_model 从父作业结果文件夹推断的模型文件
inference results_dir output_dir 当前作业结果目录
train encryption_key key 加密密钥
train model.backbone.pretrained_backbone_path ptm_if_no_resume_model 当没有恢复检查点时使用的预训练模型
train results_dir output_dir 当前作业结果目录
train train.resume_training_checkpoint_path resume_model 从当前作业结果文件夹推断的模型文件

对于parent_modelparent_model_folder,将上游train/export/AutoML子作业id作为parent_job_id传递。SDK列出父结果文件夹,过滤检查点工件,并返回所选模型文件或文件夹。不要将这些映射添加回config.json,也不要修补生成的运行器脚本来猜测检查点路径。

部署