掩码自编码器(MAE)视觉预训练与微调Skill tao-train-mask-auto-encoder

该技能提供 TAO 框架下的掩码自编码器(MAE)自监督预训练与微调能力。通过随机掩码图像补丁并重建图像来学习视觉表示,可支持视觉分类等下游任务的迁移学习,覆盖训练、评估、推理、导出与部署流程。关键词:MAE,Masked Autoencoder,掩码自编码器,自监督预训练,视觉表示学习,图像重建,微调,TAO,PyTorch,视觉模型

视觉模型训练 0 次安装 3 次浏览 更新于 9/6/2026
名称 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 命令行界面支持 trainevaluateinferenceexport。构建 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.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 都已打包时,默认通过 tao-skill-bank:tao-run-automl 将此训练操作路由到本模型的 skill_dir。保留工作流/应用程序对数据集、规范、输出目录、GPU/平台设置、父检查点和 automl_policy 的覆盖。仅在 automl_policy: off 或打包的训练模式/模板缺失时使用直接模型训练;若模式缺失,请报告 AutoML 已启用但对此模型不可运行,直到生成模式。

非训练操作(如 evaluateinferenceexport 和部署流程)保留在本模型技能中。每次运行时的 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_sourcesdataset.val_data_sourcesdataset.test_data_sources 指向解压后的 images_trainimages_valimages_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 ddpfsdp ddp
  • ddp 使用 find_unused_parameters=True
  • fsdp 强制 FP16
  • 强烈推荐使用多 GPU 进行预训练(需要大批量)

多节点环境变量(由编排器设置):WORLD_SIZENODE_RANKMASTER_ADDRMASTER_PORTNUM_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_modelparent_model_folder,请将上游训练/导出/AutoML 子作业 id 作为 parent_job_id 传入。SDK 会列出父结果文件夹,过滤检查点工件,并返回选定的模型文件或文件夹。不要将这些映射添加回 config.json,也不要通过猜测检查点路径来修补生成的运行器脚本。

当在 SDK 解析器之外解析检查点时,请准确选择所需的 epoch/step 工件,例如 model_epoch_000_step_00099.pth。仅当明确请求最新时才使用 convnextv2_atto_latest.pth 或其他最新符号链接。将 train.stagemodel.archmodel.num_classes 和导出输入大小向前传递到 evaluate、inference、export 和 deploy 规范中,以使检查点与 ONNX/引擎形状匹配。

部署