385 KiB
385 KiB
In [1]:
import torch
from d2l import torch as d2l
/home/yukun/.conda/envs/nn/lib/python3.11/site-packages/torch/cuda/__init__.py:1007: UserWarning: Can't initialize NVML raw_cnt = _raw_device_count_nvml()
In [2]:
def show_heatmaps(matrices,xlabel,ylabel,titles=None,figsize=(2.5,2.5),cmap='Reds'):
d2l.use_svg_display()
num_rows,num_cols = matrices.shape[0],matrices.shape[1]
fig,axes = d2l.plt.subplots(num_rows,num_cols,figsize=figsize,sharex=True,squeeze=False)
for i,(row_axes,row_matrices) in enumerate(zip(axes,matrices)):
for j,(ax,matrix) in enumerate(zip(row_axes,row_matrices)):
pcm = ax.imshow(matrix.detach().numpy(),cmap=cmap)
if i == num_rows - 1:
ax.set_xlabel(xlabel)
if j == 0:
ax.set_ylabel(ylabel)
if titles:
ax.set_title(titles[j])
fig.colorbar(pcm,ax=axes,shrink=0.6)In [3]:
attention_weights = torch.eye(10).reshape((1, 1, 10, 10))
show_heatmaps(attention_weights, xlabel='Keys', ylabel='Queries')In [4]:
n_train = 50
x_train,_ = torch.sort(torch.rand(n_train)*5)In [5]:
def f(x):
return 2*torch.sin(x)+x**0.8In [6]:
y_train = f(x_train)+torch.normal(0.0,0.5,(n_train,))
x_test = torch.arange(0,5,0.1)
y_truth = f(x_test)
n_test = len(x_test)
n_testOut [6]:
50
In [7]:
def plot_kernel_reg(y_hat):
d2l.plot(x_test, [y_truth, y_hat], 'x', 'y', legend=['Truth', 'Pred'],
xlim=[0, 5], ylim=[-1, 5])
d2l.plt.plot(x_train, y_train, 'o', alpha=0.5);In [8]:
y_hat = torch.repeat_interleave(y_train.mean(),n_test)
plot_kernel_reg(y_hat)In [9]:
from torch import nn
X_repeat = x_test.repeat_interleave(n_train).reshape((-1, n_train))
attention_weights = nn.functional.softmax(-(X_repeat - x_train)**2 / 2, dim=1)
y_hat = torch.matmul(attention_weights,y_train)
x_train,X_repeat,attention_weightsOut [9]:
(tensor([0.0153, 0.1180, 0.1875, 0.4215, 0.5745, 0.6163, 0.6557, 0.6917, 0.7939,
0.9230, 1.0240, 1.2932, 1.3246, 1.5167, 1.6193, 1.6496, 1.6613, 1.7104,
1.7485, 1.7840, 1.7889, 2.0187, 2.0398, 2.0413, 2.2665, 2.4287, 2.4440,
2.6026, 2.6988, 2.7922, 2.8129, 2.9356, 3.0271, 3.0442, 3.0646, 3.1996,
3.2037, 3.3222, 3.4918, 3.5313, 3.7002, 3.7714, 3.7949, 3.9609, 4.3272,
4.4440, 4.6667, 4.7213, 4.8963, 4.9971]),
tensor([[0.0000, 0.0000, 0.0000, ..., 0.0000, 0.0000, 0.0000],
[0.1000, 0.1000, 0.1000, ..., 0.1000, 0.1000, 0.1000],
[0.2000, 0.2000, 0.2000, ..., 0.2000, 0.2000, 0.2000],
...,
[4.7000, 4.7000, 4.7000, ..., 4.7000, 4.7000, 4.7000],
[4.8000, 4.8000, 4.8000, ..., 4.8000, 4.8000, 4.8000],
[4.9000, 4.9000, 4.9000, ..., 4.9000, 4.9000, 4.9000]]),
tensor([[7.9004e-02, 7.8465e-02, 7.7637e-02, ..., 1.1413e-06, 4.9195e-07,
2.9875e-07],
[7.2616e-02, 7.2865e-02, 7.2599e-02, ..., 1.6794e-06, 7.3669e-07,
4.5190e-07],
[6.6459e-02, 6.7376e-02, 6.7597e-02, ..., 2.4606e-06, 1.0985e-06,
6.8065e-07],
...,
[1.3743e-06, 2.2120e-06, 3.0334e-06, ..., 8.0086e-02, 7.8576e-02,
7.6646e-02],
[9.2161e-07, 1.4986e-06, 2.0694e-06, ..., 8.5977e-02, 8.5846e-02,
8.4585e-02],
[6.1460e-07, 1.0097e-06, 1.4040e-06, ..., 9.1792e-02, 9.3269e-02,
9.2831e-02]]))In [10]:
plot_kernel_reg(y_hat)In [11]:
d2l.show_heatmaps(attention_weights.unsqueeze(0).unsqueeze(0),
xlabel='Sorted training inputs',
ylabel='Sorted testing inputs')In [12]:
class NWKernelRegression(nn.Module):
def __init__(self,**kwargs):
super().__init__(**kwargs)
self.w = nn.Parameter(torch.randn((1,),requires_grad=True))
def forward(self,queries,keys,values):
queries = queries.repeat_interleave(keys.shape[1]).reshape(-1,keys.shape[1])
self.attention_weights = nn.functional.softmax(
-((queries - keys) * self.w)**2 / 2, dim=1)
return torch.bmm(self.attention_weights.unsqueeze(1),values.unsqueeze(-1)).reshape(-1)In [13]:
# X_tile的形状:(n_train,n_train),每一行都包含着相同的训练输入
X_tile = x_train.repeat((n_train, 1))
# Y_tile的形状:(n_train,n_train),每一行都包含着相同的训练输出
Y_tile = y_train.repeat((n_train, 1))
# keys的形状:('n_train','n_train'-1)
keys = X_tile[(1 - torch.eye(n_train)).type(torch.bool)].reshape((n_train, -1))
# values的形状:('n_train','n_train'-1)
values = Y_tile[(1 - torch.eye(n_train)).type(torch.bool)].reshape((n_train, -1))In [14]:
net = NWKernelRegression()
loss = nn.MSELoss(reduction='none')
trainer = torch.optim.SGD(net.parameters(), lr=0.5)
animator = d2l.Animator(xlabel='epoch', ylabel='loss', xlim=[1, 5])
for epoch in range(5):
trainer.zero_grad()
l = loss(net(x_train, keys, values), y_train)
l.sum().backward()
trainer.step()
print(f'epoch {epoch + 1}, loss {float(l.sum()):.6f}')
animator.add(epoch + 1, float(l.sum()))In [15]:
# keys的形状:(n_test,n_train),每一行包含着相同的训练输入(例如,相同的键)
keys = x_train.repeat((n_test, 1))
# value的形状:(n_test,n_train)
values = y_train.repeat((n_test, 1))
y_hat = net(x_test, keys, values).unsqueeze(1).detach()
plot_kernel_reg(y_hat)In [16]:
d2l.show_heatmaps(net.attention_weights.unsqueeze(0).unsqueeze(0),
xlabel='Sorted training inputs',
ylabel='Sorted testing inputs')In [17]:
def masked_softmax(X,valid_lens):
# X:3D valid_len 1D or 2D
if valid_lens is None:
return nn.functional.softmax(X,dim=-1)
else:
shape = X.shape
if valid_lens.dim()==1:
valid_lens = torch.repeat_interleave(valid_lens,shape[1])
else:
valid_lens = valid_lens.reshape(-1)
#print(valid_lens)
X = d2l.sequence_mask(X.reshape(-1,shape[-1]),valid_lens,value=-1e6)
return nn.functional.softmax(X.reshape(shape),dim=-1)In [18]:
masked_softmax(torch.rand(2,2,4),torch.tensor([2,3]))Out [18]:
tensor([[[0.5584, 0.4416, 0.0000, 0.0000],
[0.6205, 0.3795, 0.0000, 0.0000]],
[[0.4421, 0.3398, 0.2181, 0.0000],
[0.3929, 0.3625, 0.2446, 0.0000]]])In [19]:
class AdditiveAttention(nn.Module):
def __init__(self,key_size,query_size,num_hiddens,dropout,**kwargs):
super(AdditiveAttention,self).__init__(**kwargs)
self.W_k = nn.Linear(key_size,num_hiddens,bias=False)
self.W_q = nn.Linear(query_size,num_hiddens,bias=False)
self.w_v = nn.Linear(num_hiddens,1,bias=False)
self.dropout=nn.Dropout(dropout)
def forward(self,queries,keys,value,valid_lens):
queries,keys=self.W_q(queries),self.W_k(keys)
# queries (batch_size,n_q,1,num_hidden)
# key (batch_size,1,n_k,num_hiddens)
features = queries.unsqueeze(2) + keys.unsqueeze(1)
#features (batch_size,n_q,n_k,num_hidden)
features = torch.tanh(features)
scores = self.w_v(features).squeeze(-1)
#print(f"Inside AdditiveAttention: value.shape = {value.shape}") # 检查此处形状
self.attention_weights = masked_softmax(scores, valid_lens)
return torch.bmm(self.dropout(self.attention_weights), value)In [20]:
queries,keys = torch.normal(0,1,(2,1,20)) , torch.ones((2,10,2))
values = torch.arange(40,dtype=torch.float32).reshape(1,10,4).repeat(2,1,1)
valid_lens = torch.tensor([2, 6])
attention = AdditiveAttention(key_size=2, query_size=20, num_hiddens=8,
dropout=0.1)
attention.eval()
attention(queries,keys,values,valid_lens)Out [20]:
tensor([[[ 2.0000, 3.0000, 4.0000, 5.0000]],
[[10.0000, 11.0000, 12.0000, 13.0000]]], grad_fn=<BmmBackward0>)In [21]:
d2l.show_heatmaps(attention.attention_weights.reshape((1, 1, 2, 10)),
xlabel='Keys', ylabel='Queries')In [22]:
import math
class DotProductAttention(nn.Module):
"""缩放点积注意力"""
def __init__(self, dropout, **kwargs):
super(DotProductAttention, self).__init__(**kwargs)
self.dropout = nn.Dropout(dropout)
# queries的形状:(batch_size,查询的个数,d)
# keys的形状:(batch_size,“键-值”对的个数,d)
# values的形状:(batch_size,“键-值”对的个数,值的维度)
# valid_lens的形状:(batch_size,)或者(batch_size,查询的个数)
def forward(self, queries, keys, values, valid_lens=None):
d = queries.shape[-1]
# 设置transpose_b=True为了交换keys的最后两个维度
scores = torch.bmm(queries, keys.transpose(1,2)) / math.sqrt(d)
self.attention_weights = masked_softmax(scores, valid_lens)
return torch.bmm(self.dropout(self.attention_weights), values)In [23]:
queries = torch.normal(0, 1, (2, 1, 2))
attention = DotProductAttention(dropout=0.5)
attention.eval()
attention(queries, keys, values, valid_lens)Out [23]:
tensor([[[ 2.0000, 3.0000, 4.0000, 5.0000]],
[[10.0000, 11.0000, 12.0000, 13.0000]]])In [24]:
d2l.show_heatmaps(attention.attention_weights.reshape((1, 1, 2, 10)),
xlabel='Keys', ylabel='Queries')In [25]:
class AttentionDecoder(d2l.Decoder):
def __init__(self,**kwargs):
super(AttentionDecoder, self).__init__(**kwargs)
@property
def attention_weight(self):
raise NotImplementedErrorIn [26]:
class Seq2SeqAttentionDecoder(AttentionDecoder):
def __init__(self,vocab_size,embed_size,num_hiddens,num_layers,dropout=0,**kwargs):
super(Seq2SeqAttentionDecoder, self).__init__(**kwargs)
self.attention = AdditiveAttention(
num_hiddens, num_hiddens, num_hiddens,dropout)
self.embedding = nn.Embedding(vocab_size, embed_size)
self.rnn = nn.GRU(
embed_size + num_hiddens, num_hiddens, num_layers,
dropout=dropout)
self.dense = nn.Linear(num_hiddens, vocab_size)
def init_state(self, enc_outputs, enc_valid_lens, *args):
# outputs的形状为(batch_size,num_steps,num_hiddens).
# hidden_state的形状为(num_layers,batch_size,num_hiddens)
outputs, hidden_state = enc_outputs
#print(f"Encoder outputs shape before permute: {outputs.shape}") # 应为 (num_steps, batch_size, num_hiddens) 或 (batch_size, num_steps, num_hiddens)
enc_outputs_permuted = outputs.permute(1, 0, 2)
#print(f"After permute: {enc_outputs_permuted.shape}") # 期望 (batch_size, num_steps, num_hiddens)
return (enc_outputs_permuted, hidden_state, enc_valid_lens)
def forward(self, X, state):
# enc_outputs的形状为(batch_size,num_steps,num_hiddens).
# hidden_state的形状为(num_layers,batch_size,
# num_hiddens)
enc_outputs, hidden_state, enc_valid_lens = state
# 输出X的形状为(num_steps,batch_size,embed_size)
X = self.embedding(X).permute(1, 0, 2)
outputs, self._attention_weights = [], []
for x in X:
# query的形状为(batch_size,1,num_hiddens)
query = torch.unsqueeze(hidden_state[-1], dim=1)
# context的形状为(batch_size,1,num_hiddens)
#print(f"values shape before attention: {enc_outputs.shape}") # 应为 (4, 7, 16)
context = self.attention(
query, enc_outputs, enc_outputs, enc_valid_lens)
# 在特征维度上连结
x = torch.cat((context, torch.unsqueeze(x, dim=1)), dim=-1)
# 将x变形为(1,batch_size,embed_size+num_hiddens)
out, hidden_state = self.rnn(x.permute(1, 0, 2), hidden_state)
outputs.append(out)
self._attention_weights.append(self.attention.attention_weights)
# 全连接层变换后,outputs的形状为
# (num_steps,batch_size,vocab_size)
outputs = self.dense(torch.cat(outputs, dim=0))
return outputs.permute(1, 0, 2), [enc_outputs, hidden_state,
enc_valid_lens]
@property
def attention_weights(self):
return self._attention_weightsIn [27]:
encoder = d2l.Seq2SeqEncoder(vocab_size=10, embed_size=8, num_hiddens=16,
num_layers=2)
encoder.eval()
decoder = Seq2SeqAttentionDecoder(vocab_size=10, embed_size=8, num_hiddens=16,
num_layers=2)
decoder.eval()
X = torch.zeros((4, 7), dtype=torch.long) # (batch_size,num_steps)
state = decoder.init_state(encoder(X), None)
output, state = decoder(X, state)
output.shape, len(state), state[0].shape, len(state[1]), state[1][0].shapeOut [27]:
(torch.Size([4, 7, 10]), 3, torch.Size([4, 7, 16]), 2, torch.Size([4, 16]))
In [27]:
In [28]:
embed_size, num_hiddens, num_layers, dropout = 32, 32, 2, 0.1
batch_size, num_steps = 64, 10
lr, num_epochs, device = 0.005, 250, d2l.try_gpu()
train_iter, src_vocab, tgt_vocab = d2l.load_data_nmt(batch_size, num_steps)
encoder = d2l.Seq2SeqEncoder(
len(src_vocab), embed_size, num_hiddens, num_layers, dropout)
decoder = Seq2SeqAttentionDecoder(
len(tgt_vocab), embed_size, num_hiddens, num_layers, dropout)
net = d2l.EncoderDecoder(encoder, decoder)
d2l.train_seq2seq(net, train_iter, lr, num_epochs, tgt_vocab, device)[31m---------------------------------------------------------------------------[39m [31mKeyboardInterrupt[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[28][39m[32m, line 10[39m [32m 7[39m decoder = Seq2SeqAttentionDecoder( [32m 8[39m [38;5;28mlen[39m(tgt_vocab), embed_size, num_hiddens, num_layers, dropout) [32m 9[39m net = d2l.EncoderDecoder(encoder, decoder) [32m---> [39m[32m10[39m [43md2l[49m[43m.[49m[43mtrain_seq2seq[49m[43m([49m[43mnet[49m[43m,[49m[43m [49m[43mtrain_iter[49m[43m,[49m[43m [49m[43mlr[49m[43m,[49m[43m [49m[43mnum_epochs[49m[43m,[49m[43m [49m[43mtgt_vocab[49m[43m,[49m[43m [49m[43mdevice[49m[43m)[49m [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/d2l/torch.py:3421[39m, in [36mtrain_seq2seq[39m[34m(net, data_iter, lr, num_epochs, tgt_vocab, device)[39m [32m 3418[39m bos = torch.tensor([tgt_vocab[[33m'[39m[33m<bos>[39m[33m'[39m]] * Y.shape[[32m0[39m], [32m 3419[39m device=device).reshape(-[32m1[39m, [32m1[39m) [32m 3420[39m dec_input = d2l.concat([bos, Y[:, :-[32m1[39m]], [32m1[39m) [38;5;66;03m# Teacher forcing[39;00m [32m-> [39m[32m3421[39m Y_hat, _ = [43mnet[49m[43m([49m[43mX[49m[43m,[49m[43m [49m[43mdec_input[49m[43m,[49m[43m [49m[43mX_valid_len[49m[43m)[49m [32m 3422[39m l = loss(Y_hat, Y, Y_valid_len) [32m 3423[39m l.sum().backward() [38;5;66;03m# Make the loss scalar for `backward`[39;00m [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/module.py:1776[39m, in [36mModule._wrapped_call_impl[39m[34m(self, *args, **kwargs)[39m [32m 1774[39m [38;5;28;01mreturn[39;00m [38;5;28mself[39m._compiled_call_impl(*args, **kwargs) [38;5;66;03m# type: ignore[misc][39;00m [32m 1775[39m [38;5;28;01melse[39;00m: [32m-> [39m[32m1776[39m [38;5;28;01mreturn[39;00m [38;5;28;43mself[39;49m[43m.[49m[43m_call_impl[49m[43m([49m[43m*[49m[43margs[49m[43m,[49m[43m [49m[43m*[49m[43m*[49m[43mkwargs[49m[43m)[49m [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/module.py:1787[39m, in [36mModule._call_impl[39m[34m(self, *args, **kwargs)[39m [32m 1782[39m [38;5;66;03m# If we don't have any hooks, we want to skip the rest of the logic in[39;00m [32m 1783[39m [38;5;66;03m# this function, and just call forward.[39;00m [32m 1784[39m [38;5;28;01mif[39;00m [38;5;129;01mnot[39;00m ([38;5;28mself[39m._backward_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._backward_pre_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._forward_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._forward_pre_hooks [32m 1785[39m [38;5;129;01mor[39;00m _global_backward_pre_hooks [38;5;129;01mor[39;00m _global_backward_hooks [32m 1786[39m [38;5;129;01mor[39;00m _global_forward_hooks [38;5;129;01mor[39;00m _global_forward_pre_hooks): [32m-> [39m[32m1787[39m [38;5;28;01mreturn[39;00m [43mforward_call[49m[43m([49m[43m*[49m[43margs[49m[43m,[49m[43m [49m[43m*[49m[43m*[49m[43mkwargs[49m[43m)[49m [32m 1789[39m result = [38;5;28;01mNone[39;00m [32m 1790[39m called_always_called_hooks = [38;5;28mset[39m() [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/d2l/torch.py:964[39m, in [36mEncoderDecoder.forward[39m[34m(self, enc_X, dec_X, *args)[39m [32m 962[39m dec_state = [38;5;28mself[39m.decoder.init_state(enc_all_outputs, *args) [32m 963[39m [38;5;66;03m# Return decoder output only[39;00m [32m--> [39m[32m964[39m [38;5;28;01mreturn[39;00m [38;5;28;43mself[39;49m[43m.[49m[43mdecoder[49m[43m([49m[43mdec_X[49m[43m,[49m[43m [49m[43mdec_state[49m[43m)[49m [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/module.py:1776[39m, in [36mModule._wrapped_call_impl[39m[34m(self, *args, **kwargs)[39m [32m 1774[39m [38;5;28;01mreturn[39;00m [38;5;28mself[39m._compiled_call_impl(*args, **kwargs) [38;5;66;03m# type: ignore[misc][39;00m [32m 1775[39m [38;5;28;01melse[39;00m: [32m-> [39m[32m1776[39m [38;5;28;01mreturn[39;00m [38;5;28;43mself[39;49m[43m.[49m[43m_call_impl[49m[43m([49m[43m*[49m[43margs[49m[43m,[49m[43m [49m[43m*[49m[43m*[49m[43mkwargs[49m[43m)[49m [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/module.py:1787[39m, in [36mModule._call_impl[39m[34m(self, *args, **kwargs)[39m [32m 1782[39m [38;5;66;03m# If we don't have any hooks, we want to skip the rest of the logic in[39;00m [32m 1783[39m [38;5;66;03m# this function, and just call forward.[39;00m [32m 1784[39m [38;5;28;01mif[39;00m [38;5;129;01mnot[39;00m ([38;5;28mself[39m._backward_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._backward_pre_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._forward_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._forward_pre_hooks [32m 1785[39m [38;5;129;01mor[39;00m _global_backward_pre_hooks [38;5;129;01mor[39;00m _global_backward_hooks [32m 1786[39m [38;5;129;01mor[39;00m _global_forward_hooks [38;5;129;01mor[39;00m _global_forward_pre_hooks): [32m-> [39m[32m1787[39m [38;5;28;01mreturn[39;00m [43mforward_call[49m[43m([49m[43m*[49m[43margs[49m[43m,[49m[43m [49m[43m*[49m[43m*[49m[43mkwargs[49m[43m)[49m [32m 1789[39m result = [38;5;28;01mNone[39;00m [32m 1790[39m called_always_called_hooks = [38;5;28mset[39m() [36mCell[39m[36m [39m[32mIn[26][39m[32m, line 37[39m, in [36mSeq2SeqAttentionDecoder.forward[39m[34m(self, X, state)[39m [32m 35[39m x = torch.cat((context, torch.unsqueeze(x, dim=[32m1[39m)), dim=-[32m1[39m) [32m 36[39m [38;5;66;03m# 将x变形为(1,batch_size,embed_size+num_hiddens)[39;00m [32m---> [39m[32m37[39m out, hidden_state = [38;5;28;43mself[39;49m[43m.[49m[43mrnn[49m[43m([49m[43mx[49m[43m.[49m[43mpermute[49m[43m([49m[32;43m1[39;49m[43m,[49m[43m [49m[32;43m0[39;49m[43m,[49m[43m [49m[32;43m2[39;49m[43m)[49m[43m,[49m[43m [49m[43mhidden_state[49m[43m)[49m [32m 38[39m outputs.append(out) [32m 39[39m [38;5;28mself[39m._attention_weights.append([38;5;28mself[39m.attention.attention_weights) [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/module.py:1776[39m, in [36mModule._wrapped_call_impl[39m[34m(self, *args, **kwargs)[39m [32m 1774[39m [38;5;28;01mreturn[39;00m [38;5;28mself[39m._compiled_call_impl(*args, **kwargs) [38;5;66;03m# type: ignore[misc][39;00m [32m 1775[39m [38;5;28;01melse[39;00m: [32m-> [39m[32m1776[39m [38;5;28;01mreturn[39;00m [38;5;28;43mself[39;49m[43m.[49m[43m_call_impl[49m[43m([49m[43m*[49m[43margs[49m[43m,[49m[43m [49m[43m*[49m[43m*[49m[43mkwargs[49m[43m)[49m [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/module.py:1787[39m, in [36mModule._call_impl[39m[34m(self, *args, **kwargs)[39m [32m 1782[39m [38;5;66;03m# If we don't have any hooks, we want to skip the rest of the logic in[39;00m [32m 1783[39m [38;5;66;03m# this function, and just call forward.[39;00m [32m 1784[39m [38;5;28;01mif[39;00m [38;5;129;01mnot[39;00m ([38;5;28mself[39m._backward_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._backward_pre_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._forward_hooks [38;5;129;01mor[39;00m [38;5;28mself[39m._forward_pre_hooks [32m 1785[39m [38;5;129;01mor[39;00m _global_backward_pre_hooks [38;5;129;01mor[39;00m _global_backward_hooks [32m 1786[39m [38;5;129;01mor[39;00m _global_forward_hooks [38;5;129;01mor[39;00m _global_forward_pre_hooks): [32m-> [39m[32m1787[39m [38;5;28;01mreturn[39;00m [43mforward_call[49m[43m([49m[43m*[49m[43margs[49m[43m,[49m[43m [49m[43m*[49m[43m*[49m[43mkwargs[49m[43m)[49m [32m 1789[39m result = [38;5;28;01mNone[39;00m [32m 1790[39m called_always_called_hooks = [38;5;28mset[39m() [36mFile [39m[32m~/.conda/envs/nn/lib/python3.11/site-packages/torch/nn/modules/rnn.py:1415[39m, in [36mGRU.forward[39m[34m(self, input, hx)[39m [32m 1413[39m [38;5;28mself[39m.check_forward_args([38;5;28minput[39m, hx, batch_sizes) [32m 1414[39m [38;5;28;01mif[39;00m batch_sizes [38;5;129;01mis[39;00m [38;5;28;01mNone[39;00m: [32m-> [39m[32m1415[39m result = [43m_VF[49m[43m.[49m[43mgru[49m[43m([49m [32m 1416[39m [43m [49m[38;5;28;43minput[39;49m[43m,[49m [32m 1417[39m [43m [49m[43mhx[49m[43m,[49m [32m 1418[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43m_flat_weights[49m[43m,[49m[43m [49m[38;5;66;43;03m# type: ignore[arg-type][39;49;00m [32m 1419[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43mbias[49m[43m,[49m [32m 1420[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43mnum_layers[49m[43m,[49m [32m 1421[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43mdropout[49m[43m,[49m [32m 1422[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43mtraining[49m[43m,[49m [32m 1423[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43mbidirectional[49m[43m,[49m [32m 1424[39m [43m [49m[38;5;28;43mself[39;49m[43m.[49m[43mbatch_first[49m[43m,[49m [32m 1425[39m [43m [49m[43m)[49m [32m 1426[39m [38;5;28;01melse[39;00m: [32m 1427[39m result = _VF.gru( [32m 1428[39m [38;5;28minput[39m, [32m 1429[39m batch_sizes, [32m (...)[39m[32m 1436[39m [38;5;28mself[39m.bidirectional, [32m 1437[39m ) [31mKeyboardInterrupt[39m:
In [29]:
engs = ['go .', "i lost .", 'he\'s calm .', 'i\'m home .']
fras = ['va !', 'j\'ai perdu .', 'il est calme .', 'je suis chez moi .']
for eng, fra in zip(engs, fras):
translation, dec_attention_weight_seq = d2l.predict_seq2seq(
net, eng, src_vocab, tgt_vocab, num_steps, device, True)
print(f'{eng} => {translation}, ',
f'bleu {d2l.bleu(translation, fra, k=2):.3f}')go . => va !, bleu 1.000 i lost . => j'ai perdu ., bleu 1.000 he's calm . => il est riche ., bleu 0.658 i'm home . => je suis calme ., bleu 0.512
In [30]:
class MultiHeadAttention(nn.Module):
def __init__(self, key_size, query_size, value_size, num_hiddens,
num_heads, dropout, bias=False, **kwargs):
super(MultiHeadAttention, self).__init__(**kwargs)
self.num_heads = num_heads
self.attention = d2l.DotProductAttention(dropout)
self.W_q = nn.Linear(query_size, num_hiddens, bias=bias)
self.W_k = nn.Linear(key_size, num_hiddens, bias=bias)
self.W_v = nn.Linear(value_size, num_hiddens, bias=bias)
self.W_o = nn.Linear(num_hiddens, num_hiddens, bias=bias)
def forward(self, queries, keys, values, valid_lens):
# 1. 线性投影 + 变换形状以分割多头
queries = transpose_qkv(self.W_q(queries), self.num_heads)
keys = transpose_qkv(self.W_k(keys), self.num_heads)
values = transpose_qkv(self.W_v(values), self.num_heads)
# 2. 处理有效长度掩码(valid_lens)以适配多头
if valid_lens is not None:
valid_lens = torch.repeat_interleave(valid_lens, repeats=self.num_heads, dim=0)
# 3. 计算注意力(每个头独立计算)
output = self.attention(queries, keys, values, valid_lens)
# 4. 合并多头,并通过输出线性层
output_concat = transpose_output(output, self.num_heads)
return self.W_o(output_concat)
def transpose_qkv(X, num_heads):
"""为了多注意力头的并行计算而变换形状"""
# 输入X的形状:(batch_size,查询或者“键-值”对的个数,num_hiddens)
# 输出X的形状:(batch_size,查询或者“键-值”对的个数,num_heads,
# num_hiddens/num_heads)
X = X.reshape(X.shape[0], X.shape[1], num_heads, -1)
# 输出X的形状:(batch_size,num_heads,查询或者“键-值”对的个数,
# num_hiddens/num_heads)
X = X.permute(0, 2, 1, 3)
# 最终输出的形状:(batch_size*num_heads,查询或者“键-值”对的个数,
# num_hiddens/num_heads)
return X.reshape(-1, X.shape[2], X.shape[3])
def transpose_output(X, num_heads):
"""逆转transpose_qkv函数的操作"""
X = X.reshape(-1, num_heads, X.shape[1], X.shape[2])
X = X.permute(0, 2, 1, 3)
return X.reshape(X.shape[0], X.shape[1], -1)
In [31]:
num_hiddens, num_heads = 100, 5
attention = MultiHeadAttention(num_hiddens, num_hiddens, num_hiddens,
num_hiddens, num_heads, 0.5)
attention.eval()Out [31]:
MultiHeadAttention(
(attention): DotProductAttention(
(dropout): Dropout(p=0.5, inplace=False)
)
(W_q): Linear(in_features=100, out_features=100, bias=False)
(W_k): Linear(in_features=100, out_features=100, bias=False)
(W_v): Linear(in_features=100, out_features=100, bias=False)
(W_o): Linear(in_features=100, out_features=100, bias=False)
)In [32]:
batch_size, num_queries = 2, 4
num_kvpairs, valid_lens = 6, torch.tensor([3, 2])
X = torch.ones((batch_size, num_queries, num_hiddens))
Y = torch.ones((batch_size, num_kvpairs, num_hiddens))
attention(X, Y, Y, valid_lens).shapeOut [32]:
torch.Size([2, 4, 100])
In [35]:
class PositionalEncoding(nn.Module):
def __init__(self,num_hiddens,dropout,max_len=1000):
super().__init__()
self.dropout = nn.Dropout(dropout)
self.dropout = nn.Dropout(dropout)
X = torch.arange(max_len, dtype=torch.float32).reshape(
-1, 1) / torch.pow(10000, torch.arange(
0, num_hiddens, 2, dtype=torch.float32) / num_hiddens)
self.P = torch.zeros((1,max_len,num_hiddens))
self.P[:,:,0::2] = torch.sin(X)
self.P[:,:,1::2] = torch.cos(X)
def forward(self,X):
X = X + self.P[:,:X.shape[1],:].to(X.device)
return self.dropout(X)In [36]:
encoding_dim, num_steps = 32, 60
pos_encoding = PositionalEncoding(encoding_dim, 0)
pos_encoding.eval()
X = pos_encoding(torch.zeros((1, num_steps, encoding_dim)))
P = pos_encoding.P[:, :X.shape[1], :]
d2l.plot(torch.arange(num_steps), P[0, :, 6:10].T, xlabel='Row (position)',
figsize=(6, 2.5), legend=["Col %d" % d for d in torch.arange(6, 10)])In [38]:
class PositionWiseFFN(nn.Module):
def __init__(self,ffn_num_input,ffn_num_hiddens,ffn_num_outputs,**kwargs):
super().__init__(**kwargs)
self.dense1 = nn.Linear(ffn_num_input,ffn_num_hiddens)
self.relu = nn.ReLU()
self.dense2 = nn.Linear(ffn_num_hiddens,ffn_num_outputs)
def forward(self,X):
return self.dense2(self.relu(self.dense1(X)))
In [39]:
class AddNorm(nn.Module):
def __init__(self,normalized_shape,dropout,**kwargs):
super().__init__()
self.dropout=nn.Dropout()
self.ln = nn.LayerNorm(normalized_shape)
def forward(self,X,Y):
return self.ln(self.dropout(Y)+X)
In [63]:
class EncoderBlock(nn.Module):
def __init__(self,key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_heads,dropout,use_bias=False,**kwargs):
super(EncoderBlock,self).__init__(**kwargs)
self.attention = MultiHeadAttention(key_size,query_size,value_size,num_hiddens,num_heads,dropout,bias=use_bias)
self.addnorm1 = AddNorm(norm_shape,dropout)
self.ffn = PositionWiseFFN(
ffn_num_input, ffn_num_hiddens, num_hiddens)
self.addnorm2 = AddNorm(norm_shape, dropout)
def forward(self, X, valid_lens):
Y = self.addnorm1(X, self.attention(X, X, X, valid_lens))
return self.addnorm2(Y, self.ffn(Y))
In [64]:
X = torch.ones((2, 100, 24))
valid_lens = torch.tensor([3, 2])
encoder_blk = EncoderBlock(24, 24, 24, 24, [100, 24], 24, 48, 8, 0.5)
encoder_blk.eval()
encoder_blk(X, valid_lens).shapeOut [64]:
torch.Size([2, 100, 24])
In [65]:
class TransformerEncoder(d2l.Encoder):
def __init__(self,vocab_size,key_size,query_size,value_size,num_hiddens,norm_shape, ffn_num_input,ffn_num_hiddens,num_head,num_layers,dropout,use_bias=False,**kwargs):
super().__init__(**kwargs)
self.num_hiddens = num_hiddens
self.embedding = nn.Embedding(vocab_size, num_hiddens)
self.pos_encoding = PositionalEncoding(num_hiddens,dropout)
self.blks = nn.Sequential()
for i in range(num_layers):
self.blks.add_module("block"+str(i),EncoderBlock(key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_head,dropout,use_bias))
def forward(self,X,valid_lens,*args):
X = self.pos_encoding(self.embedding(X)*math.sqrt(self.num_hiddens))
self.attention_weights = [None] * len(self.blks)
for i, blk in enumerate(self.blks):
X = blk(X, valid_lens)
self.attention_weights[i] = blk.attention.attention.attention_weights
return XIn [66]:
encoder = TransformerEncoder(
200, 24, 24, 24, 24, [100, 24], 24, 48, 8, 2, 0.5)
encoder.eval()
encoder(torch.ones((2, 100), dtype=torch.long), valid_lens).shapeOut [66]:
torch.Size([2, 100, 24])
In [67]:
class DecoderBlock(nn.Module):
"""解码器中第i个块"""
def __init__(self,key_size,query_size,value_size,num_hiddens,norm_shape,ffn_num_input,ffn_num_hiddens,num_heads,dropout,i,**kwargs):
super().__init__()
self.i = i
self.attention1 = MultiHeadAttention(key_size,query_size,value_size,num_hiddens,num_heads,dropout)
self.addnorm1 = AddNorm(norm_shape, dropout)
self.attention2 = MultiHeadAttention(
key_size, query_size, value_size, num_hiddens, num_heads, dropout)
self.addnorm2 = AddNorm(norm_shape, dropout)
self.ffn = PositionWiseFFN(ffn_num_input, ffn_num_hiddens,
num_hiddens)
self.addnorm3 = AddNorm(norm_shape, dropout)
def forward(self,X,state):
enc_outputs,enc_valid_lens=state[0],state[1]
if state[2][self.i] is None:
key_values = X
else:
key_values = torch.cat((state[2][self.i],X),axis=1)
state[2][self.i] = key_values
if self.training:
batch_size, num_steps, _ = X.shape
# dec_valid_lens的开头:(batch_size,num_steps),
# 其中每一行是[1,2,...,num_steps]
dec_valid_lens = torch.arange(
1, num_steps + 1, device=X.device).repeat(batch_size, 1)
else:
dec_valid_lens = None
# 自注意力
X2 = self.attention1(X, key_values, key_values, dec_valid_lens)
Y = self.addnorm1(X, X2)
# 编码器-解码器注意力。
# enc_outputs的开头:(batch_size,num_steps,num_hiddens)
Y2 = self.attention2(Y, enc_outputs, enc_outputs, enc_valid_lens)
Z = self.addnorm2(Y, Y2)
return self.addnorm3(Z, self.ffn(Z)), stateIn [68]:
decoder_blk = DecoderBlock(24, 24, 24, 24, [100, 24], 24, 48, 8, 0.5, 0)
decoder_blk.eval()
X = torch.ones((2, 100, 24))
state = [encoder_blk(X, valid_lens), valid_lens, [None]]
decoder_blk(X, state)[0].shapeOut [68]:
torch.Size([2, 100, 24])
In [95]:
class TransformerDecoder(d2l.AttentionDecoder):
def __init__(self, vocab_size, key_size, query_size, value_size,
num_hiddens, norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, num_layers, dropout, **kwargs):
super(TransformerDecoder, self).__init__(**kwargs)
self.num_hiddens = num_hiddens
self.num_layers = num_layers
self.embedding = nn.Embedding(vocab_size, num_hiddens)
self.pos_encoding = PositionalEncoding(num_hiddens, dropout)
self.blks = nn.Sequential()
for i in range(num_layers):
self.blks.add_module("block"+str(i),
DecoderBlock(key_size, query_size, value_size, num_hiddens,
norm_shape, ffn_num_input, ffn_num_hiddens,
num_heads, dropout, i))
self.dense = nn.Linear(num_hiddens, vocab_size)
def init_state(self, enc_outputs, enc_valid_lens, *args):
return [enc_outputs, enc_valid_lens, [None] * self.num_layers]
def forward(self, X, state):
X = self.pos_encoding(self.embedding(X) * math.sqrt(self.num_hiddens))
self._attention_weights = [[None] * len(self.blks) for _ in range (2)]
for i, blk in enumerate(self.blks):
X, state = blk(X, state)
# 解码器自注意力权重
self._attention_weights[0][
i] = blk.attention1.attention.attention_weights
# “编码器-解码器”自注意力权重
self._attention_weights[1][
i] = blk.attention2.attention.attention_weights
return self.dense(X), state
@property
def attention_weights(self):
return self._attention_weightsIn [ ]:
In [104]:
num_hiddens, num_layers, dropout, batch_size, num_steps = 32, 2, 0.1, 64, 10
lr, num_epochs, device = 0.005, 200, d2l.try_gpu()
ffn_num_input, ffn_num_hiddens, num_heads = 32, 64, 4
key_size, query_size, value_size = 32, 32, 32
norm_shape = [32]
train_iter, src_vocab, tgt_vocab = d2l.load_data_nmt(batch_size, num_steps)
encoder = TransformerEncoder(
len(src_vocab), key_size, query_size, value_size, num_hiddens,
norm_shape, ffn_num_input, ffn_num_hiddens, num_heads,
num_layers, dropout)
decoder = TransformerDecoder(
len(tgt_vocab), key_size, query_size, value_size, num_hiddens,
norm_shape, ffn_num_input, ffn_num_hiddens, num_heads,
num_layers, dropout)
net = d2l.EncoderDecoder(encoder, decoder)
d2l.train_seq2seq(net, train_iter, lr, num_epochs, tgt_vocab, device)loss 0.079, 1913.5 tokens/sec on cpu
In [106]:
engs = ['go .', "i lost .", 'he\'s calm .', 'i\'m home .']
fras = ['va !', 'j\'ai perdu .', 'il est calme .', 'je suis chez moi .']
for eng, fra in zip(engs, fras):
translation, dec_attention_weight_seq = d2l.predict_seq2seq(
net, eng, src_vocab, tgt_vocab, num_steps, device, True)
print(f'{eng} => {translation}, ',
f'bleu {d2l.bleu(translation, fra, k=2):.3f}')go . => va doucement !, bleu 0.000 i lost . => je <unk> ., bleu 0.000 he's calm . => il est mouillé ., bleu 0.658 i'm home . => je suis malade ., bleu 0.512
In [107]:
enc_attention_weights = torch.cat(net.encoder.attention_weights, 0).reshape((num_layers, num_heads,
-1, num_steps))
enc_attention_weights.shapeOut [107]:
torch.Size([2, 4, 10, 10])
In [108]:
d2l.show_heatmaps(
enc_attention_weights.cpu(), xlabel='Key positions',
ylabel='Query positions', titles=['Head %d' % i for i in range(1, 5)],
figsize=(7, 3.5))In [ ]: