はじめに
この記事では、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_backbone、build_rpn、build_roi_heads は、それぞれ fcos_core.modeling.backbone、fcos_core.modeling.rpn.rpn、fcos_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
),
正規化関数(出力チャネル数)
)
# ... (その他の畳み込み層の定義)