2024/05/14

KANでMNIST

KANとは


Kolmogorov–Arnold Networkの略で、Multi-Layer Perceptron (MLP) の代わりに使えるニューラルネットワークです。

MLPでは、入力データに対して重み付き線形和を計算し、活性化関数(例えばReLU)に通す、という処理を層の数だけ繰り返します。 和をとったあとに活性化関数を適用することになります。 入力が\(N_{\rm in}\)次元、出力が\(N_{\rm out}\)次元のとき、活性化関数は\(N_{\rm out}\)回だけ実行されます。

KANでは、入力データに対してB-スプライン曲線で学習可能にした活性化関数を通し、和を取る、という処理を層の数だけ繰り返します。 入力が\(N_{\rm in}\)次元、出力が\(N_{\rm out}\)次元のとき、活性化関数は\(N_{\rm in} \times N_{\rm out}\)回だけ実行されます。 また、活性化関数は学習可能なので、\(N_{\rm in} \times N_{\rm out}\)個のそれぞれ異なる形状の活性化関数が存在します。

B-スプライン曲線で活性化関数を書くといっても、お絵描きするときのように2Dの曲線を書いて活性化関数にするのではなく、基底関数の線形和を計算するだけです。 具体的には \[ {\rm spline}(x) = \sum_i c_i B_i(x) \] となります。\(B_i\)はB-スプライン曲線で指定した点を使って曲線を描くときにどのように補間するか(どの割合で点の位置を混ぜるか)を計算する関数です。 それを学習可能な\(c_i\)で混ぜて活性化関数とするわけです。

活性化関数\(\phi(x)\)には\({\rm spline}(x)\)を直接そのまま使うのではなく \[\phi(x) = w (b(x) + {\rm spline}(x))\] を用います。ここで、 \[b(x)={\rm silu}(x)=\frac{x}{1+e^{-x}} \] です。\(w\)は学習可能な重みです。ただし、Githubで公開されているコードを読むと \[\phi(x) = w_{\rm base} b(x) + w_{\rm sp} {\rm spline}(x)\] が使われているように見えます(KANLayer.pyを参照)。

B-スプライン曲線についてはhttps://techblog.kayac.com/generate-curves-using-b-splineとかhttp://web.mit.edu/hyperbook/Patrikalakis-Maekawa-Cho/node17.htmlが分かりやすいです。

KANの面白いところは、B-スプライン曲線で作った活性化関数が \(x^2\)、\({\rm exp}(x)\)、\({\rm sin}(x)\)、\({\rm log}(x)\)、\({\rm sqrt}(x)\)、\({\rm abs}(x)\) のようなユーザーが指定できる関数に十分近い場合はそれに置き換えてしまうことができる点です。

任意の入力に対して出力が0になる活性化関数を正則化によって増やし、それらを除去していくと、有効な活性化関数が人間が理解できる程度に少なくなることがあります。このとき、入力\(x\)に対して出力を\(y=f(x)\)で計算できる場合、関数\(f\)をユーザーが指定した活性化関数を使って作った合成関数と線形和、例えば \(f(x) = 1.2 \times {\rm sin} (x^2 - 0.3) + 0.5\) のような人間がみて分かる数式で出力することができます。

論文はhttps://arxiv.org/abs/2404.19756で、 コードはhttps://github.com/KindXiaoming/pykanにあります。 ここではリビジョン e6078bc8 を使います。

MNISTで学習させてみる


KAN Layerを使ってモデルを作り、MNISTで学習させてみます。

すべてKAN Layerで作ることもできるのですが、MNISTの画像は28×28=784と大きく、これを入力として32次元のベクトルを出力するようにすると、1層だけで784×32=25088個ものB-スプライン曲線を学習することになります。実行自体はできるのですが、非常に遅いため、ここでは最初にConv2Dで次元数を減らしてからKAN Layerを利用することにします。

学習用のコードは以下のとおりです。なお、著者が公開しているpykanのKAN.pyをベースに色々書き換えているので、もとのコードのライセンスに従い、このコードの部分はMITライセンスとします。

  1
  2
  3
  4
  5
  6
  7
  8
  9
 10
 11
 12
 13
 14
 15
 16
 17
 18
 19
 20
 21
 22
 23
 24
 25
 26
 27
 28
 29
 30
 31
 32
 33
 34
 35
 36
 37
 38
 39
 40
 41
 42
 43
 44
 45
 46
 47
 48
 49
 50
 51
 52
 53
 54
 55
 56
 57
 58
 59
 60
 61
 62
 63
 64
 65
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
import matplotlib.pyplot as plt
import numpy as np
import random
import torch
import torch.nn as nn
import torch.nn.functional as F
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
from pykan.kan.KANLayer import KANLayer
import sys

def initialize_seed(seed=0):
    torch.manual_seed(seed)
    np.random.seed(seed)
    random.seed(seed)

class ConvMLP(nn.Module):
    def __init__(self, fc_layers: list[int], device):
        super().__init__()
        self.conv1 = nn.Conv2d(1, 8, kernel_size=5, stride=2, device=device)
        self.conv2 = nn.Conv2d(8, 16, kernel_size=5, stride=2, device=device)
        n_in = 16
        fcs = []
        for fc in fc_layers:
            fcs.append(nn.Linear(n_in, fc, device=device))
            n_in = fc
        self.fcs = nn.ModuleList(fcs)

    def forward(self, x):
        x = self.conv2(F.relu(F.max_pool2d(self.conv1(x), 2)))
        x = x.reshape(x.shape[0], -1)
        for fc in self.fcs:
            x = fc(F.relu(x))
        return x

    def update_grid_from_samples(self, x):
        pass

    def regularize(self, lambda_l1, lambda_entropy, lambda_coef, lambda_coefdiff, small_mag_threshold=1e-16, small_reg_factor=1.0):
        return 0.0

# Modified version of KAN in pykan/kan/KAN.py
class ConvKAN(nn.Module):
    def __init__(self,
                 width: list[int],
                 grid=5,
                 k=3,
                 noise_scale=0.1,
                 noise_scale_base=0.1,
                 base_fun=torch.nn.SiLU(),
                 bias_trainable=True,
                 grid_eps=1.0,
                 grid_range=[-1, 1],
                 sp_trainable=True,
                 sb_trainable=True,
                 device="cpu"):
        super().__init__()

        ### Initialize feature extraction layers
        self.conv1 = nn.Conv2d(1, 8, kernel_size=5, stride=2, device=device)
        self.conv2 = nn.Conv2d(8, 16, kernel_size=5, stride=2, device=device)
        width.insert(0, 16)

        ### Initialize KAN layers
        self.biases = []
        self.act_fun = []
        self.depth = len(width) - 1
        self.width = width

        for l in range(self.depth):
            # splines
            scale_base = 1 / np.sqrt(width[l]) + (torch.randn(width[l] * width[l + 1], ) * 2 - 1) * noise_scale_base
            sp_batch = KANLayer(in_dim=width[l],
                                out_dim=width[l + 1],
                                num=grid,
                                k=k,
                                noise_scale=noise_scale,
                                scale_base=scale_base,
                                scale_sp=1.0,
                                base_fun=base_fun,
                                grid_eps=grid_eps,
                                grid_range=grid_range,
                                sp_trainable=sp_trainable,
                                sb_trainable=sb_trainable,
                                device=device)
            self.act_fun.append(sp_batch)

            # bias
            bias = nn.Linear(width[l + 1], 1, bias=False, device=device).requires_grad_(bias_trainable)
            bias.weight.data *= 0.0
            self.biases.append(bias)

        self.biases = nn.ModuleList(self.biases)
        self.act_fun = nn.ModuleList(self.act_fun)

    def forward(self, x):
        # Extract features by conv
        x = self.conv2(F.relu(F.max_pool2d(self.conv1(x), 2)))
        x = x.reshape(x.shape[0], -1)

        # Run KAN layers
        self.acts = [x] # acts shape: (batch, width[l])
        self.acts_scale = []

        for l in range(self.depth):
            x, preacts, postacts, postspline = self.act_fun[l](x)
            grid_reshape = self.act_fun[l].grid.reshape(self.width[l + 1], self.width[l], -1)
            input_range = grid_reshape[:, :, -1] - grid_reshape[:, :, 0] + 1e-4
            output_range = torch.mean(torch.abs(postacts), dim=0)
            self.acts_scale.append(output_range / input_range)

            x = x + self.biases[l].weight
            self.acts.append(x)

        return x

    def update_grid_from_samples(self, x):
        for l in range(self.depth):
            self.forward(x)
            self.act_fun[l].update_grid_from_samples(self.acts[l])

    def regularize(self, lambda_l1, lambda_entropy, lambda_coef, lambda_coefdiff, small_mag_threshold=1e-16, small_reg_factor=1.0):
        def nonlinear(x, th, factor):
            return (x < th) * x * factor + (x > th) * (x + (factor - 1) * th)

        reg_ = 0.
        for i in range(len(self.acts_scale)):
            vec = self.acts_scale[i].reshape(-1, )
            vec_sum = torch.sum(vec)
            if vec_sum == 0.0:
                continue

            p = vec / vec_sum
            l1 = torch.sum(nonlinear(vec, th=small_mag_threshold, factor=small_reg_factor))
            entropy = - torch.sum(p * torch.log2(p + 1e-4))
            reg_ += lambda_l1 * l1 + lambda_entropy * entropy  # both l1 and entropy

        # regularize coefficient to encourage spline to be zero
        for i in range(len(self.act_fun)):
            coeff_l1 = torch.sum(torch.mean(torch.abs(self.act_fun[i].coef), dim=1))
            coeff_diff_l1 = torch.sum(torch.mean(torch.abs(torch.diff(self.act_fun[i].coef)), dim=1))
            reg_ += lambda_coef * coeff_l1 + lambda_coefdiff * coeff_diff_l1

        return reg_

def calc_accuracy(ys):
    rs = []
    for y, label in ys:
        rs.append((torch.argmax(y, dim=1) == label).float())
    r = torch.cat(rs, dim=0)
    return torch.mean(r)*100.0

def train(model,
          train_loader,
          test_loader,
          max_epoch,
          lamb=0.0,
          lambda_l1=1.0,
          lambda_entropy=2.0,
          lambda_coef=0.0,
          lambda_coefdiff=0.0,
          update_grid=True,
          grid_update_freq=10,
          loss_fn=torch.nn.CrossEntropyLoss(),
          lr=0.002,
          device="cpu"):

    optimizer = torch.optim.Adam(model.parameters(), lr=lr)
    for epoch in range(max_epoch):
        model.train()
        n_samples = 0
        max_samples = len(train_loader.dataset)
        ys = []
        for iter, (x, label) in enumerate(train_loader):
            x = x.to(device)
            label = label.to(device)
            if iter % grid_update_freq == 0 and update_grid:
                model.update_grid_from_samples(x)
            y = model(x)
            ys.append((y, label))
            loss = loss_fn(y, label)
            reg_ = model.regularize(lambda_l1, lambda_entropy, lambda_coef, lambda_coefdiff)
            loss = loss + lamb * reg_
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            n_samples += len(x)
            if iter % 100 == 0:
                print(f"Epoch: {epoch} [{n_samples}/{max_samples}] Loss: {loss.item():.6f}")

        # Calc train accuracy
        train_acc = calc_accuracy(ys)

        # Calc test accuracy
        model.eval()
        ys = []
        with torch.no_grad():
            for iter, (x, label) in enumerate(test_loader):
                x = x.to(device)
                label = label.to(device)
                y = model(x)
                ys.append((y, label))
        test_acc = calc_accuracy(ys)
        print(f"Epoch: {epoch} [{n_samples}/{max_samples}] Loss: {loss.item():.6f} Acc(Train): {train_acc} Acc(Test): {test_acc}")

    return

def main(mode):
    initialize_seed(123)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    reg_lambda = 0.0
    update_grid = True,
    if mode == "kan":
        model = ConvKAN(width=[20, 10], device=device)
    elif mode == "kan-no-update-grid":
        model = ConvKAN(width=[20, 10], device=device)
        update_grid = False
    elif mode == "kan-reg":
        model = ConvKAN(width=[20, 10], device=device)
        reg_lambda = 0.003
    elif mode == "mlp":
        model = ConvMLP(fc_layers=[20, 10], device=device)
    else:
        return
    train_loader = DataLoader(datasets.MNIST("./data", train=True, download=True, transform=transforms.ToTensor()), batch_size=128, shuffle=True)
    test_loader = DataLoader(datasets.MNIST("./data", train=False, download=True, transform=transforms.ToTensor()), batch_size=128, shuffle=False)
    train(model, train_loader, test_loader, max_epoch=5, lamb=reg_lambda, update_grid=update_grid, device=device)

if __name__ == "__main__":
    main(sys.argv[1])

ConvMLP


比較用としてMLPを使った単純なモデルで学習させてみた結果が以下です。
Epoch: 0 [128/60000] Loss: 2.330844
Epoch: 0 [12928/60000] Loss: 0.527279
Epoch: 0 [25728/60000] Loss: 0.264145
Epoch: 0 [38528/60000] Loss: 0.268547
Epoch: 0 [51328/60000] Loss: 0.270200
Epoch: 0 [60000/60000] Loss: 0.137060 Acc(Train): 82.13166809082031 Acc(Test): 93.66999816894531
Epoch: 1 [128/60000] Loss: 0.211759
Epoch: 1 [12928/60000] Loss: 0.172556
Epoch: 1 [25728/60000] Loss: 0.216685
Epoch: 1 [38528/60000] Loss: 0.224366
Epoch: 1 [51328/60000] Loss: 0.166860
Epoch: 1 [60000/60000] Loss: 0.297034 Acc(Train): 94.53333282470703 Acc(Test): 95.87000274658203
Epoch: 2 [128/60000] Loss: 0.169406
Epoch: 2 [12928/60000] Loss: 0.060187
Epoch: 2 [25728/60000] Loss: 0.065257
Epoch: 2 [38528/60000] Loss: 0.209189
Epoch: 2 [51328/60000] Loss: 0.137371
Epoch: 2 [60000/60000] Loss: 0.062734 Acc(Train): 95.69499969482422 Acc(Test): 96.29000091552734
Epoch: 3 [128/60000] Loss: 0.161130
Epoch: 3 [12928/60000] Loss: 0.094211
Epoch: 3 [25728/60000] Loss: 0.137475
Epoch: 3 [38528/60000] Loss: 0.143321
Epoch: 3 [51328/60000] Loss: 0.086744
Epoch: 3 [60000/60000] Loss: 0.295325 Acc(Train): 96.2300033569336 Acc(Test): 96.52999877929688
Epoch: 4 [128/60000] Loss: 0.084705
Epoch: 4 [12928/60000] Loss: 0.114923
Epoch: 4 [25728/60000] Loss: 0.071916
Epoch: 4 [38528/60000] Loss: 0.093224
Epoch: 4 [51328/60000] Loss: 0.111265
Epoch: 4 [60000/60000] Loss: 0.049419 Acc(Train): 96.69833374023438 Acc(Test): 97.1500015258789

ConvKAN


KANを使ったモデルで学習させてみた結果が以下です。Conv2Dが色々吸収してしまっているのかもしれませんが、違いがほとんどありません。 パラメータ数はConvMLPよりConvKANのほうが多いです。
Epoch: 0 [128/60000] Loss: 2.308975
Epoch: 0 [12928/60000] Loss: 0.608790
Epoch: 0 [25728/60000] Loss: 0.333690
Epoch: 0 [38528/60000] Loss: 0.214927
Epoch: 0 [51328/60000] Loss: 0.171603
Epoch: 0 [60000/60000] Loss: 0.099660 Acc(Train): 88.74166870117188 Acc(Test): 96.1500015258789
Epoch: 1 [128/60000] Loss: 0.073898
Epoch: 1 [12928/60000] Loss: 0.215137
Epoch: 1 [25728/60000] Loss: 0.109934
Epoch: 1 [38528/60000] Loss: 0.126619
Epoch: 1 [51328/60000] Loss: 0.066091
Epoch: 1 [60000/60000] Loss: 0.188135 Acc(Train): 96.20833587646484 Acc(Test): 96.88999938964844
Epoch: 2 [128/60000] Loss: 0.061085
Epoch: 2 [12928/60000] Loss: 0.078620
Epoch: 2 [25728/60000] Loss: 0.045636
Epoch: 2 [38528/60000] Loss: 0.052172
Epoch: 2 [51328/60000] Loss: 0.037537
Epoch: 2 [60000/60000] Loss: 0.149782 Acc(Train): 96.87333679199219 Acc(Test): 95.19000244140625
Epoch: 3 [128/60000] Loss: 0.040984
Epoch: 3 [12928/60000] Loss: 0.102282
Epoch: 3 [25728/60000] Loss: 0.017132
Epoch: 3 [38528/60000] Loss: 0.043684
Epoch: 3 [51328/60000] Loss: 0.126490
Epoch: 3 [60000/60000] Loss: 0.057794 Acc(Train): 97.1483383178711 Acc(Test): 97.47000122070312
Epoch: 4 [128/60000] Loss: 0.056742
Epoch: 4 [12928/60000] Loss: 0.087390
Epoch: 4 [25728/60000] Loss: 0.046712
Epoch: 4 [38528/60000] Loss: 0.058726
Epoch: 4 [51328/60000] Loss: 0.217805
Epoch: 4 [60000/60000] Loss: 0.137546 Acc(Train): 97.54500579833984 Acc(Test): 97.52999877929688

ConvKANでgridの更新なし


KANを使ったモデルでgridの更新なしで学習させてみた結果が以下です。B-スプライン曲線による活性化関数は処理できる入力値の範囲が決まっており、gridの更新なしというのは、その範囲の調整を行わないということです。今回の設定では特に効果が無いようです。
Epoch: 0 [128/60000] Loss: 2.308975
Epoch: 0 [12928/60000] Loss: 0.502888
Epoch: 0 [25728/60000] Loss: 0.330806
Epoch: 0 [38528/60000] Loss: 0.242808
Epoch: 0 [51328/60000] Loss: 0.165698
Epoch: 0 [60000/60000] Loss: 0.191306 Acc(Train): 86.51166534423828 Acc(Test): 95.06999969482422
Epoch: 1 [128/60000] Loss: 0.090638
Epoch: 1 [12928/60000] Loss: 0.238503
Epoch: 1 [25728/60000] Loss: 0.128466
Epoch: 1 [38528/60000] Loss: 0.166372
Epoch: 1 [51328/60000] Loss: 0.120421
Epoch: 1 [60000/60000] Loss: 0.113729 Acc(Train): 95.87166595458984 Acc(Test): 96.37999725341797
Epoch: 2 [128/60000] Loss: 0.094055
Epoch: 2 [12928/60000] Loss: 0.101774
Epoch: 2 [25728/60000] Loss: 0.053376
Epoch: 2 [38528/60000] Loss: 0.050028
Epoch: 2 [51328/60000] Loss: 0.049892
Epoch: 2 [60000/60000] Loss: 0.045015 Acc(Train): 96.88333129882812 Acc(Test): 96.44000244140625
Epoch: 3 [128/60000] Loss: 0.091558
Epoch: 3 [12928/60000] Loss: 0.116339
Epoch: 3 [25728/60000] Loss: 0.035658
Epoch: 3 [38528/60000] Loss: 0.047689
Epoch: 3 [51328/60000] Loss: 0.121902
Epoch: 3 [60000/60000] Loss: 0.066078 Acc(Train): 97.5 Acc(Test): 97.30999755859375
Epoch: 4 [128/60000] Loss: 0.069760
Epoch: 4 [12928/60000] Loss: 0.048372
Epoch: 4 [25728/60000] Loss: 0.042325
Epoch: 4 [38528/60000] Loss: 0.073130
Epoch: 4 [51328/60000] Loss: 0.130568
Epoch: 4 [60000/60000] Loss: 0.046441 Acc(Train): 97.72833251953125 Acc(Test): 97.43999481201172

ConvKANで正則化あり


KANを使ったモデルで正則化ありで実行してみます。正則化のロスの重み\(\lambda\)は0.003にしています。
Epoch: 0 [128/60000] Loss: 2.449388
Epoch: 0 [12928/60000] Loss: 0.862422
Epoch: 0 [25728/60000] Loss: 0.502983
Epoch: 0 [38528/60000] Loss: 0.569132
Epoch: 0 [51328/60000] Loss: 0.391648
Epoch: 0 [60000/60000] Loss: 0.360877 Acc(Train): 89.41166687011719 Acc(Test): 95.72999572753906
Epoch: 1 [128/60000] Loss: 0.345524
Epoch: 1 [12928/60000] Loss: 0.449618
Epoch: 1 [25728/60000] Loss: 0.375242
Epoch: 1 [38528/60000] Loss: 0.334639
Epoch: 1 [51328/60000] Loss: 0.313511
Epoch: 1 [60000/60000] Loss: 0.374076 Acc(Train): 96.25166320800781 Acc(Test): 96.95999908447266
Epoch: 2 [128/60000] Loss: 0.237053
Epoch: 2 [12928/60000] Loss: 0.721758
Epoch: 2 [25728/60000] Loss: 0.468227
Epoch: 2 [38528/60000] Loss: 0.489218
Epoch: 2 [51328/60000] Loss: 0.366972
Epoch: 2 [60000/60000] Loss: 0.461017 Acc(Train): 90.59166717529297 Acc(Test): 94.61000061035156
Epoch: 3 [128/60000] Loss: 0.391373
Epoch: 3 [12928/60000] Loss: 0.405158
Epoch: 3 [25728/60000] Loss: 0.296408
Epoch: 3 [38528/60000] Loss: 0.306373
Epoch: 3 [51328/60000] Loss: 0.362922
Epoch: 3 [60000/60000] Loss: 0.335211 Acc(Train): 94.8933334350586 Acc(Test): 95.55999755859375
Epoch: 4 [128/60000] Loss: 0.343693
Epoch: 4 [12928/60000] Loss: 0.300828
Epoch: 4 [25728/60000] Loss: 0.310633
Epoch: 4 [38528/60000] Loss: 0.402329
Epoch: 4 [51328/60000] Loss: 0.401533
Epoch: 4 [60000/60000] Loss: 0.272522 Acc(Train): 95.64000701904297 Acc(Test): 95.79000091552734
\(\lambda=0.005\)にすると、以下のようになり途中でモデルが崩壊しました。
Epoch: 0 [128/60000] Loss: 2.542997
Epoch: 0 [12928/60000] Loss: 0.992837
Epoch: 0 [25728/60000] Loss: 0.666648
Epoch: 0 [38528/60000] Loss: 0.642962
Epoch: 0 [51328/60000] Loss: 0.496544
Epoch: 0 [60000/60000] Loss: 0.467141 Acc(Train): 88.85499572753906 Acc(Test): 95.80999755859375
Epoch: 1 [128/60000] Loss: 0.465318
Epoch: 1 [12928/60000] Loss: 0.551555
Epoch: 1 [25728/60000] Loss: 0.480575
Epoch: 1 [38528/60000] Loss: 0.532717
Epoch: 1 [51328/60000] Loss: 0.389326
Epoch: 1 [60000/60000] Loss: 0.464731 Acc(Train): 95.5183334350586 Acc(Test): 96.31999969482422
Epoch: 2 [128/60000] Loss: 0.360698
Epoch: 2 [12928/60000] Loss: 0.528383
Epoch: 2 [25728/60000] Loss: 0.377956
Epoch: 2 [38528/60000] Loss: 0.333951
Epoch: 2 [51328/60000] Loss: 0.379329
Epoch: 2 [60000/60000] Loss: 0.398343 Acc(Train): 94.8566665649414 Acc(Test): 96.06999969482422
Epoch: 3 [128/60000] Loss: 0.374904
Epoch: 3 [12928/60000] Loss: 0.928802
Epoch: 3 [25728/60000] Loss: 0.401091
Epoch: 3 [38528/60000] Loss: 0.539629
Epoch: 3 [51328/60000] Loss: 223.694611
Epoch: 3 [60000/60000] Loss: 3.871729 Acc(Train): 76.33499908447266 Acc(Test): 10.520000457763672
Epoch: 4 [128/60000] Loss: 4.122091
Epoch: 4 [12928/60000] Loss: 3.226183
Epoch: 4 [25728/60000] Loss: 2.973773
Epoch: 4 [38528/60000] Loss: 2.831389
Epoch: 4 [51328/60000] Loss: 2.789084
Epoch: 4 [60000/60000] Loss: 2.598794 Acc(Train): 9.819999694824219 Acc(Test): 10.09999942779541

2024/04/13

Landlock

Linuxでファイルシステムへのアクセス制限ができるlandlockのサンプルプログラムを試してみました。

詳細は https://qiita.com/nekoaddict/items/39125b8cd01da08b6a91 に詳しく書かれています。

Ubuntu22.04の5.15.0-84-genericでは https://git.kernel.org/pub/scm/linux/kernel/git/torvalds/linux.git/tree/samples/landlock/sandboxer.c?id=81709f3dccacf4104a4bc2daa80bdd767a9c4c54のコードで動くことが分かりました。 このコードをローカルファイル sandboxer.c に保存して、

gcc -o sandboxer sandboxer.c
とすると、コンパイルができます。

例えば

LL_FS_RO="/usr:/etc:/home/username/.bashrc" LL_FS_RW="/dev/null" ./sandboxer bash -i
のようにすれば、
  • /usr
  • /etc
  • /home/username/.bashrc
以下のみ読み込み、
  • /dev/null
だけに書き込めるようになっている状態でCLIからいろいろ実行できるようになります。
$ ls
ls: ディレクトリ '.' を開くことが出来ません: 許可がありません
$ ls /usr
bin  games  include  lib  <以下省略>
$ echo aaa > a
bash: a: 許可がありません
最新のカーネルではネットワークアクセスも制御できるようですが、Ubuntu22.04のデフォルトではできないようです。

2023/11/12

PrivateGPTの使い方メモ

はじめに


PrivateGPTを試したのでメモ。PrivateGPTのドキュメントは
https://docs.privategpt.dev/
に公開されており、少なくともLinux環境かつGPUを利用する条件では、このドキュメント通りにインストールすると使えるようになります。

使い方


インストールが完了すると、
PGPT_PROFILES=local make run
でローカル実行できます。具体的な実行コードはMakefileに記載されています。実行後、 http://localhost:8001/ にアクセスすると利用できます。

ドキュメントをPrivateGPTに取り込むには、

PGPT_PROFILES=local make ingest /path/to/docments
を実行します。8487440a6f8d135のリビジョンのコードでは、コードを改変しない限りディレクトリしか指定できません。実行すると、ディレクトリにあるファイルが解析されてPrivateGPTに追加されます。実行するたびにファイル名が同じものを除いて追加されていきます。削除方法はPrivateGPTのドキュメントに記載されています。

ドキュメントの取り込み時は単にサーバーを動かすときよりもGPUのメモリを使用するため、もしGPUのメモリが足りない場合は

CUDA_VISIBLE_DEVICES="" PGPT_PROFILES=local make ingest /path/to/docments
のようにして、GPUを見えなくすればCPUで処理してくれます。

2023/10/26

Diffusion MNIST その3

はじめに


その1で試したDiffusion MNISTについて、 ノイズの強さを表す時刻\(t\)をニューラルネットワークに伝えないとどうなるのかを見ていきます。

方法


https://github.com/MarceloGennari/diffusion_mnist のConditionalUNetの\(t\)が関連する行、つまり、TemporalEmbedding部分をコメントアウトします。具体的には
class ConditionalUNet(UNet):
    (省略)
    def forward(self, x: Tensor, t: Tensor, label: Tensor) -> Tensor:
        x0 = x #self.embedding1(x, t)
        x1 = self.block1(x0)
        x1 = self.label_emb1(x1, label)
        #x1 = self.embedding2(x1, t)
        x2 = self.block2(self.down1(x1))
        x2 = self.label_emb2(x2, label)
        #x2 = self.embedding3(x2, t)
        crossed = self.label_emb3(self.block3(self.down2(x2)), label)
        x3 = self.up1(self.attention1(crossed))
        x4 = torch.cat([x2, x3], dim=1)
        #x4 = self.embedding4(x4, t)
        x5 = self.up2(self.label_emb4(self.block4(x4), label))
        x6 = torch.cat([x5, x1], dim=1)
        x6 = self.label_emb5(x6, label)
        #x6 = self.embedding5(x6, t)
        out = self.out(self.block5(x6))
        return out
とします。

結果


その2で試した時間刻みを100にしたバージョンをベースに比較します。左側がTemporalEmbeddingありで、右側がなしに対応します。
t=50 TemporalEmbeddingあり
t=50 TemporalEmbeddingなし

t=0 TemporalEmbeddingあり
t=0 TemporalEmbeddingなし

TemporalEmbeddingなしの場合はノイズが多いように見えるので、画像の明るさをGIMPを使って上げたものが下図です。TemporalEmbeddingありではノイズが見えませんが、TemporalEmbeddingなしではノイズがはっきり見えるケースが多くなっています。
t=0 TemporalEmbeddingあり
t=0 TemporalEmbeddingなし

まとめ


時刻の埋め込みは効果があるということを確認できました。

2023/10/22

Diffusion MNIST その2

はじめに


その1で試したDiffusion MNISTについて、 ノイズを乗せるステップの細かさを粗くするとどうなるのかを見てみます。

方法


https://github.com/MarceloGennari/diffusion_mnist をいくつか変更することで粗さを変えていきます。

スケジュール変更


DiffusionProcessの引数に渡すvariance_scheduleを変えていきます。 デフォルトでは、
variance_schedule = torch.linspace(1e-4, 0.01, steps=1000)
となっています。これをパターンAでは
variance_schedule = torch.linspace(1e-4, 0.1, steps=100)
と、パターンBでは
variance_schedule = torch.linspace(1e-4, 0.999, steps=10)
とします。

それぞれのスケジュールを使ったときのalphaは

[パターン デフォルト]
[0.99990, 0.99989, 0.99988, ... , 0.99001, 0.99000]

[パターン A]
[0.99990, 0.99889, 0.99788, ... , 0.90101, 0.90000]

[パターン B]
[0.99990, 0.88891, 0.77792, ... , 0.11199, 0.00100]
となります。

alpha_barは

[パターン デフォルト]
[0.9999, 0.9998, 0.9997, ... , 0.0064, 0.0063]

[パターン A]
[0.9999, 0.9988, 0.9967, ... , 0.0062, 0.0056]

[パターン B]
[9.9990e-01, 8.8882e-01, 6.9143e-01, ... , 9.5131e-04, 9.5130e-07]
となります。ここで重要なことは、最初の時刻(ノイズが乗っていない)をt=0、最後の時刻(完全にノイズ)をt=1とするとき、alpha_barはt=0では1に近く、t=1では0に近くなるようにvariance_scheduleを決める必要があるということです。 各時刻tにおけるノイズの強さがalpha_barで決まり、t=1のときに完全にノイズになっていないと拡散プロセスの前提が崩れてしまうためです。

実際、パターンAの

variance_schedule = torch.linspace(1e-4, 0.1, steps=100)
variance_schedule = torch.linspace(1e-4, 0.01, steps=100)
に変えると、alpha_bar の値は
0.99990, 0.99970, 0.99940, ... , 0.60857, 0.60248
となりますが、この場合、数字の画像をうまく生成できません。

学習時のtの値


デフォルトではmain.py
t = torch.randint(0, 1000, (image.shape[0],))
の1000のところを、パターンAでは100に、パターンBでは10にします。

生成時のtの値


デフォルトではinference_unet.py
for t in trange(999, -1, -1):
の999のところを、パターンAでは99に、パターンBでは9にします。刻む数が少なくなると(ステップの細かさを粗くすると)、その分だけ生成時間を短くできます。

結果


デフォルトの設定ではこのようになります(その1の再掲)
t=500
t=0

時刻を100個に刻んだパターンAでも特に変わりなく生成できています。
t=50
t=0

時刻を10個に刻んだパターンBだと、多少ノイズが残ってしまいますが、生成できないというほどではありません。
t=5
t=0

まとめ


デフォルトの1000ステップではなくても、MNIST程度なら生成できることが分かりました。

2023/10/15

Diffusion MNIST

はじめに


MNISTデータセットを使って拡散モデルの学習と、学習したモデルを使った画像を生成をしていきます。

これを実現するコードがすでに
https://github.com/MarceloGennari/diffusion_mnist
で公開されていましたので、こちらを利用して試します。

なお、拡散モデルについての説明は検索するとたくさん出てきますので理論的背景については論文や解説記事を参照ください。 本記事では具体的に何をすれば拡散モデルを動かせるのかを見ていきます。

元論文はarXiv:2006.11239です。

まずは動かす


ソースコードをgit cloneで取得します。 ここでは、コミットハッシュがdf15ee746aのものを使用します。 README.mdを読んで、必要なモジュールをpipでインストールしておきます。

この時点では学習も生成もCPUで実行するように設定されているため、GPUで処理するように変更した後に実行します。

学習


学習はmain.pyで実行できるのですが、この中の
device = "cpu"
と書かれている行を
device = "cuda"
に書き換えます。そして
$ python main.py
を実行すると、学習が始まります。しばらく待っていると完了し、モデルのパラメータがunet_mnist.pthに記録されます。

生成


こちらもGPUで動くように変更します。また、学習はConditionalUNetで行われるものの生成はUNetになっているため、その点も修正します。

修正するコードはinference_unet.pyです。 先ほどと同じように

- device = "cpu"
+ device = "cuda"
に書き換えます。-が書き換え対象の行で、+が書き換えたあとの行の内容です。さらに、
- from models import UNet
+ from models import UNet, ConditionalUNet
と書き換え、
- model = UNet().to(device)
+ model = ConditionalUNet().to(device)
と書き換えます。ConditionalUNetにすると、生成時にどの数字を生成するかを指定する必要があるため、
     model.eval()
+    labels = torch.randint(0, 10, (batch_size,))
+    print(labels)
+    labels_cpu = labels
+    labels = labels.to(device)
     with torch.no_grad():
         for t in trange(999, -1, -1):
             time = torch.ones(batch_size) * t
-            et = model(xt.to(device), time.to(device))  # predict noise
+            et = model(xt.to(device), time.to(device), labels)  # predict noise
             xt = process.inverse(xt, et.cpu(), t)
のように、書き換えます。ランダムに0〜9の値をラベルとして指定するようにしています。

出力部分を少し書き換えて

    labels = ["Generated Images"] * 9
-
-    for i in range(9):
-        plt.subplot(3, 3, i + 1)
-        plt.tight_layout()
-        plt.imshow(xt[i][0], cmap="gray", interpolation="none")
-        plt.title(labels[i])
-    plt.show()
+            if t % 10 == 0:
+                plt.figure(figsize=(10, 10))
+                for i in range(25):
+                    plt.subplot(5, 5, i + 1)
+                    plt.tight_layout()
+                    plt.imshow(xt[i][0], cmap="gray", interpolation="none")
+                    plt.title(f"{labels_cpu[i]}")
+                plt.savefig(f"images/generated_t{t}.png")
+    #plt.show()
とすると、逆拡散過程によって少しずつ数字が画像として浮かび上がるところを見ることができます。 ただし、このようにすると、生成した数字のpyplotでの描画処理のために生成時間が延びます。GPU使用率もあきらかに低下します。 途中結果を見る必要がなければ最初のコードのほうが良いでしょう。

生成結果


結果は次のようになります。ただし、乱数がいろいろなところで使われており、シードも固定されていないので、毎回結果は異なります。

[t=990] ノイズだらけで何も読み取れません。

[t=700] そこはかとなく数字があるように見えるような見えないような。

[t=500] 遠くからみれば(ぼかしてみれば)、数字が簡単に読み取れます。

[t=300] ざらざらしていますが、十分に読めるようになりました。

[t=0] 完全にノイズが取り除かれました。アンチエイリアスは残ったままです。

学習時の処理


学習時の処理がどのように実装されているのかを見ていきます。

上位ループ


モデルの学習をするmain.pyの主要な処理である学習のループ部分を抜き出して、疑似コードとして書き換えると、
for epoch in range(100):
    for image, label in 学習用画像とそのラベルの集合:
	# 画像に加えるノイズの強さをランダムに選択
        t = torch.randint(0, 1000, (image.shape[0],))
        
        # 画像に加えるノイズを作成
        epsilon = torch.randn(image.shape)
        
        # tとepsilonに基づいてノイズをimageに加える
        diffused_image = process.forward(image, t, epsilon)

        # modelを使って加えられたノイズを予測
        optimizer.zero_grad()
        output = model(diffused_image, t, label)
        
        # 予測したノイズがどれだけ正しいかを評価して、モデルの重みを更新
        loss = criterion(epsilon, output) # criterion = torch.nn.MSELoss()
        loss.backward()
        optimizer.step()
となります。ノイズが加えられた画像から、加えられたノイズを推定するモデルを学習しているだけです。 それ以外は通常のニューラルネットワークの学習と変わりがありません。

そしてこれは元論文のAlgorithm 1の通りの実装です。5行目の


\( \nabla_\theta \| \bm{\epsilon} - \bm{\epsilon}_\theta ( \sqrt{\bar{\alpha_t}}\bm{x}_0 + \sqrt{1-\bar{\alpha}_t}\bm{\epsilon}, t) \|^2 \)

と見比べると、\(\bm{\epsilon}\)はコード上のepsilonに、 \(\bm{\epsilon_{\theta}}\)はコード上のmodelに、 \(\sqrt{\bar{\alpha_t}}\bm{x}_0 + \sqrt{1-\bar{\alpha}_t}\bm{\epsilon}\)はdiffused_imageに相当していることが分かります。

Algorithm 1とは異なり、modelの引数にlabelが余計についていますが、これは生成する数字が0〜9のどれであるかをニューラルネットワークに指示するための値となります。

ノイズ付加処理


上位ループ内の
diffused_image = process.forward(image, t, epsilon)
の部分について、詳細を見ていきます。

まず、diffusion_model.pyDiffusionProcessの初期化処理部分である__init__にて\(\alpha_t\)の値が計算されています。

最初に\(\beta_t\) (コード上ではvariance_schedule) を

self.variance_schedule = torch.linspace(1e-4, 0.01, steps=1000)
のように計算しています。具体的には
self.variance_schedule = [1.0000e-04, 1.0991e-04, 1.1982e-04, ... , 0.009980, 0.009990, 0.01]
となっています。元論文では\(\beta_1=10^4\)、\(\beta_T=0.02\)、\(T=1000\)と書かれているので\(\beta_T\)の値のみ異なっています。

\(\alpha_t = 1-\beta_t\)であるので、

self.alpha = 1 - self.variance_schedule
にて計算され、具体的な値は
self.alpha = [0.999900, 0.999890, 0.999880 , ... , 0.990020, 0.990010, 0.990000]
となります。さらに、\(\bar{\alpha}_t=\prod_{s=1}^t \alpha_s\)であるので、
self.alpha_bar = torch.cumprod(self.alpha, dim=0)
にて計算され、具体的な値は
self.alpha_bar = [0.999900, 0.999790, 0.999670, ... , 0.006430, 0.006365, 0.006302]
となります。

これらの値を使って、forwardにてノイズを画像に付加します。 \(\sqrt{\bar{\alpha_t}}\bm{x}_0 + \sqrt{1-\bar{\alpha}_t}\bm{\epsilon}\)が計算できれば良いので、まず\(\sqrt{1-\bar{\alpha}_t}\)を

std_dev = torch.sqrt(1 - self.alpha_bar[time_step])
で計算します。次に\(\sqrt{\bar{\alpha_t}}\)を
mean_multiplier = torch.sqrt(self.alpha_bar[time_step])
で計算し、最後に
diffused_images = mean_multiplier * x_0 + std_dev * noise
のように足し合わせます。x_0が\(\bm{x}_0\)で、noiseが\(\bm{\epsilon}\)です。

ノイズを予測するモデル


ノイズを予測するモデルには小さめのUNetが使われています。コードからグラフに書き起こすと下図のようになります。
背景が水色のボックス内にはクラス名とインスタンス化時の引数を記載しています(画像をクリックすると拡大できます)。

降りる方向と登る方向では条件付けの場所が異なっていますが、このあたりはおおまかであっても十分に動くのでしょう。

以下では、Pytorchで提供されているクラス以外のものを見ていきます。

ResConvGroupNorm


このクラスは以下のような構成になっています。残差接続のある、よく見かける畳み込み演算を使ったブロックです。

LabelEmbedding


このクラスは、0〜9のラベルをベクトルに変換し、入力値\(x\)に埋め込む処理をします。画像のチャネル方向に埋め込むので、各チャネルの画像特徴が埋め込みによって全体的に明るくなったり暗くなったりすることになります。

TemporalEmbedding


このクラスは、加えるノイズの強さを表す時刻\(t\)を入力値\(x\)に埋め込む処理をします。こちらも画像のチャネル方向に埋め込むので、各チャネルの画像特徴が埋め込みによって全体的に明るくなったり暗くなったりすることになります。LabelEmbeddingと干渉しそうですが、モデルの学習時にうまく棲み分けしているのでしょう。
SinusoidalPositionEmbeddingsはarXiv:1706.03762の \[ PE_{(pos,2i)} = sin(pos/10000^{2i/d_{model}}) \] \[ PE_{(pos,2i+1)} = cos(pos/10000^{2i/d_{model}}) \] とほぼ同じ計算をしています。\(sin\)と\(cos\)の引数部分のみ取り出すと、 \[ pos/10000^{2i/d_{model}} \] となり、\(pos\)を時刻\(t\)、\(d_{model}\)を単に\(d\)と書き換えると、 \[ t/10000^{2i/d} \] となります。iは次元方向のインデックスです。もう少し書き換えて、 \[ t/10000^{i/(d/2)} \] \(i/(d/2)\)の範囲が[0,1]になるように、少し値は変わりますが、 \[ t/10000^{i/(d/2-1)} \] とします。さらに、 \[ \begin{aligned} \frac{t}{\exp\left(\log 10000^{i/(d/2-1)}\right)} &= t \cdot \exp \left(-\log 10000^{i/(d/2-1)} \right) \\ &= t \cdot \exp \left(-\frac{i}{\frac{d}{2}-1} \log 10000 \right) \end{aligned} \] と変形します。ここまで変形するとコードとの対応がとれ、
half_dim = self.dim // 2
で\(d/2\)を計算し、
embeddings = math.log(10000) / (half_dim - 1)
で、\(\frac{1}{\frac{d}{2}-1} \log 10000\)を計算し、続いて
embeddings = torch.exp(torch.arange(half_dim, device=device) * -embeddings)
で、\(\exp \left(-\frac{i}{\frac{d}{2}-1} \log 10000 \right)\)を計算していることが分かります。

なお、変形前の\(1/10000^{i/(d/2-1)}\)にしたがって

1/torch.pow(10000, torch.arange(half_dim, device=device)/(half_dim-1))
で計算しても、結果はほとんど変わりません。具体的には、最初の5つの値を出力すると、
式変形前 [1.0000000000, 0.9821373820, 0.9645937681, 0.9473634958, 0.9304410219]
式変形後 [1.0000000000, 0.9821373224, 0.9645937085, 0.9473634958, 0.9304410219]
となるので、どちらで計算してもよさそうです。

さて、時刻\(t\)をかけると

embeddings = time[:, None] * embeddings[None, :]
となり、最後に、
embeddings = torch.cat((embeddings.sin(), embeddings.cos()), dim=-1)
のように\(sin\)と\(cos\)を通し、それらを連結することで、埋め込みベクトルが作成できます。

LinearAttention


attention1のクラスであるLinearAttentionの実装のアイデアの元と思われる論文は のあたりのようですが、ぴったり当てはまるものを見つけることはできませんでした。

注意機構部分のみをコードから計算式に戻すと、 \[ d_{t'f'} = \frac{1}{\sqrt{32}}\frac{1}{7 \times 7} \sum_f{\frac{\exp(q_{t'f})}{\sum_{f''}{\exp(q_{t'f''})}} \sum_t{\frac{\exp(k_{tf})}{\sum_{t''}\exp(k_{t''f})} v_{tf'}}} \] となっているようです。普通の注意機構とは異なり、最初に\(K\)の特徴次元数×\(V\)の特徴次元数(今の場合は32×32)の行列を計算しているようです。 \(\frac{1}{\sqrt{32}}\)の32は\(K\)の特徴次元数で、arXiv:1706.03762の式(1)である \[ Attention(Q,K,V)=\mathrm{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V \] の\(d_k\)に相当するようです。

生成時の処理


生成時の処理がどのように実装されているのかを見ていきます。

上位ループ


生成処理をしているinference_unet.pyの主要部分を疑似コードとして書き換えると、
for t in trange(999, -1, -1):
    # 時刻tとラベルlabelの条件の下、ノイズを推定する
    et = model(xt, t, label)

    # 推定したノイズを使って画像からノイズを除去する
    xt = process.inverse(xt, et, t)

# xtを画像として描画する
draw(xt)
となります。モデルの学習時は時刻tをランダムに選んでいましたが、 生成するときはノイズから徐々にノイズを取り除いていきます。t=999ではノイズのみ、t=0ではノイズがなくなった画像になります。

ノイズ除去処理


上位ループ内の
xt = process.inverse(xt, et, t)
の部分について詳細を見ていきます。

元論文arXiv:2006.11239のAlgorithm 2の4行目の処理 \[ \bm{x}_{t-1}=\frac{1}{\sqrt{\alpha_t}}\left(\bm{x}_t-\frac{1-\alpha_t}{\sqrt{1-\bar{\alpha}_t}} \epsilon_\theta(\bm{x}_t,t)\right)+\sigma_t\bm{z} \] が実装されています。DiffusionProcessinverseを見ていくと、

scale = 1 / torch.sqrt(self.alpha[t])
では\(\frac{1}{\sqrt{\alpha_t}}\)の部分が計算されています。
noise_scale = (1 - self.alpha[t]) / torch.sqrt(1 - self.alpha_bar[t])
では\(\frac{1-\alpha_t}{\sqrt{1-\bar{\alpha}_t}}\)の部分が計算されています。
std_dev = torch.sqrt(self.variance_schedule[t])
では\(\sigma_t\)が計算されています。
z = torch.randn(xt.shape) if t > 1 else torch.Tensor([0])
では\(\bm{z} \sim N(0,\bm{I})\)が計算されています。ガウス分布に従うノイズを作っているだけです。

最後にすべてをつなげて

mu_t = scale * (xt - noise_scale * et)
xt = mu_t + std_dev * z  # remove noise from image  
を計算することで、Algorithm 2の4行目の処理を実現しています。