10
10

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

More than 5 years have passed since last update.

chainerでsinとcosから角度を学習させてみた

10
Posted at

はじめに

前回まではsinから角度(theta)のように1入力1出力しか学習していませんでしたが、今回はsinとcosから角度のように2入力1出力を学習しようと思います。

環境

  • python: 2.7.6
  • chainer: 1.8.0

学習内容

sinとcos(0〜2π)から角度(theta)を学習する。

[training data]

  • input: sin, cos(0~2π, 1000分割)
  • output: theta
学習データ
def get_dataset(N):
    theta = np.linspace(0, 2 * np.pi, N)
    sin = np.sin(theta)
    cos = np.cos(theta)

    x = np.c_[sin, cos]
    y = theta

    return x, y

実装

ユニット数の設定

入力層(in_units)と出力層(out_units)のユニット数を設定するようにしたため前回までのコードより少しだけ汎用性が増えました。

バッチ学習
class MyChain(Chain):
    def __init__(self, in_units=1, n_units=10, out_units=1):
        super(MyChain, self).__init__(
             l1=L.Linear(in_units, n_units),
             l2=L.Linear(n_units, n_units),
             l3=L.Linear(n_units, out_units))
        self.in_units = in_units
        self.out_units = out_units

学習パラメータ

  • ミニバッチサイズ(batchsize): 20
  • エポック(n_epoch): 500
  • 隠れ層の数: 2
  • 入力層のユニット数(in_units): 2
  • 隠れ層のユニット数(n_units): 100
  • 出力層のユニット数(out_units): 1
  • 活性化関数: 正規化線形関数(relu)
  • ドロップアウト(dropout): なし(0%)
  • 最適化: Adam
  • 損失誤差関数: 平均二乗誤差関数(mean_squared_error)

パラメータは全て適当。

コード全体

全体
# -*- coding: utf-8 -*-

# とりあえず片っ端からimport
import numpy as np
import chainer
from chainer import cuda, Function, gradient_check, Variable, optimizers, serializers, utils
from chainer import Link, Chain, ChainList
import chainer.functions as F
import chainer.links as L
import time
from matplotlib import pyplot as plt

# データ
def get_dataset(N):
    theta = np.linspace(0, 2 * np.pi, N)
    sin = np.sin(theta)
    cos = np.cos(theta)

    x = np.c_[sin, cos]
    y = theta

    return x, y

# ニューラルネットワーク
class MyChain(Chain):
    def __init__(self, in_units=1, n_units=10, out_units=1):
        super(MyChain, self).__init__(
             l1=L.Linear(in_units, n_units),
             l2=L.Linear(n_units, n_units),
             l3=L.Linear(n_units, out_units))
        self.in_units = in_units
        self.out_units = out_units

    def __call__(self, x_data, y_data):
        x = Variable(x_data.astype(np.float32).reshape(len(x_data),self.in_units)) # Variableオブジェクトに変換
        y = Variable(y_data.astype(np.float32).reshape(len(y_data),self.out_units)) # Variableオブジェクトに変換
        return F.mean_squared_error(self.predict(x), y)

    def  predict(self, x):
        h1 = F.relu(self.l1(x))
        h2 = F.relu(self.l2(h1))
        h3 = self.l3(h2)
        return h3

    def get_predata(self, x):
        return self.predict(Variable(x.astype(np.float32).reshape(len(x),self.in_units))).data

# main
if __name__ == "__main__":

    # 学習データ
    N = 1000
    x_train, y_train = get_dataset(N)

    # テストデータ
    N_test = 900
    x_test, y_test = get_dataset(N_test)

    # 学習パラメータ
    batchsize = 10
    n_epoch = 500
    in_units = 2
    n_units = 100
    out_units = 1

    # モデル作成
    model = MyChain(in_units, n_units, out_units)
    optimizer = optimizers.Adam()
    optimizer.setup(model)

    # 学習ループ
    print "start..."
    train_losses =[]
    test_losses =[]
    start_time = time.time()
    for epoch in range(1, n_epoch + 1):

        # training
        perm = np.random.permutation(N)
        sum_loss = 0
        for i in range(0, N, batchsize):
            x_batch = x_train[perm[i:i + batchsize]]
            y_batch = y_train[perm[i:i + batchsize]]
            model.zerograds()
            loss = model(x_batch,y_batch)
            sum_loss += loss.data * batchsize
            loss.backward()
            optimizer.update()

        average_loss = sum_loss / N
        train_losses.append(average_loss)

        # test
        loss = model(x_test, y_test)
        test_losses.append(loss.data)

        # 学習過程を出力
        if epoch % 10 == 0:
            print "epoch: {}/{} train loss: {} test loss: {}".format(epoch, n_epoch, average_loss, loss.data)

    interval = int(time.time() - start_time)
    print "実行時間(normal): {}sec".format(interval)
    print "end"

    # 誤差のグラフ作成
    plt.plot(train_losses, label = "train_loss")
    plt.plot(test_losses, label = "test_loss")
    plt.yscale('log')
    plt.legend()
    plt.grid(True)
    plt.title("loss")
    plt.xlabel("epoch")
    plt.ylabel("loss")
    plt.show()

    # 学習結果のグラフ作成
    sin, cos = np.hsplit(x_test ,in_units)
    theta = y_test
    test = model.get_predata(x_test)

    plt.subplot(3, 1, 1)
    plt.plot(theta, sin, "b", label = "sin")
    plt.legend(loc = "upper left")
    plt.grid(True)
    plt.xlim(0, 2 * np.pi)
    plt.ylim(-1.2, 1.2)

    plt.subplot(3, 1, 2)
    plt.plot(theta, cos, "g", label = "cos")
    plt.legend(loc = "upper left")
    plt.grid(True)
    plt.xlim(0, 2 * np.pi)
    plt.ylim(-1.2, 1.2)

    plt.subplot(3, 1, 3)
    plt.plot(theta, theta, "r", label = "theta")
    plt.plot(theta, test, "c", label = "test")
    plt.legend(loc = "upper left")
    plt.grid(True)
    plt.xlim(0, 2 * np.pi)
    plt.ylim(-0.5, 7)

    plt.tight_layout()
    plt.show()

実行結果

誤差

エポック数が500回では誤差が大きいです。しかも誤差の低下具合が飽和してきているので大幅にエポック数を増やさないとこれ以上の低下は望めそうにありません。

resolver_simple_loss.png

学習結果

出力が線形なこともありそれほど誤差が大きいようには見えません。

resolver_simple_test.png

入力に搬送波を掛け合わせた場合

入力のsinとcosに搬送波を掛け合わせた場合も確認しました。(搬送波も入力して3入力1出力)
これはモータの角度を検出するレゾルバと同じ原理です。

搬送波ありの学習データ
def get_dataset(N):
    theta = np.linspace(0, 2 * np.pi, N)
    ref = np.sin(40 * theta) # 搬送波
    sin = np.sin(theta) * ref
    cos = np.cos(theta) * ref

    x = np.c_[sin, cos, ref]
    y = theta

    return x, y

# in_unitsを3に変更

入力が複雑になったため見た目でも誤差が大きいことが確認できます。高調波成分が残っています。

resolver_real_test.png

まとめ

誤差は大きいもののsinとcosから角度を学習することができました。(2入力1出力)
搬送波を掛け合わしたレゾルバの波形からも角度を学習することができました。ただし高調波成分が発生しました。
大したことがない学習内容だったので、もっと精度よく学習できると思っていたのですが誤差が大きくて驚いています。今後はもっと複雑な内容を学習させたいと思っているのでこの先が思いやれます…

10
10
0

Register as a new user and use Qiita more conveniently

  1. You get articles that match your needs
  2. You can efficiently read back useful information
  3. You can use dark theme
What you can do with signing up
10
10

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?