OCRNetSkill tao-train-ocrnet

这个技能用于TAO OCRNet模型的场景文本识别,支持从裁剪文本图像中识别文字内容。它涵盖训练、评估、导出、剪枝、量化、重训练和推理等完整流程,并支持CTC与注意力两种解码器。关键词包括OCRNet、场景文本识别、Tesseract、CTC解码器、注意力解码器、训练、评估、TensorRT部署。

视觉模型训练 0 次安装 0 次浏览 更新于 9/6/2026
名称 tao-train-ocrnet
描述 OCRNet用于场景文本识别。从裁剪的文本区域图像中识别文本内容,支持CTC和基于注意力的解码器。在训练、评估、导出、剪枝、量化、重训练或对TAO OCRNet模型进行推理时使用。触发短语包括"train OCRNet"、“场景文本识别”、“OCR裁剪文本”、“CTC/attention文本解码器”。
开源协议 Apache-2.0 compatibility: 需要docker + nvidia-container-toolkit。 metadata:
版本 “0.1.0”
作者 NVIDIA Corporation allowed-tools: Read Bash tags: - 文本 - 识别

OCRNet

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

OCRNet用于场景文本识别。从裁剪的文本区域图像中识别文本内容。支持CTC和基于注意力的解码器。

对于TAO部署TensorRT操作(gen_trt_engine、TensorRT evaluate和TensorRT inference),请先阅读references/tao-deploy-ocrnet.md。部署规范模板位于此技能的references/文件夹中,前缀为spec_template_deploy_*.yaml

Dataclass 模式

生成的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覆盖不会更改模型元数据。

训练要求

  • 数据集类型: ocrnet
  • 格式: default
  • 监控指标: val_acc_1

每个操作的数据集要求

操作 规范键 来源 文件 列表?
dataset_convert dataset_convert.input_img_dir train_datasets或eval_dataset 包含裁剪文本图像的提取文件夹
dataset_convert dataset_convert.gt_file train_datasets或eval_dataset train/gt_new.txt或test/gt_new.txt
evaluate dataset.character_list_file eval_dataset character_list
evaluate evaluate.test_dataset_dir eval_dataset 提取的测试图像文件夹
evaluate evaluate.test_dataset_gt_file eval_dataset test/gt_new.txt
evaluate evaluate.checkpoint 父训练/AutoML作业 best_accuracy.pth或确切请求的周期检查点
export dataset.character_list_file eval_dataset character_list
export export.checkpoint 父训练/AutoML作业 best_accuracy.pth或确切请求的周期检查点
deploy/gen_trt_engine gen_trt_engine.tensorrt.calibration.cal_image_dir calibration_dataset 用于INT8校准的提取校准图像文件夹
deploy/gen_trt_engine gen_trt_engine.onnx_file 父导出作业 导出的.onnx工件
deploy/gen_trt_engine dataset.character_list_file eval_dataset character_list
inference dataset.character_list_file eval_dataset character_list
inference inference.inference_dataset_dir inference_dataset 提取的推理图像文件夹
inference inference.checkpoint 父训练/AutoML作业 best_accuracy.pth或确切请求的周期检查点
prune dataset.character_list_file eval_dataset character_list
prune prune.checkpoint 父训练/AutoML作业 best_accuracy.pth或确切请求的周期检查点
quantize dataset.train_dataset_dir dataset_convert训练作业 包含data.mdb和lock.mdb的LMDB文件夹
quantize dataset.val_dataset_dir dataset_convert评估作业 包含data.mdb和lock.mdb的LMDB文件夹
quantize dataset.character_list_file eval_dataset character_list
quantize dataset.quant_calibration_dataset.images_dir train_datasets 提取的校准图像文件夹
quantize quantize.model_path 父训练/AutoML作业 由解析器选择的检查点
retrain dataset.train_dataset_dir dataset_convert训练作业 包含data.mdb和lock.mdb的LMDB文件夹
retrain dataset.val_dataset_dir dataset_convert评估作业 包含data.mdb和lock.mdb的LMDB文件夹
retrain dataset.character_list_file eval_dataset character_list
retrain model.pruned_graph_path 父剪枝作业 剪枝的.pth工件
train dataset.train_dataset_dir dataset_convert训练作业 包含data.mdb和lock.mdb的LMDB文件夹
train dataset.train_gt_file train_datasets 使用原始文件夹而非LMDB时的train/gt_new.txt
train dataset.val_dataset_dir dataset_convert评估作业 包含data.mdb和lock.mdb的LMDB文件夹
train dataset.val_gt_file eval_dataset 使用原始文件夹而非LMDB时的test/gt_new.txt
train dataset.character_list_file eval_dataset character_list

检查点选择

OCRNet训练会同时写入best_accuracy.pth和类似model_epoch_000_step_00003.pth的周期步骤检查点。请使用SDK/模型检查点解析器通过references/skill_info.yaml中的spec_params映射进行解析;不要通过按最新.pth排序来猜测。

  • 对于最佳检查点的evaluateinferenceexportprune请求,使用best_accuracy.pth
  • 对于特定周期/步骤操作,使用确切请求的model_epoch_*_step_*.pth
  • 仅将train.resume_training_checkpoint_path用于恢复训练,并将model.pruned_graph_path用于从剪枝输出重训练。OCRNet在PyT镜像中没有单独的ocrnet retrain CLI子任务;模型技能的retrain操作通过设置剪枝图路径来路由到ocrnet train -e
  • OCRNetquantize通过PyTorch加载模型。对于由同一本地运行创建的可信检查点,如果PyTorch 2.6+拒绝该检查点作为仅权重的加载,请设置TORCH_FORCE_NO_WEIGHTS_ONLY_LOAD=1

典型规范覆盖

对于每个操作,数据源覆盖是强制性的。分别对训练和验证分割运行dataset_convert,然后将直接包含data.mdblock.mdb的LMDB文件夹传递给train、quantize和retrain。远程存储中的压缩包必须先解压才能用作图像目录。

TRAIN_IMAGES = "<提取的训练图像文件夹>"
TRAIN_GT = "<train gt_new.txt>"
EVAL_IMAGES = "<提取的评估图像文件夹>"
EVAL_GT = "<eval gt_new.txt>"
TRAIN_LMDB = "<train dataset_convert结果目录>"
EVAL_LMDB = "<eval dataset_convert结果目录>"
CHAR_LIST = "<character_list>"

dataset_convert(每个分割运行一次):

{
    "dataset_convert.input_img_dir": TRAIN_IMAGES,
    "dataset_convert.gt_file": TRAIN_GT,
}

train(强制数据源):

{
    "train.num_epochs": 30,
    "train.checkpoint_interval": 10,
    "train.validation_interval": 10,
    "train.num_gpus": 1,
    "dataset.batch_size": 16,
    "dataset.train_dataset_dir": [TRAIN_LMDB],
    "dataset.val_dataset_dir": EVAL_LMDB,
    "dataset.train_gt_file": "",
    "dataset.val_gt_file": "",
    "dataset.character_list_file": CHAR_LIST,
}

deploy/gen_trt_engine(强制数据源):

{
    "gen_trt_engine.onnx_file": "<选定的导出ONNX>",
    "gen_trt_engine.trt_engine": "<输出引擎路径>",
    "gen_trt_engine.tensorrt.calibration.cal_cache_file": "<输出校准缓存路径>",
    "gen_trt_engine.tensorrt.data_type": "fp16",
    "gen_trt_engine.tensorrt.calibration.cal_image_dir": [TRAIN_IMAGES],
    "dataset.character_list_file": CHAR_LIST,
}

evaluate(强制数据源):

{
    "evaluate.checkpoint": "<选定的训练/AutoML检查点>",
    "dataset.character_list_file": CHAR_LIST,
    "evaluate.test_dataset_dir": EVAL_IMAGES,
    "evaluate.test_dataset_gt_file": EVAL_GT,
}

export(强制数据源):

{
    "export.checkpoint": "<选定的训练/AutoML检查点>",
    "export.onnx_file": "<输出ONNX路径>",
    "dataset.character_list_file": CHAR_LIST,
}

inference(强制数据源):

{
    "inference.checkpoint": "<选定的训练/AutoML检查点>",
    "dataset.character_list_file": CHAR_LIST,
    "inference.inference_dataset_dir": EVAL_IMAGES,
}

prune(强制数据源):

{
    "prune.checkpoint": "<选定的训练/AutoML检查点>",
    "prune.pruned_file": "<输出剪枝PTH路径>",
    "dataset.character_list_file": CHAR_LIST,
}

quantize(强制数据源):

{
    "dataset.train_dataset_dir": [TRAIN_LMDB],
    "dataset.val_dataset_dir": EVAL_LMDB,
    "dataset.character_list_file": CHAR_LIST,
    "dataset.quant_calibration_dataset.images_dir": TRAIN_IMAGES,
    "quantize.model_path": "<选定的训练/AutoML检查点>",
}

retrain(强制数据源):

{
    "dataset.train_dataset_dir": [TRAIN_LMDB],
    "dataset.val_dataset_dir": EVAL_LMDB,
    "dataset.character_list_file": CHAR_LIST,
    "model.pruned_graph_path": "<选定的剪枝输出>",
}

评估数据集

可选。作为单独压缩包提供的测试数据。

重要参数

  • dataset.character_list_file:定义所支持字符集的字符列表文件路径。这决定了输出词汇表的大小。
  • model.backbone:默认ResNet。
  • model.prediction:解码器类型。CTC或Attn(基于注意力)。
  • train.optim.lr:学习率。默认1.0(Adadelta优化器)。高默认值特定于Adadelta。
  • dataset.batch_size:每GPU批大小。默认16。

多GPU / 多节点

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

规范键 描述 默认值
train.num_gpus GPU数量 1
train.gpu_ids GPU设备索引 [0]
train.distributed_strategy 策略名称 auto
  • 策略:单GPU时auto,多GPU时从配置读取train.distributed_strategy
  • 训练脚本中没有显式的num_nodes——面向单节点
  • 轻量级模型,单GPU通常足够

硬件

最低1个GPU,推荐1个GPU。每个GPU8GB以上VRAM。OCR文本识别是轻量级的。单GPU通常足够。

错误模式

需要dataset_convert:如果使用原始图像+gt文件,请先运行dataset_convert以生成LMDB格式。

dataset_convert输出文件夹:直接ocrnet dataset_convert会将data.mdblock.mdb写入dataset_convert.results_dir下。使用该文件夹本身作为dataset.train_dataset_dirdataset.val_dataset_dir、quantize和retrain的输入。SDK支持的运行可能会将相同的LMDB文件夹包装在作业工件目录内;解析包含data.mdblock.mdb的实际文件夹。

GT文件BOM:一些文本识别GT文件可能在第一个文件名前以UTF-8 BOM开头。如果数据集转换日志记录了一个在第一个图像名称之前带有不可见前缀的缺失路径,请在转换或评估前从GT文件的本地副本中去除BOM。

字符列表不匹配:训练数据中的所有字符都必须存在于character_list文件中。

导出/剪枝输字段必需export.onnx_fileprune.pruned_file必须是可写的输出路径。这些在references/skill_info.yaml中声明,以便SDK支持的模型运行可以自动创建路径。

TensorRT位于deploy中:PyT OCRNet CLI暴露了dataset_convertevaluateexportinferenceprunequantizetrain,但不包括gen_trt_engine。有关TensorRT引擎生成和TensorRT支持的评估/推理,请使用references/tao-deploy-ocrnet.mddeploy/skill_info.yaml

规范参数 / 父模型推断

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

来自TAO Core ocrnet.config.json的推断映射:

操作 规范字段 推断函数 含义
dataset_convert results_dir output_dir 当前作业结果目录
evaluate encryption_key key 加密密钥
evaluate evaluate.checkpoint parent_model 从父作业结果文件夹推断的模型文件
evaluate evaluate.trt_engine parent_model 从父作业结果文件夹推断的模型文件
evaluate model.pruned_graph_path pruned_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 当前作业结果目录
deploy/gen_trt_engine encryption_key key 加密密钥
deploy/gen_trt_engine gen_trt_engine.onnx_file parent_model 从父导出作业结果文件夹推断的ONNX文件
deploy/gen_trt_engine gen_trt_engine.tensorrt.calibration.cal_cache_file create_cal_cache 校准缓存路径
deploy/gen_trt_engine gen_trt_engine.trt_engine create_engine_file 输出TensorRT引擎路径
deploy/gen_trt_engine results_dir output_dir 当前作业结果目录
inference encryption_key key 加密密钥
inference inference.checkpoint parent_model 从父作业结果文件夹推断的模型文件
inference inference.trt_engine parent_model 从父作业结果文件夹推断的模型文件
inference model.pruned_graph_path pruned_model 父剪枝模型
inference results_dir output_dir 当前作业结果目录
prune encryption_key key 加密密钥
prune prune.checkpoint parent_model 从父作业结果文件夹推断的模型文件
prune prune.pruned_file create_pth_file 输出PTH路径
prune results_dir output_dir 当前作业结果目录
quantize encryption_key key 加密密钥
quantize quantize.model_path parent_model 从父作业结果文件夹推断的模型文件
quantize results_dir output_dir 当前作业结果目录
retrain encryption_key key 加密密钥
retrain model.pruned_graph_path parent_model 从父作业结果文件夹推断的模型文件
retrain 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,也不要修补生成的运行器脚本以猜测检查点路径。

部署