1
1

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

🔰PyTorchでニューラルネットワーク基礎 #36 【可変長データでのBERTタイプの事前学習】

1
Last updated at Posted at 2026-07-24

概要

個人的な備忘録を兼ねたPyTorchの基本的な解説とまとめです。BERTタイプ (Transformer Encoderタイプ) のMLM事前学習に向けて学んだことをまとめてみました。今回は、系列長の異なるID列データをそのまま扱い、BERTタイプのMasked Language Modelingによる事前学習を行う演習を行ってみたいと思います。

扱う内容

  1. BERTタイプの事前学習
  2. Datasetクラスのカスタマイズ
  3. DataLoaderクラスのカスタマイズ
  4. step数による学習の管理

方針

  1. できるだけ同じコード進行
  2. できるだけ簡潔(細かい内容は割愛)

演習用のファイル

1. データ

wikipediaのギリシア神話から抽出したテキストデータを利用しました。15万文字のとても小さなテキストファイルです。自作BERTモデルでもギリシア神話の神々の名前を<mask>して、予測することができるのか🔥

sample02.png
図:AIに描画してもらったのですが、進歩を感じますね〜。イラストはイメージであり、本文とは一切関係ありません:smile:

1.1 前処理

前処理と言っても大したことはやっていません。

  1. 「##」マークの削除
  2. 文頭にある不要な空白の削除
  3. 1文1行に変更:系列長を異なる形にしたいのであえて1文1データとなっています。

1.2 トークナイザーの作成

収集したデータを使ってトークナイザーも作成します。

  • 学習に利用するコーパスサイズは小さいですが、「アテーナー」「ポセイドーン」といった名前の頻度が多く、1トークン化して「<mask>埋め状況を確認したい」ので語彙数を3000に固定する形でunigram lmを利用して作成
  • <mask>部分の穴埋め効果を見るために、ByteLevelは使わない
  • トークナイザーの作成については第25〜27回のBPEWordPieceUnigram LMを参考にしてください。

1.3 学習データの作成

集めた文章に対して次の作業を行いデータセットを作成します。

  1. 1文1行 (短文、長文が入り交じるようにわざと1文1行になっています)
  2. トークナイザーを利用してIDに変換
  3. 文頭に<bos>のID、文末に<eos>のIDを追加

データのサンプル
text部分の文頭に<bos>のID、文末に<eos>のIDを挿入してidsを作成していきます1

No. text ids
0 アテーナー アテーナーは、知恵や戦争及び様々な技芸を司るギリシア神話の女神で、オリュンポス... [1, 131, 79, 382, 8, 144, 824, 22, 357, 2190, ...]
1 アテーネー、アターナー、アテーナイエーなどとも呼ばれる。 [1, 1398, 8, 52, 216, 23, 277, 23, 8, 468, 116...]
2 日本語では長母音を省略してアテナ、アテネと表記される場合が多い。 [1, 638, 85, 125, 1184, 52, 124, 277, 8, 52, 1...]

2. 事前学習

PyTorchによるプログラムの流れを確認します。基本的に下記の5つの流れとなります。

  1. データの読み込みとtorchテンソルへの変換 (2.1)
  2. ネットワークモデルの定義と作成 (2.2)
  3. 誤差関数と誤差最小化の手法の選択 (2.3)
  4. 変数更新のループ (2.4)
  5. モデルの保存 (2.5)

2.1 データの読み込みとtorchテンソルへの変換

毎回おなじみのライブラリーを色々読み込む部分となります。Datasetクラスをカスタマイズ、DataLoaderのcollate関数を設定するタイプとなります。

ライブラリの読み込みなど
import pandas as pd
import torch
import torch.nn as nn
import random
from tokenizers import Tokenizer

# カスタマイズする部分
from torch.utils.data import Dataset, DataLoader
from functools import partial              # collate_fnのときに使う

# 利用可能なデバイス
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

続いて、トークナイザーや学習データの読み込みです。pandasを利用して、JSONL形式を読み込みます。

pretrain_filename名前の付け方が曖昧だった!今回は事前学習が目的だから、pretrain_filenameが保存されるファイル名となります。

トークナイザーとデータファイルの読み込み
# いくつかのファイル名
data_filename = "./data/greek_data_unigram_3k.jsonl"              # 学習データ
tokenizer_filename = "tokenizer/greek_unigram_tokenizer_3k.json"  # トークナイザー
pretrain_filename = "./model/greek_unigram_3k.model"              # 保存する事前学習モデル名

# トークナイザー
tokenizer = Tokenizer.from_file(tokenizer_filename)

print("特殊トークンID:")
print(f"<pad>: {tokenizer.token_to_id('<pad>')}")
print(f"<bos>: {tokenizer.token_to_id('<bos>')}")
print(f"<eos>: {tokenizer.token_to_id('<eos>')}")
print(f"<unk>: {tokenizer.token_to_id('<unk>')}")
print(f"<mask>: {tokenizer.token_to_id('<mask>')}")
print(f"size: {tokenizer.get_vocab_size()}")

# 特殊トークンID:
# <pad>: 0
# <bos>: 1
# <eos>: 2
# <unk>: 3
# <mask>: 4
# size: 3000


#  データファイルの読み込みJSONファイルを読み込む
data = pd.read_json(data_filename, lines=True)

print(f"データの構造:{data.keys()}")
print(f"データ数: {len(data['ids'])}")
print(f"サンプル長さ: {[len(d) for d in data['ids'][:5]]}")  # 可変長を確認
# データの構造:Index(['text', 'ids'], dtype='object')
# データ数: 2979
# サンプル長さ: [23, 17, 18, 10, 47]

説明メモ

  • Tokenizerを使い、保存したトークナイザーを読み込みます。
  • サンプルの長さでids列の長さが異なることも確認
  • たまたまなのですが、idsの最大の長さmax([len(d) for d in data['ids']])が64となります。利用する最大系列長が64なので、文を切り詰めることなくすべての文字が利用されます。

2.2 ネットワークモデルの定義と作成

モデルは第31回構築したBERTタイプ (Transformer Encoder型) のネットワークモデルを利用します。トークナイザーを準備して、モデル設定クラスを利用する形です。

mlm_network.png
図1:モデルの構造

図1に準拠したモデルをコードにしていきます。詳しくは第30回を参照してください。

python モデルの定義
# (1) モデルの設定のクラス
class ModelConfig:
    def __init__(self, tokenizer):
        # モデル構造
        self.vocab_size = tokenizer.get_vocab_size()
        self.seq_len = 64
        self.d_model = 64
        self.nhead = 4
        self.dim_feedforward = 256
        self.num_layers = 6
        self.dropout = 0.1
        
        # 特殊トークンID
        self.pad_token_id = tokenizer.token_to_id("<pad>")
        self.mask_token_id = tokenizer.token_to_id("<mask>")
        self.bos_token_id = tokenizer.token_to_id("<bos>")
        self.eos_token_id = tokenizer.token_to_id("<eos>")
        self.unk_token_id = tokenizer.token_to_id("<unk>")

        # 特殊トークンのセット
        self.special_tokens_set = {
            self.pad_token_id,
            self.mask_token_id,
            self.bos_token_id,
            self.eos_token_id,
            self.unk_token_id,
        }
        
        # 通常トークンのリスト special_tokenを除くトークンのリスト(MLMランダム置換用)
        self.normal_tokens_list = [
            i for i in range(self.vocab_size) 
            if i not in self.special_tokens_set
        ]
        
        # 学習設定 (使うと便利かも、今回は一部だけ使ってみました)
        self.batch_size = 32        # 64 余裕があるときは大きめだと学習が早い
        self.learning_rate = 0.001  # config.learning_rateとして使います
        self.num_epochs = 100
        self.mask_prob = 0.15
        self.max_grad_norm = 1.0
        # PyTorchの仕様 ID= -100 は損失計算時に除外されるマスクID
        self.ignore_index = -100  # config.ignore_indexとして使います

# (2) MLM用の出力部分
class MLMHead(nn.Module):
    def __init__(self, config: ModelConfig, embedding_weight):
        super().__init__()
        self.fc = nn.Linear(config.d_model, config.d_model)
        self.act = nn.GELU()
        self.ln = nn.LayerNorm(config.d_model)

        self.classification_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
        self.classification_head.weight = embedding_weight  # weight tying

        self.bias = nn.Parameter(torch.zeros(config.vocab_size))

    def forward(self, x):
        x = self.fc(x)
        x = self.act(x)
        x = self.ln(x)
        x = self.classification_head(x) + self.bias
        return x

# (3) MLM用のモデル
class DNN(nn.Module):
    def __init__(self, config: ModelConfig):
        super().__init__()
        self.config = config
        
        # 埋め込み層
        self.token_embedding = nn.Embedding(
            num_embeddings=config.vocab_size, 
            embedding_dim=config.d_model,
            padding_idx=config.pad_token_id
        )
        self.pos_embedding = nn.Embedding(num_embeddings=config.seq_len, embedding_dim=config.d_model)
        
        self.layer_norm = nn.LayerNorm(config.d_model)
        self.dropout = nn.Dropout(config.dropout)
        
        # Transformer Encoder
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=config.d_model,
            nhead=config.nhead,
            dim_feedforward=config.dim_feedforward,
            dropout=config.dropout,
            batch_first=True,
        )
        self.transformer_encoder = nn.TransformerEncoder(encoder_layer, num_layers=config.num_layers, enable_nested_tensor=False)
        
        # MLM用の出力層:BERT風レイヤー 各トークン位置で語彙全体を予測
        self.mlm_head = MLMHead(config, self.token_embedding.weight)
    
    def forward(self, x):
        # マスクの作成
        src_key_padding_mask = (x == self.config.pad_token_id)
        
        # 埋め込み
        tok_emb = self.token_embedding(x)
        pos_emb = self.pos_embedding(torch.arange(x.size(1), device=x.device))
        x = tok_emb + pos_emb.unsqueeze(0)
        
        x = self.layer_norm(x)
        x = self.dropout(x)
        
        # Transformer Encoder
        h = self.transformer_encoder(x, src_key_padding_mask=src_key_padding_mask)
        
        # MLM予測:各位置で語彙全体での確率を出力
        logits = self.mlm_head(h)  # [batch, seq_len, vocab_size]
        
        return logits

config = ModelConfig(tokenizer)
model = DNN(config).to(device)

説明メモ

  • (1) ModelConfigクラスでモデル構造やトークナイズの属性をまとめています。
  • (2) MLMHeadクラスはマスク部分を予測する最終層のブロック。Transformer Encoder層からの特徴量を、単語予測の分類問題として扱います。classification_headの重みに、トークン埋め込み層の重みをそのまま利用する「重み共有」も活用します。
  • (3) DNNクラスに最終的な流れを記載します。マスク化されたIDベクトルを入力、トークンと位置の埋め込み層、Transformer Encoder層、mlm_head層と経由して、最終出力となります。出力は各トークごとのlogitsとなります。出力される値の形状は(系列長,語彙数)= (seq, 3000)の形になります。

DtasetクラスとDataloaderのcollate関数をカスタマイズしていきます。2つのクラスの関係は、次のようなものです。

  • Datasetクラスでデータを1件ずつ取り出す規則を設定
  • Dataloaderクラスで複数のデータをまとめて学習用に変更

Datasetクラスです。第33回で確認した方法を使います。データフレームを取得してids列をリストに変換するという単純なクラスになります。

カスタムDatasetクラス
class MLMDataset(Dataset):
    def __init__(self, data):
        # ids列をリストへ
        self.ids = data["ids"].tolist()
   
    def __len__(self):
        return len(self.ids)
    
    def __getitem__(self, idx):
        return self.ids[idx]

説明メモ

  • MLMDasetクラスに入力されるdataが
    data = pd.read_json(data_filename, lines=True)
    つまり、データフレーム形式になっていた点に注意して、クラスを定義します。
  • __getitem__()の戻り値は、素朴にidsリストだけです。
  • idsリストを受け取ってミニバッチ毎に等長化する役割がDalaloaderのcollate関数となります。
  • マスク化の処理もDataloaderのcollate関数に担わせます。

MLMDatasetクラスの確認をしてみました。

dataset = MLMDataset(data)
print(dataset[0])
# [1, 131, 79, 382, 8, 144, 824, 22, 357, 2190, 172, 446, 858, 2289, 7, 2216, 15, 147, 103, 56, 1108, 6, 2]

説明メモ

  • [dataset[0],...,dataset[n]]というdataset[i]のリストがcollate_fn (collate_wrapper) の引数となります。

DataLoaderのcollate関数を定義していきます。BERTタイプのポイントであるマスク化の部分です。マスク化の処理について、詳しくは第29回を参照してください。

今回の主目的は「等長化(パディング)」です。collate関数に、idsリストの長さをミニバッチ毎に揃える処理を追加します

マスク化関数 mlm_sampleとcollate_fnの定義
# (1) マスクの関数
def create_mlm_sample(
    original_ids: list,
    config: ModelConfig,
    mask_prob: float = 0.15,
    ):
    input_ids = original_ids.copy()
    labels = [config.ignore_index] * len(original_ids)   # config.ignore_index= -100

    special_tokens_set = config.special_tokens_set # {0,1,2,3,4}  # <pad>, <bos>, <eos>, <unk>, <mask>になる予定
    normal_tokens_list = config.normal_tokens_list  # range(5,_000) # 語彙IDのspecial_tokensを除いたもの


    # 【ポイント】句読点も除外 (マスクに句読点を入れる結果になりやすかったので今回あえて削除しました)
    no_mask_tokens = ["", "", ",", ".", "", ""]
    no_mask_ids = set(
        token_id for token_id in [tokenizer.token_to_id(tok) for tok in no_mask_tokens]
        if token_id is not None
    )
    special_tokens_set.update(no_mask_ids)


    for i in range(len(original_ids)):
        # 特殊トークンの場合は飛ばす original_ids[i]はint型なのでそのまま判定できる
        token_id = original_ids[i]
        if token_id in special_tokens_set:
            continue
        
        # 15%未満で他のトークンへ置き換え
        # 80%は<mask> mask_token_id
        # 10%は特殊トークを除く別トークン random.choice()
        # 10%はそのまま pass
        if random.random() < mask_prob:
            labels[i] = original_ids[i]
            
            rand = random.random()
            if rand < 0.8:
                input_ids[i] = config.mask_token_id
            elif rand < 0.9:
                input_ids[i] = random.choice(normal_tokens_list) 
            else:
                pass
    
    return input_ids, labels

# (2) マスクの関数を利用してさらにcollate_fn用の関数を定義
def collate_fn_mlm(batch_data, config):
    """
    バッチ内で動的にパディングを行う
    引数
        batch_data: [{'ids': [...]}, {'ids': [...]}, ...]
        config: ModelConfig (モデルの設定)
    戻り値
        辞書形式
        input_idsが入力されるデータ、<mask>化されたid列
        labelsが教師データ、<mask>化される前のid列
        {input_ids: [batch_size, max_len_in_batch],
         labels: [batch_size, max_len_in_batch]}
    """

    mask_token_id = config.mask_token_id
    pad_token_id = config.pad_token_id
    
    # MLMマスクを適用
    input_ids_list = []
    labels_list = []
    # (2-1) マスク化したidsと対応するラベルを作成
    for item in batch_data:
        original_ids = item
        input_ids, labels = create_mlm_sample(original_ids, config=config, mask_prob=config.mask_prob)
        input_ids_list.append(input_ids)
        labels_list.append(labels)
    
    # (2-2) バッチ内の最大長を取得, config.seq_len=64がmax
    max_len = min(max(len(ids) for ids in input_ids_list), config.seq_len)

    # (2-3) パディング
    padded_input_ids = []
    padded_labels = []
    
    for input_ids, labels in zip(input_ids_list, labels_list):
        input_ids = input_ids[:max_len]   # max_lenで切り詰め
        labels    = labels[:max_len]      # max_lenで切り詰め
        padding_length = max_len - len(input_ids)
        
        # (2-4) input_idsをパディング
        padded_input = input_ids + [pad_token_id] * padding_length
        padded_input_ids.append(padded_input)
        
        # (2-5) labelsをパディング(ignore_index=-100でパディング)
        padded_label = labels + [config.ignore_index] * padding_length
        padded_labels.append(padded_label)
    
    # (2-6) Tensorに変換
    input_ids_tensor = torch.LongTensor(padded_input_ids)
    labels_tensor = torch.LongTensor(padded_labels)
    
    return {
        "input_ids": input_ids_tensor,
        "label_ids": labels_tensor
    }

# (3) DataLoaderで利用するためのラップ処理
# collate_fn_mlmのままだと使えない
collate_wrapper = partial(collate_fn_mlm, config=config)

説明メモ

  • (1) id列の15%をマスクの対象とします。対象部分について「80/10/10」の割合で、「マスク化/別のトークン/そのまま」に置き換えます。

  • (2) 今回の最大のポイント部分

  • (2-1) create_mlm_sampleで<mask>文を作成

  • (2-2) 内包表記で入り組んでいますが、input_ids_list、つまり、バッチ毎のデータに対しての処理となります。まず、max()部分でバッチ内の系列長の最大値を求めます。min()部分でモデルで決めた最大系列長 config.seq_len=64とバッチ内最大長を比較します。

  • (2-3) input_idsとlabelsに対してリストのスライス機能 list[:max_len] を利用して、系列長がmax_lenを超えないように調整します。
    「max_len − input_idsの長さ」がpadで埋める量padding_lengthになります。

  • (2-4) input_idsにpadding_length分の<pad>のIDを追加します。

  • (2-5) labelsに対しても、<pad>のIDを追加します。

  • この時点でバッチ内の系列長が同じ長さに揃うことになります。

  • (2-6) pad付きのinput_idsとlabelsをテンソルに変換、戻り値とします。

  • (3) 地味に重要:smile:
    partialを利用して、collate_fnとして利用できる引数の形へ変形します。functools.partialを利用してcollate_fn_mlm(batch_data, config)の引数configを固定して、batch_dataだけにします。

どうしても、マスク化の処理のほうが、パディングよりも目立ってしまう:sweat:

DataLoaderを使って、マスク化、等長化したミニバッチを作成します。

dataloader = DataLoader(
    dataset,
    batch_size=config.batch_size,
    shuffle=True,
    drop_last=True,
    collate_fn=collate_wrapper, # 動的パディング
    num_workers=2
)

dataloaderの出力を確認してみました。input_idsが「0」でパティングされていることが確認できます。labels部分も「-100」以外のトークンIDが省略された部分に表示されているはずです。詳しい挙動は第29回のマスク化を参照してください。

for data in dataloader:
    print(data)
    break

#  出力結果
# {'input_ids': tensor([[   1,  346,  552,  ...,    0,    0,    0],
#         [   1,   88,   33,  ...,    0,    0,    0],
#         [   1, 1357,   13,  ...,    0,    0,    0],
#         ...,
#         [   1,   33, 1173,  ...,    0,    0,    0],
#         [   1,   30,    4,  ...,    0,    0,    0],
#         [   1,  468,    7,  ...,    0,    0,    0]]),
#  'label_ids': tensor([[-100, -100, -100,  ..., -100, -100, -100],
#         [-100, -100, -100,  ..., -100, -100, -100],
#         [-100, -100, -100,  ..., -100, -100, -100],
#         ...,
#         [-100, -100, -100,  ..., -100, -100, -100],
#         [-100, -100,  168,  ..., -100, -100, -100],
#         [-100, -100, -100,  ..., -100, -100, -100]])}

やや長かったですが、これで、可変長を扱うタイプのモデルが完成しました🐾

2.3 誤差関数と誤差最小化の手法の選択

誤差関数と最適化の手法
criterion = torch.nn.CrossEntropyLoss(ignore_index=config.ignore_index)
optimizer = torch.optim.AdamW(model.parameters(), lr=config.learning_rate)

説明メモ

  • ignore_index = config.ignore_index = -100で損失計算時に除外するIDを指定できます。「ー100」がデフォルト値となります。
  • 学習率についても徐々に減少するタイプの方が効率的だと思います。今回そのままですが:sweat:

2.4 変数更新のループ

今回は、変数を更新する数(ステップ数)で繰り返しの回数をコントロールしてみたいと思います2

dataloaderをそのまま使うと、見た目はfor文の2重ループの形になります。

epoch数でループ
for epoch in range(LOOP):   
    for batch in dataloader:
        y = model(x)

epoch数ではなく、変数更新数でforループを使う形に修正してみます。雰囲気としては、次のような形になる予定です。

step数でループ
for step in range(max_iters):
    y = model(x)

ステップで管理する準備1
dataloader からバッチを取り出し続け、1エポック終了したら最初に戻って、またバッチを供給し続けるジェネレータを作成します。

データ供給関数
# step数で管理するためのデータローダー関数
def infinite_loader(dataloader):
    while True:
        for batch in dataloader:
            yield batch

ステップで管理する準備2 (不要 :sweat:)
epoch数とstep数の関係を確認しておきます。

step数 (更新回数) = epoch数 (LOOPの回数)× ミニバッチ分割数

準備を踏まえた変数更新ループとなります3。DataLoaderのオプションに「drop_last=True」をつけているので、ステップ数とエポック数とできれいに整合性が取れているはずです。

変数更新ループ
# step数で管理するためのデータローダー関数
def infinite_loader(dataloader):
    while True:
        for batch in dataloader:
            yield batch

# (1) 精度を求める関数 -100の部分をマスクして精度計算から除外
def accuracy(y, t, ignore_index=config.ignore_index):
    """
    マスク位置での予測精度を計算
    ※この関数内では勾配計算を行わない
    """
    with torch.no_grad():
        preds = torch.argmax(y, dim=-1)
        mask = (t != ignore_index)
        correct = (preds == t) & mask

        num_correct = correct.sum().item()
        num_total = mask.sum().item()
        acc = (num_correct / num_total) if num_total > 0 else 0.0

        return acc

# --- 変数更新部分 〜ここまで、長かった^^;〜
# (2) step数による繰り返し部分
data_iter = infinite_loader(dataloader)  # epochではなく、step数で計測
max_iters = 70_000                       # step数 (更新回数) = epoch数 (LOOPの回数)× ミニバッチ分割数、1500エポックくらい
model.train()                            # trainモードを明示

for step in range(max_iters):
    # (2-1)
    batch = next(data_iter)             # infinite_loaderによるデータ取り出し
    x = batch["input_ids"].to(device)
    t = batch["label_ids"].to(device)

    optimizer.zero_grad()
    y = model(x)
    # (2-2)
    loss = criterion(y.view(-1, config.vocab_size), t.view(-1))
    acc = accuracy(y, t)    
    loss.backward()
    # (2-3)
    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=config.max_grad_norm)  # 勾配クリップ
    optimizer.step()

    if (step+1)%1000 == 0:
        print(f"{step+1}-step:\tloss:{loss.item():.3f}\tacc:{acc:.3f}")

説明メモ

  • (1) 精度の関数:IDがconfig.ignore_indexつまりID番号「−100」になっている部分を除外します。if num_total > 0としています。系列長が64と短めなので、<mask>が1箇所も登場しないこともあり得ます。予防的につけておきました。
  • (2) forループ1重に見えますが、実際は
        data_iter = infinite_loader(dataloader)
    でループ処理を代用しています。
  • (2-1) batch = next(data_iter)でdata_iterからミニバッチサイズのデータが供給されます。
  • collate_fn_mlmの戻り値が辞書形式だったので、input_idsキーとlabel_idsキーを利用して入力データと教師データに分割します。
  • (2-2) 予測値 y の形状が (bs=32, seq_len, vocab_size=3000)、教師ラベル t の形状が(bs=32, seq_len)となります。CrossEntropyLossを利用してseq_len トークン × 32 個のトークンについて、logits と教師ラベル t と比較すればよいので、y と t の形状を下記のように変形することになります。

    loss = criterion(y.reshape(-1, config.vocab_size), t.reshape(-1))

  • (2-3) パラメータ$\theta_t$を更新する部分。基本的には、「$\theta_{t+1}=\theta_{t}-\eta \nabla L(\theta_t)$」という雰囲気の式でした。式からわかるように、パラメータを更新する時、微分した値が大きいとパラメータの更新が大きくなり挙動が不安定になることがあります。loss.backward()で求めた勾配について、「勾配が大きい場合は小さくなるように計算せよ!」という設定がclip_grad_normになります。

マスク部分のトークンを正しく予測できているか確認したいので、精度が0.4〜0.5くらいまで学習を頑張ります。1時間ほど待ちましょう。バッチサイズを小さくしてあるのでGPUメモリー(VRAM)は4GBもあれば十分動作します。

2.5 モデルの保存方法

pretrain_filenameに指定したファイル名にモデルを保存します。

モデルの保存

torch.save({
        "model_state_dict": model.state_dict(),
        "config": config.__dict__,  # configも一緒に保存
        }, pretrain_filename)

モデルの復元

モデルの読み込み(復元方法)
checkpoint = torch.load(pretrain_filename)
config = ModelConfig(tokenizer)                  # tokenizerも事前に読み込んでおきます
config.__dict__.update(checkpoint["config"])     # checkpoint["config"]辞書を使って、configをアップデート

model = DNN(config).to(device)
model.load_state_dict(checkpoint["model_state_dict"])

torch.load以外に、設定も保存してあるのでその部分をupdateする必要があるので注意が必要です。

説明メモ

  • checkpoint["config"]:保存したconfig
  • checkpoint["config"]は辞書形式なので、updateを使い中身を保存したものへ置き換えます。
  • checkpoipnt["model_state_dict"]: 保存したモデルの重み

地味です:sweat_smile: よく見かける from_dict(辞書名)みたいなメソッドを利用する

config = ModelConfig.from_dict(model_id["config"], tokenizer=tokenizer)

というスマート方法はModelConfig少し改良する必要があるのと、多分、ニューラルネットワークとは別の部分になるので今回はパス。

2.6 検証と精度について

最後に穴埋め学習によって<mask>部分が正しく予測できているか確認してみましょう。このためだけに、頑張りました。
トークナイザーによってギリシャ神話の神様の名前(ゼウスやアルテミスなど)が1トークン扱いになっています。

<mask>予測の関数は第31回の関数をそのままコピーして使いました:sweat:

import torch.nn.functional as F
# (1) テキストをtokenizerに従い、idsへ変換する関数
def encode_text_with_special_tokens(text, tokenizer, config):
    """
    テキストをtokenizerに従い、idsへ変換する関数
    text: 例 織田信<mask>は桶狭間の戦いで... : <mask>は1個に限定
    tokenizer: 学習時に使った tokenizer
    戻り値: (1, seq_len) = (1, 64) の torch.LongTensor input_ids
    """
    ids = tokenizer.encode(text).ids  # 文字列をID化
    ids = [config.bos_token_id] + ids + [config.eos_token_id]  # <bos>, <eos> を追加
    ids = ids[:config.seq_len]        # 長すぎる場合は切る
    if len(ids) < config.seq_len:     # 足りない場合は<pad>
        ids = ids + [config.pad_token_id] * (config.seq_len - len(ids))

    # LongTensorに変換して出力
    input_ids = torch.tensor([ids], dtype=torch.long, device=device)
    return input_ids


# (2) text中の <mask> 位置を見つけて,上位候補を返す
@torch.inference_mode()
def predict_mask_topk(model, tokenizer, config, text, topk=5):
    """
    text中の <mask> 位置を見つけて,上位候補を返す
    """
    model.eval()
    input_ids = encode_text_with_special_tokens(text, tokenizer, config)  # textをエンコード化<bos>や<eos>もついているぞ
    logits = model(input_ids).to(device)           # モデル出力: (1, seq_len, vocab_size) = (1, 64, 2000)
    mask_positions = (input_ids[0] == config.mask_token_id).nonzero(as_tuple=True)[0]  # <mask> の位置を探す
    pos = mask_positions.item()                    # mask_positionはtorch.Tensorなので数値に戻す
    mask_logits = logits[0, pos]                   # maski_position位置の語彙方向のスコア logitsを取得 (vocab_size)
    probs = F.softmax(mask_logits, dim=-1)         # 確率化 (自分が解釈するため)
    top_probs, top_ids = torch.topk(probs, k=topk) # 上位 topk 個

    # topkの(ID,トークン,確率)を戻り値に指定
    candidates = []
    for prob, token_id in zip(top_probs.tolist(), top_ids.tolist()):
        token_str = tokenizer.id_to_token(token_id)
        candidates.append((token_id, token_str, prob))
    return {"position": pos,"candidates": candidates}

<mask>が一箇所だけある文のみ扱います。2個以上はNGだよ😱(エラーになります)
基本的に、

  1. 予測させたい<mask>付きの文をIDに変換、
  2. <mask>位置のmodel出力について、大きい順に5つ(topk=5)

を求めればOKです。

説明メモ

  • (1) ベタな関数ですが順番にID化、<bos>などの追加、長さの調整を行います。
  • (2) (1)の関数を利用して、文章中の<mask>部分の予想語彙を出力させます。こちらもベタな関数です。順番に、文をID化、学習したモデルで予測、出力値 (logits) の<mask>部分の確率をsoftmaxで計算して、上位5つまでのトークンと確率を出力させます。

具体例で検証
学習した内容と無関係のものに対しては全く異なる予測をします。利用した文章から適切に学習が行われたのかを確認する意味で、類似の文章・同一文章を使い<mask>を予測させてみます。

# 例文1〜5
# 正解:黄金の
text = "神話ではトロイア戦争のきっかけは<mask>林檎を巡る"
# 正解:アルテミス
#text = "神話の中ではオレステースがイーピゲネイアと共にもたらした<mask>の神像は人身御供を要求する神であった"
# 正解:鷹
#text = "神々は変身してエジプトへ逃げた時、アポローンは<mask>に、アレースは魚に、ヘーパイストスは雄牛に変身した。"
# 正解:ハーデス
#text = "アルテミスは後に神となるほどの腕前の医師アスクレーピオスを訪ね、オーリーオーンの復活を依頼したが、冥府の王<mask>がそれに異を唱えた。"
# 正解:神
#text = "アルテミスは後に<mask>となるほどの腕前の医師アスクレーピオスを訪ね、オーリーオーンの復活を依頼した"#が、冥府の王<mask>がそれに異を唱えた。"


results = predict_mask_topk(model, tokenizer, config, text, topk=5)

import pandas as pd
print(f"mask position = {results['position']}")
df = pd.DataFrame(results["candidates"] , columns=["id", "token", "prob"])
df

結果
ギリシア神話以外では<mask>予測するのは不可能なので試していません。

例1. 正解「黄金の」

text = "神話ではトロイア戦争のきっかけは<mask>林檎を巡る"

予測結果:正解

index id token prob
0 756 黄金の 0.973537
1 637 巨大な 0.004002
2 1007 ここで 0.002844

この例文は、何度か学習を試しましたが、毎回正解でした。黄金の林檎は有名なのか?

例2. 正解「アルテミス」

text = "神話の中ではオレステースがイーピゲネイアと共にもたらした<mask>の神像は人身御供を要求する神であった"

予測結果:正解

index id token prob
0 202 アルテミス 0.325990
1 931 性格 0.064852
2 530 ほど 0.055759

「神様の名前を予測できるのか!」が今回の目標の一つでした。しっかり予測できているようです:smile:

例3. 正解「鷹」

text = "神々は変身してエジプトへ逃げた時、アポローンは<mask>に、アレースは魚に、ヘーパイストスは雄牛に変身した。"

予測結果:不正解

index id token prob
0 196 ヘーラクレース 0.099213
1 204 ヘルメース 0.071102
2 316 0.058443
3 1151 アルクメーネー 0.051095
4 1714 0.048138

この例文が曲者😂正解する結果もあれば、今回のように不正解もありました。確率見るともう少し頑張れば正解になりそうな予感ですね:sweat_smile:

例4. 正解「ハーデス」

text = "アルテミスは後に神となるほどの腕前の医師アスクレーピオスを訪ね、オーリーオーンの復活を依頼したが、冥府の王<mask>がそれに異を唱えた。"

予測結果:正解

index id token prob
0 283 ハーデース 0.480711
1 183 デーメーテール 0.048887
2 165 ヘーラー 0.035613

冥府の王といえば、「ハーデス」?こちらも正解になりやすかったです。

例5. 正解「神」

text = "アルテミスは後に<mask>となるほどの腕前の医師アスクレーピオスを訪ね、オーリーオーンの復活を依頼した"

予測結果:不正解

index id token prob
0 128 英雄 0.143521
1 1000 アスクレーピオス 0.069757
2 1155 ケイローン 0.053887

最後の例は、難問でした。「神」が正解なのですが、これはなかなか正解になりませんでした。テキトーに学習させている時、一度だけ正解していた気がします。

次回

HuggingFaceのライブラリーを利用してDataset・DataLoaderクラスのカスタマイズ部分をスマートにしてみたいと思います。

扱いたいテーマ

  1. Datasetクラスのカスタマイズ方法 (第33回
  2. DataLoaderのcollate_fnを使ってみる(第34回
  3. 可変長データでの文章分類 (第35回
  4. 可変長データでのBERTタイプの事前学習 (今回・第36回)
  5. HuggingFaceのライブラリーとつなげてみる(次回・第37回)

目次

  1. 事前にpaddingせず使う形にするのが今回の目的となります。id化もDatasetクラスに任せることができます。第33回【Datasetクラスのカスタマイズ】を参考。

  2. データが5万件あるとします。バッチサイズを100にすると1エポックで500回更新されます。バッチサイズを1000だと、50回しか更新されません:scream: 同じエポック数でも、前者は後者の10倍もパラメータ更新を行っていることになります。ということは、「バッチサイズ=1」最強😆そうなのですが、今度は小さすぎで学習が不安定になりそうですね。

  3. githubのサンプルコードでは、from tqdm import tqdmを使って進捗状況を示すバーを表示させちゃっています:wink:

1
1
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
1
1

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?