練指南:用 references/classification 腳本訓(xùn)練字符分類器與方向分類器(PyTorch))
人工智能深度學(xué)習(xí)計(jì)算機(jī)視覺OCR【免費(fèi)下載鏈接】doctrdocTR (Document Text Recognition) - a seamless, high-performing accessible library for OCR-related tasks powered by Deep Learning. Ongoing development and maintenance by t2k.項(xiàng)目地址https://gitcode.com/gh_mirrors/do/doctr點(diǎn)擊查看免費(fèi)下載導(dǎo)讀本文基于 docTRDocument Text Recognition倉庫中 references/classification 目錄下的官方訓(xùn)練腳本完整講解如何用 PyTorch 訓(xùn)練兩類圖像分類模型字符分類器Character Classification識(shí)別單個(gè)字符屬于哪個(gè)字符類別與方向分類器Orientation Classification判斷文檔頁面或單詞裁剪圖被旋轉(zhuǎn)了 0°、90°、180° 還是 -90°。讀完本文后你將掌握從環(huán)境安裝、數(shù)據(jù)組織、命令行參數(shù)含義、設(shè)備與混合精度配置到 checkpoint 管理、訓(xùn)練監(jiān)控、ONNX 導(dǎo)出與 HuggingFace Hub 推送的完整實(shí)戰(zhàn)流程并了解這些腳本在 docTR 源碼層面的底層實(shí)現(xiàn)邏輯。一、環(huán)境準(zhǔn)備訓(xùn)練腳本依賴 docTR 本體以及若干訓(xùn)練輔助庫官方 README 給出的安裝方式如下pip install -e . --upgrade pip install -r references/requirements.txt第一條命令以可編輯模式安裝當(dāng)前倉庫根目錄下的 docTR-e .指向倉庫根目錄的 setup.py / pyproject.toml第二條命令安裝訓(xùn)練輔助依賴內(nèi)容見 references/requirements.txt包括依賴用途tqdm訓(xùn)練進(jìn)度條且支持通過tqdm.contrib.slack將進(jìn)度推送到 Slackslack-sdkSlack 日志所需的 SDKwandb0.10.31Weights Biases 實(shí)驗(yàn)跟蹤配合--wb使用clearml1.11.1ClearML 實(shí)驗(yàn)跟蹤配合--clearml使用matplotlib3.1.0樣本可視化與 LR 搜索曲線繪制安裝完成后兩個(gè)入口腳本分別位于 references/classification/train_character.py 與 references/classification/train_orientation.py。二、字符分類訓(xùn)練train_character.py字符分類任務(wù)的定位是為 OCR 流水線中的單個(gè)字符訓(xùn)練分類頭可配合識(shí)別模型或獨(dú)立使用。最簡單的訓(xùn)練命令為python references/classification/train_character.py mobilenet_v3_large --epochs 5 --device 0第一個(gè)位置參數(shù)arch指定分類骨干網(wǎng)絡(luò)例如mobilenet_v3_large--epochs 5表示訓(xùn)練 5 個(gè) epoch--device 0表示使用 CUDA 設(shè)備 0設(shè)備選擇規(guī)則詳見下文。從源碼 train_character.py 的parse_args()可以拿到完整的參數(shù)清單逐一說明如下參數(shù)默認(rèn)值說明arch必填分類模型架構(gòu)名如mobilenet_v3_small、resnet18、vit_s等--output_dir.checkpoint 與元數(shù)據(jù)保存目錄--nameNone實(shí)驗(yàn)名缺省時(shí)自動(dòng)生成為{arch}_{時(shí)間戳}--epochs10訓(xùn)練輪數(shù)-b/--batch_size64訓(xùn)練批量大小--device自動(dòng)選擇CUDA index如0、cuda:N、mps、cpu--input_size32輸入尺寸模型輸入為(input_size, input_size)的方形圖--lr0.001Adam / AdamW 學(xué)習(xí)率--wd/--weight-decay0權(quán)重衰減AdamW 下若為 0 會(huì)自動(dòng)落到1e-4-j/--workers自動(dòng)DataLoader 工作進(jìn)程數(shù)缺省為min(16, cpu_count)--resumeNone從指定 checkpoint 恢復(fù)權(quán)重model.from_pretrained--fontFreeMono.ttf,FreeSans.ttf,FreeSerif.ttf合成字符圖像使用的字體族逗號(hào)分隔--vocabfrench訓(xùn)練詞表取自doctr.datasets.VOCABS--train-samples1000每個(gè)字符的合成訓(xùn)練樣本數(shù)總樣本數(shù) train_samples × vocab 長度--val-samples20每個(gè)字符的合成驗(yàn)證樣本數(shù)--test-only關(guān)閉只跑驗(yàn)證循環(huán)不訓(xùn)練--show-samples關(guān)閉展示一批未歸一化的訓(xùn)練樣本--wb關(guān)閉啟用 Weights Biases 日志--clearml關(guān)閉啟用 ClearML 日志--push-to-hub關(guān)閉訓(xùn)練結(jié)束后推送模型到 HuggingFace Hub--pretrained關(guān)閉訓(xùn)練前加載 ImageNet 預(yù)訓(xùn)練權(quán)重--export-onnx關(guān)閉訓(xùn)練結(jié)束后導(dǎo)出 ONNX 模型--optimadam優(yōu)化器可選adam/adamw--schedcosine學(xué)習(xí)率調(diào)度可選cosine/onecycle/poly--amp關(guān)閉開啟自動(dòng)混合精度僅 CUDA 支持--amp-dtypefloat16autocast 精度可選float16/bfloat16--find-lr關(guān)閉學(xué)習(xí)率網(wǎng)格搜索LR Finder--early-stop關(guān)閉啟用早停--early-stop-epochs5早停耐心值patience--early-stop-delta0.01早停最小改善閾值字符分類的數(shù)據(jù)來源在線合成字符分類腳本不需要外部數(shù)據(jù)集而是使用 docTR 的CharacterGenerator在線合成單字符圖像。核心邏輯在 doctr/datasets/generator/base.pysynthesize_text_img使用 PIL 以默認(rèn)字號(hào) 32 渲染單個(gè)字符自動(dòng)裁出字形包圍盒并把單字符渲染為方形圖_fonts_per_char會(huì)檢查每個(gè)字體是否真的能渲染詞表中的字符避免字體缺字形導(dǎo)致.notdef空框污染訓(xùn)練數(shù)據(jù)無法渲染的字符會(huì)嘗試系統(tǒng)字體回退腳本中train_set CharacterGenerator(vocabvocab, num_samplesargs.train_samples * len(vocab), ...)見 train_character.py即總樣本數(shù)為train_samples × 詞表長度每個(gè)字符在其可渲染字體間輪換取樣。合成之后腳本通過img_transforms施加了一整套數(shù)據(jù)增強(qiáng)見 train_character.pyRandomApply(T.ColorInversion(), 0.9)90% 概率反色保證樣本中 90% 是白底黑字RandomGrayscale(p0.1)、RandomPhotometricDistort(p0.1)灰度化與光度擾動(dòng)RandomApply(T.RandomShadow(), p0.4)隨機(jī)陰影RandomApply(T.GaussianNoise(mean0, std0.1), 0.1)與RandomApply(T.GaussianBlur(sigma(0.5, 1.5)), 0.3)噪聲與模糊RandomPerspective(distortion_scale0.2, p0.3)與RandomRotation(15)透視畸變與 ±15° 旋轉(zhuǎn)。喂入網(wǎng)絡(luò)前統(tǒng)一做Normalize(mean(0.694, 0.695, 0.693), std(0.299, 0.296, 0.301))train_character.py。詞表選擇--vocab參數(shù)取值來自 doctr/datasets/vocabs.py 的VOCABS字典例如latin數(shù)字 ASCII 字母 標(biāo)點(diǎn)englishlatin 基礎(chǔ)上加°與貨幣符號(hào)frenchenglish 基礎(chǔ)上再疊加法語變音字符àaéèê????ù?ü?...這也是默認(rèn)值此外還有g(shù)erman、spanish、portuguese、chinese、japanese、korean、cyrillic等數(shù)十種語言詞表可供選擇。分類模型的輸出類別數(shù)等于len(vocab)類別順序即詞表字符順序源碼classification.__dict__args.arch, classeslist(vocab))。三、方向分類訓(xùn)練train_orientation.py方向分類器解決的是 OCR 前處理中的文檔旋轉(zhuǎn)校正問題輸入可能是整頁文檔圖像page也可能是單詞裁剪圖crop模型輸出四類旋轉(zhuǎn)角度[0, -90, 180, 90]見 train_orientation.py 的CLASSES常量。典型命令python references/classification/train_orientation.py resnet18 --type page \ --train_path path/to/your/train_set --val_path path/to/your/val_set --epochs 5--type為必填參數(shù)只能取page文檔整頁或crop單詞裁剪圖并直接決定輸入尺寸input_size (512, 512) if args.type page else (256, 256)見 train_orientation.py。與字符分類不同方向分類必須提供真實(shí)圖片數(shù)據(jù)--train_path與--val_path為必填參數(shù)分別指向訓(xùn)練/驗(yàn)證圖片文件夾。腳本內(nèi)部使用doctr.datasets.OrientationDataset實(shí)現(xiàn)見 doctr/datasets/orientation.py讀取文件夾內(nèi)全部圖片初始目標(biāo)統(tǒng)一記為 0°。旋轉(zhuǎn)標(biāo)簽的在線生成訓(xùn)練/驗(yàn)證時(shí)腳本通過sample_transforms中的rnd_rotate函數(shù)train_orientation.py為每張圖動(dòng)態(tài)生成旋轉(zhuǎn)標(biāo)簽先從CLASSES [0, -90, 180, 90]中隨機(jī)選一個(gè)基準(zhǔn)角度以 50% 概率疊加一個(gè)從-25°到25°步長 5°的隨機(jī)微調(diào)模擬真實(shí)掃描中并非精確旋轉(zhuǎn)的情況用torchvision.transforms.functional.rotate執(zhí)行旋轉(zhuǎn)并填充背景。因此模型學(xué)到的是對大致屬于某類角度的魯棒判別而非對精確角度的記憶。方向分類訓(xùn)練同樣帶有一套增強(qiáng)ColorInversion 0.1、GaussianNoise、RandomShadow 0.2、GaussianBlur 0.3、RandomPhotometricDistort、RandomGrayscale、RandomPerspective 等見 train_orientation.py并先做Resize(input_size, preserve_aspect_ratioTrue, symmetric_padTrue)等比縮放 對稱填充。注意方向分類腳本默認(rèn)--batch_size 2比字符分類的 64 小得多原因是 page 類型的輸入是 512×512 的大圖。四、數(shù)據(jù)目錄格式兩個(gè)腳本對用戶數(shù)據(jù)的組織方式非常寬容只要給到圖片文件夾路徑即可。官方 README 給出的目錄結(jié)構(gòu)為images ├── sample_img_01.png ├── sample_img_02.png ├── sample_img_03.png └── ...對方向分類--train_path/--val_path指向這樣的圖片目錄腳本內(nèi)部讀取os.path.join(path, images)子目錄見 train_orientation.py 與 train_orientation.py字符分類則完全無需外部數(shù)據(jù)。五、設(shè)備選擇與混合精度設(shè)備解析規(guī)則--device支持多種寫法CUDA 索引如0、cuda:N、mpsApple Silicon GPU或cpu。設(shè)備解析邏輯在 references/classification/utils.py 的resolve_device()中實(shí)現(xiàn)不傳該參數(shù)時(shí)腳本按CUDA → MPS → CPU的順序自動(dòng)選擇可用設(shè)備傳入純數(shù)字索引時(shí)若 CUDA 不可用或索引越界會(huì)直接報(bào)錯(cuò)在分布式訓(xùn)練torchrun場景下該參數(shù)會(huì)被忽略每個(gè)進(jìn)程自動(dòng)使用自己的 GPU。另外當(dāng)設(shè)備為 CUDA 時(shí)腳本會(huì)開啟torch.backends.cudnn.benchmark True以加速卷積train_character.py。自動(dòng)混合精度AMP--amp開啟自動(dòng)混合精度僅支持 CUDA源碼在非 CUDA 設(shè)備上開啟--amp會(huì)拋出ValueError(--amp (automatic mixed precision) is only supported on CUDA devices)--amp-dtype可選float16默認(rèn)或bfloat16。官方 README 特別說明bfloat16 適用于 Ampere 或更新的 GPU它擁有與 float32 相同的指數(shù)范圍因此無需 loss scaling也能避免 float16 在某些損失函數(shù)中出現(xiàn)的上溢問題。對應(yīng)源碼實(shí)現(xiàn)amp_dtype()utils.py把字符串映射為torch.bfloat16/torch.float16_autocast()與_scaler()train_character.py中GradScaler僅在float16時(shí)啟用bfloat16時(shí)自動(dòng)禁用。官方 README 給出了兩個(gè)針對性示例# Apple Silicon: 開啟回退讓 MPS 缺失的少數(shù)算子回落到 CPU 執(zhí)行 PYTORCH_ENABLE_MPS_FALLBACK1 python references/classification/train_character.py mobilenet_v3_small --epochs 5 --device mps # NVIDIA GPU bfloat16 混合精度 python references/classification/train_character.py mobilenet_v3_small --epochs 5 --device 0 --amp --amp-dtype bfloat16六、Checkpoint 與運(yùn)行元數(shù)據(jù)每次訓(xùn)練會(huì)產(chǎn)出兩類文件均保存在--output_dir文件名以實(shí)驗(yàn)名exp_name為前綴權(quán)重文件exp_name.pt由save_checkpoint()utils.py通過torch.save(model.state_dict(), ...)保存。注意訓(xùn)練循環(huán)只在驗(yàn)證損失下降時(shí)才覆蓋保存if val_loss min_loss見 train_character.py因此目錄下始終是當(dāng)前最優(yōu)權(quán)重。元數(shù)據(jù)文件exp_name.json由save_run_metadata()utils.py每次運(yùn)行只寫一次而非隨每個(gè) checkpoint 重復(fù)寫。元數(shù)據(jù)由run_metadata()utils.py生成內(nèi)容包括官方 README 所述的全部要素架構(gòu)名與任務(wù)設(shè)置classes/vocab_name使用本地?cái)?shù)據(jù)時(shí)數(shù)據(jù)集的哈希值docTR 版本doctr.__version__與 PyTorch 版本torch.__version__Git revisiongit rev-parse HEAD獲取失敗時(shí)為None本次運(yùn)行的全部命令行參數(shù)dict(vars(args))。這份 JSON 的價(jià)值在于它記錄了重建模型與復(fù)現(xiàn)訓(xùn)練所需的全部信息即使事后忘了當(dāng)初怎么訓(xùn)練的也能據(jù)此還原出可推理的模型。七、訓(xùn)練監(jiān)控tqdm、Slack、WB 與 ClearML進(jìn)度條與 Slack 推送訓(xùn)練循環(huán)統(tǒng)一使用 tqdm 進(jìn)度條train_character.py。如果同時(shí)設(shè)置了以下兩個(gè)環(huán)境變量進(jìn)度條會(huì)切換到tqdm.contrib.slack把訓(xùn)練信息直接推送到 Slack 頻道TQDM_SLACK_TOKENSlack Bot TokenTQDM_SLACK_CHANNEL在 Slack 頻道上右鍵 → Copy → Copy link可得到形如https://xxxxxx.slack.com/archives/yyyyyyyy的鏈接只需保留最后的yyyyyyyy部分作為頻道 ID。腳本中還會(huì)對 tqdm 的write方法做 monkey patch讓進(jìn)度消息直接通過 Slack API 發(fā)送train_character.py。WB 與 ClearML傳入--wb時(shí)腳本調(diào)用wandb.init(nameexp_name, projectcharacter-classification)初始化實(shí)驗(yàn)并記錄每步的train_loss_step/val_loss_step/step_lr以及每 epoch的train_loss/val_loss/learning_rate/acctrain_character.py傳入--clearml時(shí)使用 ClearMLTaskLogger.report_scalar記錄相同指標(biāo)項(xiàng)目名為docTR/character-classification方向分類腳本為docTR/orientation-classification。八、進(jìn)階選項(xiàng)LR Finder、早停與調(diào)度器腳本內(nèi)置了若干訓(xùn)練技巧開關(guān)均源自 train_character.pyLR Finder--find-lrrecord_lr()從1e-7到1指數(shù)增長地網(wǎng)格搜索學(xué)習(xí)率plot_recorder()utils.py以對數(shù)橫軸繪制學(xué)習(xí)率-損失曲線并做 EMA 平滑幫助你確定合適的學(xué)習(xí)率區(qū)間。該實(shí)現(xiàn)改編自 Holocron 訓(xùn)練器。早停--early-stop配合--early-stop-epochs默認(rèn) 5與--early-stop-delta默認(rèn) 0.01使用EarlyStopperutils.py在驗(yàn)證損失連續(xù)patience個(gè) epoch 未能比歷史最優(yōu)降低超過min_delta時(shí)終止訓(xùn)練。優(yōu)化器與調(diào)度器--optim可選adambetas(0.95, 0.999), eps1e-6或adamwbetas(0.9, 0.999)weight_decay 缺省 1e-4--sched可選cosineCosineAnnealingLReta_minlr/25e4、onecycleOneCycleLR或polyPolynomialLR三者均按epochs × len(train_loader)的總步數(shù)調(diào)度。九、可用骨干架構(gòu)與推理入口可訓(xùn)練的架構(gòu)清單分類腳本的arch位置參數(shù)從doctr.models.classification命名空間取模型完整清單見 doctr/models/classification/zoo.py 的ARCHS包括magc_resnet31、mobilenet_v3_small/large含_r變體、resnet18/31/34/50、resnet34_wide、textnet_tiny/small/base、vgg16_bn_r、vit_s/b、vip_tiny/base、vit_det_s/m、starnet_s3等。官方 README 示例中字符分類用mobilenet_v3_large、方向分類用resnet18。方向分類另有兩條專門預(yù)訓(xùn)練入口同為分類骨干見 zoo.py 的ORIENTATION_ARCHScrop_orientation_predictor(archmobilenet_v3_small_crop_orientation)處理單詞裁剪圖默認(rèn)batch_size128page_orientation_predictor(archmobilenet_v3_small_page_orientation)處理整頁默認(rèn)batch_size4。兩者均在PreProcessor中使用preserve_aspect_ratioTrue, symmetric_padTrue預(yù)處理輸入尺寸與訓(xùn)練腳本的(256,256)/(512,512)對應(yīng)。訓(xùn)練后的自定義模型可類比這兩個(gè)入口在推理階段用相同預(yù)處理加載權(quán)重。訓(xùn)練后導(dǎo)出與共享兩個(gè)訓(xùn)練腳本都支持--test-only僅加載模型跑驗(yàn)證循環(huán)并打印Validation loss (Acc: ...)--show-samples可視化一批未歸一化的增強(qiáng)樣本plot_samplesutils.py--export-onnx訓(xùn)練結(jié)束后用export_model_to_onnx導(dǎo)出 ONNX 模型train_character.py--push-to-hub訓(xùn)練結(jié)束后通過push_to_hf_hub(model, exp_name, taskclassification, run_configargs)推送到 HuggingFace Hub需要先執(zhí)行l(wèi)ogin_to_hub()登錄--resume checkpoint用model.from_pretrained()恢復(fù)權(quán)重繼續(xù)訓(xùn)練。十、性能基準(zhǔn)測試latency.pyreferences/classification/latency.py 提供了分類模型的延遲基準(zhǔn)腳本隨機(jī)生成(batch_size, 3, size, size)的輸入先做 10 次 warmup再運(yùn)行--it默認(rèn) 100次計(jì)時(shí)輸出平均與標(biāo)準(zhǔn)差延遲毫秒。典型用法python references/classification/latency.py mobilenet_v3_small --size 32 --batch-size 64 --gpu --it 100 --pretrained--gpu在 CUDA:0 上評(píng)測否則回落到 CPU--pretrained加載模型庫預(yù)訓(xùn)練權(quán)重。該腳本有助于在訓(xùn)練前后快速評(píng)估不同架構(gòu)的推理成本為部署選型提供依據(jù)。十一、完整參數(shù)速查需要查看任意腳本的完整幫助時(shí)官方 README 建議直接使用python references/classification/train_character.py --help python references/classification/train_orientation.py --help由于腳本使用argparse.ArgumentDefaultsHelpFormatter--help輸出會(huì)附帶每個(gè)參數(shù)的默認(rèn)值是查詢參數(shù)語義最權(quán)威的途徑。小結(jié)docTR 的 references/classification 目錄提供了兩套開箱即用的 PyTorch 分類訓(xùn)練管線字符分類完全基于字體在線合成數(shù)據(jù)無需外部數(shù)據(jù)集方向分類基于真實(shí)圖片 在線旋轉(zhuǎn)標(biāo)簽增強(qiáng)兩者共享同一套設(shè)備解析、混合精度、checkpoint 元數(shù)據(jù)、實(shí)驗(yàn)跟蹤與模型導(dǎo)出機(jī)制。配合 doctr/models/classification/zoo.py 中十余種骨干架構(gòu)開發(fā)者可以快速訓(xùn)練出面向自己語種、版式或掃描質(zhì)量的定制分類模型再通過 ONNX 導(dǎo)出或 HuggingFace Hub 推送接入生產(chǎn)推理鏈路。贊分享人工智能深度學(xué)習(xí)計(jì)算機(jī)視覺OCR【免費(fèi)下載鏈接】doctrdocTR (Document Text Recognition) - a seamless, high-performing accessible library for OCR-related tasks powered by Deep Learning. Ongoing development and maintenance by t2k.項(xiàng)目地址https://gitcode.com/gh_mirrors/do/doctr點(diǎn)擊查看免費(fèi)下載相關(guān)推薦視頻分類模型訓(xùn)練指南用Ludwig實(shí)現(xiàn)分類訓(xùn)練視頻分類模型訓(xùn)練指南用Ludwig實(shí)現(xiàn)分類訓(xùn)練 1. 視頻分類的技術(shù)挑戰(zhàn)與解決方案 你是否在構(gòu)建視頻分類系統(tǒng)時(shí)面臨以下痛點(diǎn)標(biāo)注數(shù)據(jù)不足導(dǎo)致模型泛化能力差復(fù)人工智能深度學(xué)習(xí)機(jī)器學(xué)習(xí)大模型預(yù)訓(xùn)練微調(diào)LoRA多模態(tài)NLP計(jì)算機(jī)視覺模型推理服務(wù)PaddleOCR文本方向分類模型配置與訓(xùn)練指南PaddleOCR文本方向分類模型配置與訓(xùn)練指南 文本方向分類模型現(xiàn)狀分析 PaddleOCR項(xiàng)目中提供的文本方向分類功能主要用于識(shí)別文本行的方向0度、90度人工智能計(jì)算機(jī)視覺OCR深度學(xué)習(xí)大模型RAGccv ConvNet 深度卷積網(wǎng)絡(luò)指南從 ImageNet 預(yù)訓(xùn)練模型分類到自訓(xùn)練圖像分類器ccv ConvNet 深度卷積網(wǎng)絡(luò)指南從 ImageNet 預(yù)訓(xùn)練模型分類到自訓(xùn)練圖像分類器 ccvC based/Cached/Core Compute計(jì)算機(jī)視覺深度學(xué)習(xí)上一篇如何構(gòu)建自定義nfs-subdir-external-provisioner鏡像完整構(gòu)建流程詳解下一篇Nav2路徑規(guī)劃器對比SMAC、Theta*、NavFn哪個(gè)更適合你的應(yīng)用場景創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考