FCOS 目標検出モデルのコード詳細解説

はじめに

この記事では、FCOS(Fully Convolutional One-Stage Object Detection)という目標検出モデルのトレーニングプロセスとモデル構築の詳細なコードを解説します。特に、モデルの構築を担う `build_detection_model` 関数の内部処理に焦点を当てます。

トレーニングの実行

FCOSモデルのトレーニングは、以下のコマンドラインから開始されます。このコマンドは、PyTorchの分散トレーニング機能を利用して `train_net.py` スクリプトを実行します。

python -m torch.distributed.launch \
    --nproc_per_node=1 \
    --master_port=$((RANDOM + 10000)) \
    D:\Project\FCOS-master\tools\train_net.py \
    --config-file D:\Project\FCOS-master\configs\fcos\fcos_R_50_FPN_1x.yaml \
    DATALOADER.NUM_WORKERS 1 \
    OUTPUT_DIR D:\Project\FCOS-master\FCOS_imprv_R_50_FPN_1x.pth

train_net.py の役割

`train_net.py` スクリプトは、モデルのトレーニングとテストのエントリーポイントです。このスクリプトには、主に以下の関数が定義されています。

  • main(): メインのエントリーポイント関数。コマンドライン引数を解析し、分散トレーニングの設定を判断します。トレーニングかテストかを判定し、それぞれに対応する関数を呼び出します。
  • train(): 実際のトレーニングループを含む関数。
  • test(): モデルの評価を行う関数。

main() 関数内で、以下のように train() 関数が呼び出されます。

model = train(cfg, args.local_rank, args.distributed)

train() 関数の内部

train() 関数の最初の重要なステップは、設定ファイル(`cfg`)に基づいてモデルを構築することです。これは build_detection_model(cfg) 関数によって実行されます。

model = build_detection_model(cfg)

この build_detection_model 関数は、fcos_core.modeling.detector モジュールからインポートされ、GeneralizedRCNN クラスのインスタンスを生成します。

GeneralizedRCNN クラス

GeneralizedRCNN クラスは、PyTorchの nn.Module を継承した、一般的なRCNNベースのモデルの基底クラスです。そのコンストラクタは、モデルの主要な3つのコンポーネントを初期化します。

class GeneralizedRCNN(nn.Module):
    def __init__(self, cfg):
        super(GeneralizedRCNN, self).__init__()
        # バックボーン(特徴抽出ネットワーク)
        self.backbone = build_backbone(cfg)
        # RPN(領域提案ネットワーク)
        self.rpn = build_rpn(cfg, self.backbone.out_channels)
        # ROIヘッド(領域ごとのタスク処理)
        self.roi_heads = build_roi_heads(cfg, self.backbone.out_channels)

ここで、build_backbonebuild_rpnbuild_roi_heads は、それぞれ fcos_core.modeling.backbonefcos_core.modeling.rpn.rpnfcos_core.modeling.roi_head.roi_heads からインポートされます。

バックボーンの構築: build_backbone()

build_backbone() 関数は、モデルのバックボーン(例: ResNet)を構築する責任を持ちます。この関数は、Registry というデザインパターンを利用して、設定に応じて適切なバックボーン構築関数を選択します。

Registry クラスは、異なるモジュール(バックボーン、ROIヘッド、損失関数など)を登録し、管理するためのヘルパークラスです。以下にその実装の一部を示します。

class Registry(dict):
    def __init__(self, *args, **kwargs):
        super(Registry, self).__init__(*args, **kwargs)

    def register(self, module_name, module=None):
        if module is not None:
            self[module_name] = module
            return module
        # デコレータとして使用する場合
        def register_fn(fn):
            self[module_name] = fn
            return fn
        return register_fn

この Registry を使って、特定のバックボーン構築関数が登録されます。例えば、ResNetバックボーンの場合は以下のように登録されています。

@registry.BACKBONES.register("ResNet50-C4")
@registry.BACKBONES.register("ResNet50-C5")
def build_resnet_backbone(設定):
    # ResNetの本体を構築
    body = resnet.ResNet(設定)
    # nn.Sequentialにラップ
    model = nn.Sequential(OrderedDict([("body", body)]))
    # 出力チャネル数を設定
    model.out_channels = 設定.MODEL.RESNETS.BACKBONE_OUT_CHANNELS
    return model

このように、設定で指定された名前(例: "ResNet50-C4")に基づいて、対応する build_resnet_backbone 関数が呼び出されます。

ResNet バックボーンの詳細

build_resnet_backbone 関数は、resnet.py にある ResNet クラスを呼び出して、ResNetのボディ部分を構築します。

ResNet クラスのコンストラクタは、まず BaseStem という基本のステージを初期化します。これは、入力画像に対する最初の畳み込み処理を担当します。

class BaseStem(nn.Module):
    def __init__(self, 設定, 正規化関数):
        super(BaseStem, self).__init__()
        # 出力チャネル数を設定から取得
        出力チャネル数 = 設定.MODEL.RESNETS.STEM_OUT_CHANNELS
        # 7x7の畳み込み層
        self.conv1 = Conv2d(
            3, 出力チャネル数, kernel_size=7, stride=2, padding=3, bias=False
        )
        # バッチ正規化層
        self.bn1 = 正規化関数(出力チャネル数)
        # 重みの初期化(Heの初期化)
        for l in [self.conv1,]:
            nn.init.kaiming_uniform_(l.weight, a=1)

    def forward(self, x):
        x = self.conv1(x)
        x = self.bn1(x)
        x = F.relu_(x)
        # 最大プーリング
        x = F.max_pool2d(x, kernel_size=3, stride=2, padding=1)
        return x

この BaseStem クラス内で使用されている Conv2d は、通常の畳み込み層ですが、ここでは FrozenBatchNorm2d という特殊なバッチ正規化層が使用される場合があります。この層は、トレーニング中でも統計量を固定し、更新しない特徴があります。

class FrozenBatchNorm2d(nn.Module):
    def __init__(self, n):
        super(FrozenBatchNorm2d, self).__init__()
        self.register_buffer("weight", torch.ones(n))
        self.register_buffer("bias", torch.zeros(n))
        self.register_buffer("running_mean", torch.zeros(n))
        self.register_buffer("running_var", torch.ones(n))

    def forward(self, x):
        # 固定された統計量を使ってスケールとバイアスを計算
        scale = self.weight * self.running_var.rsqrt()
        bias = self.bias - self.running_mean * scale
        # 入力テンソルに適用
        return x * scale.view(1, -1, 1, 1) + bias.view(1, -1, 1, 1)

最後に、ResNet クラスは、Bottleneck という基本ブロックを繰り返し使用して、深いネットワークを構築します。この Bottleneck ブロックは、入力チャネル数、ボトルネックチャネル数、出力チャネル数などのパラメータを持ち、ダウンサンプリングなどの処理を行います。

class Bottleneck(nn.Module):
    def __init__(
        self,
        入力チャネル数,
        ボトルネックチャネル数,
        出力チャネル数,
        グループ数,
        stride_in_1x1,
        ストライド,
        ダイレーション,
        正規化関数,
        dcn_config
    ):
        super(Bottleneck, self).__init__()
        # ダウンサンプリング層の初期化
        self.downsample = None
        if 入力チャネル数 != 出力チャネル数:
            down_stride = ストライド if ダイレーション == 1 else 1
            self.downsample = nn.Sequential(
                Conv2d(
                    入力チャネル数, 出力チャネル数,
                    kernel_size=1, stride=down_stride, bias=False
                ),
                正規化関数(出力チャネル数)
            )
        # ... (その他の畳み込み層の定義)

タグ: FCOS PyTorch 物体検出 モデル構築

7月23日 01:49 投稿