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])