はじめに
今回私は最近はやりのchatGPTに興味を持ち、深層学習について学んでみたいと思い立ちました!
深層学習といえばPythonということなので、最終的にはPythonを使って深層学習ができるとこまでコツコツと学習していくことにしました。
ただ、勉強するだけではなく少しでもアウトプットをしようということで、備忘録として学習した内容をまとめていこうと思います。
この記事が少しでも誰かの糧になることを願っております!
※投稿主の環境はWindowsなのでMacの方は多少違う部分が出てくると思いますが、ご了承ください。
最初の記事:Python初心者の備忘録 #01
前の記事:Python初心者の備忘録 #28 ~深層学習超入門編04~
次の記事:まだ
今回は正則化、バッチ正規化、レイヤー正規化についてまとめております。
■学習に使用している資料
Udemy:②米国AI開発者がやさしく教える深層学習超入門第二弾【Pythonで実践】
■正則化(regularization)
基本的にparameter(layer)を増やしていけば精度は上がっていくのだが、深層学習でもbias-variance trade offという考え方が適用でき、モデルが複雑になるほどbiasは下がり、varianceは高くなる。
正則化とは、深層学習における過学習への対策の1つである。
本記事ではL2正則とドロップアウトについて紹介する
▶L2正則(重み減衰:Weight Decay)
- 正則化項といえばRidgeとLassoが有名だが、L2正則はRidgeとほぼ同じである
- 損失関数にパラメータのL2ノルムの正則化項を追加することで、パラメータの値(重み)を小さくすることで過学習を防ぐ
- ハイパーパラメータ$\lambda$で正則化の強さを調整する
- バイアスは正則化項には含まない
L2正則をスクラッチ実装
- モデルのパラメータを取得し、重みに対してフロベニウス正則(L2正則)を損失関数に加える
-
torch.linalg.norm()でフロベニウスノルムを計算
-
- 正則化項の強さを調整する係数
weight_decayを使う - 学習ループにあることを想定して、lossに追加していく(実際のlossの計算等は省略する)
import matplotlib.pyplot as plt
import torch
from torch.nn import functional as F
from torch import nn, optim
from torch.utils.data import DataLoader
import torchvision
from torchvision import transforms
%load_ext autoreload
%autoreload 2
import utils # 本記事で用意しているモジュール
# tensor生成
X = torch.tensor([[1., 2., 3.], [4., 5., 6.]])
# フロベニウスノルムを計算
torch.linalg.norm(X) # torch.sqrt(torch.sum(X**2))と同じ
# -> tensor(9.5394)
# CNN
def get_conv_model():
return nn.Sequential(
# 3x28x28
nn.Conv2d(3, 4, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 4x14x14
nn.Conv2d(4, 8, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 8x7x7
nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 16x4x4
nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 32x2x2 -> GAP -> 32 x 1 x 1
nn.Flatten(),
# # 128 -> 32
nn.Linear(128, 10)
# nn.Linear(32, 10)
# 10
)
# 実際には学習ループに組み込む
l2_reg = torch.tensor(0.)
for name, param in get_conv_model().named_parameters():
# print(name, param)
if 'weight' in name:
l2_reg += torch.linalg.norm(param)**2
# loss += (weight_decay / (2*m)) * l2_reg
PytorchモジュールでL2正則
-
torch.nn.optim.<class>()において、waight_decay引数を指定することでL2正則を利用できる
例:torch.nn.optim.SGD(model.parameters(), lr=0.03, weight_decay=0.01) -
waight_decayの値は一般的に0.01、0.001、0.0001などの小さな値が使用される
▶L2正則の有無での学習結果を比較する
例として下記条件で比較を行う。
- CIFAR10データで学習
- GAP層を使ったCNNを使用
- 20epochで、重み減衰有りのOptimizerと無しのOptimizerのlearning curveを比較
# model準備
conv_model = get_conv_model()
conv_model_l2 = get_conv_model()
transform = transforms.Compose([
transforms.ToTensor(),
# transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))])
# CIFAR10用Normalize
transforms.Normalize(mean=[0.4914, 0.4822, 0.4465], std=[0.2023, 0.1994, 0.2010])
])
classes = ('plane', 'car', 'bird', 'cat',
'deer', 'dog', 'frog', 'horse', 'ship', 'truck')
train_dataset = torchvision.datasets.CIFAR10('./cifar10_data', train=True, download=True, transform=transform)
val_dataset = torchvision.datasets.CIFAR10('./cifar10_data', train=False, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=128, shuffle=True, num_workers=4)
val_loader = DataLoader(val_dataset, batch_size=128, num_workers=4)
# L2正則有り/無しを用意
opt = optim.SGD(conv_model.parameters(), lr=0.03)
opt_l2 = optim.SGD(conv_model_l2.parameters(), lr=0.03, weight_decay=0.001)
# 学習
num_epoch = 20
train_losses, val_losses, val_accuracies = utils.learn(conv_model, train_loader, val_loader, opt, F.cross_entropy, num_epoch=num_epoch)
train_losses_l2, val_losses_l2, val_accuracies_l2 = utils.learn(conv_model_l2, train_loader, val_loader, opt_l2, F.cross_entropy, num_epoch=num_epoch)
# 学習曲線出力
plt.plot(train_losses, label='train loss without weight decay')
plt.plot(train_losses_l2, label='train loss with weight decay')
plt.plot(val_losses_l2, label='val loss with weight decay')
plt.plot(val_losses, label='val loss without weight decay')
plt.xlabel('epoch')
plt.ylabel('loss')
plt.legend()

L2正則無しのほうが学習は速く進んでいるが、L2正則有りのほうがval lossが安定している。
▶ドロップアウト
- モデルの過学習を抑制し、汎化性能を向上させるためのテクニックで、主に全結合層で用いられる
- 学習時にランダムで一部のニューロンを確率𝑝でドロップアウト(無効化)する
- アンサンブルの効果が得られるとともに一般に汎化性能の向上を期待できる
- Dropoutは層ごとに指定でき、それぞれ異なる確率𝑝を指定できる
※層が深くなるにつれ確率𝑝を高くしたり、特定の層にのみdropoutを適用する等 - バッチ毎に非活性ニューロンをランダムに選択する
※各バッチで異なるモデルを学習することになり、アンサンブル効果が得られる
- 非活性ニューロンにより学習時の層の出力の総和が小さくなってしまう、一方で予測時には全てのニューロンを使用するので、スケーリングが必要
-> 学習時のdropoutがある層の出力を$\frac{1}{1-p}$倍する
Dropoutをスクラッチ実装
- 入力tensorと確率𝑝を受け取りランダムにtensorの要素を0にする
- 出力時には$(1-p)$の逆数をかけてスケーリングする
# Dropout関数
def dropout(X, drop_p):
keep_p = 1 - drop_p # (1 - p)を設定
mask = torch.rand(X.shape) < keep_p # Dropout用のマスクを生成
return X * mask / keep_p
# 実行
X = torch.randn((100, 100))
droped_X = dropout(X, 0.3)
"""実行結果(ところどころが0になっている)
tensor([[ 0.0000, -3.2176, 0.0000, ..., -1.3100, 1.0123, -1.5714],
[-1.5627, -1.7712, 0.1377, ..., 0.6681, 0.0000, -1.2267],
[-0.6961, 0.8658, 2.1287, ..., -0.0000, -0.0000, -1.0496],
..,
[ 0.6516, -0.0000, -1.3543, ..., -0.0000, -1.0073, 0.9672],
[-0.4615, 0.6114, 0.0000, ..., -0.0099, 3.4280, -0.1678],
[ 0.5435, 0.2993, -1.6568, ..., -0.0000, -0.9031, -1.9652]])
"""
# 現状だと+-が混在しているので、先にReLuを行う
# ReLUを適用した後にDropoutをすると、学習時と予測時でスケールが変わるので/keep_pでリスケールしスケールを合わせる
def relu(X):
return torch.clamp(X, min=0)
relu_out = relu(X)
keep_p = 0.5
mask = torch.rand(X.shape) < keep_p
# スケーリング(Dropout後でも似た値に戻すことができている)
print((relu_out * mask).sum() / keep_p) # -> tensor(3958.6870)
print(relu_out.sum()) # -> tensor(3983.1404)
PytorchモジュールでDropout
- nn.Dropout()
-
p:dropする確率を指定
-
model = nn.Sequential(
nn.Linear(64, 20),
nn.ReLU(),
nn.Dropout(p=0.4), # Dropoutを指定
nn.Linear(20, 10)
)
■正規化層(バッチ正規化とレイヤー正規化)
▶バッチ正規化(Batch Normalization)
- 深層学習モデルの学習を安定化させてより高速にする必須テクニックで、各層の活性化関数の入力($Z^{[l]}$)の平均と分散を正規化する
- ネットワークの各活性化層(主にReLU層)の入力を正規化することで、学習を安定化させ、学習率を大きくできるので、学習の高速化ができる
- 内部共変量シフトの問題を緩和する
▶Activationの分布
- 各層の活性化関数の出力(Activation:$A^{[l]}$)の分布を確認することで、学習の進行状況を理解できる
この分布が下記のようなトラブルシューティングの際に重要な情報となる- 内部共変量シフト
- 勾配消失/爆発問題
- 不活性なニューロン
- 重みの初期化問題
▶PythonでActivationの情報を取得
⇒ Hookを使って情報を取得する
- NNモデルの層(
nn.Module)やtensorに対して勾配計算時などのタイミングで特定の関数を実行することが可能 - tensorに対しては
.register_hook()を使用する
-> tensorの勾配計算の直後に引数で渡した関数を実行する -
nn.Moduleに対しては以下のメソッドでhookを登録する-
Forward Hook
-
nn.Moduleのforwardメソッドが呼び出された直後に実行される関数を定義 - 特定の層の出力を記録できる
-
nn.Moduleオブジェクトに対して.register_forward_hook()に関数を渡す
-
-
Backward Hook
- 逆伝播の間に勾配の計算が行われた直後に実行される関数を定義
- 勾配の値を記録したり、勾配に変更を加えたりすることが可能
-
nn.Moduleオブジェクトに対して.register_full_backward_hook()に関数を渡す
※.register_backward_hook()はPytorch1.9以降非推奨
-
Forward Hook
tensorにHookを使用
-
.register_hook(func)メソッドでそのtensorの.gradが計算された時に、引数に渡した関数を実行する - 引数の関数funcは、そのtensorの勾配(
.grad)を引数に受け取る
from functools import partial
import matplotlib.pyplot as plt
import numpy as np
import torch
from torch.utils.data import DataLoader
from torch.nn import functional as F
from torch import nn, optim
import torchvision
from torchvision import transforms
%load_ext autoreload
%autoreload 2
import utils
a = torch.ones(5, requires_grad=True)
b = 2*a
b.retain_grad()
b.register_hook(lambda grad: print(grad))
c = b.mean()
c.backward() # -> tensor([0.2000, 0.2000, 0.2000, 0.2000, 0.2000])
モデルのレイヤーにforward hookを登録する
-
nn.Moduleに対して.register_forward_hook(func)でhookを登録する -
model.named_modules()等でモデルのモジュールをイテレーションすることで、モデルの各層(モジュール)にhookをつけることができる - 自動で3つの引数を受け取る
-
module:hookが登録されているモジュールそのもの(nn.Module) -
inp:そのモジュールの入力 -
out:そのモジュールの出力
-
partial関数で、引数の固定
-
functools.partialを使って、関数の一部の引数を固定して新しい関数を作成できる
partial(func, *args, **keywords)の形で使用
from functools import partial
# 元となる関数の定義
def power(base, exponent):
return base ** exponent
# 部分適⽤した関数の作成 -> power関数のexponent引数が2に固定された関数が作られる
square = partial(power, exponent=2)
# 部分適⽤した関数の利⽤
print(square(5)) # 出⼒:25
# モデル準備
conv_model = nn.Sequential(
# 1x28x28
nn.Conv2d(1, 4, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 4x14x14
nn.Conv2d(4, 8, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 8x7x7
nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 16x4x4
nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 32x2x2 -> GAP -> 32 x 1 x 1
nn.Flatten(),
# # 128 -> 32
nn.Linear(128, 10)
# nn.Linear(32, 10)
# 10
)
# Hookを定義
outputs = {}
def save_output(name, module, inp, out):
module_name = f'{name}_{str(module)}'
outputs[module_name] = out.shape
for name, module in conv_model.named_modules():
if name: # 自分自身のmoduleにはhookをつけない
module.register_forward_hook(partial(save_output, name)) # nameを固定
# Hookがあるか確認する関数
def print_hooks(model):
for name, module in model.named_modules():
if hasattr(module, "_forward_hooks"):
for hook in module._forward_hooks.values():
print(f'Module {name} has forward hook: {hook}')
if hasattr(module, "_backward_hooks"):
for hook in module._backward_hooks.values():
print(f'Module {name} has backward hook: {hook}')
# forwardでhook発動
X = torch.randn((1, 1, 28, 28))
output = conv_model(X)
"""実行結果(outputsの値)
{'0_Conv2d(1, 4, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': torch.Size([1, 4, 14, 14]),
'1_ReLU()': torch.Size([1, 4, 14, 14]),
'2_Conv2d(4, 8, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': torch.Size([1, 8, 7, 7]),
'3_ReLU()': torch.Size([1, 8, 7, 7]),
'4_Conv2d(8, 16, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': torch.Size([1, 16, 4, 4]),
'5_ReLU()': torch.Size([1, 16, 4, 4]),
'6_Conv2d(16, 32, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': torch.Size([1, 32, 2, 2]),
'7_ReLU()': torch.Size([1, 32, 2, 2]),
'8_Flatten(start_dim=1, end_dim=-1)': torch.Size([1, 128]),
'9_Linear(in_features=128, out_features=10, bias=True)': torch.Size([1, 10])}
"""
モデルのレイヤーにbackward hookを登録する
-
nn.Moduleに対して.register_full_backward_hook(func)でhookを登録する- モジュールの入力と出力の勾配に対するhookを設定する
- 勾配の変更や監視などの操作を実行したいときに使用
-
model.named_modules()等でモデルのモジュールをイテレーションすることで、モデルの各層(モジュール)にhookをつけることができる -
.backward()で逆伝播をし、各モジュールの勾配を計算した後にhookに登録した関数が実行される - 自動で3つの引数を受け取る
-
module:hookが登録されているモジュールそのもの(nn.Module) -
grad_in:そのモジュールの入力の勾配 -
grad_out:そのモジュールの出力の勾配
-
# Hookを定義
grads = {}
def save_grad_in(name, module, grad_in, grad_out):
module_name = f'{name}_{str(module)}'
grads[module_name] = grad_in
for name, module in conv_model.named_modules():
if name: # 自分自身のmoduleにはhookをつけない
module.register_full_backward_hook(partial(save_grad_in, name))
# backward
X = torch.randn((1, 1, 28, 28))
output = conv_model(X)
loss = output.mean()
loss.backward()
実行結果
{'9_Linear(in_features=128, out_features=10, bias=True)': (tensor([[ 1.0921e-02, -4.3990e-03, 7.2159e-03, 3.1243e-02, 1.6620e-02,
-7.0275e-03, 9.4797e-03, -1.7395e-03, -2.3359e-03, -2.1701e-05,
4.7168e-04, 2.6008e-02, 5.9918e-03, -2.0172e-02, 4.6917e-03,
2.7294e-02, 1.0300e-03, -6.0820e-03, -1.1718e-02, 6.9794e-03,
8.5153e-03, 2.6091e-02, -1.9366e-03, 3.1181e-03, -4.0697e-03,
2.0126e-02, -5.8546e-03, 1.9624e-02, -2.0346e-03, 2.1318e-02,
9.7808e-03, -6.4242e-03, 5.0385e-03, 1.7358e-02, 2.1928e-03,
-3.4724e-03, 2.3878e-02, 1.9980e-02, -9.4096e-03, 7.2701e-03,
-9.2114e-04, -1.2644e-02, 1.4586e-02, -1.1472e-02, -3.0500e-02,
1.1921e-02, -1.1921e-02, 1.0437e-02, 2.4066e-02, 8.0657e-03,
-1.4798e-02, 2.4862e-02, -1.8980e-02, 1.8410e-02, -3.4257e-04,
1.4696e-02, -2.1735e-02, 2.0870e-03, 1.5663e-02, 4.9493e-03,
4.2591e-03, -9.8904e-03, 2.7069e-02, -4.9831e-03, 1.0082e-02,
1.8965e-02, 8.3023e-03, 1.2518e-02, 4.3696e-03, -1.0321e-04,
-2.5166e-02, 3.6042e-03, 1.0916e-02, -4.6339e-03, -2.8633e-03,
-1.6319e-02, 1.4807e-02, -8.8942e-03, 2.9575e-04, 3.9900e-03,
-1.0500e-03, -2.4952e-03, -2.1930e-02, -1.2692e-02, -2.9128e-02,
-1.5449e-02, -2.4347e-03, 2.1425e-02, 1.8583e-02, 2.5633e-02,
4.0231e-03, 3.5832e-03, 2.2667e-02, -3.5905e-02, 1.7062e-02,
3.0661e-02, 2.8080e-02, 2.9406e-03, 5.4189e-04, 1.4014e-02,
1.0859e-02, -2.5285e-03, 4.2081e-02, 9.2557e-03, -1.6036e-02,
1.6470e-02, 1.4494e-02, -8.3996e-04, 1.5382e-02, -1.7236e-02,
1.5295e-02, -1.8159e-02, -2.8988e-02, -2.5866e-02, -1.1155e-02,
4.0374e-03, 1.3646e-02, 6.7126e-04, 1.8978e-02, 1.1644e-02,
-1.6774e-02, -2.1379e-02, -1.8311e-02, 3.0744e-02, -1.7930e-02,
-3.0645e-02, 9.1215e-03, 1.9649e-02]]),),
'8_Flatten(start_dim=1, end_dim=-1)': (tensor([[[[ 1.0921e-02, -4.3990e-03],
[ 7.2159e-03, 3.1243e-02]],
[[ 1.6620e-02, -7.0275e-03],
[ 9.4797e-03, -1.7395e-03]],
[[-2.3359e-03, -2.1701e-05],
[ 4.7168e-04, 2.6008e-02]],
[[ 5.9918e-03, -2.0172e-02],
[ 4.6917e-03, 2.7294e-02]],
[[ 1.0300e-03, -6.0820e-03],
[-1.1718e-02, 6.9794e-03]],
[[ 8.5153e-03, 2.6091e-02],
[-1.9366e-03, 3.1181e-03]],
[[-4.0697e-03, 2.0126e-02],
[-5.8546e-03, 1.9624e-02]],
[[-2.0346e-03, 2.1318e-02],
[ 9.7808e-03, -6.4242e-03]],
[[ 5.0385e-03, 1.7358e-02],
[ 2.1928e-03, -3.4724e-03]],
[[ 2.3878e-02, 1.9980e-02],
[-9.4096e-03, 7.2701e-03]],
[[-9.2114e-04, -1.2644e-02],
[ 1.4586e-02, -1.1472e-02]],
[[-3.0500e-02, 1.1921e-02],
[-1.1921e-02, 1.0437e-02]],
[[ 2.4066e-02, 8.0657e-03],
[-1.4798e-02, 2.4862e-02]],
[[-1.8980e-02, 1.8410e-02],
[-3.4257e-04, 1.4696e-02]],
[[-2.1735e-02, 2.0870e-03],
[ 1.5663e-02, 4.9493e-03]],
[[ 4.2591e-03, -9.8904e-03],
[ 2.7069e-02, -4.9831e-03]],
[[ 1.0082e-02, 1.8965e-02],
[ 8.3023e-03, 1.2518e-02]],
[[ 4.3696e-03, -1.0321e-04],
[-2.5166e-02, 3.6042e-03]],
[[ 1.0916e-02, -4.6339e-03],
[-2.8633e-03, -1.6319e-02]],
[[ 1.4807e-02, -8.8942e-03],
[ 2.9575e-04, 3.9900e-03]],
[[-1.0500e-03, -2.4952e-03],
[-2.1930e-02, -1.2692e-02]],
[[-2.9128e-02, -1.5449e-02],
[-2.4347e-03, 2.1425e-02]],
[[ 1.8583e-02, 2.5633e-02],
[ 4.0231e-03, 3.5832e-03]],
[[ 2.2667e-02, -3.5905e-02],
[ 1.7062e-02, 3.0661e-02]],
[[ 2.8080e-02, 2.9406e-03],
[ 5.4189e-04, 1.4014e-02]],
[[ 1.0859e-02, -2.5285e-03],
[ 4.2081e-02, 9.2557e-03]],
[[-1.6036e-02, 1.6470e-02],
[ 1.4494e-02, -8.3996e-04]],
[[ 1.5382e-02, -1.7236e-02],
[ 1.5295e-02, -1.8159e-02]],
[[-2.8988e-02, -2.5866e-02],
[-1.1155e-02, 4.0374e-03]],
[[ 1.3646e-02, 6.7126e-04],
[ 1.8978e-02, 1.1644e-02]],
[[-1.6774e-02, -2.1379e-02],
[-1.8311e-02, 3.0744e-02]],
[[-1.7930e-02, -3.0645e-02],
[ 9.1215e-03, 1.9649e-02]]]]),),
'7_ReLU()': (tensor([[[[ 0.0109, -0.0044],
[ 0.0072, 0.0312]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0010, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[-0.0020, 0.0213],
[ 0.0098, -0.0064]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0239, 0.0200],
[-0.0094, 0.0073]],
[[-0.0009, -0.0126],
[ 0.0146, -0.0115]],
[[-0.0305, 0.0119],
[-0.0119, 0.0104]],
[[ 0.0241, 0.0081],
[-0.0148, 0.0249]],
[[-0.0190, 0.0000],
[-0.0003, 0.0147]],
[[-0.0217, 0.0021],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0101, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0044, -0.0001],
[-0.0252, 0.0036]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0148, 0.0000],
[ 0.0003, 0.0040]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0000, 0.0000]],
[[ 0.0000, 0.0256],
[ 0.0040, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0171, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0005, 0.0000]],
[[ 0.0000, 0.0000],
[ 0.0421, 0.0093]],
[[-0.0160, 0.0165],
[ 0.0145, -0.0008]],
[[ 0.0154, -0.0172],
[ 0.0153, -0.0182]],
[[-0.0290, -0.0259],
[-0.0112, 0.0040]],
[[ 0.0000, 0.0000],
[ 0.0190, 0.0116]],
[[-0.0168, -0.0214],
[-0.0183, 0.0000]],
[[-0.0179, -0.0306],
[ 0.0091, 0.0196]]]]),),
'6_Conv2d(16, 32, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': (tensor([[[[-8.1224e-04, 8.8051e-03, -3.7924e-03, 4.0243e-03],
[ 2.9598e-03, 3.3052e-03, 4.0327e-03, 3.5059e-03],
[-4.5375e-03, 2.6876e-03, -3.0632e-03, 2.4276e-03],
[-2.9871e-03, 5.3796e-03, -5.2571e-03, 3.7799e-03]],
[[-1.5733e-03, 3.3068e-03, -2.1497e-03, 4.0445e-03],
[ 1.7453e-03, -3.6990e-03, 8.4699e-04, 8.3692e-05],
[-2.3996e-03, 3.1891e-03, 1.5152e-03, 2.4021e-03],
[ 4.2043e-03, 1.0123e-03, -2.2671e-03, 1.2533e-03]],
[[ 5.9415e-03, 2.5239e-03, 6.5090e-03, 1.2484e-03],
[-2.8309e-03, 1.1945e-02, -6.3376e-04, 9.9478e-03],
[-6.5318e-03, -1.4537e-04, -1.2008e-03, -1.8552e-03],
[ 1.8780e-04, 8.7959e-04, 6.6143e-04, -1.8889e-03]],
[[-1.9355e-03, 1.7704e-04, -3.9984e-03, -1.4748e-03],
[-7.7303e-03, 9.2516e-03, -8.6722e-03, 1.1642e-03],
[-3.6055e-03, -1.0900e-02, 1.5369e-03, 3.4651e-03],
[ 1.5549e-03, 3.5049e-03, 1.8408e-03, 3.7570e-03]],
[[-2.6236e-03, -2.7631e-04, 7.2431e-03, 3.4406e-03],
[-2.8467e-03, -2.1614e-03, 3.2836e-03, -5.6541e-03],
[ 6.8491e-04, 7.4490e-03, -3.2392e-03, -1.9207e-03],
[ 7.3717e-04, 2.6391e-03, 1.3052e-03, 2.0603e-03]],
[[-4.7392e-03, 9.7611e-04, -2.2489e-03, -1.1619e-03],
[-8.4127e-03, -6.5658e-03, 5.6114e-03, -6.5003e-03],
[-2.1920e-03, -8.4643e-03, -7.7131e-04, -2.1159e-03],
[ 2.2930e-03, -5.3064e-03, 3.9402e-03, 4.2334e-03]],
[[-8.8421e-03, 3.2870e-03, -3.3622e-03, -2.6443e-03],
[-5.0419e-03, -6.6077e-04, -1.4780e-03, -8.6279e-03],
[ 1.2909e-03, 3.2973e-03, -5.4123e-04, 3.5995e-03],
[-5.0128e-03, 1.9425e-03, 1.0865e-03, 1.4871e-03]],
[[-3.2767e-03, 6.7580e-03, 1.9946e-04, -5.2361e-03],
[-7.3424e-04, 9.9075e-03, 4.2868e-03, -2.0631e-03],
[ 5.6509e-03, 7.1259e-03, 4.6860e-04, 6.1161e-03],
[-8.6625e-04, 7.1809e-03, -4.7639e-04, 4.6767e-03]],
[[ 3.4179e-03, -1.4861e-04, -3.2047e-03, 7.4699e-04],
[-7.8380e-03, -3.4127e-03, -9.3033e-03, -5.3440e-03],
[ 3.1743e-03, -2.9292e-05, -3.5800e-03, -5.8630e-04],
[ 5.7492e-03, -7.4535e-03, -1.6904e-03, 3.4651e-04]],
[[-1.8648e-03, 3.7695e-03, -4.1276e-03, 3.7306e-03],
[-1.2048e-03, 2.3942e-03, -1.2842e-03, -3.9905e-03],
[ 5.7733e-03, 3.8128e-03, 2.1531e-03, -3.7171e-03],
[ 2.6650e-03, -4.4391e-03, -4.7052e-04, -4.0273e-03]],
[[-3.1481e-03, 4.7381e-03, 4.0787e-04, 5.7281e-03],
[-4.6669e-03, 4.5235e-03, 4.3917e-04, 2.7792e-03],
[ 6.2606e-03, 5.6273e-03, 6.0640e-03, -2.4458e-04],
[ 9.2786e-04, -2.4128e-03, 2.1609e-03, 9.9979e-04]],
[[-1.3908e-03, 5.4902e-03, -3.0880e-03, -2.7354e-03],
[ 6.3066e-04, -4.3386e-03, 5.8105e-03, 1.5834e-03],
[-1.7890e-03, 8.0994e-03, 2.0549e-03, 1.2166e-03],
[-2.7448e-03, 5.0549e-03, -9.6755e-04, -3.6717e-03]],
[[ 1.0410e-02, 4.0270e-03, 4.3084e-05, 3.2104e-03],
[-1.9653e-03, 1.7322e-03, -4.4630e-03, 2.3962e-03],
[ 2.7378e-04, 3.6787e-03, 1.1893e-03, 6.3945e-04],
[ 5.2412e-05, -6.7765e-04, -1.6120e-03, -2.1347e-03]],
[[ 3.7880e-03, -2.0208e-03, 2.2090e-03, -1.1848e-03],
[-3.0157e-03, -9.1156e-04, -6.3857e-03, -1.7833e-04],
[-4.8676e-03, 2.3498e-03, 1.4701e-03, 9.2037e-04],
[ 3.0809e-03, -1.0096e-02, -2.0238e-04, 5.4446e-04]],
[[ 1.4401e-04, -2.9455e-03, -1.3885e-03, 2.1430e-03],
[ 6.6681e-03, 2.9532e-03, 4.2780e-03, 6.9571e-03],
[ 1.2775e-03, 8.6522e-03, 6.3525e-03, -6.9465e-03],
[ 4.8444e-03, 6.0539e-03, 2.4884e-03, 4.9810e-03]],
[[-5.3359e-03, -3.6717e-03, -1.3821e-03, -2.7377e-03],
[ 1.0219e-03, 6.8303e-03, -1.0258e-04, -1.9569e-03],
[-2.9363e-03, 2.0021e-03, -1.1634e-03, 5.8919e-03],
[ 9.3405e-04, 3.8363e-04, -4.2537e-03, -9.1564e-04]]]]),),
'5_ReLU()': (tensor([[[[-8.1224e-04, 0.0000e+00, -3.7924e-03, 4.0243e-03],
[ 2.9598e-03, 0.0000e+00, 4.0327e-03, 3.5059e-03],
[-4.5375e-03, 0.0000e+00, -3.0632e-03, 2.4276e-03],
[-2.9871e-03, 0.0000e+00, 0.0000e+00, 3.7799e-03]],
[[-1.5733e-03, 3.3068e-03, -2.1497e-03, 4.0445e-03],
[ 1.7453e-03, -3.6990e-03, 8.4699e-04, 8.3692e-05],
[ 0.0000e+00, 0.0000e+00, 1.5152e-03, 2.4021e-03],
[ 0.0000e+00, 1.0123e-03, -2.2671e-03, 1.2533e-03]],
[[ 5.9415e-03, 2.5239e-03, 6.5090e-03, 1.2484e-03],
[ 0.0000e+00, 1.1945e-02, -6.3376e-04, 0.0000e+00],
[-6.5318e-03, 0.0000e+00, 0.0000e+00, -1.8552e-03],
[ 1.8780e-04, 8.7959e-04, 6.6143e-04, -1.8889e-03]],
[[-1.9355e-03, 1.7704e-04, -3.9984e-03, -1.4748e-03],
[-7.7303e-03, 9.2516e-03, -8.6722e-03, 1.1642e-03],
[-3.6055e-03, 0.0000e+00, 0.0000e+00, 3.4651e-03],
[ 1.5549e-03, 3.5049e-03, 1.8408e-03, 3.7570e-03]],
[[-2.6236e-03, -2.7631e-04, 7.2431e-03, 3.4406e-03],
[-2.8467e-03, -2.1614e-03, 3.2836e-03, 0.0000e+00],
[ 6.8491e-04, 7.4490e-03, -3.2392e-03, -1.9207e-03],
[ 7.3717e-04, 2.6391e-03, 1.3052e-03, 2.0603e-03]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[-8.4127e-03, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 2.2930e-03, -5.3064e-03, 0.0000e+00, 0.0000e+00]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 1.2909e-03, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[-5.0128e-03, 0.0000e+00, 1.0865e-03, 0.0000e+00]],
[[-3.2767e-03, 6.7580e-03, 1.9946e-04, 0.0000e+00],
[-7.3424e-04, 9.9075e-03, 4.2868e-03, -2.0631e-03],
[ 5.6509e-03, 7.1259e-03, 4.6860e-04, 6.1161e-03],
[-8.6625e-04, 7.1809e-03, -4.7639e-04, 4.6767e-03]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[-7.8380e-03, -3.4127e-03, -9.3033e-03, -5.3440e-03],
[ 3.1743e-03, 0.0000e+00, 0.0000e+00, -5.8630e-04],
[ 5.7492e-03, -7.4535e-03, 0.0000e+00, 3.4651e-04]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 3.7306e-03],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00]],
[[-3.1481e-03, 0.0000e+00, 4.0787e-04, 0.0000e+00],
[-4.6669e-03, 4.5235e-03, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, -2.4128e-03, 2.1609e-03, 0.0000e+00]],
[[ 0.0000e+00, 5.4902e-03, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 5.8105e-03, 0.0000e+00],
[-1.7890e-03, 8.0994e-03, 2.0549e-03, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00]],
[[ 1.0410e-02, 0.0000e+00, 0.0000e+00, 3.2104e-03],
[ 0.0000e+00, 0.0000e+00, -4.4630e-03, 0.0000e+00],
[ 0.0000e+00, 3.6787e-03, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, -2.1347e-03]],
[[ 3.7880e-03, -2.0208e-03, 2.2090e-03, -1.1848e-03],
[-3.0157e-03, -9.1156e-04, -6.3857e-03, -1.7833e-04],
[-4.8676e-03, 2.3498e-03, 1.4701e-03, 0.0000e+00],
[ 3.0809e-03, 0.0000e+00, -2.0238e-04, 5.4446e-04]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 6.6681e-03, 2.9532e-03, 0.0000e+00, 0.0000e+00],
[ 1.2775e-03, 0.0000e+00, 6.3525e-03, 0.0000e+00],
[ 4.8444e-03, 6.0539e-03, 2.4884e-03, 0.0000e+00]],
[[ 0.0000e+00, 0.0000e+00, -1.3821e-03, -2.7377e-03],
[ 1.0219e-03, 0.0000e+00, -1.0258e-04, 0.0000e+00],
[-2.9363e-03, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00]]]]),),
'4_Conv2d(8, 16, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': (tensor([[[[-1.0656e-03, 7.5747e-04, -1.6185e-03, -5.3223e-04, -1.0706e-03,
9.7928e-05, -9.7779e-04],
[-8.1075e-04, -2.6244e-03, -1.5192e-04, 1.5267e-03, -1.2031e-03,
-2.1059e-03, -7.7150e-04],
[-6.0472e-04, -1.3894e-04, -2.1541e-03, 2.2548e-04, 3.5344e-05,
1.6608e-03, 1.9653e-04],
[-2.8736e-03, -1.4244e-03, -8.6384e-04, 1.4071e-03, -1.1771e-03,
-9.8173e-04, -6.2930e-05],
[ 1.0392e-03, 9.1389e-04, -2.9733e-03, 8.2200e-04, -6.7711e-04,
-8.8036e-04, -7.0756e-04],
[ 1.5740e-04, 2.1206e-03, -6.0328e-04, -4.9052e-04, -1.8407e-04,
9.4428e-05, -3.8516e-06],
[-5.3994e-04, -1.2651e-03, -2.3370e-03, -4.3497e-04, -1.7912e-04,
-3.1361e-04, -8.4835e-04]],
[[-4.1640e-04, -2.4980e-03, -9.0657e-04, -6.9413e-04, -2.2370e-04,
3.3716e-04, -3.8910e-04],
[ 5.1844e-04, -1.3270e-04, 5.4790e-04, -2.1540e-04, 1.3245e-03,
7.4849e-04, -1.0925e-05],
[-7.6038e-04, -5.7810e-04, -4.0637e-05, -3.8023e-04, -1.3220e-03,
4.0301e-04, 3.2908e-05],
[ 7.3732e-04, -2.6721e-03, 1.1793e-03, 1.5071e-03, 2.0445e-04,
-5.4433e-04, -3.7882e-05],
[-3.1092e-04, -6.8968e-04, -2.3050e-03, 2.2061e-03, -7.7050e-04,
9.7026e-05, -5.0178e-04],
[ 4.4780e-05, 3.3589e-04, 2.5528e-03, -2.5554e-03, 5.9973e-04,
-5.7766e-05, 2.7670e-04],
[-8.2545e-05, -7.7172e-04, -1.1081e-03, 2.8213e-03, -4.0030e-05,
1.5509e-04, -1.3812e-04]],
[[ 1.0352e-03, 1.2007e-03, 7.5558e-04, 1.8508e-04, 8.6633e-04,
-3.6132e-04, 5.7727e-04],
[-2.1258e-03, 5.3812e-04, 2.3194e-03, -4.3054e-04, 8.3277e-04,
1.0553e-03, -1.1102e-03],
[-1.6536e-06, 6.5862e-04, 4.9397e-04, 2.2869e-03, -1.0077e-03,
-1.0656e-03, -9.1100e-04],
[-5.6003e-04, -1.2753e-03, 2.2697e-03, -1.5830e-03, 5.2504e-04,
-1.7667e-03, -1.5750e-04],
[ 4.5050e-04, -9.7666e-04, 9.9400e-04, 1.9136e-04, 3.7191e-04,
1.1241e-03, 1.5986e-04],
[ 4.1770e-04, -1.2715e-04, 1.3569e-03, 2.4989e-04, 4.5238e-04,
-8.3924e-06, 8.1631e-04],
[ 7.1249e-04, 9.7514e-04, 3.6068e-04, -1.1837e-03, -2.4222e-04,
6.3730e-05, 6.1609e-05]],
[[-1.0615e-03, -1.0752e-04, 6.9042e-04, -1.0681e-03, 1.0034e-03,
9.0056e-04, 9.0564e-04],
[-2.9002e-04, -3.4739e-03, -4.6412e-04, 6.8232e-05, -1.5989e-03,
-3.1644e-03, 5.9167e-04],
[-1.0635e-03, 2.9646e-04, -1.6517e-05, -2.4679e-03, 1.0296e-03,
-8.1664e-06, -5.2583e-04],
[ 1.1080e-03, 1.4427e-03, -8.3323e-04, -2.3484e-03, 6.8837e-04,
9.0577e-04, 1.2139e-03],
[ 5.6764e-04, 4.2960e-04, 4.7950e-04, 3.5519e-04, -5.6601e-04,
-1.0193e-04, -1.6735e-04],
[-1.6584e-04, 1.0465e-03, 1.0939e-03, -1.7193e-03, -9.2489e-04,
-1.2682e-03, 1.2618e-03],
[ 7.1151e-04, 6.6203e-04, -7.8366e-04, -1.4712e-03, -1.8660e-04,
-2.1273e-04, 2.2677e-04]],
[[-5.5757e-05, 9.9819e-04, -1.2479e-04, 1.0773e-03, 6.8591e-04,
-4.7962e-04, -1.2547e-04],
[-9.3380e-04, -2.8332e-04, 2.6574e-03, -2.4873e-03, 9.0571e-05,
3.5035e-05, -9.0065e-04],
[ 3.9473e-05, 1.6200e-03, -2.4472e-04, 2.1644e-03, -1.2779e-04,
-3.7209e-04, 1.2569e-06],
[-1.2878e-03, -2.8467e-03, 6.4494e-04, 2.2062e-03, 8.4960e-04,
-4.1877e-04, 2.3373e-04],
[-2.7420e-04, -1.7249e-03, -5.3101e-05, -6.2599e-04, 5.9828e-06,
6.3329e-04, -6.2544e-04],
[-2.4746e-05, -1.7150e-03, 2.2267e-03, -2.9928e-03, -3.8832e-04,
3.0428e-04, 1.6922e-04],
[ 6.9785e-04, -4.9448e-04, 3.8832e-04, 1.7894e-03, -9.8686e-06,
2.6351e-04, -4.5929e-04]],
[[ 1.4509e-04, -3.1800e-04, -3.7117e-04, -1.3032e-04, -5.6349e-04,
8.4643e-04, -6.1448e-04],
[ 2.4576e-04, -2.6615e-03, 2.4249e-03, -1.0139e-03, 6.5292e-04,
1.6201e-03, 1.1232e-05],
[ 9.2611e-04, 1.7564e-03, 1.5166e-04, -3.4335e-05, -8.1387e-04,
-9.0837e-06, 7.4637e-05],
[ 1.6329e-03, -1.6931e-03, 9.6722e-05, 4.4217e-03, 2.0745e-03,
-1.0868e-03, 1.0959e-03],
[ 3.9892e-04, -1.4426e-03, -1.1159e-03, -2.3341e-03, 1.2100e-03,
4.4169e-04, 4.5841e-04],
[-5.7219e-04, -2.3464e-03, 9.5252e-05, -6.3566e-05, -4.8333e-05,
-1.1582e-03, 5.7359e-04],
[-2.4290e-05, -9.1920e-04, 1.0833e-03, 1.2837e-03, -1.2491e-04,
-2.5292e-05, -1.5558e-04]],
[[ 7.9668e-04, 1.2505e-03, 2.4611e-04, 7.6998e-04, 3.3443e-04,
-6.0057e-04, -3.9938e-04],
[-7.9400e-04, 1.7224e-04, -8.7273e-04, 2.4792e-03, -5.1523e-04,
-2.8066e-03, 6.3521e-04],
[-1.6282e-04, 1.7916e-03, -9.5535e-05, 5.2070e-04, 1.0430e-03,
-1.5612e-04, 3.6284e-04],
[-1.8953e-03, -1.0941e-03, -6.6283e-04, -2.1561e-03, -1.9229e-03,
3.5086e-04, -7.6511e-04],
[-2.0394e-04, -2.2606e-04, 2.1148e-03, 1.0084e-03, 2.6319e-04,
2.2455e-04, -4.2830e-05],
[ 5.9260e-05, -1.2758e-03, -1.5653e-03, 4.9268e-05, -7.0527e-04,
-6.6742e-04, -1.2761e-04],
[-3.3747e-04, 1.7132e-03, 5.8338e-04, 4.7087e-04, 3.8362e-04,
4.2629e-04, -9.0944e-06]],
[[-6.3080e-04, 9.4132e-04, 6.8788e-04, -4.7812e-04, 3.0448e-04,
1.8219e-04, 3.6464e-06],
[-2.5405e-03, -8.5279e-04, 9.8303e-04, -3.8908e-03, -1.5163e-03,
-7.8586e-05, -1.5794e-04],
[-1.0955e-04, 1.7112e-03, -4.0635e-04, -3.6135e-03, 5.7530e-04,
-8.6006e-04, -5.4358e-04],
[ 3.5968e-03, -2.9882e-03, -3.5051e-04, -2.0188e-03, 1.5229e-03,
-2.3937e-03, 7.5154e-04],
[ 9.1789e-06, 2.3676e-03, -8.5528e-04, -1.8046e-03, -4.0570e-04,
1.6788e-03, -2.3612e-05],
[-1.3230e-03, 2.6724e-04, 7.1319e-04, 1.5617e-03, 2.7488e-04,
-1.1601e-03, 1.0248e-03],
[-2.0661e-05, 8.9161e-04, -1.2851e-03, -1.4811e-03, -6.6366e-04,
1.9848e-04, 9.3297e-06]]]]),),
'3_ReLU()': (tensor([[[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, -5.3223e-04, 0.0000e+00,
9.7928e-05, -9.7779e-04],
[ 0.0000e+00, -2.6244e-03, 0.0000e+00, 1.5267e-03, 0.0000e+00,
0.0000e+00, -7.7150e-04],
[-6.0472e-04, 0.0000e+00, -2.1541e-03, 2.2548e-04, 0.0000e+00,
0.0000e+00, 1.9653e-04],
[-2.8736e-03, 0.0000e+00, -8.6384e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, -6.2930e-05],
[ 0.0000e+00, 0.0000e+00, -2.9733e-03, 8.2200e-04, 0.0000e+00,
-8.8036e-04, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, -4.9052e-04, 0.0000e+00,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, -2.3370e-03, 0.0000e+00, 0.0000e+00,
-3.1361e-04, 0.0000e+00]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, -2.2370e-04,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 5.4790e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
4.0301e-04, 3.2908e-05],
[ 7.3732e-04, 0.0000e+00, 1.1793e-03, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00],
[-3.1092e-04, 0.0000e+00, 0.0000e+00, 2.2061e-03, 0.0000e+00,
0.0000e+00, -5.0178e-04],
[ 0.0000e+00, 3.3589e-04, 0.0000e+00, -2.5554e-03, 5.9973e-04,
0.0000e+00, 2.7670e-04],
[-8.2545e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, -1.3812e-04]],
[[ 1.0352e-03, 1.2007e-03, 7.5558e-04, 0.0000e+00, 8.6633e-04,
-3.6132e-04, 5.7727e-04],
[-2.1258e-03, 5.3812e-04, 2.3194e-03, 0.0000e+00, 8.3277e-04,
1.0553e-03, -1.1102e-03],
[-1.6536e-06, 6.5862e-04, 0.0000e+00, 2.2869e-03, -1.0077e-03,
-1.0656e-03, -9.1100e-04],
[-5.6003e-04, -1.2753e-03, 2.2697e-03, -1.5830e-03, 0.0000e+00,
-1.7667e-03, 0.0000e+00],
[ 4.5050e-04, -9.7666e-04, 9.9400e-04, 0.0000e+00, 3.7191e-04,
0.0000e+00, 1.5986e-04],
[ 4.1770e-04, -1.2715e-04, 1.3569e-03, 2.4989e-04, 4.5238e-04,
-8.3924e-06, 8.1631e-04],
[ 7.1249e-04, 9.7514e-04, 3.6068e-04, -1.1837e-03, 0.0000e+00,
6.3730e-05, 6.1609e-05]],
[[-1.0615e-03, -1.0752e-04, 6.9042e-04, -1.0681e-03, 1.0034e-03,
9.0056e-04, 9.0564e-04],
[-2.9002e-04, -3.4739e-03, 0.0000e+00, 6.8232e-05, -1.5989e-03,
-3.1644e-03, 5.9167e-04],
[-1.0635e-03, 2.9646e-04, -1.6517e-05, -2.4679e-03, 1.0296e-03,
-8.1664e-06, -5.2583e-04],
[ 1.1080e-03, 1.4427e-03, -8.3323e-04, 0.0000e+00, 6.8837e-04,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 4.2960e-04, 4.7950e-04, 3.5519e-04, -5.6601e-04,
-1.0193e-04, -1.6735e-04],
[ 0.0000e+00, 1.0465e-03, 1.0939e-03, 0.0000e+00, -9.2489e-04,
-1.2682e-03, 1.2618e-03],
[ 7.1151e-04, 6.6203e-04, -7.8366e-04, 0.0000e+00, -1.8660e-04,
-2.1273e-04, 2.2677e-04]],
[[-5.5757e-05, 0.0000e+00, -1.2479e-04, 0.0000e+00, 0.0000e+00,
-4.7962e-04, 0.0000e+00],
[-9.3380e-04, -2.8332e-04, 2.6574e-03, -2.4873e-03, 9.0571e-05,
0.0000e+00, 0.0000e+00],
[ 3.9473e-05, 0.0000e+00, 0.0000e+00, 2.1644e-03, 0.0000e+00,
-3.7209e-04, 1.2569e-06],
[-1.2878e-03, -2.8467e-03, 6.4494e-04, 2.2062e-03, 8.4960e-04,
-4.1877e-04, 2.3373e-04],
[ 0.0000e+00, -1.7249e-03, 0.0000e+00, -6.2599e-04, 5.9828e-06,
6.3329e-04, -6.2544e-04],
[-2.4746e-05, -1.7150e-03, 2.2267e-03, -2.9928e-03, -3.8832e-04,
0.0000e+00, 1.6922e-04],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
2.6351e-04, 0.0000e+00]],
[[ 1.4509e-04, -3.1800e-04, -3.7117e-04, -1.3032e-04, -5.6349e-04,
8.4643e-04, -6.1448e-04],
[ 2.4576e-04, -2.6615e-03, 2.4249e-03, 0.0000e+00, 6.5292e-04,
1.6201e-03, 1.1232e-05],
[ 9.2611e-04, 0.0000e+00, 1.5166e-04, -3.4335e-05, -8.1387e-04,
-9.0837e-06, 7.4637e-05],
[ 1.6329e-03, -1.6931e-03, 0.0000e+00, 4.4217e-03, 2.0745e-03,
-1.0868e-03, 1.0959e-03],
[ 3.9892e-04, -1.4426e-03, -1.1159e-03, -2.3341e-03, 1.2100e-03,
4.4169e-04, 4.5841e-04],
[-5.7219e-04, -2.3464e-03, 9.5252e-05, -6.3566e-05, -4.8333e-05,
0.0000e+00, 5.7359e-04],
[-2.4290e-05, -9.1920e-04, 1.0833e-03, 1.2837e-03, -1.2491e-04,
-2.5292e-05, -1.5558e-04]],
[[ 7.9668e-04, 1.2505e-03, 2.4611e-04, 7.6998e-04, 3.3443e-04,
0.0000e+00, -3.9938e-04],
[-7.9400e-04, 1.7224e-04, -8.7273e-04, 0.0000e+00, 0.0000e+00,
-2.8066e-03, 6.3521e-04],
[-1.6282e-04, 1.7916e-03, -9.5535e-05, 5.2070e-04, 1.0430e-03,
0.0000e+00, 0.0000e+00],
[-1.8953e-03, 0.0000e+00, 0.0000e+00, 0.0000e+00, -1.9229e-03,
0.0000e+00, -7.6511e-04],
[-2.0394e-04, -2.2606e-04, 2.1148e-03, 1.0084e-03, 2.6319e-04,
2.2455e-04, -4.2830e-05],
[ 5.9260e-05, -1.2758e-03, 0.0000e+00, 4.9268e-05, 0.0000e+00,
-6.6742e-04, -1.2761e-04],
[-3.3747e-04, 1.7132e-03, 5.8338e-04, 4.7087e-04, 3.8362e-04,
0.0000e+00, -9.0944e-06]],
[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 3.6464e-06],
[ 0.0000e+00, -8.5279e-04, 0.0000e+00, 0.0000e+00, -1.5163e-03,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, -3.6135e-03, 0.0000e+00,
0.0000e+00, -5.4358e-04],
[ 3.5968e-03, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
1.6788e-03, 0.0000e+00],
[-1.3230e-03, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
1.9848e-04, 0.0000e+00]]]]),),
'2_Conv2d(4, 8, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': (tensor([[[[ 2.3937e-05, 1.1012e-04, -1.6486e-04, 1.6532e-05, -1.3814e-04,
2.7187e-05, -5.4272e-05, -1.0197e-04, -1.6643e-04, -1.5402e-04,
-6.0708e-05, -2.6618e-04, -2.2534e-04, -1.4727e-04],
[ 9.5885e-06, 3.0527e-04, 6.3060e-04, 1.1652e-04, -1.1160e-04,
1.3917e-04, 5.1844e-04, -4.0658e-04, 3.2414e-04, 8.1624e-04,
2.2539e-04, 3.1220e-04, -1.5987e-04, 2.4545e-04],
[ 1.1413e-04, 1.1815e-05, -2.4900e-04, -7.2038e-04, 2.8631e-04,
1.2080e-03, 9.8164e-05, -2.3473e-04, 4.4620e-05, 6.9611e-04,
7.9264e-04, 4.3300e-04, -2.4726e-04, -1.7247e-04],
[-3.3721e-04, -1.0692e-03, 3.4237e-04, -1.5182e-03, -2.3796e-04,
7.2905e-04, -5.6431e-05, 5.0506e-05, -2.3128e-04, -6.7149e-04,
-6.1109e-04, -3.6258e-04, 1.3754e-04, 1.0892e-04],
[ 1.1079e-04, 3.1046e-06, -2.5770e-04, -1.0208e-04, -2.7592e-04,
-4.5945e-05, -8.6948e-05, 8.9369e-04, -2.8710e-04, -1.5136e-04,
-5.4339e-05, -2.6684e-05, 2.9127e-05, -3.9120e-05],
[-6.2364e-04, 2.0763e-04, 3.8731e-04, -8.5275e-05, 3.4847e-04,
-1.1107e-04, -4.2422e-04, 8.3567e-04, -7.9133e-05, 9.6508e-04,
-2.0009e-06, 2.1964e-04, -1.6144e-04, 4.1951e-05],
[ 8.2500e-05, -1.6533e-04, -3.7370e-04, -1.2196e-03, -1.3410e-04,
-4.0664e-05, 3.4612e-04, 5.7151e-04, 3.1037e-04, 6.0903e-04,
-7.0793e-05, -5.2193e-04, 1.5663e-04, 1.6343e-04],
[-3.8814e-04, 1.1547e-03, 3.4974e-05, -7.9794e-04, -2.5181e-06,
-2.3280e-04, 3.2494e-04, -7.4044e-04, -3.9292e-04, 2.6597e-04,
-1.7974e-04, -3.4615e-04, -2.1567e-04, -1.1748e-05],
[ 7.9036e-05, 4.2049e-04, -1.7023e-04, -1.1428e-03, -7.9333e-04,
1.2891e-04, -4.4567e-04, -7.0718e-04, 9.8013e-05, 1.1661e-04,
1.0632e-04, 1.1011e-04, 7.6861e-05, 6.5782e-05],
[-3.1710e-06, 4.8413e-04, 2.4798e-04, -4.3924e-04, 3.6353e-04,
8.0010e-04, 4.8942e-04, 7.4723e-05, 1.3952e-04, -3.7216e-04,
2.5528e-04, -1.8998e-04, -3.3490e-04, 2.2057e-04],
[-1.9393e-04, 2.3483e-04, -2.0061e-04, -1.2143e-03, -3.0736e-05,
5.0260e-04, 7.0531e-05, -6.8975e-05, 1.9606e-05, 2.2458e-04,
2.2688e-04, -2.5099e-04, -1.2218e-04, 7.4733e-05],
[ 6.4238e-05, 1.6851e-04, 2.5903e-05, -6.2447e-04, 2.1598e-05,
5.3430e-04, -9.1805e-04, 1.1443e-04, 4.5163e-05, -1.9687e-04,
-1.5153e-04, -2.4613e-04, 9.0351e-07, 3.6472e-04],
[-3.4823e-05, -1.9172e-04, -3.4215e-04, -3.2354e-04, -2.4731e-04,
2.3442e-04, 1.8285e-05, 9.5674e-05, -3.3052e-05, -1.0156e-05,
1.3121e-05, 2.3594e-05, -1.8721e-05, -2.3318e-05],
[-1.6353e-05, 1.4439e-04, 3.5971e-04, -8.1760e-05, 5.5780e-05,
3.2592e-04, -1.1950e-04, 9.1855e-05, 5.7093e-05, -1.1070e-04,
4.7344e-05, -1.9019e-05, 6.5746e-06, 1.9975e-06]],
[[-2.6704e-04, 1.1212e-05, -2.5228e-04, 9.5089e-06, 7.5840e-06,
2.8798e-05, -2.4484e-04, 3.3928e-06, 1.1820e-05, 2.6876e-05,
2.0486e-04, 6.1986e-04, 2.3585e-04, -1.4135e-05],
[-2.8994e-04, 8.3559e-06, -3.2679e-04, -5.4188e-04, 4.6399e-04,
-2.5265e-04, -1.7093e-04, 3.5683e-04, -4.6983e-05, 3.8496e-04,
5.2358e-04, -6.9416e-04, -4.9376e-04, -6.5162e-05],
[ 4.8207e-06, 5.6033e-04, -6.3901e-04, -1.3562e-03, 8.2737e-04,
-1.9954e-04, -4.2241e-04, 7.4009e-05, -6.6958e-05, -1.3321e-03,
2.0348e-04, 7.1236e-05, 7.6279e-05, -6.5804e-05],
[ 1.7157e-04, -5.4876e-04, -1.0512e-03, 8.2409e-04, 1.1303e-04,
-1.1919e-04, 4.3505e-04, -5.9100e-04, -3.4649e-04, 8.1442e-04,
-3.0242e-04, 2.2382e-04, -4.1207e-05, 2.5229e-04],
[ 7.5319e-05, 1.5205e-04, -2.3221e-04, -2.6131e-05, 2.9503e-04,
-5.5548e-04, -1.0027e-04, -7.1619e-04, -1.3592e-04, -5.6776e-05,
-2.1956e-05, -2.4653e-04, -7.6074e-05, -4.6931e-05],
[ 1.2626e-04, -2.5248e-04, -1.4998e-04, -6.9698e-05, -1.6767e-05,
-7.8133e-04, -4.3178e-04, 7.5472e-04, 5.4500e-04, -4.5319e-04,
-1.9419e-04, 2.8705e-04, 2.6846e-04, 5.6569e-05],
[ 8.1981e-04, 7.9898e-04, -3.2506e-04, 7.8140e-04, 1.3980e-04,
-1.6165e-03, 8.7030e-04, -3.1482e-04, 8.0963e-04, 3.8212e-04,
-2.2113e-04, -1.7861e-04, 3.1704e-04, 1.3941e-04],
[ 2.8836e-05, -3.2592e-04, -3.3883e-04, -9.8969e-04, -6.2059e-04,
1.9970e-04, 6.4372e-04, 1.4320e-03, 1.8642e-04, 1.3202e-04,
-6.1798e-05, 4.0974e-04, 1.1952e-04, 1.2719e-05],
[ 7.5060e-05, 2.7094e-04, -2.9502e-04, 1.2047e-03, -6.0954e-05,
-8.1272e-05, -4.6556e-04, -5.8528e-04, 5.5247e-05, -1.7957e-05,
1.0831e-04, -1.6150e-04, -4.6134e-05, 9.7511e-05],
[ 1.4366e-04, 3.2599e-04, 2.2366e-08, -1.5219e-03, -1.3583e-04,
-1.4898e-04, -3.8998e-04, -6.1777e-04, 8.1882e-05, 6.0266e-05,
8.7902e-05, 3.9386e-04, 3.1381e-04, -1.9107e-04],
[-3.4162e-05, 2.9535e-04, -1.3569e-04, 5.4000e-04, 4.2075e-04,
-1.0431e-04, -4.3081e-04, 4.6695e-04, -1.4563e-04, -5.2346e-04,
-9.2486e-05, 1.4801e-04, 3.4713e-04, 1.2845e-04],
[ 8.7399e-06, -2.7850e-04, -6.0840e-04, -7.7733e-04, -1.1043e-04,
8.3233e-05, -2.1357e-04, 5.8537e-05, -2.2692e-04, 1.5011e-04,
-6.5490e-05, 5.5975e-05, -3.3625e-06, -2.8705e-04],
[ 1.6210e-04, 7.1103e-04, -2.9768e-04, -1.1943e-04, 2.3233e-04,
-8.6740e-04, 1.0715e-04, 8.0686e-05, -1.0888e-04, -7.5533e-05,
1.8711e-05, -3.0638e-05, 4.6634e-06, 4.8491e-05],
[-1.0344e-04, -6.6617e-04, -1.3288e-04, -2.7363e-04, -1.5321e-04,
5.8064e-04, 3.6950e-04, 2.5912e-05, 6.5170e-06, -3.0432e-05,
-3.0982e-05, 5.8315e-05, -1.6274e-05, -1.6815e-05]],
[[-4.0228e-05, 1.4272e-04, -6.4145e-05, -1.3240e-04, -8.4810e-05,
6.4997e-06, 6.6723e-05, 1.4506e-05, -9.0218e-05, -4.3075e-04,
-6.1727e-05, -4.4069e-05, -1.4123e-04, -7.1466e-05],
[-3.2086e-04, -1.5973e-04, -7.7067e-04, 7.5304e-04, 1.6558e-04,
2.6055e-04, -8.7314e-05, -1.1949e-05, -2.4506e-04, 1.3508e-04,
-2.5712e-04, 5.0259e-04, 1.0112e-04, 1.2568e-04],
[ 2.4292e-04, 5.9650e-04, 2.3969e-04, 8.7696e-04, -6.0283e-04,
4.0869e-04, 1.6449e-04, -4.3470e-04, -2.2458e-05, 9.3307e-04,
-1.5938e-04, 1.0204e-03, 6.4942e-05, -1.2595e-04],
[-2.5315e-04, 2.6695e-04, -2.4080e-04, 1.2067e-03, -4.4358e-04,
-1.2772e-03, -8.7892e-04, 2.1312e-04, -5.4002e-05, -2.9787e-04,
-8.0498e-04, 6.4100e-05, 1.9813e-04, -2.4368e-05],
[-5.7082e-05, 2.5399e-04, -3.4804e-05, 1.9550e-04, -1.4575e-04,
4.9889e-04, -6.4127e-05, 9.1943e-04, 1.5602e-04, -3.9264e-04,
1.1528e-04, -5.1565e-05, 1.5230e-04, 6.5153e-05],
[-1.1543e-04, 2.5606e-04, 1.5928e-04, -6.8074e-05, 4.9274e-05,
1.2956e-03, -5.5321e-05, 2.0192e-03, 5.5524e-04, -3.2656e-05,
7.0507e-05, 1.3275e-04, 7.3171e-05, 1.2044e-05],
[-4.5353e-04, -3.7591e-04, 2.8801e-04, -4.9969e-04, -3.3958e-04,
8.0113e-04, -2.5325e-04, 8.8753e-04, -3.0150e-04, 4.1176e-04,
3.0548e-04, -2.7025e-04, -1.2688e-04, 2.4008e-04],
[-5.0538e-04, -2.5669e-04, 1.7525e-04, 5.4144e-05, -7.8949e-04,
-1.3130e-04, 4.6083e-04, -1.7226e-03, -2.0996e-04, 1.0125e-04,
1.9874e-04, 6.2382e-05, -2.0128e-04, -1.4744e-04],
[-7.6448e-05, -1.5684e-04, 2.5149e-04, -2.3964e-04, -1.6167e-04,
5.9585e-05, 1.6871e-04, -2.9720e-04, -1.0129e-04, 5.6402e-04,
-1.3528e-04, -3.4259e-05, -4.7456e-06, -5.0054e-05],
[-8.6366e-05, -2.0974e-04, 4.1264e-04, 1.0381e-05, 6.7187e-04,
8.4585e-04, -6.2139e-04, 1.9965e-04, -3.3008e-04, -3.2754e-04,
-2.2731e-04, 1.6206e-04, 9.0695e-05, -9.0865e-05],
[ 3.3175e-05, -1.2354e-04, 1.4219e-04, -4.4020e-04, -2.9590e-04,
-4.4689e-04, 1.6233e-04, -5.8411e-04, -1.0068e-05, 3.6661e-04,
5.4367e-05, 1.5689e-04, -2.3851e-04, -4.0392e-06],
[ 3.0935e-04, 8.2410e-05, 1.6129e-05, 5.4838e-04, -2.4528e-04,
8.7201e-05, -1.8823e-04, 5.9477e-04, -3.7406e-04, 1.8751e-04,
-2.4077e-04, 1.4221e-06, 7.3875e-05, -1.3567e-04],
[-1.2874e-04, -2.0046e-04, -1.8863e-05, 7.6757e-05, -2.0815e-04,
4.5474e-04, 4.4234e-05, 7.9381e-05, 3.3896e-05, 1.0218e-04,
-2.4316e-05, 4.9032e-05, -1.2110e-06, -7.3315e-05],
[ 2.6147e-05, 1.4361e-04, 1.1021e-04, 3.7724e-04, -2.9279e-05,
3.7330e-04, 5.8772e-05, -2.6094e-04, -5.7822e-06, 5.9071e-05,
-4.2728e-06, 3.2576e-05, 5.5953e-05, 6.4956e-06]],
[[ 1.6470e-04, -2.7624e-04, 1.7499e-04, -1.2259e-04, 6.9882e-07,
4.9388e-05, 1.5995e-04, -1.5952e-04, 1.4625e-05, 3.7111e-05,
4.3648e-05, 1.1004e-04, -1.3988e-04, 1.4384e-04],
[-1.7775e-04, 4.5784e-05, 1.1488e-03, 7.7995e-06, 3.3354e-04,
4.6698e-04, -8.5577e-05, -5.0907e-04, 6.2724e-04, 8.6607e-04,
4.3535e-04, 2.7672e-04, 3.1742e-05, 1.6021e-04],
[ 3.5031e-05, 4.1418e-04, 1.3538e-04, 2.5667e-04, -5.0345e-04,
-2.3831e-04, 2.9470e-04, 7.5360e-05, 1.1585e-04, 3.3693e-04,
-2.8915e-04, -5.5372e-04, 8.7012e-05, 3.2053e-05],
[-2.3625e-04, 5.0858e-04, -2.6603e-04, -1.7864e-03, 4.2650e-04,
4.1888e-04, 9.1864e-04, 6.4201e-04, -1.2718e-04, 7.6525e-05,
-5.0092e-04, 2.5285e-04, 1.2927e-05, -5.0823e-05],
[ 3.1950e-05, -4.0170e-04, 2.6143e-04, 7.5459e-05, -3.5227e-05,
3.1396e-04, 3.8865e-05, -6.0907e-04, 1.1948e-04, 1.7387e-05,
-6.4454e-07, -4.2278e-05, 5.7739e-05, -1.0788e-05],
[-2.5086e-04, -7.0095e-04, 2.9502e-04, -1.8155e-04, 8.2582e-04,
1.0952e-04, -6.1037e-04, 2.6006e-04, -1.5226e-04, 1.5867e-04,
-3.5271e-04, 1.5612e-04, -2.2940e-04, -9.6533e-05],
[-4.3130e-04, -3.4917e-04, 2.3871e-04, 4.9014e-04, -2.6723e-04,
-1.2312e-04, -1.1016e-04, -4.1635e-05, -3.8083e-04, -2.5270e-04,
5.1527e-05, 3.7086e-04, -1.2186e-04, -7.2179e-05],
[ 4.5180e-05, 2.5626e-04, 4.1634e-05, -7.3664e-04, 5.5430e-04,
-2.6693e-04, 4.8458e-04, -5.3410e-04, 1.9737e-05, 8.2474e-04,
-4.9684e-04, -4.8043e-04, -2.3831e-05, 2.4546e-04],
[ 1.6957e-05, -2.6987e-04, 1.2108e-04, 2.8456e-04, 2.3749e-04,
-7.2097e-04, -1.7035e-04, 5.5760e-04, 8.9456e-05, -2.6423e-04,
-9.1239e-05, 6.6347e-05, 1.5629e-04, -2.0357e-04],
[ 4.2304e-04, -1.7646e-04, 1.1746e-04, -4.5712e-04, 6.0311e-04,
-8.7996e-04, -5.7582e-04, 6.8055e-04, 3.5117e-04, 2.5471e-04,
8.1237e-05, -1.0538e-04, 2.1575e-04, -1.5488e-04],
[ 3.8732e-05, -5.7904e-06, -1.6344e-04, 5.6333e-04, -3.2018e-04,
6.0812e-04, 7.2461e-04, -5.3970e-04, -1.7153e-05, 2.0488e-05,
-5.1389e-05, -1.5774e-06, -1.3674e-04, 5.4398e-05],
[ 1.4171e-04, -6.2153e-04, -5.7166e-05, -4.4995e-04, 4.5814e-04,
1.1146e-03, 3.7766e-04, -8.4037e-05, -1.2185e-04, 1.4746e-04,
-2.1026e-04, -4.4517e-04, 2.3162e-04, 3.9549e-04],
[-8.5389e-05, -8.0975e-05, 2.0450e-04, 3.0029e-04, 1.2117e-04,
-5.9433e-04, 1.2826e-04, -2.0950e-04, 6.6218e-05, 3.4343e-06,
-3.3028e-05, 6.6400e-07, 4.8704e-06, 3.2811e-05],
[ 1.4581e-04, 2.1354e-04, 2.9725e-04, 3.1627e-04, 3.2114e-04,
-3.6244e-04, 3.8799e-05, -4.7348e-05, 7.5557e-06, -2.5671e-05,
-2.0939e-05, -8.2872e-05, 4.0288e-05, 1.2124e-05]]]]),),
'1_ReLU()': (tensor([[[[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 1.6532e-05, 0.0000e+00,
2.7187e-05, -5.4272e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, -2.6618e-04, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, -4.0658e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 1.1413e-04, 0.0000e+00, -2.4900e-04, 0.0000e+00, 0.0000e+00,
1.2080e-03, 9.8164e-05, 0.0000e+00, 4.4620e-05, 6.9611e-04,
7.9264e-04, 4.3300e-04, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, -2.3128e-04, 0.0000e+00,
0.0000e+00, -3.6258e-04, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 3.1046e-06, -2.5770e-04, -1.0208e-04, 0.0000e+00,
0.0000e+00, -8.6948e-05, 8.9369e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 2.9127e-05, -3.9120e-05],
[ 0.0000e+00, 0.0000e+00, 3.8731e-04, -8.5275e-05, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 9.6508e-04,
0.0000e+00, 0.0000e+00, -1.6144e-04, 4.1951e-05],
[ 8.2500e-05, 0.0000e+00, 0.0000e+00, -1.2196e-03, -1.3410e-04,
0.0000e+00, 3.4612e-04, 5.7151e-04, 0.0000e+00, 0.0000e+00,
-7.0793e-05, 0.0000e+00, 1.5663e-04, 0.0000e+00],
[-3.8814e-04, 0.0000e+00, 0.0000e+00, -7.9794e-04, -2.5181e-06,
0.0000e+00, 0.0000e+00, -7.4044e-04, 0.0000e+00, 2.6597e-04,
0.0000e+00, 0.0000e+00, -2.1567e-04, 0.0000e+00],
[ 7.9036e-05, 0.0000e+00, 0.0000e+00, -1.1428e-03, -7.9333e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 1.1661e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 2.4798e-04, 0.0000e+00, 3.6353e-04,
8.0010e-04, 4.8942e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[-1.9393e-04, 0.0000e+00, -2.0061e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 2.2458e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 7.4733e-05],
[ 6.4238e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00, 2.1598e-05,
0.0000e+00, -9.1805e-04, 0.0000e+00, 0.0000e+00, -1.9687e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 3.6472e-04],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, -3.2354e-04, 0.0000e+00,
2.3442e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
1.3121e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 1.4439e-04, 0.0000e+00, -8.1760e-05, 0.0000e+00,
3.2592e-04, -1.1950e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, -1.9019e-05, 0.0000e+00, 1.9975e-06]],
[[ 0.0000e+00, 0.0000e+00, -2.5228e-04, 9.5089e-06, 0.0000e+00,
2.8798e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
2.0486e-04, 6.1986e-04, 2.3585e-04, -1.4135e-05],
[ 0.0000e+00, 0.0000e+00, -3.2679e-04, 0.0000e+00, 0.0000e+00,
-2.5265e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, -4.9376e-04, 0.0000e+00],
[ 4.8207e-06, 0.0000e+00, 0.0000e+00, 0.0000e+00, 8.2737e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, -6.6958e-05, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, -6.5804e-05],
[ 1.7157e-04, -5.4876e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00,
-1.1919e-04, 4.3505e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00,
-3.0242e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 7.5319e-05, 1.5205e-04, 0.0000e+00, -2.6131e-05, 0.0000e+00,
-5.5548e-04, -1.0027e-04, 0.0000e+00, 0.0000e+00, -5.6776e-05,
-2.1956e-05, 0.0000e+00, -7.6074e-05, -4.6931e-05],
[ 0.0000e+00, 0.0000e+00, -1.4998e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 7.8140e-04, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 3.8212e-04,
-2.2113e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, -3.2592e-04, -3.3883e-04, -9.8969e-04, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 1.1952e-04, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 1.2047e-03, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
1.0831e-04, -1.6150e-04, 0.0000e+00, 9.7511e-05],
[ 1.4366e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00, -1.3583e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 8.1882e-05, 6.0266e-05,
8.7902e-05, 0.0000e+00, 0.0000e+00, -1.9107e-04],
[ 0.0000e+00, 2.9535e-04, -1.3569e-04, 5.4000e-04, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, -1.4563e-04, -5.2346e-04,
0.0000e+00, 1.4801e-04, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, -6.0840e-04, -7.7733e-04, -1.1043e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, -1.0888e-04, 0.0000e+00,
1.8711e-05, -3.0638e-05, 4.6634e-06, 4.8491e-05],
[-1.0344e-04, 0.0000e+00, 0.0000e+00, -2.7363e-04, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00]],
[[-4.0228e-05, 0.0000e+00, 0.0000e+00, -1.3240e-04, 0.0000e+00,
6.4997e-06, 0.0000e+00, 0.0000e+00, -9.0218e-05, 0.0000e+00,
0.0000e+00, 0.0000e+00, -1.4123e-04, -7.1466e-05],
[-3.2086e-04, -1.5973e-04, -7.7067e-04, 7.5304e-04, 0.0000e+00,
0.0000e+00, -8.7314e-05, 0.0000e+00, -2.4506e-04, 1.3508e-04,
0.0000e+00, 0.0000e+00, 1.0112e-04, 1.2568e-04],
[ 0.0000e+00, 5.9650e-04, 2.3969e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 1.6449e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, -2.4080e-04, 1.2067e-03, 0.0000e+00,
0.0000e+00, -8.7892e-04, 2.1312e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 0.0000e+00, -2.4368e-05],
[-5.7082e-05, 0.0000e+00, -3.4804e-05, 0.0000e+00, -1.4575e-04,
0.0000e+00, -6.4127e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00,
1.1528e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 4.9274e-05,
1.2956e-03, -5.5321e-05, 2.0192e-03, 0.0000e+00, 0.0000e+00,
7.0507e-05, 0.0000e+00, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
8.0113e-04, 0.0000e+00, 8.8753e-04, -3.0150e-04, 0.0000e+00,
0.0000e+00, -2.7025e-04, 0.0000e+00, 2.4008e-04],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 4.6083e-04, 0.0000e+00, -2.0996e-04, 0.0000e+00,
0.0000e+00, 6.2382e-05, -2.0128e-04, 0.0000e+00],
[ 0.0000e+00, -1.5684e-04, 2.5149e-04, 0.0000e+00, 0.0000e+00,
0.0000e+00, 1.6871e-04, 0.0000e+00, -1.0129e-04, 0.0000e+00,
0.0000e+00, -3.4259e-05, -4.7456e-06, -5.0054e-05],
[ 0.0000e+00, -2.0974e-04, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 0.0000e+00, 1.9965e-04, -3.3008e-04, 0.0000e+00,
-2.2731e-04, 1.6206e-04, 9.0695e-05, -9.0865e-05],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00, 0.0000e+00,
0.0000e+00, 1.6233e-04, -5.8411e-04, 0.0000e+00, 3.6661e-04,
5.4367e-05, 1.5689e-04, -2.3851e-04, -4.0392e-06],
[ 0.0000e+00, 0.0000e+00, 1.6129e-05, 5.4838e-04, 0.0000e+00,
0.0000e+00, 0.0000e+00, 5.9477e-04, -3.7406e-04, 0.0000e+00,
-2.4077e-04, 1.4221e-06, 7.3875e-05, 0.0000e+00],
[ 0.0000e+00, -2.0046e-04, 0.0000e+00, 7.6757e-05, 0.0000e+00,
0.0000e+00, 0.0000e+00, 7.9381e-05, 3.3896e-05, 0.0000e+00,
0.0000e+00, 4.9032e-05, 0.0000e+00, 0.0000e+00],
[ 0.0000e+00, 1.4361e-04, 0.0000e+00, 0.0000e+00, -2.9279e-05,
0.0000e+00, 0.0000e+00, -2.6094e-04, -5.7822e-06, 0.0000e+00,
-4.2728e-06, 0.0000e+00, 0.0000e+00, 6.4956e-06]],
[[ 1.6470e-04, -2.7624e-04, 1.7499e-04, 0.0000e+00, 6.9882e-07,
0.0000e+00, 1.5995e-04, -1.5952e-04, 1.4625e-05, 3.7111e-05,
4.3648e-05, 0.0000e+00, -1.3988e-04, 0.0000e+00],
[-1.7775e-04, 4.5784e-05, 1.1488e-03, 0.0000e+00, 3.3354e-04,
0.0000e+00, 0.0000e+00, -5.0907e-04, 0.0000e+00, 8.6607e-04,
4.3535e-04, 2.7672e-04, 0.0000e+00, 1.6021e-04],
[ 3.5031e-05, 0.0000e+00, 1.3538e-04, 2.5667e-04, -5.0345e-04,
-2.3831e-04, 2.9470e-04, 7.5360e-05, 1.1585e-04, 3.3693e-04,
-2.8915e-04, -5.5372e-04, 8.7012e-05, 3.2053e-05],
[-2.3625e-04, 0.0000e+00, -2.6603e-04, -1.7864e-03, 0.0000e+00,
4.1888e-04, 0.0000e+00, 0.0000e+00, -1.2718e-04, 7.6525e-05,
-5.0092e-04, 0.0000e+00, 1.2927e-05, 0.0000e+00],
[ 3.1950e-05, -4.0170e-04, 2.6143e-04, 7.5459e-05, -3.5227e-05,
3.1396e-04, 0.0000e+00, -6.0907e-04, 1.1948e-04, 0.0000e+00,
-6.4454e-07, 0.0000e+00, 0.0000e+00, -1.0788e-05],
[-2.5086e-04, -7.0095e-04, 2.9502e-04, -1.8155e-04, 8.2582e-04,
0.0000e+00, -6.1037e-04, 2.6006e-04, -1.5226e-04, 1.5867e-04,
0.0000e+00, 1.5612e-04, -2.2940e-04, -9.6533e-05],
[-4.3130e-04, -3.4917e-04, 2.3871e-04, 4.9014e-04, 0.0000e+00,
-1.2312e-04, -1.1016e-04, -4.1635e-05, -3.8083e-04, -2.5270e-04,
5.1527e-05, 3.7086e-04, -1.2186e-04, -7.2179e-05],
[ 0.0000e+00, 2.5626e-04, 4.1634e-05, -7.3664e-04, 5.5430e-04,
-2.6693e-04, 4.8458e-04, -5.3410e-04, 1.9737e-05, 8.2474e-04,
0.0000e+00, 0.0000e+00, -2.3831e-05, 2.4546e-04],
[ 0.0000e+00, 0.0000e+00, 0.0000e+00, 2.8456e-04, 2.3749e-04,
0.0000e+00, 0.0000e+00, 5.5760e-04, 8.9456e-05, -2.6423e-04,
-9.1239e-05, 6.6347e-05, 1.5629e-04, 0.0000e+00],
[ 0.0000e+00, -1.7646e-04, 1.1746e-04, -4.5712e-04, 6.0311e-04,
-8.7996e-04, -5.7582e-04, 6.8055e-04, 3.5117e-04, 2.5471e-04,
8.1237e-05, -1.0538e-04, 2.1575e-04, -1.5488e-04],
[ 3.8732e-05, -5.7904e-06, 0.0000e+00, 5.6333e-04, 0.0000e+00,
0.0000e+00, 7.2461e-04, -5.3970e-04, -1.7153e-05, 0.0000e+00,
0.0000e+00, -1.5774e-06, -1.3674e-04, 0.0000e+00],
[ 1.4171e-04, 0.0000e+00, -5.7166e-05, 0.0000e+00, 4.5814e-04,
0.0000e+00, 3.7766e-04, -8.4037e-05, 0.0000e+00, 1.4746e-04,
0.0000e+00, 0.0000e+00, 0.0000e+00, 3.9549e-04],
[ 0.0000e+00, -8.0975e-05, 2.0450e-04, 3.0029e-04, 1.2117e-04,
-5.9433e-04, 1.2826e-04, -2.0950e-04, 6.6218e-05, 0.0000e+00,
-3.3028e-05, 6.6400e-07, 4.8704e-06, 0.0000e+00],
[ 1.4581e-04, 0.0000e+00, 2.9725e-04, 0.0000e+00, 3.2114e-04,
-3.6244e-04, 3.8799e-05, -4.7348e-05, 7.5557e-06, 0.0000e+00,
-2.0939e-05, -8.2872e-05, 4.0288e-05, 1.2124e-05]]]]),),
'0_Conv2d(1, 4, kernel_size=(3, 3), stride=(2, 2), padding=(1, 1))': (None,)}
Activation分布の可視化
- 今までのHookの設定を活用して、内部共変量シフトの有無を確認する
- register_forward_hookを使って、各ReLU層の出力の平均と標準偏差の記録し、それぞれがどのように推移しているかを可視化する
- CNNを作成し、fashion MNISTデータセットで学習
# model準備
conv_model = nn.Sequential(
# 1x28x28
nn.Conv2d(1, 4, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 4x14x14
nn.Conv2d(4, 8, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 8x7x7
nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 16x4x4
nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1),
nn.ReLU(),
# 32x2x2 -> GAP -> 32 x 1 x 1
nn.Flatten(),
# # 128 -> 32
nn.Linear(128, 10)
# nn.Linear(32, 10)
# 10
)
# forward hook
# def save_out_stats(i, module, inp, out):
# act_means[i].append(out.mean().item())
# act_stds[i].append(out.std().item())
# relu_layers = [module for module in conv_model if isinstance(module, nn.ReLU)]
# for i, relu in enumerate(relu_layers):
# relu.register_forward_hook(partial(save_out_stats, i))
# Activation用クラス(上記forward hookをクラス化)
class ActivationStatistics:
def __init__(self, model):
self.model = model
self.act_means = [[] for module in self.model if isinstance(module, nn.ReLU)]
self.act_stds = [[] for module in self.model if isinstance(module, nn.ReLU)]
self.register_hook()
def register_hook(self):
relu_layers = [module for module in self.model if isinstance(module, nn.ReLU)]
for i, relu in enumerate(relu_layers):
relu.register_forward_hook(partial(self.save_out_stats, i))
def save_out_stats(self, i, module, inp, out):
# 学習データに対してのみActivationをtrackする
if self.model.training:
self.act_means[i].append(out.detach().cpu().mean().item())
self.act_stds[i].append(out.detach().cpu().std().item())
def get_statistics(self):
return self.act_means, self.act_stds
def plot_statistics(self):
fig, axs = plt.subplots(1, 2, figsize=(15, 5))
for act_mean in self.act_means:
axs[0].plot(act_mean)
axs[0].set_title('Activation means')
axs[0].legend(range(len(self.act_means)))
for act_std in self.act_stds:
axs[1].plot(act_std)
axs[1].set_title('Activation stds')
axs[1].legend(range(len(self.act_stds)))
plt.show()
act_stats = ActivationStatistics(conv_model)
# データ準備
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,) )
])
train_dataset = torchvision.datasets.FashionMNIST('./fmnist_data', train=True, download=True, transform=transform)
val_dataset = torchvision.datasets.FashionMNIST('./fmnist_data', train=False, download=True, transform=transform)
train_loader = DataLoader(train_dataset, batch_size=1024, shuffle=True, num_workers=4)
val_loader = DataLoader(val_dataset, batch_size=1024, num_workers=4)
opt = optim.SGD(conv_model.parameters(), lr=0.6)
train_losses, val_losses, val_accuracies = utils.learn(conv_model, train_loader, val_loader, opt, F.cross_entropy, 3)
# 結果の描画
act_stats.plot_statistics()
▶畳み込み層におけるバッチ正規化
- 畳み込み層にもバッチ正規化を用いるが、全結合層と少し異なる
- 各特徴マップの平均と分散は、全てのサンプルに対し、さらに全ての空間位置(H, W)にわたって計算する -> channel軸で平均、標準偏差を計算する
つまり、バッチ毎で正規化するのではなく、バッチをまたいで同じchannelについて正規化を行う
▶予測時のバッチ正規化
- 予測時はバッチサイズ=1になるので、どのように正規化するかを考える必要がある
-> 平均/標準偏差を計算できない - 学習全体を通して得られた全体の平均/標準偏差を使って正規化する
-> 移動平均と移動分散を使用し、学習データ全体の平均と分散を推定する

※$\gamma=0.7$とした場合、一つ前の結果70%に、新しい結果30%を追加(累積)するようなイメージ
バッチ正規化をスクラッチで実装
- 入力tensorに対して平均と標準偏差を計算する
- スケーリングとシフト用の係数(gamma, beta)も引数として受け取る
- 本来は学習可能なパラメータだが、今回の実装では引数として受け取る形を取る
- 入力tensorは[B, C, H, W]を想定する
- 予測時の対応は無視してよい
# 全結合の場合はX.shape = [b, out_features]なので、dim=(0)で平均,分散を求める
# バッチ正規化関数
def batch_norm(X, gamma=1, beta=0, eps=1e-5):
mean = X.mean(dim=(0, 2, 3), keepdim=True)
var = X.var(dim=(0, 2, 3), keepdim=True)
X_norm = (X - mean) / (torch.sqrt(var) + eps)
return gamma*X_norm + beta
# データ準備
X, y = train_dataset[0]
X.shape # -> torch.Size([1, 28, 28])
X = X /2 + 0.5
# 描画
plt.imshow(np.transpose(X, (1, 2, 0)), cmap='gray')
# バッチ正規化したデータに畳み込みのフィルターを適用して、特徴量を抽出
def apply_filter(im, filter):
im_h, im_w = im.shape
f_h, f_w = filter.shape
output_data = []
for i in range(im_h - f_h + 1):
row = []
for j in range(im_w - f_w + 1):
row.append((im[i:i+f_h, j:j+f_w] * filter).sum().item())
output_data.append(row)
return torch.tensor(output_data)
left_edge_filter = torch.tensor([[-1, 0, 1],
[-1, 0, 1],
[-1, 0, 1]])
def relu(X):
return torch.clamp(X, min=0)
X_ = X[0, :, :]
conv_out = apply_filter(X_, left_edge_filter)
# 平均と標準偏差を確認
print(conv_out.mean(), conv_out.std()) # -> tensor(0.0442) tensor(0.6041)
# 描画
plt.imshow(conv_out, cmap='gray')
# 次元を整えて、バッチ正則化
conv_out = conv_out[None, None, :, :]
norm_out = batch_norm(conv_out, )
norm_out.shape # -> torch.Size([1, 1, 26, 26])
# 平均と標準偏差を確認
print(norm_out.mean(), norm_out.std()) # -> tensor(-1.6224e-08) tensor(1.0000)
# ReLUで無駄な値を削除
relu_out = relu(norm_out)
relu_out = relu_out[0, 0, :, :]
# 描画
plt.imshow(relu_out, cmap='gray')
Pytorchでバッチ正規化層
- 畳み込み層と全結合層で異なる
- 畳み込み層では
nn.BatchNorm2d()を畳み込み層の後に配置する-
num_features:正規化される特徴マップの数で、通常直前の畳み込み層の出力channel数を指定する
-
- 全結合層の場合は
nn.BatchNorm1d()を全結合層の後に配置する-
num_features:正規化される特徴量の数で通常直前の全結合層の出力次元を指定する
-
X, y = next(iter(train_loader))
conv_out = nn.Conv2d(1, 8, kernel_size=3, stride=2, padding=1)(X)
norm_out = nn.BatchNorm2d(8)(conv_out)
X.shape # -> torch.Size([1024, 1, 28, 28])
# バッチ正規化のパラメータ(スケーリングとシフト)
list(nn.BatchNorm2d(8).parameters())
"""
[Parameter containing:
tensor([1., 1., 1., 1., 1., 1., 1., 1.], requires_grad=True),
Parameter containing:
tensor([0., 0., 0., 0., 0., 0., 0., 0.], requires_grad=True)]
index0がスケーリングで、index1がシフト
初期値はそれぞれ1.0と0.0で学習が進めば、この値も変動していく
"""
# CNNへ組み込み
def get_conv_model():
return nn.Sequential(
# 1x28x28
nn.Conv2d(1, 4, kernel_size=3, stride=2, padding=1),
nn.BatchNorm2d(4), # バッチ正規化
nn.ReLU(),
# 4x14x14
nn.Conv2d(4, 8, kernel_size=3, stride=2, padding=1),
nn.BatchNorm2d(8), # バッチ正規化
nn.ReLU(),
# 8x7x7
nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=1),
nn.BatchNorm2d(16), # バッチ正規化
nn.ReLU(),
# 16x4x4
nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1),
nn.BatchNorm2d(32), # バッチ正規化
nn.ReLU(),
# 32x2x2 -> GAP -> 32 x 1 x 1
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(32, 10)
# 10
)
# インスタンス化
conv_model_w_bn = get_conv_model()
# Activation
act_stats = ActivationStatistics(conv_model_w_bn)
# 学習
opt = optim.SGD(conv_model_w_bn.parameters(), lr=0.6)
train_losses, val_losses, val_accuracies = utils.learn(conv_model_w_bn, train_loader, val_loader, opt, F.cross_entropy, 3)
# 描画
act_stats.plot_statistics()

前回のバッチ正規化を使用していないものと比べるとかなり安定していることがわかる
▶バッチ正規化と正則化の効果
- バッチ正規化も副次的に正則化の効果がある
- バッチ単位で平均/標準偏差を計算する際にノイズが乗ることになり、そのノイズがそれぞれのactivationにのり正則化の効果を生む
- 目的はあくまでも学習の安定/高速化なので、正則化のために使用するものではない
▶バッチ正規化と重み減衰
- 理論的にはバッチ正規化が完全に正しくできていれば、重み減衰は意味をなさない
- バッチ正規化により結局0~1に正規化するので、重みが0.001でも1000でもスケールが出力に与える影響が軽減する
- 完全に適切に機能している場合、重みのスケールの影響を無視できるので、重み減衰をする必要がなくなる
- 実際にはバッチ正規化を完全に正しく機能させることは不可能で、バッチ正規化と重み減衰は組み合わせて使われる事は非常に一般的
▶レイヤー正規化(Layer Normalization)
- バッチ正規化と同じく正規化手法の一種で、バッチ正規化と異なり,一つのデータに対して正規化を行う
- バッチサイズに依存せずに適用できる
- 予測時の特別な対応は不要(学習時と同じ)
- RNNなどの再帰型NNに使用されることが多い
バッチ正規化は列について、レイヤー正規化は行について正規化を行う
上記画像の右側の図形を例にとって、画像について言うと下記のようなことになる
・バッチ正規化:全画像のR, G, Bごとに正規化を行う
・レイヤー正規化:それぞれの画像ごとに正規化を行う
layer Normalizationをスクラッチで実装
- 入力tensorに対して平均と標準偏差を計算する
- スケーリングとシフト用の係数(gamma, beta)も引数として受け取る
- 本来は学習可能なパラメータだが、今回の実装では引数として受け取る形を取る
- 入力tensorは[B, C, H, W]を想定する
def layer_norm(X, gamma=1, beta=0, eps=1e-5):
mean = X.mean(dim=(1, 2, 3), keepdim=True)
var = X.var(dim=(1, 2, 3), keepdim=True)
X_norm = (X - mean) / (torch.sqrt(var) + eps)
return gamma*X_norm + beta
X = torch.randn(5, 3, 3, 3) * 3. + 10
norm_out = layer_norm(X)
norm_out.shape # -> torch.Size([5, 3, 3, 3])
PytorchモジュールでLayer Normalization
-
nn.LayerNorm()を畳み込み層や全結合の直後に配置する-
normalized_shape:正規化するtensorのshapeを指定
※通常直前の畳み込み層の出力サイズ([C, H, W])を指定、全結合の場合は出力の次元を指定する
-
# CNNにLayer Normalizationを組み込み
def get_conv_model_ln():
return nn.Sequential(
# 1x28x28
nn.Conv2d(1, 4, kernel_size=3, stride=2, padding=1),
nn.LayerNorm([4, 14, 14]), # レイヤー正規化
nn.ReLU(),
# 4x14x14
nn.Conv2d(4, 8, kernel_size=3, stride=2, padding=1),
nn.LayerNorm([8, 7, 7]), # レイヤー正規化
nn.ReLU(),
# 8x7x7
nn.Conv2d(8, 16, kernel_size=3, stride=2, padding=1),
nn.LayerNorm([16, 4, 4]), # レイヤー正規化
nn.ReLU(),
# 16x4x4
nn.Conv2d(16, 32, kernel_size=3, stride=2, padding=1),
nn.LayerNorm([32, 2, 2]), # レイヤー正規化
nn.ReLU(),
# 32x2x2 -> GAP -> 32 x 1 x 1
nn.AdaptiveAvgPool2d(1),
nn.Flatten(),
nn.Linear(32, 10)
# 10
)
# インスタンス化
conv_model_w_ln = get_conv_model_ln()
# Activation
act_stats = ActivationStatistics(conv_model_w_ln)
# 学習
opt = optim.SGD(conv_model_w_ln.parameters(), lr=0.6)
train_losses, val_losses, val_accuracies = utils.learn(conv_model_w_ln, train_loader, val_loader, opt, F.cross_entropy, 3)
# 描画
act_stats.plot_statistics()

バッチ正規化に比べるとばらけているが、何もしていないものよりはよくなっている
次の記事
まだ












