HuggingFace TransformersのBertEmbeddingsクラスの内部構造解説

BertEmbeddingsクラスはBERTモデルの中心的なコンポーネントです。このクラスは単語埋め込み、位置埋め込み、セグメント埋め込みを統合し、各入力トークンに対して最終的な埋め込み表現を生成します。これにより、後続のエンコーダ層への入力が提供されます。

BertEmbeddingsクラスの処理フロー

ソースコード位置:transformers/src/transformers/models/bert/modeling_bert.py

forward()メソッドの主要な処理フロー:

1)input_shapeパラメータ:input_idsまたはinputs_embedsから取得可能

input_shape: [batch_size, seq_length]

2)position_idsパラメータ:パラメータに含まれていればそのまま使用、なければバッファから取得

例:text = "私は学習が好きです", position_ids = tensor([[0, 1, 2, 3, 4, 5]])

3)token_type_idsパラメータ:パラメータに含まれていればそのまま使用、なければバッファから取得

例:text = "私は学習が好きです", token_type_ids = tensor([[0, 0, 0, 0, 0, 0]])

4)inputs_embedsパラメータ:通常input_idsから取得

self.word_embeddings = torch.nn.Embedding(vocab_size, hidden_size, padding_idx=pad_token_id)
inputs_embeds = self.word_embeddings(input_ids)

5)token_type_embeddingsパラメータ:

self.token_type_embeddings = torch.nn.Embedding(type_vocab_size, hidden_size)
token_type_embeddings = self.token_type_embeddings(token_type_ids)

6)position_embeddingsパラメータ:position_embedding_type == "absolute"に依存

- True: position_embeddings = self.position_embeddings(position_ids) # 絶対位置エンコーディングを使用
- False: # 絶対エンコーディングを使用しない

7)最終的なembeddings結果:

# 絶対位置エンコーディングを使用:embeddings = inputs_embeds + token_type_embeddings + position_embeddings
# 絶対位置エンコーディングを使用しない:embeddings = inputs_embeds + token_type_embeddings
embeddings = inputs_embeds + token_type_embeddings (+ position_embeddings)
embeddings = self.dropout(self.LayerNorm(embeddings))

BertEmbeddingsクラスのソースコード解説

# -*- coding: utf-8 -*-
# @time: 2024/7/12 11:35

import torch

from torch import nn
from typing import Optional


class TokenEmbeddingLayer(nn.Module):
    """単語、位置、およびトークンタイプの埋め込みを構築します。"""

    def __init__(self, config):
        super().__init__()
        # 単語埋め込み層:語彙表の単語を固定サイズのベクトル空間にマッピング
        self.word_embeddings = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id)

        # 位置埋め込み層:入力シーケンス内の各位置の埋め込みを表現
        self.position_embeddings = nn.Embedding(config.max_position_embeddings, config.hidden_size)

        # トークンタイプ埋め込み層:異なるタイプのトークンを区別(通常は文対処理で使用)
        self.segment_embeddings = nn.Embedding(config.type_vocab_size, config.hidden_size)

        # 層正規化:隠れ層の出力を正規化して訓練の安定性を向上
        self.normalization = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)

        # ドロップアウト:過学習を防ぐために一部のニューロンの出力をランダムに0に設定
        self.dropout_layer = nn.Dropout(config.hidden_dropout_prob)

        # 位置埋め込みタイプ:デフォルトは"absolute"
        self.embedding_type = getattr(config, "position_embedding_type", "absolute")

        # バッファの登録:位置IDを0からmax_position_embeddingsまで保存
        self.register_buffer("position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False)

        # バッファの登録:トークンタイプIDを初期化(すべて0)
        self.register_buffer("segment_ids", torch.zeros(self.position_ids.size(), dtype=torch.long), persistent=False)

    def forward(
        self,
        input_ids: Optional[torch.LongTensor] = None,
        segment_ids: Optional[torch.LongTensor] = None,
        position_ids: Optional[torch.LongTensor] = None,
        token_embeddings: Optional[torch.FloatTensor] = None,
        past_length: int = 0,
    ) -> torch.Tensor:
        if input_ids is not None:
            input_shape = input_ids.size()
        else:
            input_shape = token_embeddings.size()[:-1]

        # input_shape: [batch_size, seq_length]
        seq_length = input_shape[1]

        # position_idsがNoneの場合、バッファから取得
        if position_ids is None:
            position_ids = self.position_ids[:, past_length: seq_length + past_length]

        # トークンタイプIDの設定
        if segment_ids is None:
            if hasattr(self, "segment_ids"):
                buffered_segment_ids = self.segment_ids[:, :seq_length]
                buffered_segment_ids_expanded = buffered_segment_ids.expand(input_shape[0], seq_length)
                segment_ids = buffered_segment_ids_expanded
            else:
                segment_ids = torch.zeros(input_shape, dtype=torch.long, device=self.position_ids.device)

        # 入力埋め込みの取得
        if token_embeddings is None:
            token_embeddings = self.word_embeddings(input_ids)

        # セグメント埋め込みの取得
        segment_embeddings = self.segment_embeddings(segment_ids)

        # 埋め込みの初期合成
        combined_embeddings = token_embeddings + segment_embeddings

        # 絶対位置エンコーディングの追加
        if self.embedding_type == "absolute":
            position_embeddings = self.position_embeddings(position_ids)
            combined_embeddings += position_embeddings

        # 層正規化とドロップアウトの適用
        normalized_embeddings = self.normalization(combined_embeddings)
        final_embeddings = self.dropout_layer(normalized_embeddings)

        return final_embeddings

BertEmbeddingsクラスのテスト

# -*- coding: utf-8 -*-
# @time: 2024/7/12 16:21

import torch

from transformers import BertTokenizer, BertConfig
from TokenEmbeddingLayer import TokenEmbeddingLayer

# BERTトークナイザーのロード
model_name = 'google-bert/bert-base-uncased'
tokenizer = BertTokenizer.from_pretrained(model_name)

# 二つの文を含むサンプルテキスト
text = [("my dog is cute", "he likes playing")]
inputs = tokenizer(text, truncation=True, padding=True, return_tensors='pt')

input_ids = inputs["input_ids"]
print(input_ids)

segment_ids = inputs["token_type_ids"]
print(segment_ids)

config = BertConfig()
embedding_layer = TokenEmbeddingLayer(config)

with torch.no_grad():
    output_embeddings = embedding_layer(input_ids=input_ids, segment_ids=segment_ids)

print(output_embeddings)
print(output_embeddings.shape)

出力

tensor([[  101,  2026,  3899,  2003, 10140,   102,  2002,  7777,  2377, 13749,
           102]])
tensor([[0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1]])
tensor([[[ 1.1078,  1.1187,  0.8171,  ..., -1.7894, -0.3242, -0.3227],
         [-1.1724,  0.6758, -0.7380,  ..., -2.0156,  0.6501,  0.0658],
         [ 2.4001,  1.4040,  0.3180,  ...,  0.2173,  0.3123, -0.3133],
         ...,
         [-0.2485,  0.0995,  1.1544,  ...,  0.6161,  0.6230, -0.8850],
         [ 0.2799, -1.1622,  0.0000,  ...,  1.9079,  0.0000, -0.8867],
         [-0.1645, -0.0000, -0.3955,  ...,  2.5497, -1.2822, -2.1249]]])
torch.Size([1, 11, 768])

タグ: HuggingFace transformers BERT PyTorch 自然言語処理

8月3日 20:48 投稿