GenRec式

Netflixの推薦モデル GenRec は、LLMの Transformer を利用しながら、通常のLLMのように推薦結果を1トークンずつ生成するのではなく、入力から得られたユーザーの表現を使って、多数の推薦候補を一度にスコアリングするモデルである。

通常の生成型LLMでは、Transformerによって入力を処理した後、「次にどのトークンを出力するか」を予測する。そのため、1回の推論で得られるのは基本的に次の1トークンであり、複数のトークンを生成するには、この処理を繰り返す必要がある。

一方、推薦では、推薦対象となるアイテムがあらかじめ候補として存在している。したがって、「推薦文を1トークンずつ生成する」必要はない。ユーザーの履歴から「このユーザーは何を好みそうか」という表現を一度作り、それを各候補と比較して、どの候補が適しているかをまとめて計算すればよい。

ここで重要なのがTransformerである。GenRecもTransformerを使うため、内部では通常のLLMと同じようにQuery(Q)、Key(K)、Value(V)によるAttentionを計算する。\[Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V\]\[\operatorname{Attention}(Q,K,V)=\operatorname{softmax}\left(\frac{QK^{T}}{\sqrt{d_k}}\right)V\]

この処理によって、ユーザーの過去の行動など、入力された要素同士の関係を考慮した表現が得られる。つまり、Q・K・Vは GenRec でも必要であり、通常のLLMとの違いはTransformerそのものではなく、Transformerの出力を何に使うかにある。

通常の生成型LLMでは、

Transformer(Q,K,V) $\rightarrow$ 語彙ごとのlogits $\rightarrow$ 1トークン生成 $\rightarrow$ 再び推論

となる。

これに対してGenRecでは、

Transformer(Q,K,V) $\rightarrow$ ユーザー表現 $\rightarrow$ 推薦候補を一括スコアリング

となる。

例えば、Transformerから得られたユーザー表現を $\mathbf h\in\mathbb{R}^{D}$、推薦候補の埋め込みを $E\in\mathbb{R}^{D\times M}$ とすると、\[\mathbf{s}=E^{T}\mathbf h\]

によって$M$個の候補に対するスコアを一度に計算できる。

この方式が効率的なのは、推薦では「次のトークンを何個も生成する」という問題を解く必要がなく、「既存の候補の中からどれが適切か」を判断すればよいからである。候補数 $M$ に対するスコア計算は、ユーザー表現と候補表現の行列演算としてまとめて実行できる。そのため、推薦結果を長い系列として逐次生成する場合に発生する生成ステップを省略できる。

したがって、GenRecの基本的な考え方は、

Q・K・Vでユーザーの文脈を理解する $\rightarrow$ ユーザー表現を作る $\rightarrow$ 既存の推薦候補を一括評価する

というものである。

さらに、LLMの事前学習によって得られた言語的・意味的な表現能力を推薦に利用できるため、推薦専用モデルとして大量のデータを一から学習する必要性も小さくなる。GenRecでは、このようにLLMの「文脈を表現する能力」を利用しながら、LLM本来の「文章を逐次生成する処理」を推薦には必要な部分だけに置き換えることで、推論と学習の効率化を図っている。

サンプルコード

Netflixの推薦モデル GenRec を概念的に試してみる。

import time
import torch
import torch.nn as nn
import torch.nn.functional as F

# 1. 環境の設定(GPUが使えればGPUを使用、なければCPU)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")

# ==========================================
# 設定パラメータ(GenRecと通常のLLMの比較用)
# ==========================================
BATCH_SIZE = 1 # 1ユーザー
HIDDEN_DIM = 2048 # LLMの隠れ状態ベクトルの次元数 (D)
PREFILL_LEN = 512 # ユーザー履歴などの入力トークン数 (N)
DECODE_LEN = 20 # 通常のLLMが推薦作品名を生成するのにかかるトークン数 (L)
NUM_CANDIDATES = 5000 # カタログ全体の推薦候補作品数 (M)

# ダミーの重み(LLMの1レイヤー分をシミュレーションするための軽量な線形層)
# 1B~10Bモデルの全レイヤーを載せるとメモリを圧迫するため、主要なボトルネック特性を再現します
llm_layer = nn.Linear(HIDDEN_DIM, HIDDEN_DIM).to(device)

print(f"\n--- 実験設定 ---")
print(f"入力トークン数 (N): {PREFILL_LEN}")
print(f"候補作品数 (M): {NUM_CANDIDATES}")
print(f"デコード長 (L): {DECODE_LEN}\n")

# ==========================================
# パターンA: 通常のLLMによる自己回帰デコード (Auto-regressive Decode)
# ==========================================
print("A. 通常のLLM(自己回帰デコード)による推薦をシミュレート中...")
torch.cuda.synchronize() if torch.cuda.is_available() else None
start_time = time.time()

# ① Prefill相: 入力コンテキストを一括処理
prefill_input = torch.randn(BATCH_SIZE, PREFILL_LEN, HIDDEN_DIM).to(device)
with torch.no_grad():
hidden_states = llm_layer(prefill_input)

# ② Decode相: 1トークンずつ自己回帰的にループ処理 (L回)
# ※ デコード時は毎回モデルの重みにアクセスする(メモリ帯域ボトルネックをループで表現)
current_input = hidden_states[:, -1:, :] # 最後のトークン
for t in range(DECODE_LEN):
with torch.no_grad():
# 次のトークンを予測するためにLLMを1ステップ実行
current_input = llm_layer(current_input)

torch.cuda.synchronize() if torch.cuda.is_available() else None
time_autoregressive = time.time() - start_time
print(f"⇒ 通常LLM 処理時間: {time_autoregressive:.6f} 秒\n")


# ==========================================
# パターンB: GenRec方式 (Prefill + ランキングヘッド一括計算)
# ==========================================
print("B. GenRec方式(Prefill + ランキングヘッド)による推薦をシミュレート中...")

# ランキングヘッドの定義 (隠れ状態ベクトル D から 候補数 M のスコアへ一括変換する軽量な行列積)
ranking_head = nn.Linear(HIDDEN_DIM, NUM_CANDIDATES, bias=False).to(device)

torch.cuda.synchronize() if torch.cuda.is_available() else None
start_time = time.time()

# ① Prefill相: 入力コンテキストを一括処理 (パターンAと同じ)
with torch.no_grad():
hidden_states_genrec = llm_layer(prefill_input)

# ② ランキングヘッド相: 最終位置の隠れ状態を取り出し、全カタログ(M個)のスコアを一括で計算
# 自己回帰ループは「ゼロ」
user_vector = hidden_states_genrec[:, -1, :] # (BATCH_SIZE, HIDDEN_DIM)
with torch.no_grad():
all_scores = ranking_head(user_vector) # (BATCH_SIZE, NUM_CANDIDATES) の行列積を一発で計算

torch.cuda.synchronize() if torch.cuda.is_available() else None
time_genrec = time.time() - start_time
print(f"⇒ GenRec方式 処理時間: {time_genrec:.6f} 秒\n")


# ==========================================
# 結果の比較
# ==========================================
speedup = time_autoregressive / time_genrec
print("=== 結果のまとめ ===")
print(f"通常のLLM (自己回帰): {time_autoregressive:.6f} 秒")
print(f"GenRec方式 (一括): {time_genrec:.6f} 秒")
print(f"→ **GenRec方式は通常のLLMに比べて約 {speedup:.1f} 倍 高速です**")

上記の nn.Linear(HIDDEN_DIM, HIDDEN_DIM) は単なる線形変換。そのため、Q・K・Vを実際に計算する簡易Transformerに置き換える。

import time
import torch
import torch.nn as nn
import torch.nn.functional as F

# ==========================================
# 1. 環境
# ==========================================
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")


# ==========================================
# 2. 設定
# ==========================================
BATCH_SIZE = 1
HIDDEN_DIM = 256 # 2048でもよいが、Colabでは256程度が軽い
NUM_HEADS = 4
PREFILL_LEN = 64 # 入力トークン数 N
DECODE_LEN = 20 # 自己回帰生成する長さ L
NUM_CANDIDATES = 5000 # 推薦候補数 M


# ==========================================
# 3. 簡易Transformer
# Q / K / V を明示的に計算する
# ==========================================
class SimpleSelfAttention(nn.Module):

def __init__(self, hidden_dim, num_heads):
super().__init__()

assert hidden_dim % num_heads == 0

self.hidden_dim = hidden_dim
self.num_heads = num_heads
self.head_dim = hidden_dim // num_heads

# Q, K, V を作る線形変換
self.W_Q = nn.Linear(hidden_dim, hidden_dim, bias=False)
self.W_K = nn.Linear(hidden_dim, hidden_dim, bias=False)
self.W_V = nn.Linear(hidden_dim, hidden_dim, bias=False)

# Attention後の出力変換
self.W_O = nn.Linear(hidden_dim, hidden_dim, bias=False)

def forward(self, x):

# --------------------------------------
# Q / K / V
# --------------------------------------
Q = self.W_Q(x)
K = self.W_K(x)
V = self.W_V(x)

# shape:
# (B, N, D)
# ↓
# (B, heads, N, head_dim)

B, N, D = Q.shape

Q = Q.view(B, N, self.num_heads, self.head_dim)
K = K.view(B, N, self.num_heads, self.head_dim)
V = V.view(B, N, self.num_heads, self.head_dim)

Q = Q.transpose(1, 2)
K = K.transpose(1, 2)
V = V.transpose(1, 2)

# --------------------------------------
# Attention
#
# Q K^T
# --------------------------------------
attention_scores = torch.matmul(
Q,
K.transpose(-2, -1)
) / (self.head_dim ** 0.5)

# --------------------------------------
# Softmax
# --------------------------------------
attention_weights = F.softmax(
attention_scores,
dim=-1
)

# --------------------------------------
# Attention × V
# --------------------------------------
context = torch.matmul(
attention_weights,
V
)

# --------------------------------------
# Headを結合
# --------------------------------------
context = context.transpose(1, 2)

context = context.contiguous().view(
B,
N,
D
)

# 出力射影
output = self.W_O(context)

return output


# ==========================================
# 4. Transformer本体
# ==========================================
class SimpleTransformer(nn.Module):

def __init__(self, hidden_dim, num_heads):
super().__init__()

self.attention = SimpleSelfAttention(
hidden_dim,
num_heads
)

self.ffn = nn.Sequential(
nn.Linear(hidden_dim, hidden_dim * 4),
nn.ReLU(),
nn.Linear(hidden_dim * 4, hidden_dim)
)

def forward(self, x):

# Self-Attention
x = x + self.attention(x)

# Feed Forward
x = x + self.ffn(x)

return x


# ==========================================
# 5. Transformerを作成
# ==========================================
transformer = SimpleTransformer(
HIDDEN_DIM,
NUM_HEADS
).to(device)

transformer.eval()


# ==========================================
# 6. 入力
# ==========================================
prefill_input = torch.randn(
BATCH_SIZE,
PREFILL_LEN,
HIDDEN_DIM
).to(device)

print("\n--- 実験設定 ---")
print(f"入力トークン数 N : {PREFILL_LEN}")
print(f"隠れ状態次元 D : {HIDDEN_DIM}")
print(f"Attention Head数 : {NUM_HEADS}")
print(f"候補数 M : {NUM_CANDIDATES}")
print(f"生成長 L : {DECODE_LEN}")


# ==========================================
# 7. パターンA
# 通常のLLM:自己回帰的に生成
# ==========================================
print("\nA. 通常のLLM(自己回帰生成)")

if torch.cuda.is_available():
torch.cuda.synchronize()

start_time = time.time()

with torch.no_grad():

# --------------------------------------
# Prefill
# --------------------------------------
hidden_states = transformer(prefill_input)

# 最後の位置を次の生成に利用
current_input = hidden_states[:, -1:, :]

# --------------------------------------
# Decode
# --------------------------------------
for t in range(DECODE_LEN):

# Transformerをもう一度実行
current_output = transformer(current_input)

# 次のトークンの入力になる
current_input = current_output

if torch.cuda.is_available():
torch.cuda.synchronize()

time_autoregressive = time.time() - start_time

print(
f"通常LLM 処理時間: "
f"{time_autoregressive:.6f} 秒"
)


# ==========================================
# 8. パターンB
# GenRec:ランキング
# ==========================================
print("\nB. GenRec方式(ランキング)")


# 候補アイテムのEmbedding
candidate_embeddings = nn.Parameter(
torch.randn(
NUM_CANDIDATES,
HIDDEN_DIM
)
).to(device)

# --------------------------------------
# Prefill
# --------------------------------------
if torch.cuda.is_available():
torch.cuda.synchronize()

start_time = time.time()

with torch.no_grad():

hidden_states_genrec = transformer(
prefill_input
)

# ユーザーを表すベクトル
user_vector = hidden_states_genrec[:, -1, :]

# ----------------------------------
# 候補を一括スコアリング
#
# user_vector:
# (1, D)
#
# candidate_embeddings:
# (M, D)
#
# 結果:
# (1, M)
# ----------------------------------
all_scores = torch.matmul(
user_vector,
candidate_embeddings.T
)

if torch.cuda.is_available():
torch.cuda.synchronize()

time_genrec = time.time() - start_time


print(
f"GenRec方式 処理時間: "
f"{time_genrec:.6f} 秒"
)


# ==========================================
# 9. 結果
# ==========================================
print("\n=== 結果 ===")

print(
f"通常LLM(自己回帰): "
f"{time_autoregressive:.6f} 秒"
)

print(
f"GenRec(ランキング): "
f"{time_genrec:.6f} 秒"
)

if time_genrec > 0:
speedup = time_autoregressive / time_genrec

print(
f"\nGenRec方式 / 通常LLM "
f"処理時間比: {speedup:.1f} 倍"
)


# ==========================================
# 10. 形状を確認
# ==========================================
print("\n=== Tensor Shape ===")

print(f"入力 X : {prefill_input.shape}")
print(f"ユーザー表現 h : {user_vector.shape}")
print(f"推薦スコア : {all_scores.shape}")

print(
"\n推薦スコアの形状は "
"(BATCH_SIZE, NUM_CANDIDATES) です。"
)

参考文献

Mathematics is the language with which God has written the universe.





















ELFプログラムヘッダーのデータ構造と型 マルチヘッド・アテンション 学習式位置エンコーディング Dockerのネットワーク構成 kindのインストール Network Namespaceによる最小ネットワーク構成