| 名称 | tao-train-mask-auto-encoder |
| 描述 | 掩码自编码器(Masked Auto-Encoder,MAE)用于自监督预训练和微调。通过掩码随机块并重建它们来学习视觉表示;支持预训练和微调阶段。当需要训练、评估、导出或对 TAO MAE 骨干网络进行推理时使用。触发短语包括“pretrain MAE”、“self-supervised vision pretraining”、“Masked Autoencoder”、“Mask Auto-Encoder”、“MAE fine-tune”。 |
| 开源协议 | Apache-2.0 compatibility: 需要 docker + nvidia-container-toolkit。 metadata: |
| 版本 | ‘0.1.0’ |
| 作者 | NVIDIA Corporation allowed-tools: 读取 Bash tags: - 自监督 - 监督 - 学习 |
掩码自编码器(MAE)
独立安装? 如果本会话不是由 TAO 技能库插件初始化的,请先运行
tao-setup技能(主机预检、凭据、跨技能发现)。
MAE(掩码自编码器)用于自监督预训练和微调。通过掩码随机块并重建它们来学习视觉表示。支持预训练和微调两个阶段。
微调时,请将 train.pretrained_model_path 设置为预训练的 MAE 权重。
对于 TAO Deploy TensorRT 操作(gen_trt_engine),请先阅读 references/tao-deploy-mask-auto-encoder.md。部署规范模板位于本技能的 references/ 文件夹中,文件前缀为 spec_template_deploy_*.yaml。
父级 PyTorch mae 命令行界面支持 train、evaluate、inference 和 export。构建 TensorRT 引擎应通过部署工作流,而非模型技能。
数据类模式
生成的 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.json 和 references/spec_template_<action>.yaml 且可解析。请使用打包好的所选操作模式来获取 automl_default_parameters、automl_disabled_parameters、默认值、最小/最大边界、枚举、选项权重、数学条件、依赖关系和常用参数。运行时不依赖 ~/tao-core;维护者在打包技能库之前重新生成模式/模板。
训练操作策略
该模型在模型层启用了 AutoML。在处理任何训练阶段请求前,请阅读 references/skill_info.yaml,并从显式的 automl_policy 值或用户的工作流请求中解析出本次运行的覆盖设置。默认使用 automl_policy: on,并且只在新启动提示中暴露 on / off。将“关闭 AutoML”、“禁用 AutoML”、“无 HPO”或“普通训练”等短语视为本次运行的 automl_policy: off。当 automl_policy: on、automl_enabled: true,并且 schemas/train.schema.json 与 references/spec_template_train.yaml 都已打包时,默认通过 tao-skill-bank:tao-run-automl 将此训练操作路由到本模型的 skill_dir。保留工作流/应用程序对数据集、规范、输出目录、GPU/平台设置、父检查点和 automl_policy 的覆盖。仅在 automl_policy: off 或打包的训练模式/模板缺失时使用直接模型训练;若模式缺失,请报告 AutoML 已启用但对此模型不可运行,直到生成模式。
非训练操作(如 evaluate、inference、export 和部署流程)保留在本模型技能中。每次运行时的 automl_policy 覆盖不会改变模型元数据。
训练要求
- 数据集类型: image_classification
- 格式: ssl
- 接受的数据集用途: training、evaluation、testing
- 监控指标: train_loss
各操作的数据集要求
| 操作 | 规范键 | 来源 | 文件 | 列表? |
|---|---|---|---|---|
| train | dataset.train_data_sources | train_datasets | images_train.tar.gz | 否 |
| train | dataset.val_data_sources | eval_dataset | images_val.tar.gz | 否 |
| evaluate | dataset.val_data_sources | eval_dataset | images_val.tar.gz | 否 |
| inference | dataset.test_data_sources | inference_dataset | images_test.tar.gz | 否 |
对于 SDK/应用作业输入,images_*.tar.gz 归档作为操作输入上传。对于针对宿主机挂载数据的直接本地 Docker 运行,请先解压这些归档,然后将 dataset.train_data_sources、dataset.val_data_sources 和 dataset.test_data_sources 指向解压后的 images_train、images_val 和 images_test 文件夹。直接将本地 tar 路径传给 MAE CLI 可能导致零样本数据加载器,因为本地数据加载器不会解包该归档路径。
典型规范覆盖
数据源覆盖对每个操作都是强制的 —— 代理必须根据上面的“各操作的数据集要求”表构造数据源路径,并将其包含在 spec_overrides 中。
S3_TRAIN = 's3://bucket/data/train'
S3_EVAL = 's3://bucket/data/eval'
训练(强制数据源):
{
'dataset.train_data_sources': f'{S3_TRAIN}/images_train.tar.gz',
'dataset.val_data_sources': f'{S3_EVAL}/images_val.tar.gz',
'train.num_epochs': 10,
'train.optim.lr': 2e-4,
}
评估(强制数据源):
{
'dataset.val_data_sources': f'{S3_EVAL}/images_val.tar.gz',
'evaluate.checkpoint': '<selected train/AutoML checkpoint>',
'train.stage': 'finetune',
}
推理(强制数据源):
{
'dataset.test_data_sources': f'{S3_EVAL}/images_test.tar.gz',
'inference.checkpoint': '<selected train/AutoML checkpoint>',
'train.stage': 'finetune',
}
评估数据集
可选。预训练不需要评估数据。微调可选择使用验证集。
重要参数
- train.stage:训练阶段。可选值:pretrain、finetune。预训练通过掩码学习表示。微调添加分类头。
- model.arch:架构。默认
convnextv2_base。对于本地冒烟 AutoML,使用convnextv2_atto而不是不支持的名称(如vit_tiny_patch16)。支持的系列包括vit_base_patch16及更大的 ViT、ConvNeXtV2 atto/femto/pico/nano/tiny/base/large/huge,以及 Hiera tiny/small/base/large/huge。 - model.num_classes:微调时的类别数。默认1000(ImageNet)。仅在微调阶段相关。
- model.mask_ratio:预训练期间掩码的块比例。通常为 0.75。
- model.norm_pix_loss:是否在重建损失中对像素值进行归一化。
- dataset.augmentation.input_size:保持本地冒烟配置文件为 224 以适应 ConvNeXtV2 MAE。减小到 112 可能会使 MAE 掩码网格与特征图尺寸不兼容。
- MAE 不公开
dataset.workers规范字段。不要将其添加到冒烟测试覆盖中;Hydra 在训练前会拒绝未知的数据集键。 - train.optim.lr:学习率。默认 2e-4。
- dataset.augmentation:增强设置,包括微调的 mixup、cutmix。
多 GPU / 多节点
启动方法: Lightning 托管(单个 python 进程,Lightning 派生子进程)。
| 规范键 | 描述 | 默认 |
|---|---|---|
train.num_gpus |
GPU 数量 | 1 |
train.gpu_ids |
GPU 设备索引 | [0] |
train.num_nodes |
节点数量 | 1 |
train.distributed_strategy |
ddp 或 fsdp |
ddp |
ddp使用find_unused_parameters=Truefsdp强制 FP16- 强烈推荐使用多 GPU 进行预训练(需要大批量)
多节点环境变量(由编排器设置):WORLD_SIZE、NODE_RANK、MASTER_ADDR、MASTER_PORT、NUM_GPU_PER_NODE。
硬件
最低 2 个 GPU,推荐 8 个 GPU。每个 GPU 需 24GB+ 显存(推荐 A100)。MAE 预训练受益于跨多个 GPU 的大批量。微调的资源需求较低。
错误模式
阶段不匹配:确保 train.stage 与您的意图(预训练 vs 微调)一致。在没有 pretrained_model_path 的情况下进行微调会从零开始训练。
使用预训练检查点进行推理:MAE 预测数据加载器在 train.stage: pretrain 时会抛出 NotImplementedError。对于推理和分类式评估,请使用 finetune 检查点;或仅将纯预训练运行限于 train/evaluate/export。
num_classes 不匹配(仅微调):微调时确保 model.num_classes 与数据集类别数一致。
规范参数 / 父模型推断
特定于模型的推断映射记录在此 MD 文件中,而不是 config.json。生成的运行器应读取本部分,并在 create_job() 之前通过 SDK 助手应用这些映射。这类似于旧微服务的 infer_params.py 流程。
来自 TAO Core mae.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 |
当前作业结果目录 |
| 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 | results_dir |
output_dir |
当前作业结果目录 |
| train | train.pretrained_model_path |
ptm_if_no_resume_model |
如果没有恢复检查点则使用 PTM |
| train | train.resume_training_checkpoint_path |
resume_model |
从当前作业结果文件夹推断出的模型文件 |
对于 parent_model 或 parent_model_folder,请将上游训练/导出/AutoML 子作业 id 作为 parent_job_id 传入。SDK 会列出父结果文件夹,过滤检查点工件,并返回选定的模型文件或文件夹。不要将这些映射添加回 config.json,也不要通过猜测检查点路径来修补生成的运行器脚本。
当在 SDK 解析器之外解析检查点时,请准确选择所需的 epoch/step 工件,例如 model_epoch_000_step_00099.pth。仅当明确请求最新时才使用 convnextv2_atto_latest.pth 或其他最新符号链接。将 train.stage、model.arch、model.num_classes 和导出输入大小向前传递到 evaluate、inference、export 和 deploy 规范中,以使检查点与 ONNX/引擎形状匹配。