Files
nn/chapter5-9.ipynb
T
2026-04-01 23:34:58 +08:00

463 KiB

In [1]:

import torch
import d2l
import numpy
import torch.nn as nn
import torch.nn.functional as F
In [2]:
net = nn.Sequential(nn.Linear(20, 256), nn.ReLU(), nn.Linear(256, 10))
X = torch.rand(2, 20)
net(X)
Out [2]:
tensor([[-0.0620,  0.3494, -0.2706,  0.0934,  0.1624,  0.0507,  0.0276,  0.1200,
          0.1418,  0.1422],
        [-0.2129,  0.3243, -0.3727, -0.0147,  0.1620, -0.0534,  0.0918,  0.0569,
          0.0515, -0.0515]], grad_fn=<AddmmBackward0>)
In [3]:
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.hidden=nn.Linear(20,256)
        self.out=nn.Linear(256,10)
    def forward(self,X):
        return self.out(F.relu(self.hidden(X)))
In [4]:
net=MLP()
net(X)
Out [4]:
tensor([[-0.0426, -0.0041,  0.0686,  0.0151, -0.0754,  0.0269,  0.2757,  0.0227,
          0.2260,  0.0424],
        [ 0.0319,  0.0394,  0.0179,  0.0704, -0.1369, -0.0294,  0.2276,  0.0702,
          0.1313,  0.2124]], grad_fn=<AddmmBackward0>)
In [5]:
class FixedHiddenMLP(nn.Module):
    def __init__(self):
        super().__init__()
        # 不计算梯度的随机权重参数。因此其在训练期间保持不变
        self.rand_weight = torch.rand((20, 20), requires_grad=False)
        self.linear = nn.Linear(20, 20)
    def forward(self, X):
        X = self.linear(X)
        # 使用创建的常量参数以及relu和mm函数
        X = F.relu(torch.mm(X, self.rand_weight) + 1)
        # 复用全连接层。这相当于两个全连接层共享参数
        X = self.linear(X)
        # 控制流
        while X.abs().sum() > 1:
            X /= 2
        return X.sum()
In [6]:
net = FixedHiddenMLP()
net(X)
Out [6]:
tensor(-0.0847, grad_fn=<SumBackward0>)
In [7]:
class NestMLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.net = nn.Sequential(nn.Linear(20, 64), nn.ReLU(),
                                    nn.Linear(64, 32), nn.ReLU())
        self.linear = nn.Linear(32, 16)
    def forward(self, X):
        return self.linear(self.net(X))
        chimera = nn.Sequential(NestMLP(), nn.Linear(16, 20), FixedHiddenMLP())
        chimera(X)
In [8]:
net = nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 1))
X = torch.rand(size=(2, 4))
net(X)
Out [8]:
tensor([[-0.5641],
        [-0.5857]], grad_fn=<AddmmBackward0>)
In [9]:
print(net[2].state_dict())
OrderedDict([('weight', tensor([[-0.3334,  0.0416, -0.3501, -0.1285,  0.0512,  0.2549, -0.2154, -0.2633]])), ('bias', tensor([-0.0772]))])
In [10]:
net[2].state_dict()
Out [10]:
OrderedDict([('weight',
              tensor([[-0.3334,  0.0416, -0.3501, -0.1285,  0.0512,  0.2549, -0.2154, -0.2633]])),
             ('bias', tensor([-0.0772]))])
In [11]:
print(type(net[2].bias))
<class 'torch.nn.parameter.Parameter'>
In [12]:
print(net[2].bias)
print(net[2].bias.data)
Parameter containing:
tensor([-0.0772], requires_grad=True)
tensor([-0.0772])
In [13]:
net[2].weight.grad==None
Out [13]:
True
In [14]:
print(*[(name, param.shape) for name, param in net[0].named_parameters()])
print(*[(name, param.shape) for name, param in net.named_parameters()])
('weight', torch.Size([8, 4])) ('bias', torch.Size([8]))
('0.weight', torch.Size([8, 4])) ('0.bias', torch.Size([8])) ('2.weight', torch.Size([1, 8])) ('2.bias', torch.Size([1]))
In [15]:
net.state_dict()['2.bias'].data
Out [15]:
tensor([-0.0772])
In [16]:
def block1():
    return nn.Sequential(nn.Linear(4, 8), nn.ReLU(), nn.Linear(8, 4),nn.ReLU())
def block2():
    net = nn.Sequential()
    for i in range(4):
        net.add_module(f'block{i}', block1())
    return net
In [17]:
rgnet = nn.Sequential(block2(),nn.Linear(4,1))
rgnet(X)
Out [17]:
tensor([[0.2190],
        [0.2190]], grad_fn=<AddmmBackward0>)
In [18]:
print(rgnet)
Sequential(
  (0): Sequential(
    (block0): Sequential(
      (0): Linear(in_features=4, out_features=8, bias=True)
      (1): ReLU()
      (2): Linear(in_features=8, out_features=4, bias=True)
      (3): ReLU()
    )
    (block1): Sequential(
      (0): Linear(in_features=4, out_features=8, bias=True)
      (1): ReLU()
      (2): Linear(in_features=8, out_features=4, bias=True)
      (3): ReLU()
    )
    (block2): Sequential(
      (0): Linear(in_features=4, out_features=8, bias=True)
      (1): ReLU()
      (2): Linear(in_features=8, out_features=4, bias=True)
      (3): ReLU()
    )
    (block3): Sequential(
      (0): Linear(in_features=4, out_features=8, bias=True)
      (1): ReLU()
      (2): Linear(in_features=8, out_features=4, bias=True)
      (3): ReLU()
    )
  )
  (1): Linear(in_features=4, out_features=1, bias=True)
)
In [19]:
rgnet[0][1][0].bias.data
Out [19]:
tensor([-0.3955,  0.0030, -0.0100,  0.3198, -0.4639, -0.4023, -0.3653,  0.0766])
In [20]:
def init_normal(m):
    if type(m) == nn.Linear:
        nn.init.normal_(m.weight, mean=0, std=0.01)
        nn.init.zeros_(m.bias)
net.apply(init_normal)
net[0].weight.data[0], net[0].bias.data[0]
Out [20]:
(tensor([-0.0059, -0.0004, -0.0091,  0.0014]), tensor(0.))
In [21]:
def init_xavier(m):
    if type(m) == nn.Linear:
        nn.init.xavier_uniform_(m.weight)
def init_42(m):
    if type(m) == nn.Linear:
        nn.init.constant_(m.weight, 42)

net[0].apply(init_xavier)
net[2].apply(init_42)
print(net[0].weight.data[0])
print(net[2].weight.data)
tensor([ 0.1297, -0.3070, -0.2955,  0.3630])
tensor([[42., 42., 42., 42., 42., 42., 42., 42.]])
In [22]:
x = torch.arange(4)
torch.save(x, 'x-file')
In [23]:
x2 = torch.load('x-file')
x2
Out [23]:
tensor([0, 1, 2, 3])
In [24]:
class MLP(nn.Module):
    def __init__(self):
        super().__init__()
        self.hidden = nn.Linear(20, 256)
        self.output = nn.Linear(256, 10)
    def forward(self, x):
        return self.output(F.relu(self.hidden(x)))

net = MLP()
X = torch.randn(size=(2, 20))
Y = net(X)
In [25]:
torch.save(net.state_dict(), 'mlp.params')
In [26]:
clone = MLP()
clone.load_state_dict(torch.load('mlp.params'))
clone.eval()
Out [26]:
MLP(
  (hidden): Linear(in_features=20, out_features=256, bias=True)
  (output): Linear(in_features=256, out_features=10, bias=True)
)
In [27]:
Y_clone = clone(X)
Y_clone == Y
Out [27]:
tensor([[True, True, True, True, True, True, True, True, True, True],
        [True, True, True, True, True, True, True, True, True, True]])
In [28]:
def corr2d(X,K):
    h,w=K.shape
    Y=torch.ones((X.shape[0]-h+1,X.shape[1]-w+1))
    for i in range(Y.shape[0]):
        for j in range(Y.shape[1]):
            Y[i,j]=(X[i:i+h,j:j+w]*K).sum()
    return Y
In [29]:
X = torch.tensor([[0.0, 1.0, 2.0], [3.0, 4.0, 5.0], [6.0, 7.0, 8.0]])
K = torch.tensor([[0.0, 1.0], [2.0, 3.0]])
corr2d(X,K)
Out [29]:
tensor([[19., 25.],
        [37., 43.]])
In [30]:
class Conv2D(nn.Module):
    def __init__(self, kernel_size):
        super().__init__()
        self.weight = nn.Parameter(torch.rand(kernel_size))
        self.bias = nn.Parameter(torch.zeros(1))
    def forward(self, x):
        return corr2d(x, self.weight) + self.bias
In [31]:
X = torch.ones((6, 8))
X[:, 2:6] = 0
X
Out [31]:
tensor([[1., 1., 0., 0., 0., 0., 1., 1.],
        [1., 1., 0., 0., 0., 0., 1., 1.],
        [1., 1., 0., 0., 0., 0., 1., 1.],
        [1., 1., 0., 0., 0., 0., 1., 1.],
        [1., 1., 0., 0., 0., 0., 1., 1.],
        [1., 1., 0., 0., 0., 0., 1., 1.]])
In [32]:
K = torch.tensor([[1.0, -1.0]])
Y = corr2d(X, K)
Y
Out [32]:
tensor([[ 0.,  1.,  0.,  0.,  0., -1.,  0.],
        [ 0.,  1.,  0.,  0.,  0., -1.,  0.],
        [ 0.,  1.,  0.,  0.,  0., -1.,  0.],
        [ 0.,  1.,  0.,  0.,  0., -1.,  0.],
        [ 0.,  1.,  0.,  0.,  0., -1.,  0.],
        [ 0.,  1.,  0.,  0.,  0., -1.,  0.]])
In [33]:
corr2d(X.t(), K)
Out [33]:
tensor([[0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.],
        [0., 0., 0., 0., 0.]])
In [34]:
conv2d = nn.Conv2d(1,1, kernel_size=(1, 2), bias=False)
In [35]:
X = X.reshape((1, 1, 6, 8))
Y = Y.reshape((1, 1, 6, 7))
lr = 3e-2
In [36]:
for i in range(100):
    Y_hat = conv2d(X)
    l = (Y_hat - Y) ** 2
    conv2d.zero_grad()
    l.sum().backward()
    # 迭代卷积核
    conv2d.weight.data[:] -= lr * conv2d.weight.grad
    if (i + 1) % 20 == 0:
        print(f'epoch {i+1}, loss {l.sum():.3f}')
epoch 20, loss 0.000
epoch 40, loss 0.000
epoch 60, loss 0.000
epoch 80, loss 0.000
epoch 100, loss 0.000
In [37]:
conv2d.weight.data.reshape((1, 2))
Out [37]:
tensor([[ 1.0000, -1.0000]])
In [38]:

# 为了方便起见,我们定义了一个计算卷积层的函数。
# 此函数初始化卷积层权重,并对输入和输出提高和缩减相应的维数
def comp_conv2d(conv2d, X):
# 这里的(1,1)表示批量大小和通道数都是1
    X = X.reshape((1, 1) + X.shape)
    Y = conv2d(X)
    # 省略前两个维度:批量大小和通道
    return Y.reshape(Y.shape[2:])
# 请注意,这里每边都填充了1行或1列,因此总共添加了2行或2列
conv2d = nn.Conv2d(1, 1, kernel_size=3, padding=1)
In [39]:
X = torch.rand(size=(8, 8))
comp_conv2d(conv2d, X).shape
Out [39]:
torch.Size([8, 8])
In [40]:
conv2d = nn.Conv2d(1, 1, kernel_size=(5, 3), padding=(2, 1))
comp_conv2d(conv2d, X).shape
Out [40]:
torch.Size([8, 8])
In [41]:
conv2d = nn.Conv2d(1, 1, kernel_size=3, padding=1, stride=2)
comp_conv2d(conv2d, X).shape
Out [41]:
torch.Size([4, 4])
In [42]:
conv2d = nn.Conv2d(1, 1, kernel_size=(3, 5), padding=(0, 1), stride=(3, 4))
comp_conv2d(conv2d, X).shape
Out [42]:
torch.Size([2, 2])
In [43]:
def corr2d_multi_in(X,K):
    return sum(corr2d(x,k) for x,k in zip(X,K))
X = torch.tensor([[[0.0, 1.0, 2.0], [3.0, 4.0, 5.0], [6.0, 7.0, 8.0]],
[[1.0, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 9.0]]])
K = torch.tensor([[[0.0, 1.0], [2.0, 3.0]], [[1.0, 2.0], [3.0, 4.0]]])
corr2d_multi_in(X, K)
Out [43]:
tensor([[ 56.,  72.],
        [104., 120.]])
In [44]:
def corr2d_multi_in_out(X,K) ->torch.Tensor :
    return torch.stack([corr2d_multi_in(X,k) for k in K],0)
In [45]:
K = torch.stack((K, K + 1, K + 2), 0)
K.shape
Out [45]:
torch.Size([3, 2, 2, 2])
In [46]:
corr2d_multi_in_out(X, K)
Out [46]:
tensor([[[ 56.,  72.],
         [104., 120.]],

        [[ 76., 100.],
         [148., 172.]],

        [[ 96., 128.],
         [192., 224.]]])
In [47]:
def corr2d_multi_in_out_1x1(X, K):
    h_i,h,w=X.shape
    h_o=K.shape[0]
    X=X.reshape((h_i,h*w))
    print(X.shape)
    K=K.reshape((h_o,h_i))
    print(K.shape)
    Y=torch.matmul(K,X)
    return Y.reshape((h_o,h,w))
In [48]:
X = torch.normal(0, 1, (3, 3, 3))
K = torch.normal(0, 1, (2, 3, 1, 1))
In [49]:
Y1 = corr2d_multi_in_out_1x1(X, K)
torch.Size([3, 9])
torch.Size([2, 3])
In [50]:
Y2 = corr2d_multi_in_out(X, K)
assert float(torch.abs(Y1 - Y2).sum()) < 1e-6
In [51]:
def pool2d(X,pool_size,mode='max'):
    p_h,p_w =pool_size
    Y = torch.zeros((X.shape[0]-p_h+1,X.shape[1]-p_w+1))
    for i in range(Y.shape[0]):
        for j in range(Y.shape[1]):
            match mode:
                case 'max':
                    Y[i,j]=X[i:i+p_h,j:j+p_w].max()
                case 'avg':
                    Y[i,j]=X[i:i+p_h,j:j+p_w].mean()

    return Y
In [52]:
X = torch.tensor([[0.0, 1.0, 2.0], [3.0, 4.0, 5.0], [6.0, 7.0, 8.0]])
pool2d(X, (2, 2))
Out [52]:
tensor([[4., 5.],
        [7., 8.]])
In [53]:
pool2d(X, (2, 2), 'avg')
Out [53]:
tensor([[2., 3.],
        [5., 6.]])
In [54]:
X = torch.arange(16, dtype=torch.float32).reshape((1, 1, 4, 4))
X
Out [54]:
tensor([[[[ 0.,  1.,  2.,  3.],
          [ 4.,  5.,  6.,  7.],
          [ 8.,  9., 10., 11.],
          [12., 13., 14., 15.]]]])
In [55]:
pool2d=nn.MaxPool2d(3)
pool2d(X)
Out [55]:
tensor([[[[10.]]]])
In [56]:
pool2d = nn.MaxPool2d(3, padding=1, stride=2)
pool2d(X)
Out [56]:
tensor([[[[ 5.,  7.],
          [13., 15.]]]])
In [57]:
pool2d = nn.MaxPool2d((2, 3), stride=(2, 3), padding=(0, 1))
pool2d(X)
Out [57]:
tensor([[[[ 5.,  7.],
          [13., 15.]]]])
In [58]:
X = torch.cat((X, X + 1), 1)
X
Out [58]:
tensor([[[[ 0.,  1.,  2.,  3.],
          [ 4.,  5.,  6.,  7.],
          [ 8.,  9., 10., 11.],
          [12., 13., 14., 15.]],

         [[ 1.,  2.,  3.,  4.],
          [ 5.,  6.,  7.,  8.],
          [ 9., 10., 11., 12.],
          [13., 14., 15., 16.]]]])
In [59]:
pool2d = nn.MaxPool2d(3, padding=1, stride=2)
pool2d(X)
Out [59]:
tensor([[[[ 5.,  7.],
          [13., 15.]],

         [[ 6.,  8.],
          [14., 16.]]]])
In [60]:
net = nn.Sequential(
    nn.Conv2d(1,6,kernel_size=5,padding=2), #1*1*28*28 -> 1*6*28*28
    nn.Sigmoid(),
    nn.AvgPool2d(kernel_size=2, stride=2),  #1*6*28*28 -> 1*6*14*14
    nn.Conv2d(6, 16, kernel_size=5), nn.Sigmoid(), #1*6*14*14 -> 1*16*10*10
    nn.AvgPool2d(kernel_size=2, stride=2), #1*16*10*10 -> 1*16*5*5
    nn.Flatten(),
    nn.Linear(16 * 5 * 5, 120), nn.Sigmoid(),
    nn.Linear(120, 84), nn.Sigmoid(),
    nn.Linear(84, 10)
)
X = torch.rand(size=(1,1,28,28),dtype=torch.float32)
for layer in net:
    X=layer(X)
    print(layer.__class__.__name__,'output shape: \t',X.shape)
Conv2d output shape: 	 torch.Size([1, 6, 28, 28])
Sigmoid output shape: 	 torch.Size([1, 6, 28, 28])
AvgPool2d output shape: 	 torch.Size([1, 6, 14, 14])
Conv2d output shape: 	 torch.Size([1, 16, 10, 10])
Sigmoid output shape: 	 torch.Size([1, 16, 10, 10])
AvgPool2d output shape: 	 torch.Size([1, 16, 5, 5])
Flatten output shape: 	 torch.Size([1, 400])
Linear output shape: 	 torch.Size([1, 120])
Sigmoid output shape: 	 torch.Size([1, 120])
Linear output shape: 	 torch.Size([1, 84])
Sigmoid output shape: 	 torch.Size([1, 84])
Linear output shape: 	 torch.Size([1, 10])
In [61]:
import d2l.torch as d2l
batch_size = 256
train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size=batch_size)
In [62]:
lr, num_epochs = 0.9, 10
#d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())
In [63]:
class Inception(nn.Module):
    def __init__(self,in_channels,c1,c2,c3,c4,**kwargs):
        super(Inception,self).__init__(**kwargs)
        self.p1_1 = nn.Conv2d(in_channels,c1,kernel_size=1)
        self.p2_1 = nn.Conv2d(in_channels,c2[0],kernel_size=1)
        self.p2_2 = nn.Conv2d(c2[0],c2[1],kernel_size=3,padding=1)
        self.p3_1 = nn.Conv2d(in_channels,c3[0],kernel_size=1)
        self.p3_2 = nn.Conv2d(c3[0],c3[1],kernel_size=5,padding=2)
        self.p4_1 = nn.MaxPool2d(kernel_size=3, stride=1, padding=1)
        self.p4_2 = nn.Conv2d(in_channels, c4, kernel_size=1)
    def forward(self,x):
        p1 = F.relu(self.p1_1(x))
        p2 = F.relu(self.p2_2(F.relu(self.p2_1(x))))
        p3 = F.relu(self.p3_2(F.relu(self.p3_1(x))))
        p4 = F.relu(self.p4_2(self.p4_1(x)))
        return torch.cat((p1,p2,p3,p4),dim=1)
In [64]:
b1 = nn.Sequential(nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3),
                    nn.ReLU(),
                    nn.MaxPool2d(kernel_size=3, stride=2, padding=1))
b2 = nn.Sequential(nn.Conv2d(64, 64, kernel_size=1),
                    nn.ReLU(),
                    nn.Conv2d(64, 192, kernel_size=3, padding=1),
                    nn.ReLU(),
                    nn.MaxPool2d(kernel_size=3, stride=2, padding=1))
b3 = nn.Sequential(Inception(192, 64, (96, 128), (16, 32), 32),
                    Inception(256, 128, (128, 192), (32, 96), 64),
                    nn.MaxPool2d(kernel_size=3, stride=2, padding=1))
b4 = nn.Sequential(Inception(480, 192, (96, 208), (16, 48), 64),
                    Inception(512, 160, (112, 224), (24, 64), 64),
                    Inception(512, 128, (128, 256), (24, 64), 64),
                    Inception(512, 112, (144, 288), (32, 64), 64),
                    Inception(528, 256, (160, 320), (32, 128), 128),
                    nn.MaxPool2d(kernel_size=3, stride=2, padding=1))
b5 = nn.Sequential(Inception(832, 256, (160, 320), (32, 128), 128),
                    Inception(832, 384, (192, 384), (48, 128), 128),
                    nn.AdaptiveAvgPool2d((1,1)),
                    nn.Flatten())
net = nn.Sequential(b1, b2, b3, b4, b5, nn.Linear(1024, 10))
X = torch.rand(size=(1, 1, 96, 96))
for layer in net:
    X = layer(X)
    print(layer.__class__.__name__,'output shape:\t', X.shape)
Sequential output shape:	 torch.Size([1, 64, 24, 24])
Sequential output shape:	 torch.Size([1, 192, 12, 12])
Sequential output shape:	 torch.Size([1, 480, 6, 6])
Sequential output shape:	 torch.Size([1, 832, 3, 3])
Sequential output shape:	 torch.Size([1, 1024])
Linear output shape:	 torch.Size([1, 10])
In [65]:
import torchinfo
torchinfo.summary(net,(1,1,96,96))
Out [65]:
==========================================================================================
Layer (type:depth-idx)                   Output Shape              Param #
==========================================================================================
Sequential                               [1, 10]                   --
├─Sequential: 1-1                        [1, 64, 24, 24]           --
│    └─Conv2d: 2-1                       [1, 64, 48, 48]           3,200
│    └─ReLU: 2-2                         [1, 64, 48, 48]           --
│    └─MaxPool2d: 2-3                    [1, 64, 24, 24]           --
├─Sequential: 1-2                        [1, 192, 12, 12]          --
│    └─Conv2d: 2-4                       [1, 64, 24, 24]           4,160
│    └─ReLU: 2-5                         [1, 64, 24, 24]           --
│    └─Conv2d: 2-6                       [1, 192, 24, 24]          110,784
│    └─ReLU: 2-7                         [1, 192, 24, 24]          --
│    └─MaxPool2d: 2-8                    [1, 192, 12, 12]          --
├─Sequential: 1-3                        [1, 480, 6, 6]            --
│    └─Inception: 2-9                    [1, 256, 12, 12]          --
│    │    └─Conv2d: 3-1                  [1, 64, 12, 12]           12,352
│    │    └─Conv2d: 3-2                  [1, 96, 12, 12]           18,528
│    │    └─Conv2d: 3-3                  [1, 128, 12, 12]          110,720
│    │    └─Conv2d: 3-4                  [1, 16, 12, 12]           3,088
│    │    └─Conv2d: 3-5                  [1, 32, 12, 12]           12,832
│    │    └─MaxPool2d: 3-6               [1, 192, 12, 12]          --
│    │    └─Conv2d: 3-7                  [1, 32, 12, 12]           6,176
│    └─Inception: 2-10                   [1, 480, 12, 12]          --
│    │    └─Conv2d: 3-8                  [1, 128, 12, 12]          32,896
│    │    └─Conv2d: 3-9                  [1, 128, 12, 12]          32,896
│    │    └─Conv2d: 3-10                 [1, 192, 12, 12]          221,376
│    │    └─Conv2d: 3-11                 [1, 32, 12, 12]           8,224
│    │    └─Conv2d: 3-12                 [1, 96, 12, 12]           76,896
│    │    └─MaxPool2d: 3-13              [1, 256, 12, 12]          --
│    │    └─Conv2d: 3-14                 [1, 64, 12, 12]           16,448
│    └─MaxPool2d: 2-11                   [1, 480, 6, 6]            --
├─Sequential: 1-4                        [1, 832, 3, 3]            --
│    └─Inception: 2-12                   [1, 512, 6, 6]            --
│    │    └─Conv2d: 3-15                 [1, 192, 6, 6]            92,352
│    │    └─Conv2d: 3-16                 [1, 96, 6, 6]             46,176
│    │    └─Conv2d: 3-17                 [1, 208, 6, 6]            179,920
│    │    └─Conv2d: 3-18                 [1, 16, 6, 6]             7,696
│    │    └─Conv2d: 3-19                 [1, 48, 6, 6]             19,248
│    │    └─MaxPool2d: 3-20              [1, 480, 6, 6]            --
│    │    └─Conv2d: 3-21                 [1, 64, 6, 6]             30,784
│    └─Inception: 2-13                   [1, 512, 6, 6]            --
│    │    └─Conv2d: 3-22                 [1, 160, 6, 6]            82,080
│    │    └─Conv2d: 3-23                 [1, 112, 6, 6]            57,456
│    │    └─Conv2d: 3-24                 [1, 224, 6, 6]            226,016
│    │    └─Conv2d: 3-25                 [1, 24, 6, 6]             12,312
│    │    └─Conv2d: 3-26                 [1, 64, 6, 6]             38,464
│    │    └─MaxPool2d: 3-27              [1, 512, 6, 6]            --
│    │    └─Conv2d: 3-28                 [1, 64, 6, 6]             32,832
│    └─Inception: 2-14                   [1, 512, 6, 6]            --
│    │    └─Conv2d: 3-29                 [1, 128, 6, 6]            65,664
│    │    └─Conv2d: 3-30                 [1, 128, 6, 6]            65,664
│    │    └─Conv2d: 3-31                 [1, 256, 6, 6]            295,168
│    │    └─Conv2d: 3-32                 [1, 24, 6, 6]             12,312
│    │    └─Conv2d: 3-33                 [1, 64, 6, 6]             38,464
│    │    └─MaxPool2d: 3-34              [1, 512, 6, 6]            --
│    │    └─Conv2d: 3-35                 [1, 64, 6, 6]             32,832
│    └─Inception: 2-15                   [1, 528, 6, 6]            --
│    │    └─Conv2d: 3-36                 [1, 112, 6, 6]            57,456
│    │    └─Conv2d: 3-37                 [1, 144, 6, 6]            73,872
│    │    └─Conv2d: 3-38                 [1, 288, 6, 6]            373,536
│    │    └─Conv2d: 3-39                 [1, 32, 6, 6]             16,416
│    │    └─Conv2d: 3-40                 [1, 64, 6, 6]             51,264
│    │    └─MaxPool2d: 3-41              [1, 512, 6, 6]            --
│    │    └─Conv2d: 3-42                 [1, 64, 6, 6]             32,832
│    └─Inception: 2-16                   [1, 832, 6, 6]            --
│    │    └─Conv2d: 3-43                 [1, 256, 6, 6]            135,424
│    │    └─Conv2d: 3-44                 [1, 160, 6, 6]            84,640
│    │    └─Conv2d: 3-45                 [1, 320, 6, 6]            461,120
│    │    └─Conv2d: 3-46                 [1, 32, 6, 6]             16,928
│    │    └─Conv2d: 3-47                 [1, 128, 6, 6]            102,528
│    │    └─MaxPool2d: 3-48              [1, 528, 6, 6]            --
│    │    └─Conv2d: 3-49                 [1, 128, 6, 6]            67,712
│    └─MaxPool2d: 2-17                   [1, 832, 3, 3]            --
├─Sequential: 1-5                        [1, 1024]                 --
│    └─Inception: 2-18                   [1, 832, 3, 3]            --
│    │    └─Conv2d: 3-50                 [1, 256, 3, 3]            213,248
│    │    └─Conv2d: 3-51                 [1, 160, 3, 3]            133,280
│    │    └─Conv2d: 3-52                 [1, 320, 3, 3]            461,120
│    │    └─Conv2d: 3-53                 [1, 32, 3, 3]             26,656
│    │    └─Conv2d: 3-54                 [1, 128, 3, 3]            102,528
│    │    └─MaxPool2d: 3-55              [1, 832, 3, 3]            --
│    │    └─Conv2d: 3-56                 [1, 128, 3, 3]            106,624
│    └─Inception: 2-19                   [1, 1024, 3, 3]           --
│    │    └─Conv2d: 3-57                 [1, 384, 3, 3]            319,872
│    │    └─Conv2d: 3-58                 [1, 192, 3, 3]            159,936
│    │    └─Conv2d: 3-59                 [1, 384, 3, 3]            663,936
│    │    └─Conv2d: 3-60                 [1, 48, 3, 3]             39,984
│    │    └─Conv2d: 3-61                 [1, 128, 3, 3]            153,728
│    │    └─MaxPool2d: 3-62              [1, 832, 3, 3]            --
│    │    └─Conv2d: 3-63                 [1, 128, 3, 3]            106,624
│    └─AdaptiveAvgPool2d: 2-20           [1, 1024, 1, 1]           --
│    └─Flatten: 2-21                     [1, 1024]                 --
├─Linear: 1-6                            [1, 10]                   10,250
==========================================================================================
Total params: 5,977,530
Trainable params: 5,977,530
Non-trainable params: 0
Total mult-adds (Units.MEGABYTES): 276.66
==========================================================================================
Input size (MB): 0.04
Forward/backward pass size (MB): 4.74
Params size (MB): 23.91
Estimated Total Size (MB): 28.69
==========================================================================================
In [66]:
lr, num_epochs, batch_size = 0.1, 10, 128
train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=96)
#d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())
In [67]:
class Residual(nn.Module):
    def __init__(self,input_channels,num_channels,use_1x1conv=False,strides=1):
        super().__init__()
        self.conv1 = nn.Conv2d(input_channels,num_channels,kernel_size=3,padding=1,stride=strides)
        self.conv2 = nn.Conv2d(num_channels,num_channels,kernel_size=3,padding=1)
        if use_1x1conv:
            self.conv3 = nn.Conv2d(input_channels,num_channels,kernel_size=1,stride=strides)
        else:
            self.conv3= None
        self.bn1=nn.BatchNorm2d(num_channels)
        self.bn2=nn.BatchNorm2d(num_channels)
    def forward(self,X):
        Y=F.relu(self.bn1(self.conv1(X)))
        Y=self.bn2(self.conv2(Y))
        if self.conv3:
            X = self.conv3(X)
        Y+=X
        return F.relu(Y)
In [68]:
blk = Residual(3,3)
X = torch.rand(4, 3, 6, 6)
In [69]:
blk = Residual(3,6, use_1x1conv=True, strides=2)
blk(X).shape
Out [69]:
torch.Size([4, 6, 3, 3])
In [70]:
b1 = nn.Sequential(nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3),
nn.BatchNorm2d(64), nn.ReLU(),
nn.MaxPool2d(kernel_size=3, stride=2, padding=1))
In [71]:
def resnet_block(input_channels, num_channels, num_residuals,
                first_block=False):
    blk = []
    for i in range(num_residuals):
        if i == 0 and not first_block:
            blk.append(Residual(input_channels, num_channels,
                        use_1x1conv=True, strides=2))
        else:
            blk.append(Residual(num_channels, num_channels))
    return blk
In [72]:
b2 = nn.Sequential(*resnet_block(64, 64, 2, first_block=True))
b3 = nn.Sequential(*resnet_block(64, 128, 2))
b4 = nn.Sequential(*resnet_block(128, 256, 2))
b5 = nn.Sequential(*resnet_block(256, 512, 2))
In [73]:
net = nn.Sequential(b1, b2, b3, b4, b5,
nn.AdaptiveAvgPool2d((1,1)),
nn.Flatten(), nn.Linear(512, 10))
In [74]:
X = torch.rand(size=(1, 1, 224, 224))
for layer in net:
    X = layer(X)
    print(layer.__class__.__name__,'output shape:\t', X.shape)
Sequential output shape:	 torch.Size([1, 64, 56, 56])
Sequential output shape:	 torch.Size([1, 64, 56, 56])
Sequential output shape:	 torch.Size([1, 128, 28, 28])
Sequential output shape:	 torch.Size([1, 256, 14, 14])
Sequential output shape:	 torch.Size([1, 512, 7, 7])
AdaptiveAvgPool2d output shape:	 torch.Size([1, 512, 1, 1])
Flatten output shape:	 torch.Size([1, 512])
Linear output shape:	 torch.Size([1, 10])
In [75]:
lr, num_epochs, batch_size = 0.05, 10, 256
train_iter, test_iter = d2l.load_data_fashion_mnist(batch_size, resize=96)
#d2l.train_ch6(net, train_iter, test_iter, num_epochs, lr, d2l.try_gpu())
In [76]:
import torch
import d2l.torch as d2l
import numpy
import torch.nn as nn
import torch.nn.functional as F
print(torch.version.__version__)
2.10.0+cu128
In [77]:
A=torch.Tensor([[1,2,0,0],[0,2,0,0],[0,0,2,1],[0,0,0,3]])
C=torch.Tensor([[1,0,0,0],[0,1,0,0],[0,0,-2,3],[0,0,0,-3]])
In [78]:
B=torch.Tensor([[2,0,0,0],[-2,1,0,0],[0,0,-3,0],[0,0,0,-3]])
In [79]:
torch.mm(A,C)
Out [79]:
tensor([[ 1.,  2.,  0.,  0.],
        [ 0.,  2.,  0.,  0.],
        [ 0.,  0., -4.,  3.],
        [ 0.,  0.,  0., -9.]])
In [80]:
torch.det(torch.mm(torch.mm(A,C),B))
Out [80]:
tensor(1296.)
In [81]:
1296**5
Out [81]:
3656158440062976
In [82]:
torch.mm(C,B)
Out [82]:
tensor([[ 2.,  0.,  0.,  0.],
        [-2.,  1.,  0.,  0.],
        [ 0.,  0.,  6., -9.],
        [ 0.,  0.,  0.,  9.]])
In [83]:
T = 1000 # 总共产生1000个点
time = torch.arange(1, T + 1, dtype=torch.float32)
x = torch.sin(0.01 * time) + torch.normal(0, 0.2, (T,))
d2l.plot(time, [x], 'time', 'x', xlim=[1, 1000], figsize=(6, 3))
In [84]:
tau = 4
features = torch.zeros((T - tau, tau))
for i in range(tau):
    features[:, i] = x[i: T - tau + i]
labels = x[tau:].reshape((-1, 1))
x,features,labels
Out [84]:
(tensor([ 6.4438e-02, -2.8849e-02,  2.2463e-01, -1.7219e-01,  6.6353e-02,
         -1.0363e-01,  8.8868e-02,  1.5603e-01,  2.0316e-01,  2.2209e-01,
          4.0233e-01,  1.8191e-01,  9.4271e-02,  1.9833e-01,  2.7132e-01,
         -2.6344e-02,  1.3314e-01, -9.5498e-02,  4.2949e-01,  2.9735e-01,
          2.6208e-01,  2.5798e-01,  2.6224e-01,  4.0028e-01,  1.6453e-01,
         -3.6497e-03,  4.5941e-02,  2.9152e-01,  2.8247e-01,  3.1230e-01,
          4.0011e-01,  1.4096e-01,  4.1744e-01,  3.4225e-01, -4.5256e-02,
          2.8870e-01,  3.8852e-01,  3.4837e-01,  6.1889e-01,  6.2549e-01,
          2.3834e-01,  4.8642e-01,  3.5614e-01,  1.1784e-01,  3.8346e-01,
          4.5669e-01,  3.6588e-01,  2.6488e-01,  6.0995e-01,  6.9697e-01,
          7.5780e-01,  5.8101e-01,  3.5400e-01,  2.4635e-01,  4.7288e-01,
          6.6484e-01,  6.3196e-01,  5.6758e-01,  3.1575e-01,  7.3676e-01,
          8.1908e-01,  8.8408e-01,  7.6086e-01,  4.3549e-01,  7.9157e-01,
          4.1029e-01,  3.4122e-01,  1.0624e+00,  9.8399e-01,  7.3473e-01,
          6.9833e-01,  3.4119e-01,  7.3251e-01,  6.7880e-01,  6.3626e-01,
          1.0105e+00,  7.0007e-01,  1.1702e+00,  6.0600e-01,  8.9456e-01,
          5.1218e-01,  7.3733e-01,  6.1851e-01,  7.6468e-01,  6.5189e-01,
          1.0688e+00,  1.0419e+00,  9.7937e-01,  1.0725e+00,  5.4258e-01,
          9.2915e-01,  4.3405e-01,  4.7934e-01,  1.1528e+00,  7.5340e-01,
          5.4904e-01,  5.4025e-01,  5.1751e-01,  2.9075e-01,  5.5143e-01,
          8.6160e-01,  9.6728e-01,  6.0795e-01,  7.0219e-01,  1.0551e+00,
          7.9270e-01,  9.2103e-01,  8.7458e-01,  9.9153e-01,  6.0989e-01,
          7.4993e-01,  6.9077e-01,  5.6804e-01,  7.0561e-01,  7.8830e-01,
          9.7916e-01,  9.9039e-01,  7.9061e-01,  9.6164e-01,  8.5340e-01,
          8.2899e-01,  7.3213e-01,  6.8678e-01,  1.2765e+00,  1.2545e+00,
          1.1249e+00,  1.3865e+00,  8.9114e-01,  8.0419e-01,  1.2773e+00,
          8.9564e-01,  6.2510e-01,  1.1143e+00,  9.5270e-01,  9.5466e-01,
          7.9755e-01,  1.0294e+00,  5.8184e-01,  1.2175e+00,  1.0392e+00,
          9.4017e-01,  1.1067e+00,  9.4888e-01,  8.5048e-01,  9.3845e-01,
          1.2021e+00,  9.6893e-01,  1.0378e+00,  1.1524e+00,  1.0356e+00,
          1.2582e+00,  9.6289e-01,  1.1062e+00,  6.3397e-01,  8.6299e-01,
          8.4336e-01,  6.3310e-01,  1.3410e+00,  9.9408e-01,  7.2785e-01,
          1.2686e+00,  8.6769e-01,  1.2330e+00,  1.0719e+00,  7.1942e-01,
          1.1396e+00,  9.7066e-01,  1.1835e+00,  9.9540e-01,  7.0843e-01,
          1.0240e+00,  9.7269e-01,  8.0341e-01,  1.0116e+00,  8.1968e-01,
          7.5934e-01,  9.6155e-01,  1.1158e+00,  9.6646e-01,  7.0824e-01,
          1.2226e+00,  7.9395e-01,  9.5238e-01,  1.1413e+00,  8.7932e-01,
          9.0194e-01,  8.6707e-01,  1.2885e+00,  1.2093e+00,  1.2944e+00,
          6.3165e-01,  5.5781e-01,  1.0286e+00,  1.3556e+00,  8.9051e-01,
          1.1046e+00,  9.9294e-01,  1.2333e+00,  8.1726e-01,  3.1740e-01,
          7.6761e-01,  6.5092e-01,  8.9621e-01,  6.8024e-01,  1.0196e+00,
          7.5904e-01,  9.4622e-01,  8.6102e-01,  8.3147e-01,  7.4786e-01,
          9.6059e-01,  8.6270e-01,  9.4766e-01,  5.2817e-01,  1.0704e+00,
          9.7133e-01,  8.4503e-01,  8.8452e-01,  5.0776e-01,  9.2430e-01,
          5.9001e-01,  8.0198e-01,  7.8422e-01,  5.8749e-01,  7.0924e-01,
          7.5607e-01,  4.9436e-01,  7.7539e-01,  6.4349e-01,  7.5043e-01,
          7.3827e-01,  8.6847e-01,  5.4753e-01,  6.5105e-01,  1.0554e+00,
          7.8901e-01,  8.9882e-01,  7.0064e-01,  5.4479e-01,  8.2511e-01,
          5.1943e-01,  1.5267e-01,  7.0765e-01,  5.7810e-01,  6.0173e-01,
          4.8342e-01,  6.4010e-01,  9.0313e-01,  3.0786e-01,  1.0283e+00,
          2.3870e-01,  5.7824e-01, -3.3643e-02,  4.2503e-01,  5.7349e-01,
          4.4148e-01,  6.8640e-01,  6.1931e-01,  2.5912e-01,  2.2371e-01,
          6.0460e-01,  3.7744e-01,  5.9038e-01,  3.7926e-01,  5.2749e-01,
          6.2748e-01,  6.6149e-01,  4.3518e-01,  4.0026e-01,  2.9409e-01,
          4.5821e-01,  3.7015e-01,  3.4187e-01,  1.8859e-01,  6.9215e-01,
          3.2195e-01,  2.4332e-02,  5.2798e-01,  9.0723e-02,  1.6245e-01,
          2.7128e-01,  2.0240e-01, -8.3513e-02,  3.9523e-01,  5.8745e-01,
          3.2908e-01,  2.4919e-01,  2.8691e-01,  1.1735e-01,  4.3031e-01,
          2.6840e-01,  2.5892e-01,  1.7928e-01,  4.2978e-01,  5.2001e-02,
          2.7463e-01, -1.3417e-01, -1.6025e-02,  2.9625e-01,  2.3450e-01,
          3.7070e-01,  7.5755e-02,  1.6683e-01,  1.1036e-01, -6.8264e-03,
          7.2137e-03,  2.7841e-01,  9.1316e-02,  9.1231e-03, -3.5385e-01,
          1.1431e-01, -2.0163e-01,  2.0756e-01, -9.4054e-02, -1.2446e-01,
          2.4384e-01, -2.3242e-01, -2.0931e-01, -1.8707e-01, -1.3957e-01,
         -1.8903e-01, -8.1507e-02, -3.4759e-01, -1.1257e-02, -1.9703e-01,
          2.7359e-02,  7.6564e-04, -6.1846e-01, -7.9818e-02, -3.6661e-01,
          6.5931e-02, -5.8843e-01, -1.7423e-01, -2.1431e-01,  6.6695e-02,
         -2.9555e-01, -1.3760e-01, -5.6146e-02, -1.7448e-02, -3.2177e-01,
         -1.8931e-01, -3.3209e-01, -2.6944e-01, -3.6146e-01, -3.5334e-01,
         -3.3019e-01, -2.8488e-02, -3.6981e-01, -3.1455e-01, -4.2320e-01,
         -5.8333e-01, -4.5083e-01, -2.6372e-01, -6.6177e-01, -5.4376e-01,
         -2.1988e-01, -7.4067e-02, -3.6120e-01, -7.7958e-01, -2.7244e-01,
         -3.3669e-01, -6.1547e-01, -7.1691e-01, -3.2713e-01, -3.2994e-01,
         -4.0011e-01, -4.5194e-01, -5.5936e-01, -4.8557e-01, -7.0421e-01,
         -1.8149e-01, -6.7299e-01, -4.4816e-01, -5.8107e-01, -4.7465e-01,
         -4.2599e-01, -8.3749e-01, -6.6348e-01, -6.9997e-01, -5.8357e-01,
         -3.5789e-01, -7.1656e-01, -8.7729e-01, -6.3883e-01, -6.9591e-01,
         -9.3954e-01, -4.9190e-01, -7.3111e-01, -4.3942e-01, -6.2770e-01,
         -9.2674e-01, -8.6653e-01, -9.6315e-01, -6.9102e-01, -5.9326e-01,
         -8.5505e-01, -8.5113e-01, -6.2499e-01, -9.9391e-01, -8.4853e-01,
         -7.8337e-01, -6.2505e-01, -8.0748e-01, -8.2683e-01, -6.9701e-01,
         -7.8696e-01, -6.9023e-01, -6.2324e-01, -9.8813e-01, -8.8023e-01,
         -7.4747e-01, -9.1390e-01, -1.1208e+00, -1.3740e+00, -1.0556e+00,
         -9.5917e-01, -7.6300e-01, -1.0235e+00, -1.0120e+00, -7.2330e-01,
         -1.0387e+00, -5.4913e-01, -5.8775e-01, -9.5260e-01, -7.9546e-01,
         -6.3363e-01, -7.6522e-01, -1.0495e+00, -1.1376e+00, -9.7715e-01,
         -1.0136e+00, -1.0233e+00, -4.2618e-01, -9.1957e-01, -1.1356e+00,
         -9.9493e-01, -6.0729e-01, -6.6857e-01, -8.6124e-01, -6.8043e-01,
         -7.9954e-01, -9.4514e-01, -9.6364e-01, -7.6007e-01, -7.9190e-01,
         -9.0256e-01, -1.0012e+00, -8.1351e-01, -1.0263e+00, -7.2729e-01,
         -8.8496e-01, -1.2832e+00, -8.1455e-01, -7.6072e-01, -9.7935e-01,
         -1.0354e+00, -1.1628e+00, -8.2264e-01, -7.7549e-01, -1.6201e+00,
         -6.8463e-01, -1.0060e+00, -8.6768e-01, -7.4937e-01, -1.0767e+00,
         -5.6351e-01, -9.3733e-01, -1.2270e+00, -9.3480e-01, -1.1192e+00,
         -1.2890e+00, -1.3016e+00, -9.8122e-01, -1.2527e+00, -8.9113e-01,
         -9.9467e-01, -7.3336e-01, -1.2790e+00, -1.2139e+00, -7.6957e-01,
         -8.9945e-01, -1.2749e+00, -7.1642e-01, -1.0271e+00, -1.3689e+00,
         -8.8271e-01, -9.3780e-01, -1.0709e+00, -1.0079e+00, -1.2095e+00,
         -8.3435e-01, -1.1892e+00, -5.8446e-01, -1.0576e+00, -7.8082e-01,
         -9.9774e-01, -1.0047e+00, -9.4661e-01, -7.9260e-01, -7.8298e-01,
         -8.1630e-01, -1.1429e+00, -9.0614e-01, -1.2286e+00, -1.0185e+00,
         -9.2398e-01, -9.3490e-01, -1.1074e+00, -7.6938e-01, -7.7835e-01,
         -7.9201e-01, -8.3866e-01, -5.0138e-01, -1.0518e+00, -1.1464e+00,
         -8.3545e-01, -6.3239e-01, -8.6411e-01, -1.0649e+00, -8.3904e-01,
         -9.3103e-01, -9.5688e-01, -1.3042e+00, -7.8724e-01, -8.9785e-01,
         -5.8319e-01, -9.7922e-01, -9.8292e-01, -1.0255e+00, -6.4694e-01,
         -8.0609e-01, -6.8586e-01, -1.0256e+00, -6.2613e-01, -5.2035e-01,
         -8.2406e-01, -6.1214e-01, -6.3858e-01, -7.9211e-01, -8.4110e-01,
         -8.7759e-01, -1.0926e+00, -4.8413e-01, -8.8961e-01, -8.6125e-01,
         -8.4024e-01, -7.4395e-01, -7.6605e-01, -7.7586e-01, -6.6531e-01,
         -8.7798e-01, -5.3314e-01, -3.8761e-01, -5.2371e-01, -4.9831e-01,
         -4.7124e-01, -4.0311e-01, -3.9151e-01, -6.0217e-01, -3.5831e-01,
         -8.1952e-01, -3.7521e-01, -4.1182e-01, -6.9520e-01, -1.3176e-01,
         -3.5725e-01, -6.3746e-01, -7.1734e-01, -6.0116e-01, -4.7620e-01,
         -1.3156e-01, -6.7144e-01, -4.4765e-01, -7.9655e-01, -4.6641e-01,
         -5.1395e-01, -5.9736e-01, -7.0441e-02, -2.9234e-01, -6.3963e-01,
         -5.2695e-01, -8.9920e-01, -3.9060e-01, -3.3070e-01, -5.8975e-02,
         -3.4768e-01, -3.9728e-01, -5.0462e-01, -6.9052e-01, -2.6284e-01,
         -5.3189e-01, -2.8471e-01, -3.8808e-01, -2.5389e-01, -1.0635e-01,
         -4.4742e-01, -2.8809e-01, -4.6124e-01, -1.8804e-01, -6.5422e-01,
         -2.9021e-01, -1.6320e-01, -2.7098e-01, -1.8750e-01,  2.4683e-01,
         -1.2878e-01, -2.1855e-01, -5.1811e-01,  5.4305e-02, -1.7425e-01,
         -2.6757e-01,  1.3357e-01, -3.1198e-01,  8.1655e-02, -2.8527e-01,
         -1.3569e-01, -1.4587e-01,  1.7095e-01, -2.3103e-02,  2.1838e-01,
         -1.6752e-01,  3.1579e-01,  2.7031e-01, -2.4856e-01,  7.6009e-03,
         -1.1322e-03, -2.0730e-01, -1.2818e-01,  1.4944e-01,  9.0087e-02,
          4.0082e-01,  2.9144e-01, -1.4654e-01,  8.8202e-02, -1.7362e-01,
         -9.1909e-03, -8.0324e-02, -5.5343e-02,  5.8975e-01,  1.6623e-01,
          3.3504e-01,  2.4773e-02,  8.7046e-02, -1.6163e-01,  5.1961e-01,
          1.7512e-01,  1.0362e-01,  2.0362e-01,  1.9839e-01,  4.1845e-01,
          4.6793e-01, -1.1853e-01,  1.2487e-01,  1.9344e-01,  3.0220e-01,
          8.3792e-02,  3.1021e-01,  3.2869e-01,  3.0734e-01,  5.7626e-01,
          4.3928e-01,  3.1897e-01,  2.8437e-01,  5.5234e-01,  6.2213e-01,
          6.1580e-01,  3.9967e-01,  4.5823e-01,  4.3247e-01,  5.0114e-01,
          8.3448e-01,  4.9888e-01,  5.0631e-01,  2.0848e-01,  3.6072e-01,
          2.7618e-01,  4.0099e-01,  5.4027e-01,  2.4210e-01,  1.2701e-01,
          4.4325e-01,  3.0193e-01,  3.6690e-01,  5.7623e-01,  5.2195e-01,
          6.5280e-01,  5.7883e-01,  2.9837e-01,  2.5124e-01,  3.4579e-01,
          2.2099e-01,  3.4217e-01,  8.5317e-01,  7.0991e-01,  3.0030e-01,
          7.5253e-01,  7.0718e-01,  7.5546e-01,  8.3272e-01,  8.2167e-01,
          6.8525e-01,  8.5421e-01,  3.8577e-01,  6.1654e-01,  6.7905e-01,
          9.9523e-01,  7.6051e-01,  8.6416e-01,  6.0249e-01,  1.2840e+00,
          6.4849e-01,  5.6504e-01,  6.7845e-01,  3.4798e-01,  6.4645e-01,
          7.8018e-01,  9.8716e-01,  7.3607e-01,  7.5667e-01,  9.4265e-01,
          8.0938e-01,  2.6675e-01,  6.1355e-01,  9.0162e-01,  4.2799e-01,
          8.3804e-01,  4.6611e-01,  8.9841e-01,  9.1118e-01,  6.0615e-01,
          8.8064e-01,  1.1570e+00,  9.5360e-01,  6.1671e-01,  7.3730e-01,
          1.1604e+00,  7.5139e-01,  1.0676e+00,  8.5293e-01,  1.2444e+00,
          7.3487e-01,  7.2214e-01,  1.3246e+00,  9.9059e-01,  1.1876e+00,
          8.9527e-01,  1.1861e+00,  1.1246e+00,  1.4331e+00,  1.0912e+00,
          1.0825e+00,  1.1025e+00,  9.9692e-01,  7.3992e-01,  6.6441e-01,
          9.1873e-01,  9.2798e-01,  1.3192e+00,  1.0102e+00,  1.3542e+00,
          1.0066e+00,  1.0593e+00,  1.0544e+00,  1.0728e+00,  8.1399e-01,
          1.2641e+00,  1.0079e+00,  1.0601e+00,  1.0796e+00,  1.0223e+00,
          1.1092e+00,  7.7240e-01,  1.3857e+00,  9.1455e-01,  8.6129e-01,
          8.0728e-01,  6.1754e-01,  1.0502e+00,  9.3687e-01,  9.5378e-01,
          9.7054e-01,  1.2540e+00,  1.1463e+00,  1.1405e+00,  1.2536e+00,
          8.8310e-01,  1.3811e+00,  1.1267e+00,  7.9463e-01,  1.2574e+00,
          1.0988e+00,  1.3334e+00,  1.2709e+00,  1.0338e+00,  8.9485e-01,
          8.5191e-01,  6.2941e-01,  8.1570e-01,  1.1244e+00,  1.0805e+00,
          9.9755e-01,  1.0758e+00,  1.1607e+00,  1.0960e+00,  9.6049e-01,
          1.0852e+00,  9.1462e-01,  9.4122e-01,  9.8505e-01,  7.3513e-01,
          1.0134e+00,  8.3373e-01,  7.4578e-01,  1.1270e+00,  1.0679e+00,
          8.9848e-01,  9.9106e-01,  9.4795e-01,  1.0659e+00,  8.2919e-01,
          8.6020e-01,  1.3219e+00,  1.0991e+00,  1.0899e+00,  1.1484e+00,
          1.0549e+00,  8.9757e-01,  1.2341e+00,  7.1129e-01,  7.8177e-01,
          7.1453e-01,  9.2287e-01,  5.1673e-01,  7.2670e-01,  5.6472e-01,
          1.0603e+00,  5.5677e-01,  7.6662e-01,  5.9738e-01,  8.7946e-01,
          7.2365e-01,  1.1941e+00,  9.4780e-01,  5.6618e-01,  5.3710e-01,
          6.8202e-01,  1.0785e+00,  7.5097e-01,  7.3525e-01,  7.4950e-01,
          7.1948e-01,  8.9217e-01,  3.9244e-01,  7.3835e-01,  4.3247e-01,
          7.5097e-01,  7.1474e-01,  8.1818e-01,  6.3685e-01,  1.0300e+00,
          5.9656e-01,  1.0586e+00,  8.1963e-01,  4.9452e-01,  1.0996e+00,
          5.0523e-01,  9.3571e-01,  5.5205e-01,  6.1644e-01,  6.4985e-01,
          6.3577e-01,  8.5211e-01,  9.2536e-01,  4.3236e-01,  5.6647e-01,
          4.7429e-01,  8.5065e-01,  5.0285e-01,  1.0053e+00,  4.8989e-01,
          1.1755e-01,  8.2002e-01,  7.0019e-01,  6.1519e-02,  6.4777e-01,
          3.1640e-01,  2.8206e-01,  7.1172e-01,  7.0953e-01,  6.4411e-01,
          4.9723e-01,  4.9160e-01,  7.2991e-01,  4.4568e-01,  4.7622e-01,
          2.6474e-01,  6.0209e-01,  5.5910e-01,  4.3042e-01,  5.2249e-01,
          2.1004e-01,  5.4428e-01,  1.2475e-01,  4.2799e-01,  7.4566e-02,
          5.3251e-01,  6.1238e-01,  3.2354e-01,  1.3797e-02,  2.1109e-01,
          5.6343e-01,  3.2116e-01,  5.0386e-01,  9.1126e-02,  4.6912e-01,
          3.4669e-02,  4.0979e-01,  1.4810e-02,  3.8405e-01,  2.2161e-01,
          1.9445e-01, -3.5447e-01,  1.5456e-01,  1.6863e-01,  2.0110e-01,
          1.5556e-01,  2.2514e-02,  1.6489e-01,  1.6907e-01, -9.4499e-02,
          1.3021e-01,  2.4134e-01,  9.6924e-02,  1.5037e-01,  3.9969e-02,
         -2.2726e-01,  2.8770e-01, -1.7184e-01, -1.3635e-01, -8.5396e-02,
         -9.3818e-02, -4.1428e-02, -4.6396e-01, -1.7805e-01,  3.6114e-01,
          1.5889e-01, -2.7120e-01,  2.0932e-01, -4.9246e-01, -1.9852e-02,
         -9.9432e-02, -3.6289e-01,  2.1602e-01, -1.5902e-01,  2.5226e-01,
         -4.1119e-01,  7.3532e-03, -2.6737e-01, -9.9375e-02, -5.8365e-01,
         -3.8112e-01,  1.0808e-02, -6.2558e-01, -4.5019e-01, -3.2798e-01,
         -7.1162e-02, -2.6805e-01, -2.4978e-01, -3.4975e-01, -2.8487e-01,
         -2.4127e-01, -5.3032e-01, -5.4788e-01, -6.5170e-01, -3.3645e-01,
         -3.3031e-01, -2.5862e-01, -4.1498e-01, -3.1122e-01, -4.6045e-01,
         -5.4418e-01, -2.6614e-01, -4.7850e-01, -3.8730e-01, -3.8611e-01,
         -4.1716e-01, -4.9462e-01, -6.8122e-01, -4.3859e-01, -4.6447e-01,
         -2.6121e-01, -6.4777e-01, -2.9884e-01, -2.7754e-01, -3.8261e-01,
         -5.6598e-01, -1.7966e-01, -8.3324e-01, -5.7268e-01, -5.2891e-01]),
 tensor([[ 0.0644, -0.0288,  0.2246, -0.1722],
         [-0.0288,  0.2246, -0.1722,  0.0664],
         [ 0.2246, -0.1722,  0.0664, -0.1036],
         ...,
         [-0.2775, -0.3826, -0.5660, -0.1797],
         [-0.3826, -0.5660, -0.1797, -0.8332],
         [-0.5660, -0.1797, -0.8332, -0.5727]]),
 tensor([[ 6.6353e-02],
         [-1.0363e-01],
         [ 8.8868e-02],
         [ 1.5603e-01],
         [ 2.0316e-01],
         [ 2.2209e-01],
         [ 4.0233e-01],
         [ 1.8191e-01],
         [ 9.4271e-02],
         [ 1.9833e-01],
         [ 2.7132e-01],
         [-2.6344e-02],
         [ 1.3314e-01],
         [-9.5498e-02],
         [ 4.2949e-01],
         [ 2.9735e-01],
         [ 2.6208e-01],
         [ 2.5798e-01],
         [ 2.6224e-01],
         [ 4.0028e-01],
         [ 1.6453e-01],
         [-3.6497e-03],
         [ 4.5941e-02],
         [ 2.9152e-01],
         [ 2.8247e-01],
         [ 3.1230e-01],
         [ 4.0011e-01],
         [ 1.4096e-01],
         [ 4.1744e-01],
         [ 3.4225e-01],
         [-4.5256e-02],
         [ 2.8870e-01],
         [ 3.8852e-01],
         [ 3.4837e-01],
         [ 6.1889e-01],
         [ 6.2549e-01],
         [ 2.3834e-01],
         [ 4.8642e-01],
         [ 3.5614e-01],
         [ 1.1784e-01],
         [ 3.8346e-01],
         [ 4.5669e-01],
         [ 3.6588e-01],
         [ 2.6488e-01],
         [ 6.0995e-01],
         [ 6.9697e-01],
         [ 7.5780e-01],
         [ 5.8101e-01],
         [ 3.5400e-01],
         [ 2.4635e-01],
         [ 4.7288e-01],
         [ 6.6484e-01],
         [ 6.3196e-01],
         [ 5.6758e-01],
         [ 3.1575e-01],
         [ 7.3676e-01],
         [ 8.1908e-01],
         [ 8.8408e-01],
         [ 7.6086e-01],
         [ 4.3549e-01],
         [ 7.9157e-01],
         [ 4.1029e-01],
         [ 3.4122e-01],
         [ 1.0624e+00],
         [ 9.8399e-01],
         [ 7.3473e-01],
         [ 6.9833e-01],
         [ 3.4119e-01],
         [ 7.3251e-01],
         [ 6.7880e-01],
         [ 6.3626e-01],
         [ 1.0105e+00],
         [ 7.0007e-01],
         [ 1.1702e+00],
         [ 6.0600e-01],
         [ 8.9456e-01],
         [ 5.1218e-01],
         [ 7.3733e-01],
         [ 6.1851e-01],
         [ 7.6468e-01],
         [ 6.5189e-01],
         [ 1.0688e+00],
         [ 1.0419e+00],
         [ 9.7937e-01],
         [ 1.0725e+00],
         [ 5.4258e-01],
         [ 9.2915e-01],
         [ 4.3405e-01],
         [ 4.7934e-01],
         [ 1.1528e+00],
         [ 7.5340e-01],
         [ 5.4904e-01],
         [ 5.4025e-01],
         [ 5.1751e-01],
         [ 2.9075e-01],
         [ 5.5143e-01],
         [ 8.6160e-01],
         [ 9.6728e-01],
         [ 6.0795e-01],
         [ 7.0219e-01],
         [ 1.0551e+00],
         [ 7.9270e-01],
         [ 9.2103e-01],
         [ 8.7458e-01],
         [ 9.9153e-01],
         [ 6.0989e-01],
         [ 7.4993e-01],
         [ 6.9077e-01],
         [ 5.6804e-01],
         [ 7.0561e-01],
         [ 7.8830e-01],
         [ 9.7916e-01],
         [ 9.9039e-01],
         [ 7.9061e-01],
         [ 9.6164e-01],
         [ 8.5340e-01],
         [ 8.2899e-01],
         [ 7.3213e-01],
         [ 6.8678e-01],
         [ 1.2765e+00],
         [ 1.2545e+00],
         [ 1.1249e+00],
         [ 1.3865e+00],
         [ 8.9114e-01],
         [ 8.0419e-01],
         [ 1.2773e+00],
         [ 8.9564e-01],
         [ 6.2510e-01],
         [ 1.1143e+00],
         [ 9.5270e-01],
         [ 9.5466e-01],
         [ 7.9755e-01],
         [ 1.0294e+00],
         [ 5.8184e-01],
         [ 1.2175e+00],
         [ 1.0392e+00],
         [ 9.4017e-01],
         [ 1.1067e+00],
         [ 9.4888e-01],
         [ 8.5048e-01],
         [ 9.3845e-01],
         [ 1.2021e+00],
         [ 9.6893e-01],
         [ 1.0378e+00],
         [ 1.1524e+00],
         [ 1.0356e+00],
         [ 1.2582e+00],
         [ 9.6289e-01],
         [ 1.1062e+00],
         [ 6.3397e-01],
         [ 8.6299e-01],
         [ 8.4336e-01],
         [ 6.3310e-01],
         [ 1.3410e+00],
         [ 9.9408e-01],
         [ 7.2785e-01],
         [ 1.2686e+00],
         [ 8.6769e-01],
         [ 1.2330e+00],
         [ 1.0719e+00],
         [ 7.1942e-01],
         [ 1.1396e+00],
         [ 9.7066e-01],
         [ 1.1835e+00],
         [ 9.9540e-01],
         [ 7.0843e-01],
         [ 1.0240e+00],
         [ 9.7269e-01],
         [ 8.0341e-01],
         [ 1.0116e+00],
         [ 8.1968e-01],
         [ 7.5934e-01],
         [ 9.6155e-01],
         [ 1.1158e+00],
         [ 9.6646e-01],
         [ 7.0824e-01],
         [ 1.2226e+00],
         [ 7.9395e-01],
         [ 9.5238e-01],
         [ 1.1413e+00],
         [ 8.7932e-01],
         [ 9.0194e-01],
         [ 8.6707e-01],
         [ 1.2885e+00],
         [ 1.2093e+00],
         [ 1.2944e+00],
         [ 6.3165e-01],
         [ 5.5781e-01],
         [ 1.0286e+00],
         [ 1.3556e+00],
         [ 8.9051e-01],
         [ 1.1046e+00],
         [ 9.9294e-01],
         [ 1.2333e+00],
         [ 8.1726e-01],
         [ 3.1740e-01],
         [ 7.6761e-01],
         [ 6.5092e-01],
         [ 8.9621e-01],
         [ 6.8024e-01],
         [ 1.0196e+00],
         [ 7.5904e-01],
         [ 9.4622e-01],
         [ 8.6102e-01],
         [ 8.3147e-01],
         [ 7.4786e-01],
         [ 9.6059e-01],
         [ 8.6270e-01],
         [ 9.4766e-01],
         [ 5.2817e-01],
         [ 1.0704e+00],
         [ 9.7133e-01],
         [ 8.4503e-01],
         [ 8.8452e-01],
         [ 5.0776e-01],
         [ 9.2430e-01],
         [ 5.9001e-01],
         [ 8.0198e-01],
         [ 7.8422e-01],
         [ 5.8749e-01],
         [ 7.0924e-01],
         [ 7.5607e-01],
         [ 4.9436e-01],
         [ 7.7539e-01],
         [ 6.4349e-01],
         [ 7.5043e-01],
         [ 7.3827e-01],
         [ 8.6847e-01],
         [ 5.4753e-01],
         [ 6.5105e-01],
         [ 1.0554e+00],
         [ 7.8901e-01],
         [ 8.9882e-01],
         [ 7.0064e-01],
         [ 5.4479e-01],
         [ 8.2511e-01],
         [ 5.1943e-01],
         [ 1.5267e-01],
         [ 7.0765e-01],
         [ 5.7810e-01],
         [ 6.0173e-01],
         [ 4.8342e-01],
         [ 6.4010e-01],
         [ 9.0313e-01],
         [ 3.0786e-01],
         [ 1.0283e+00],
         [ 2.3870e-01],
         [ 5.7824e-01],
         [-3.3643e-02],
         [ 4.2503e-01],
         [ 5.7349e-01],
         [ 4.4148e-01],
         [ 6.8640e-01],
         [ 6.1931e-01],
         [ 2.5912e-01],
         [ 2.2371e-01],
         [ 6.0460e-01],
         [ 3.7744e-01],
         [ 5.9038e-01],
         [ 3.7926e-01],
         [ 5.2749e-01],
         [ 6.2748e-01],
         [ 6.6149e-01],
         [ 4.3518e-01],
         [ 4.0026e-01],
         [ 2.9409e-01],
         [ 4.5821e-01],
         [ 3.7015e-01],
         [ 3.4187e-01],
         [ 1.8859e-01],
         [ 6.9215e-01],
         [ 3.2195e-01],
         [ 2.4332e-02],
         [ 5.2798e-01],
         [ 9.0723e-02],
         [ 1.6245e-01],
         [ 2.7128e-01],
         [ 2.0240e-01],
         [-8.3513e-02],
         [ 3.9523e-01],
         [ 5.8745e-01],
         [ 3.2908e-01],
         [ 2.4919e-01],
         [ 2.8691e-01],
         [ 1.1735e-01],
         [ 4.3031e-01],
         [ 2.6840e-01],
         [ 2.5892e-01],
         [ 1.7928e-01],
         [ 4.2978e-01],
         [ 5.2001e-02],
         [ 2.7463e-01],
         [-1.3417e-01],
         [-1.6025e-02],
         [ 2.9625e-01],
         [ 2.3450e-01],
         [ 3.7070e-01],
         [ 7.5755e-02],
         [ 1.6683e-01],
         [ 1.1036e-01],
         [-6.8264e-03],
         [ 7.2137e-03],
         [ 2.7841e-01],
         [ 9.1316e-02],
         [ 9.1231e-03],
         [-3.5385e-01],
         [ 1.1431e-01],
         [-2.0163e-01],
         [ 2.0756e-01],
         [-9.4054e-02],
         [-1.2446e-01],
         [ 2.4384e-01],
         [-2.3242e-01],
         [-2.0931e-01],
         [-1.8707e-01],
         [-1.3957e-01],
         [-1.8903e-01],
         [-8.1507e-02],
         [-3.4759e-01],
         [-1.1257e-02],
         [-1.9703e-01],
         [ 2.7359e-02],
         [ 7.6564e-04],
         [-6.1846e-01],
         [-7.9818e-02],
         [-3.6661e-01],
         [ 6.5931e-02],
         [-5.8843e-01],
         [-1.7423e-01],
         [-2.1431e-01],
         [ 6.6695e-02],
         [-2.9555e-01],
         [-1.3760e-01],
         [-5.6146e-02],
         [-1.7448e-02],
         [-3.2177e-01],
         [-1.8931e-01],
         [-3.3209e-01],
         [-2.6944e-01],
         [-3.6146e-01],
         [-3.5334e-01],
         [-3.3019e-01],
         [-2.8488e-02],
         [-3.6981e-01],
         [-3.1455e-01],
         [-4.2320e-01],
         [-5.8333e-01],
         [-4.5083e-01],
         [-2.6372e-01],
         [-6.6177e-01],
         [-5.4376e-01],
         [-2.1988e-01],
         [-7.4067e-02],
         [-3.6120e-01],
         [-7.7958e-01],
         [-2.7244e-01],
         [-3.3669e-01],
         [-6.1547e-01],
         [-7.1691e-01],
         [-3.2713e-01],
         [-3.2994e-01],
         [-4.0011e-01],
         [-4.5194e-01],
         [-5.5936e-01],
         [-4.8557e-01],
         [-7.0421e-01],
         [-1.8149e-01],
         [-6.7299e-01],
         [-4.4816e-01],
         [-5.8107e-01],
         [-4.7465e-01],
         [-4.2599e-01],
         [-8.3749e-01],
         [-6.6348e-01],
         [-6.9997e-01],
         [-5.8357e-01],
         [-3.5789e-01],
         [-7.1656e-01],
         [-8.7729e-01],
         [-6.3883e-01],
         [-6.9591e-01],
         [-9.3954e-01],
         [-4.9190e-01],
         [-7.3111e-01],
         [-4.3942e-01],
         [-6.2770e-01],
         [-9.2674e-01],
         [-8.6653e-01],
         [-9.6315e-01],
         [-6.9102e-01],
         [-5.9326e-01],
         [-8.5505e-01],
         [-8.5113e-01],
         [-6.2499e-01],
         [-9.9391e-01],
         [-8.4853e-01],
         [-7.8337e-01],
         [-6.2505e-01],
         [-8.0748e-01],
         [-8.2683e-01],
         [-6.9701e-01],
         [-7.8696e-01],
         [-6.9023e-01],
         [-6.2324e-01],
         [-9.8813e-01],
         [-8.8023e-01],
         [-7.4747e-01],
         [-9.1390e-01],
         [-1.1208e+00],
         [-1.3740e+00],
         [-1.0556e+00],
         [-9.5917e-01],
         [-7.6300e-01],
         [-1.0235e+00],
         [-1.0120e+00],
         [-7.2330e-01],
         [-1.0387e+00],
         [-5.4913e-01],
         [-5.8775e-01],
         [-9.5260e-01],
         [-7.9546e-01],
         [-6.3363e-01],
         [-7.6522e-01],
         [-1.0495e+00],
         [-1.1376e+00],
         [-9.7715e-01],
         [-1.0136e+00],
         [-1.0233e+00],
         [-4.2618e-01],
         [-9.1957e-01],
         [-1.1356e+00],
         [-9.9493e-01],
         [-6.0729e-01],
         [-6.6857e-01],
         [-8.6124e-01],
         [-6.8043e-01],
         [-7.9954e-01],
         [-9.4514e-01],
         [-9.6364e-01],
         [-7.6007e-01],
         [-7.9190e-01],
         [-9.0256e-01],
         [-1.0012e+00],
         [-8.1351e-01],
         [-1.0263e+00],
         [-7.2729e-01],
         [-8.8496e-01],
         [-1.2832e+00],
         [-8.1455e-01],
         [-7.6072e-01],
         [-9.7935e-01],
         [-1.0354e+00],
         [-1.1628e+00],
         [-8.2264e-01],
         [-7.7549e-01],
         [-1.6201e+00],
         [-6.8463e-01],
         [-1.0060e+00],
         [-8.6768e-01],
         [-7.4937e-01],
         [-1.0767e+00],
         [-5.6351e-01],
         [-9.3733e-01],
         [-1.2270e+00],
         [-9.3480e-01],
         [-1.1192e+00],
         [-1.2890e+00],
         [-1.3016e+00],
         [-9.8122e-01],
         [-1.2527e+00],
         [-8.9113e-01],
         [-9.9467e-01],
         [-7.3336e-01],
         [-1.2790e+00],
         [-1.2139e+00],
         [-7.6957e-01],
         [-8.9945e-01],
         [-1.2749e+00],
         [-7.1642e-01],
         [-1.0271e+00],
         [-1.3689e+00],
         [-8.8271e-01],
         [-9.3780e-01],
         [-1.0709e+00],
         [-1.0079e+00],
         [-1.2095e+00],
         [-8.3435e-01],
         [-1.1892e+00],
         [-5.8446e-01],
         [-1.0576e+00],
         [-7.8082e-01],
         [-9.9774e-01],
         [-1.0047e+00],
         [-9.4661e-01],
         [-7.9260e-01],
         [-7.8298e-01],
         [-8.1630e-01],
         [-1.1429e+00],
         [-9.0614e-01],
         [-1.2286e+00],
         [-1.0185e+00],
         [-9.2398e-01],
         [-9.3490e-01],
         [-1.1074e+00],
         [-7.6938e-01],
         [-7.7835e-01],
         [-7.9201e-01],
         [-8.3866e-01],
         [-5.0138e-01],
         [-1.0518e+00],
         [-1.1464e+00],
         [-8.3545e-01],
         [-6.3239e-01],
         [-8.6411e-01],
         [-1.0649e+00],
         [-8.3904e-01],
         [-9.3103e-01],
         [-9.5688e-01],
         [-1.3042e+00],
         [-7.8724e-01],
         [-8.9785e-01],
         [-5.8319e-01],
         [-9.7922e-01],
         [-9.8292e-01],
         [-1.0255e+00],
         [-6.4694e-01],
         [-8.0609e-01],
         [-6.8586e-01],
         [-1.0256e+00],
         [-6.2613e-01],
         [-5.2035e-01],
         [-8.2406e-01],
         [-6.1214e-01],
         [-6.3858e-01],
         [-7.9211e-01],
         [-8.4110e-01],
         [-8.7759e-01],
         [-1.0926e+00],
         [-4.8413e-01],
         [-8.8961e-01],
         [-8.6125e-01],
         [-8.4024e-01],
         [-7.4395e-01],
         [-7.6605e-01],
         [-7.7586e-01],
         [-6.6531e-01],
         [-8.7798e-01],
         [-5.3314e-01],
         [-3.8761e-01],
         [-5.2371e-01],
         [-4.9831e-01],
         [-4.7124e-01],
         [-4.0311e-01],
         [-3.9151e-01],
         [-6.0217e-01],
         [-3.5831e-01],
         [-8.1952e-01],
         [-3.7521e-01],
         [-4.1182e-01],
         [-6.9520e-01],
         [-1.3176e-01],
         [-3.5725e-01],
         [-6.3746e-01],
         [-7.1734e-01],
         [-6.0116e-01],
         [-4.7620e-01],
         [-1.3156e-01],
         [-6.7144e-01],
         [-4.4765e-01],
         [-7.9655e-01],
         [-4.6641e-01],
         [-5.1395e-01],
         [-5.9736e-01],
         [-7.0441e-02],
         [-2.9234e-01],
         [-6.3963e-01],
         [-5.2695e-01],
         [-8.9920e-01],
         [-3.9060e-01],
         [-3.3070e-01],
         [-5.8975e-02],
         [-3.4768e-01],
         [-3.9728e-01],
         [-5.0462e-01],
         [-6.9052e-01],
         [-2.6284e-01],
         [-5.3189e-01],
         [-2.8471e-01],
         [-3.8808e-01],
         [-2.5389e-01],
         [-1.0635e-01],
         [-4.4742e-01],
         [-2.8809e-01],
         [-4.6124e-01],
         [-1.8804e-01],
         [-6.5422e-01],
         [-2.9021e-01],
         [-1.6320e-01],
         [-2.7098e-01],
         [-1.8750e-01],
         [ 2.4683e-01],
         [-1.2878e-01],
         [-2.1855e-01],
         [-5.1811e-01],
         [ 5.4305e-02],
         [-1.7425e-01],
         [-2.6757e-01],
         [ 1.3357e-01],
         [-3.1198e-01],
         [ 8.1655e-02],
         [-2.8527e-01],
         [-1.3569e-01],
         [-1.4587e-01],
         [ 1.7095e-01],
         [-2.3103e-02],
         [ 2.1838e-01],
         [-1.6752e-01],
         [ 3.1579e-01],
         [ 2.7031e-01],
         [-2.4856e-01],
         [ 7.6009e-03],
         [-1.1322e-03],
         [-2.0730e-01],
         [-1.2818e-01],
         [ 1.4944e-01],
         [ 9.0087e-02],
         [ 4.0082e-01],
         [ 2.9144e-01],
         [-1.4654e-01],
         [ 8.8202e-02],
         [-1.7362e-01],
         [-9.1909e-03],
         [-8.0324e-02],
         [-5.5343e-02],
         [ 5.8975e-01],
         [ 1.6623e-01],
         [ 3.3504e-01],
         [ 2.4773e-02],
         [ 8.7046e-02],
         [-1.6163e-01],
         [ 5.1961e-01],
         [ 1.7512e-01],
         [ 1.0362e-01],
         [ 2.0362e-01],
         [ 1.9839e-01],
         [ 4.1845e-01],
         [ 4.6793e-01],
         [-1.1853e-01],
         [ 1.2487e-01],
         [ 1.9344e-01],
         [ 3.0220e-01],
         [ 8.3792e-02],
         [ 3.1021e-01],
         [ 3.2869e-01],
         [ 3.0734e-01],
         [ 5.7626e-01],
         [ 4.3928e-01],
         [ 3.1897e-01],
         [ 2.8437e-01],
         [ 5.5234e-01],
         [ 6.2213e-01],
         [ 6.1580e-01],
         [ 3.9967e-01],
         [ 4.5823e-01],
         [ 4.3247e-01],
         [ 5.0114e-01],
         [ 8.3448e-01],
         [ 4.9888e-01],
         [ 5.0631e-01],
         [ 2.0848e-01],
         [ 3.6072e-01],
         [ 2.7618e-01],
         [ 4.0099e-01],
         [ 5.4027e-01],
         [ 2.4210e-01],
         [ 1.2701e-01],
         [ 4.4325e-01],
         [ 3.0193e-01],
         [ 3.6690e-01],
         [ 5.7623e-01],
         [ 5.2195e-01],
         [ 6.5280e-01],
         [ 5.7883e-01],
         [ 2.9837e-01],
         [ 2.5124e-01],
         [ 3.4579e-01],
         [ 2.2099e-01],
         [ 3.4217e-01],
         [ 8.5317e-01],
         [ 7.0991e-01],
         [ 3.0030e-01],
         [ 7.5253e-01],
         [ 7.0718e-01],
         [ 7.5546e-01],
         [ 8.3272e-01],
         [ 8.2167e-01],
         [ 6.8525e-01],
         [ 8.5421e-01],
         [ 3.8577e-01],
         [ 6.1654e-01],
         [ 6.7905e-01],
         [ 9.9523e-01],
         [ 7.6051e-01],
         [ 8.6416e-01],
         [ 6.0249e-01],
         [ 1.2840e+00],
         [ 6.4849e-01],
         [ 5.6504e-01],
         [ 6.7845e-01],
         [ 3.4798e-01],
         [ 6.4645e-01],
         [ 7.8018e-01],
         [ 9.8716e-01],
         [ 7.3607e-01],
         [ 7.5667e-01],
         [ 9.4265e-01],
         [ 8.0938e-01],
         [ 2.6675e-01],
         [ 6.1355e-01],
         [ 9.0162e-01],
         [ 4.2799e-01],
         [ 8.3804e-01],
         [ 4.6611e-01],
         [ 8.9841e-01],
         [ 9.1118e-01],
         [ 6.0615e-01],
         [ 8.8064e-01],
         [ 1.1570e+00],
         [ 9.5360e-01],
         [ 6.1671e-01],
         [ 7.3730e-01],
         [ 1.1604e+00],
         [ 7.5139e-01],
         [ 1.0676e+00],
         [ 8.5293e-01],
         [ 1.2444e+00],
         [ 7.3487e-01],
         [ 7.2214e-01],
         [ 1.3246e+00],
         [ 9.9059e-01],
         [ 1.1876e+00],
         [ 8.9527e-01],
         [ 1.1861e+00],
         [ 1.1246e+00],
         [ 1.4331e+00],
         [ 1.0912e+00],
         [ 1.0825e+00],
         [ 1.1025e+00],
         [ 9.9692e-01],
         [ 7.3992e-01],
         [ 6.6441e-01],
         [ 9.1873e-01],
         [ 9.2798e-01],
         [ 1.3192e+00],
         [ 1.0102e+00],
         [ 1.3542e+00],
         [ 1.0066e+00],
         [ 1.0593e+00],
         [ 1.0544e+00],
         [ 1.0728e+00],
         [ 8.1399e-01],
         [ 1.2641e+00],
         [ 1.0079e+00],
         [ 1.0601e+00],
         [ 1.0796e+00],
         [ 1.0223e+00],
         [ 1.1092e+00],
         [ 7.7240e-01],
         [ 1.3857e+00],
         [ 9.1455e-01],
         [ 8.6129e-01],
         [ 8.0728e-01],
         [ 6.1754e-01],
         [ 1.0502e+00],
         [ 9.3687e-01],
         [ 9.5378e-01],
         [ 9.7054e-01],
         [ 1.2540e+00],
         [ 1.1463e+00],
         [ 1.1405e+00],
         [ 1.2536e+00],
         [ 8.8310e-01],
         [ 1.3811e+00],
         [ 1.1267e+00],
         [ 7.9463e-01],
         [ 1.2574e+00],
         [ 1.0988e+00],
         [ 1.3334e+00],
         [ 1.2709e+00],
         [ 1.0338e+00],
         [ 8.9485e-01],
         [ 8.5191e-01],
         [ 6.2941e-01],
         [ 8.1570e-01],
         [ 1.1244e+00],
         [ 1.0805e+00],
         [ 9.9755e-01],
         [ 1.0758e+00],
         [ 1.1607e+00],
         [ 1.0960e+00],
         [ 9.6049e-01],
         [ 1.0852e+00],
         [ 9.1462e-01],
         [ 9.4122e-01],
         [ 9.8505e-01],
         [ 7.3513e-01],
         [ 1.0134e+00],
         [ 8.3373e-01],
         [ 7.4578e-01],
         [ 1.1270e+00],
         [ 1.0679e+00],
         [ 8.9848e-01],
         [ 9.9106e-01],
         [ 9.4795e-01],
         [ 1.0659e+00],
         [ 8.2919e-01],
         [ 8.6020e-01],
         [ 1.3219e+00],
         [ 1.0991e+00],
         [ 1.0899e+00],
         [ 1.1484e+00],
         [ 1.0549e+00],
         [ 8.9757e-01],
         [ 1.2341e+00],
         [ 7.1129e-01],
         [ 7.8177e-01],
         [ 7.1453e-01],
         [ 9.2287e-01],
         [ 5.1673e-01],
         [ 7.2670e-01],
         [ 5.6472e-01],
         [ 1.0603e+00],
         [ 5.5677e-01],
         [ 7.6662e-01],
         [ 5.9738e-01],
         [ 8.7946e-01],
         [ 7.2365e-01],
         [ 1.1941e+00],
         [ 9.4780e-01],
         [ 5.6618e-01],
         [ 5.3710e-01],
         [ 6.8202e-01],
         [ 1.0785e+00],
         [ 7.5097e-01],
         [ 7.3525e-01],
         [ 7.4950e-01],
         [ 7.1948e-01],
         [ 8.9217e-01],
         [ 3.9244e-01],
         [ 7.3835e-01],
         [ 4.3247e-01],
         [ 7.5097e-01],
         [ 7.1474e-01],
         [ 8.1818e-01],
         [ 6.3685e-01],
         [ 1.0300e+00],
         [ 5.9656e-01],
         [ 1.0586e+00],
         [ 8.1963e-01],
         [ 4.9452e-01],
         [ 1.0996e+00],
         [ 5.0523e-01],
         [ 9.3571e-01],
         [ 5.5205e-01],
         [ 6.1644e-01],
         [ 6.4985e-01],
         [ 6.3577e-01],
         [ 8.5211e-01],
         [ 9.2536e-01],
         [ 4.3236e-01],
         [ 5.6647e-01],
         [ 4.7429e-01],
         [ 8.5065e-01],
         [ 5.0285e-01],
         [ 1.0053e+00],
         [ 4.8989e-01],
         [ 1.1755e-01],
         [ 8.2002e-01],
         [ 7.0019e-01],
         [ 6.1519e-02],
         [ 6.4777e-01],
         [ 3.1640e-01],
         [ 2.8206e-01],
         [ 7.1172e-01],
         [ 7.0953e-01],
         [ 6.4411e-01],
         [ 4.9723e-01],
         [ 4.9160e-01],
         [ 7.2991e-01],
         [ 4.4568e-01],
         [ 4.7622e-01],
         [ 2.6474e-01],
         [ 6.0209e-01],
         [ 5.5910e-01],
         [ 4.3042e-01],
         [ 5.2249e-01],
         [ 2.1004e-01],
         [ 5.4428e-01],
         [ 1.2475e-01],
         [ 4.2799e-01],
         [ 7.4566e-02],
         [ 5.3251e-01],
         [ 6.1238e-01],
         [ 3.2354e-01],
         [ 1.3797e-02],
         [ 2.1109e-01],
         [ 5.6343e-01],
         [ 3.2116e-01],
         [ 5.0386e-01],
         [ 9.1126e-02],
         [ 4.6912e-01],
         [ 3.4669e-02],
         [ 4.0979e-01],
         [ 1.4810e-02],
         [ 3.8405e-01],
         [ 2.2161e-01],
         [ 1.9445e-01],
         [-3.5447e-01],
         [ 1.5456e-01],
         [ 1.6863e-01],
         [ 2.0110e-01],
         [ 1.5556e-01],
         [ 2.2514e-02],
         [ 1.6489e-01],
         [ 1.6907e-01],
         [-9.4499e-02],
         [ 1.3021e-01],
         [ 2.4134e-01],
         [ 9.6924e-02],
         [ 1.5037e-01],
         [ 3.9969e-02],
         [-2.2726e-01],
         [ 2.8770e-01],
         [-1.7184e-01],
         [-1.3635e-01],
         [-8.5396e-02],
         [-9.3818e-02],
         [-4.1428e-02],
         [-4.6396e-01],
         [-1.7805e-01],
         [ 3.6114e-01],
         [ 1.5889e-01],
         [-2.7120e-01],
         [ 2.0932e-01],
         [-4.9246e-01],
         [-1.9852e-02],
         [-9.9432e-02],
         [-3.6289e-01],
         [ 2.1602e-01],
         [-1.5902e-01],
         [ 2.5226e-01],
         [-4.1119e-01],
         [ 7.3532e-03],
         [-2.6737e-01],
         [-9.9375e-02],
         [-5.8365e-01],
         [-3.8112e-01],
         [ 1.0808e-02],
         [-6.2558e-01],
         [-4.5019e-01],
         [-3.2798e-01],
         [-7.1162e-02],
         [-2.6805e-01],
         [-2.4978e-01],
         [-3.4975e-01],
         [-2.8487e-01],
         [-2.4127e-01],
         [-5.3032e-01],
         [-5.4788e-01],
         [-6.5170e-01],
         [-3.3645e-01],
         [-3.3031e-01],
         [-2.5862e-01],
         [-4.1498e-01],
         [-3.1122e-01],
         [-4.6045e-01],
         [-5.4418e-01],
         [-2.6614e-01],
         [-4.7850e-01],
         [-3.8730e-01],
         [-3.8611e-01],
         [-4.1716e-01],
         [-4.9462e-01],
         [-6.8122e-01],
         [-4.3859e-01],
         [-4.6447e-01],
         [-2.6121e-01],
         [-6.4777e-01],
         [-2.9884e-01],
         [-2.7754e-01],
         [-3.8261e-01],
         [-5.6598e-01],
         [-1.7966e-01],
         [-8.3324e-01],
         [-5.7268e-01],
         [-5.2891e-01]]))
In [85]:
batch_size, n_train = 16, 600
# 只有前n_train个样本用于训练
train_iter = d2l.load_array((features[:n_train], labels[:n_train]),
batch_size, is_train=True)
In [86]:
def init_weights(m):
    if type(m) == nn.Linear:
        nn.init.xavier_uniform_(m.weight)
In [87]:
def get_net():
    net = nn.Sequential(nn.Linear(4, 10),
    nn.ReLU(),
    nn.Linear(10, 1))
    net.apply(init_weights)
    return net
loss = nn.MSELoss(reduction='none')
In [88]:
def train(net, train_iter, loss, epochs, lr):
    trainer = torch.optim.Adam(net.parameters(), lr)
    for epoch in range(epochs):
        for X, y in train_iter:
            trainer.zero_grad()
            l = loss(net(X), y)
            l.sum().backward()
            trainer.step()
        print(f'epoch {epoch + 1}, '
                f'loss: {d2l.evaluate_loss(net, train_iter, loss):f}')
net = get_net()
train(net, train_iter, loss, 5, 0.01)
epoch 1, loss: 0.057278
epoch 2, loss: 0.051284
epoch 3, loss: 0.047806
epoch 4, loss: 0.048928
epoch 5, loss: 0.049009
/home/yukun/.conda/envs/nn/lib/python3.11/site-packages/d2l/torch.py:3179: UserWarning: Converting a tensor with requires_grad=True to a scalar may lead to unexpected behavior.
Consider using tensor.detach() first. (Triggered internally at /pytorch/torch/csrc/autograd/generated/python_variable_methods.cpp:836.)
  self.data = [a + float(b) for a, b in zip(self.data, args)]
In [89]:
onestep_preds = net(features)
d2l.plot([time, time[tau:]],
[x.detach().numpy(), onestep_preds.detach().numpy()], 'time',
'x', legend=['data', '1-step preds'], xlim=[1, 1000],
figsize=(6, 3))
In [90]:
multistep_preds = torch.zeros(T)
multistep_preds[: n_train + tau] = x[: n_train + tau]
for i in range(n_train + tau, T):
    multistep_preds[i] = net(
        multistep_preds[i - tau:i].reshape((1, -1)))
d2l.plot([time, time[tau:], time[n_train + tau:]],
            [x.detach().numpy(), onestep_preds.detach().numpy(),
            multistep_preds[n_train + tau:].detach().numpy()], 'time',
            'x', legend=['data', '1-step preds', 'multistep preds'],
                xlim=[1, 1000], figsize=(6, 3))
In [91]:
import collections
import re
In [92]:
d2l.DATA_HUB['time_machine'] = (d2l.DATA_URL + 'timemachine.txt',
'090b5e7e70c295757f55df93cb0a180b9691891a')
def read_time_machine(): #@save
    """将时间机器数据集加载到文本行的列表中"""
    with open(d2l.download('time_machine'), 'r') as f:
        lines = f.readlines()
    return [re.sub('[^A-Za-z]+', ' ', line).strip().lower() for line in lines]
lines = read_time_machine()
print(f'# 文本总行数: {len(lines)}')
print(lines[0])
print(lines[10])
# 文本总行数: 3221
the time machine by h g wells
twinkled and his usually pale face was flushed and animated the
In [93]:
def tokenize(lines, token='word'): #@save
    """将文本行拆分为单词或字符词元"""
    if token == 'word':
        return [line.split() for line in lines]
    elif token == 'char':
        return [list(line) for line in lines]
    else:
        print('错误:未知词元类型:' + token)
tokens = tokenize(lines)
for i in range(11):
    print(tokens[i])
['the', 'time', 'machine', 'by', 'h', 'g', 'wells']
[]
[]
[]
[]
['i']
[]
[]
['the', 'time', 'traveller', 'for', 'so', 'it', 'will', 'be', 'convenient', 'to', 'speak', 'of', 'him']
['was', 'expounding', 'a', 'recondite', 'matter', 'to', 'us', 'his', 'grey', 'eyes', 'shone', 'and']
['twinkled', 'and', 'his', 'usually', 'pale', 'face', 'was', 'flushed', 'and', 'animated', 'the']
In [94]:
def count_corpus(tokens): #@save
    """统计词元的频率"""
    # 这里的tokens是1D列表或2D列表
    if len(tokens) == 0 or isinstance(tokens[0], list):
    # 将词元列表展平成一个列表
        tokens = [token for line in tokens for token in line]
    return collections.Counter(tokens)
class Vocab: #@save
    """文本词表"""
    def __init__(self, tokens=None, min_freq=0, reserved_tokens=None):
        if tokens is None:
            tokens = []
        if reserved_tokens is None:
            reserved_tokens = []
        # 按出现频率排序
        counter = count_corpus(tokens)
        self._token_freqs = sorted(counter.items(), key=lambda x: x[1],
                                    reverse=True)
        # 未知词元的索引为0
        self.idx_to_token = ['<unk>'] + reserved_tokens
        self.token_to_idx = {token: idx
                                for idx, token in enumerate(self.idx_to_token)}
        for token, freq in self._token_freqs:
            if freq < min_freq:
                break
            if token not in self.token_to_idx:
                self.idx_to_token.append(token)
                self.token_to_idx[token] = len(self.idx_to_token) - 1
    def __len__(self):
        return len(self.idx_to_token)
    def __getitem__(self, tokens):
        if not isinstance(tokens, (list, tuple)):
            return self.token_to_idx.get(tokens, self.unk)
        return [self.__getitem__(token) for token in tokens]
    def to_tokens(self, indices):
        if not isinstance(indices, (list, tuple)):
            return self.idx_to_token[indices]
        return [self.idx_to_token[index] for index in indices]
    @property
    def unk(self): # 未知词元的索引为0
        return 0
    @property
    def token_freqs(self):
        return self._token_freqs
In [95]:
vocab = Vocab(tokens)
print(list(vocab.token_to_idx.items())[:10])
[('<unk>', 0), ('the', 1), ('i', 2), ('and', 3), ('of', 4), ('a', 5), ('to', 6), ('was', 7), ('in', 8), ('that', 9)]
In [96]:
for i in [0, 100]:
    print('文本:', tokens[i])
    print('索引:', vocab[tokens[i]])
文本: ['the', 'time', 'machine', 'by', 'h', 'g', 'wells']
索引: [1, 19, 50, 40, 2183, 2184, 400]
文本: ['were', 'three', 'dimensional', 'representations', 'of', 'his', 'four', 'dimensioned']
索引: [20, 175, 1452, 2250, 4, 25, 262, 2251]
In [97]:
def load_corpus_time_machine(max_tokens=-1): #@save
    """返回时光机器数据集的词元索引列表和词表"""
    lines = read_time_machine()
    tokens = tokenize(lines, 'char')
    vocab = Vocab(tokens)
    # 因为时光机器数据集中的每个文本行不一定是一个句子或一个段落,
    # 所以将所有文本行展平到一个列表中
    corpus = [vocab[token] for line in tokens for token in line]
    if max_tokens > 0:
        corpus = corpus[:max_tokens]
    return corpus, vocab
corpus, vocab = load_corpus_time_machine()

len(corpus), len(vocab)
Out [97]:
(170580, 28)
In [98]:
tokens = d2l.tokenize(read_time_machine())
# 因为每个文本行不一定是一个句子或一个段落,因此我们把所有文本行拼接到一起
corpus = [token for line in tokens for token in line]
vocab = d2l.Vocab(corpus)
vocab.token_freqs[:10]
Out [98]:
[('the', 2261),
 ('i', 1267),
 ('and', 1245),
 ('of', 1155),
 ('a', 816),
 ('to', 695),
 ('was', 552),
 ('in', 541),
 ('that', 443),
 ('my', 440)]
In [99]:
freqs = [freq for token, freq in vocab.token_freqs]
d2l.plot(freqs, xlabel='token: x', ylabel='frequency: n(x)',
xscale='log', yscale='log')
In [100]:
bigram_tokens = [pair for pair in zip(corpus[:-1], corpus[1:])]
bigram_vocab = Vocab(bigram_tokens)
bigram_vocab.token_freqs[:10]
Out [100]:
[(('of', 'the'), 309),
 (('in', 'the'), 169),
 (('i', 'had'), 130),
 (('i', 'was'), 112),
 (('and', 'the'), 109),
 (('the', 'time'), 102),
 (('it', 'was'), 99),
 (('to', 'the'), 85),
 (('as', 'i'), 78),
 (('of', 'a'), 73)]
Warning:
Output truncated. This notebook contains too many cells to display efficiently.