🔖

RF-DETRを追加学習してCPUで動くドキュメントレイアウト検出器を作ってみた

に公開

はじめに

RAGなど用いてAIに文書を読み込ませる際に、PDFをうまくパースする技術がよく使われます。
例えば、論文のようなドキュメントだと、タイトル、テキスト、図、表、といったパーツを理解して、それらをマークダウンなどの形式に変換することで、AIの文書読み込みの精度を上げることができます。

一方で、このようなレイアウト検出器は深層学習をベースにした手法が精度が高く、それらが使われることも多いため、GPU環境が必要になってきます。
GPU環境作ってデプロイするのは金銭的にも時間的にも少し面倒な感じなので、この辺がCPUだけでサクッと動かせるようにならないかなと思っていました。

ドキュメントパースで有名なものとしてdoclingというパッケージがありますが、doclingではRT-DETRという比較的軽量で高速な物体検出モデルを使っており、CPUでもそれなりの速度で動くのですが、ちょっと遅いかなと言う感じでした。
https://zenn.dev/kun432/scraps/2a1e2456fc68bf

昨年に、RF-DETRというRT-DETRよりさらに高速で精度の高いアルゴリズムが発表されました。
これを使うことで、よりCPUで高速に動くレイアウト検出器が作れるのではと思い、今回は、RF-DETRを使ってドキュメントレイアウト分析モデルを作れるか試してみたいと思います。
https://roboflow.com/model/rf-detr

その他のCPUで動くドキュメントレイアウトモデル

DocLayout-YOLO

YOLOをDocLayNetで追加学習したモデルです。
ライセンスがAGPLなので商用では使いづらい場面ありそうです。
https://github.com/opendatalab/DocLayout-YOLO

Yomitoku

日本語の文書に特化したドキュメント解析ソフトです。
レイアウト検出だけでなく、OCRや表の構造解析も行うことができます。
https://github.com/kotaro-kinoshita/yomitoku

RF-DETRとは?

物体検出は自動運転などのリアルタイム性が高いシステムで使用されることが多く、しばしば精度と計算速度の両立が求められます。
これまで両者のバランスに優れたモデルといえばCNNを用いたYOLOが一般的だったかと思います。
一方で、Transformerを用いた物体検出モデルも研究されてきた中で、これまで処理が比較的重いとされてきたTransformerベースの物体検出モデルがYOLOに匹敵する速度と精度を持つようになってきました。
有名所でいくとRT-DETRやLW-DETRなどが登場し、YOLOの代替としても使えるレベルのモデルが出てきたのではないかと思います。
そんな中で、RF-DETRは昨年発表された新しいモデルで、 LW-DETRと事前学習済みのDINOv2バックボーンを組み合わせたアーキテクチャになっています。

ドキュメントレイアウト分析用データセット

データセットはdoclingでも使用しているDocLayNetを今回使用しました。
https://github.com/DS4SD/DocLayNet
DocLayNetは論文に限らず比較的幅広い分野の文書を扱っており、文書内の以下のようなカテゴリがアノテーションされています。

CLASS_NAMES = [
    "Caption","Footnote","Formula","List-item","Page-footer",
    "Page-header","Picture","Section-header","Table","Text","Title"
]

学習方法

RF-DETRの公式ドキュメントに則って、学習用のコードを準備します。
COCOデータセットであればかなり簡単に学習を行うことができます。
DocLayNetもCOCOデータセットなのですが、RF-DETRの学習で期待されているディレクトリ構成と微妙に異なるので、簡単なディレクトリ構造の変換を行うコードを用意しました。

RF-DETRはパッケージがかなりちゃんと作られていて、学習用のコードも簡単に準備できました。
以下のような感じで、train関数を呼ぶだけで学習を開始できます。

from rfdetr import RFDETRMedium

CLASS_NAMES = [
    "Caption","Footnote","Formula","List-item","Page-footer",
    "Page-header","Picture","Section-header","Table","Text","Title"
]


def parse_args():
    import argparse
    import os

    dataset_default = os.environ.get("SM_CHANNEL_TRAIN", "dataset")
    output_default = os.environ.get("SM_MODEL_DIR", "models/rfdetr-doclayout")
    parser = argparse.ArgumentParser()
    parser.add_argument("--dataset_dir", type=str, default=dataset_default)
    parser.add_argument("--output_dir", type=str, default=output_default)
    parser.add_argument("--epochs", type=int, default=30)
    parser.add_argument("--batch_size", type=int, default=4)
    parser.add_argument("--grad_accum_steps", type=int, default=4)
    parser.add_argument("--lr", type=float, default=1e-4)
    parser.add_argument("--resolution", type=int, default=1120)
    parser.add_argument("--tensorboard", type=bool, default=True)
    parser.add_argument("--early_stopping", type=bool, default=True)
    parser.add_argument("--early_stopping_patience", type=int, default=5)
    args = parser.parse_args()
    return args


def train_rfdetr(args):
    model = RFDETRMedium(class_names=CLASS_NAMES)

    model.train(
        dataset_dir=args.dataset_dir,
        epochs=args.epochs,
        batch_size=args.batch_size,
        grad_accum_steps=args.grad_accum_steps,
        lr=args.lr,
        output_dir=args.output_dir,
        resolution=args.resolution,
        tensorboard=args.tensorboard,
        early_stopping=args.early_stopping,
        early_stopping_patience=args.early_stopping_patience,
    )


if __name__ == "__main__":
    args = parse_args()
    train_rfdetr(args)

SageMakerを使った学習

実際に学習を回すとそこそこのデータ量があって、1GPUで1エポックに2時間以上かかるので、私の作業PCでこれをやると学習で一杯になって他の作業ができなくなってしまうのもあり、今回はAWS SageMakerを使って学習を行いました。
RF-DETRはモデルのサイズ違いでいくつか種類があるのですが、今回はMediumモデルを使用しています。
学習はml.g5.2xlargeを使ってだいたい3日間くらいで、1学習で1万円いくかいかないかくらいになってると思います。

以下がSageMaker用のデプロイコードです。上記の学習用のコードを呼び出すような形になってます。
データはS3に保存してそのURLを指定することで、SageMakerが立ち上げたインスタンスにマウントされるようになっています。

import os
from datetime import datetime

import boto3
from dotenv import load_dotenv
from sagemaker.experiments.run import Run
from sagemaker.pytorch import PyTorch
from sagemaker.inputs import TrainingInput

load_dotenv()

NUM = r"([+-]?\d+(?:\.\d+)?(?:[eE][+-]?\d+)?)"

bucket = os.getenv("AWS_BUCKET_NAME")
role_name = os.getenv("AWS_SAGEMAKER_ROLE_NAME")
iam = boto3.client("iam")
role = iam.get_role(RoleName=role_name)["Role"]["Arn"]

train_input = TrainingInput(
    s3_data=f"s3://{bucket}/dataset/",
    input_mode="FastFile",
    distribution="ShardedByS3Key",
)

est = PyTorch(
    entry_point="doclaynet_train.py",
    source_dir="scripts",
    role=role,
    framework_version="2.2",
    py_version="py310",
    instance_type="ml.g5.2xlarge",
    instance_count=1,
    use_spot_instances=False,
    enable_sagemaker_metrics=True,
    metric_definitions=[
        {"Name": "epoch", "Regex": r"Epoch:\s*\[(\d+)\]"},
        {"Name": "class_error", "Regex": rf"class_error:\s*{NUM}"},
        {"Name": "loss", "Regex": rf"loss:\s*{NUM}\s*\("},
        {"Name": "time", "Regex": rf"time:\s*{NUM}"},
    ],
    # max_wait=12*60*60,
    max_run=72*60*60,
    checkpoint_s3_uri=f"s3://{bucket}/checkpoints/rfdetr/",
    checkpoint_local_path="/opt/ml/checkpoints",
    output_path=f"s3://{bucket}/models/rfdetr-doclayout/",
    hyperparameters={
        "dataset_dir": "/opt/ml/input/data/train",
        "output_dir": "/opt/ml/model",
    },
)

run_name = f"rfdetr-{datetime.now():%Y%m%d-%H%M%S}"
with Run(experiment_name="doclaynet-exp", run_name=run_name) as run:
    est.fit({"train": train_input}, wait=False)
    job_name = est.latest_training_job.name
    print("Training Job Started.")
    print("Please check the training job at https://console.aws.amazon.com/sagemaker/home")
    print(f"Job name: {job_name}")

できたモデルを試す

学習されたモデルをONNXに変換し、CPUで動かしてテストしてみました。
以下のようなPDFを画像にしたものをテストで使用しました。検出結果を重ねて描画しています。

CPUで実行して画像一枚あたり、0.3秒位で実行できました。
これくらいであればCPU上でも使えそうかなと思います。

今回学習したモデルをONNXに変換したものを以下のhuggingfaceに置いています。
https://huggingface.co/neka-nat/rfdetr-doclaynet-onnx

また今回使用したコードは以下のリポジトリに置いています。
今回はSageMakerで学習させましたが、ローカルでも実行できるコードと別れているので、ローカルで学習させることも可能です。
https://github.com/neka-nat/rfdetr-doclayout

まとめ

本来であれば、doclingやDocLayout-YOLO/Yomitokuといったモデルとのベンチマークをとってみたかったのですが、一旦記事にまとめたかったので、追々、時間あればやってみようかなと思ってます。
とりあえず、CPUでもそこそこ速く動くモデルができてよかったです。

Discussion