背景

書籍「ゼロから作るDeep Learning」(以降,「ゼロつく」と呼びます)のRust実装を研究室の勉強会で行うことになりました. 基本的に勉強会なので,どこまでやるかは自分次第です. 私は,自作モデルでMNISTデータセットの手書き数字認識の推論精度を95%以上にすることを目標としました.

目的

本記事では,後学のため,学習済みモデルを動作させて精度を計算する部分までをまとめます. つまり,第3章の手書き数字認識を学習済みモデルで解くところまでです.

手法

実験には研究室から貸与されたノートPCを使用しました. スペック詳細は割愛します. OSにはWSL(Windows Subsystem for Linux)上のUbuntu 24.04.4 LTSを用い,ビルドと実行はすべてこの環境で行いました. また,実験に際してエディタにはVSCodeを利用し,Rust公式が開発している拡張機能であるrust-analyzerを導入しました1

プロジェクトにはRustのワークスペース機能を利用しました2. ワークスペース機能は,複数のCargoプロジェクトをまとめて管理するための機能です. 採用した理由は,推論用のコードと学習用のコードを別々のプロジェクトに分けたかったためです. また,ワークスペースとして構成しておくことで,先述のrust-analyzerが各プロジェクトを正しく読み込めるという利点もあります. ディレクトリ構造は以下のとおりです.

$ tree -L 1
.
├── Cargo.lock
├── Cargo.toml
├── README.md
├── common # 共通で使用するコード(MNISTデータセットを読み出すコード等)
├── data # MNISTデータセット置き場
├── learn # 学習に使用するコード(Cargoプロジェクト)
├── predict # 推論を行うコード(Cargoプロジェクト)
├── target # ビルド済みバイナリ
└── weight # 重みとバイアスの保存場所

7 directories, 3 files

このうちlearnは,今後の学習フェーズの実装を見越して用意したもので,今回は中身が空の状態です. 本記事で実装したのは推論のみであるため,実際に使用したのは以下の4つのディレクトリです.

  • predict : ニューラルネットワークの構造体と推論ロジックの実装
  • common : シグモイド関数・ソフトマックス関数・MNISTデータ読み出しコード
  • weight : 学習済みの重み(pickle形式)と,pickle形式データをnpy形式に変換するためのPythonスクリプト
  • data : MNISTデータセット

以降では,各ディレクトリの構造と,含まれるソースコードについて説明します.

dataディレクトリ

dataディレクトリの構造は以下のとおりです.

$ tree
.
├── Makefile # unfreeze.shを呼び出すエイリアスとその説明等を書いた
├── README.md # dataディレクトリの説明
├── t10k-images-idx3-ubyte # テストデータ(10000件)
├── t10k-images-idx3-ubyte.gz 
├── t10k-labels-idx1-ubyte # テストデータのラベル(10000件)
├── t10k-labels-idx1-ubyte.gz 
├── train-images-idx3-ubyte # 訓練データ(60000件)
├── train-images-idx3-ubyte.gz 
├── train-labels-idx1-ubyte # 訓練データのラベル(60000件)
├── train-labels-idx1-ubyte.gz 
└── unfreeze.sh # gz形式のファイルを解凍するbashスクリプト

1 directory, 11 files

これらのうち,gz形式のファイルはCVDF(Common Visual Data Foundation)がGCS上で公開しているものです3. それぞれ以下のURLからダウンロードできます.

weightディレクトリ

weightディレクトリの構造は以下のとおりです.

$ tree
.
├── README.md
├── main.py # pickle形式からnpy形式にデータを変換するPythonスクリプト
├── output_npy # 変換後のnpy形式のファイル
│   ├── W1.npy
│   ├── W2.npy
│   ├── W3.npy
│   ├── b1.npy
│   ├── b2.npy
│   └── b3.npy
├── pyproject.toml
├── sample_weight.pkl
└── uv.lock

2 directories, 11 files

学習済みの重みデータには,ゼロつくのGitHubリポジトリで配布されているsample_weight.pklを使用しました4. pickle形式のファイルであるため,本来ならRustから直接読み込めるのが理想です. 実際,Rustにもpickle形式を扱うserde_pickleというクレートが存在します5. しかし,私が調べた限りでは,今回のsample_weight.pklの形式はserde_pickleのサポート範囲外でした.

サポート範囲外である理由は以下のとおりです. 書籍のPythonスクリプトを見る限り,重みデータはDict[str, np.ndarray]の形式で保存されていました. 素直に考えればHashMap<String, ndarray::Array2>形式で読み込めそうです. しかし,公式ドキュメントにはCurrently, this crate only supports Python’s built-in types that map easily to Rust constructs.と記述されています. つまり以下の図に示すように,Rustの型へ素直に対応づけられるPythonの組み込み型しか扱えません6

serde_pickleのドキュメント

np.ndarrayはPythonの組み込み型ではないため,それを値に持つDict[str, np.ndarray]serde_pickleでは読み込めませんでした.

そこで解決策として,Pythonでpickleファイルをnpyファイルに分解して書き出し,Rust側ではnpyファイルを1つずつ読み込む方式を採用しました. 先のディレクトリ構造にあるmain.pyが変換スクリプト,output_npyがその出力先にあたります.

commonディレクトリ

commonディレクトリの構造は以下のとおりです.

$ tree
.
├── Cargo.toml
└── src
    ├── lib.rs
    └── operation
        ├── functions.rs # シグモイド関数やソフトマックス関数など
        ├── mnist_data.rs # MNISTデータセットをndarray::Array2型に読み出すコード
        └── mod.rs

3 directories, 5 files

commonディレクトリはライブラリクレートとして作成しました7

lib.rsmod.rsはライブラリの公開設定のみを行っているため,詳細は割愛します.

functions.rsには,推論処理で使用するシグモイド関数とソフトマックス関数を記述しました.

// functions.rs
use ndarray::Array2;

// シグモイド関数(要素ごとに計算)
pub fn sigmoid(x: &Array2<f32>) -> Array2<f32> {
    x.mapv(|v| 1.0 / (1.0 + (-v).exp()))
}

// ソフトマックス関数 (バッチ対応)
// 分類問題を解くため
pub fn softmax(x: &Array2<f32>) -> Array2<f32> {
    let mut res = x.clone();
    // 画像ごとに処理
    for mut row in res.outer_iter_mut() {
        // 最大値を求める
        // NOTE: NEG_INFINITYを使う理由は,負の最大値と比較することで,必ずrowの中身の値で更新されるようにするため
        let max = row.fold(f32::NEG_INFINITY, |acc, &v| acc.max(v));
        // オーバーフロー対策
        // NOTE: 任意の数を引いてもソフトマックス関数の出力結果は変わらないため
        row.mapv_inplace(|v| (v - max).exp());
        // 合計値を求める
        let sum = row.sum();
        // 各行に対してソフトマックス関数の計算をする
        row.mapv_inplace(|v| v / sum);
    }
    res
}

mnist_data.rsには,MNISTデータセットをndarray::Array2型に読み出すコードを記述しました. 当初は独自定義型を用意していませんでしたが,型が複雑すぎるとclippyに指摘されたため,リファクタリングの過程で定義しました.

// mnist_data.rs
#![allow(unused)]

use mnist::{Mnist, MnistBuilder};
use ndarray::Array2;

// MnistDatasetの独自定義型
pub struct MnistDataset {
    pub train: Dataset,
    pub test: Dataset,
}
pub struct Dataset {
    pub data: Array2<f32>,
    pub label: Vec<f32>,
}

// MnistDatasetの読み出し
pub fn load_mnist(base_path: &str) -> MnistDataset {
    let Mnist {
        trn_img,
        trn_lbl,
        tst_img,
        tst_lbl,
        ..
    } = MnistBuilder::new()
        .base_path(base_path)
        .label_format_digit()
        .training_set_length(60_000)
        .test_set_length(10_000)
        .finalize();
    let train_data = Array2::from_shape_vec((60_000, 784), trn_img)
        .expect("訓練用データ変換エラー")
        .map(|x| *x as f32 / 255.0);
    let train_label: Vec<f32> = trn_lbl.into_iter().map(|x| x as f32).collect();
    let test_data = Array2::from_shape_vec((10_000, 784), tst_img)
        .expect("テスト用データ変換エラー")
        .map(|x| *x as f32 / 255.0);
    let test_label: Vec<f32> = tst_lbl.into_iter().map(|x| x as f32).collect();
    MnistDataset {
        train: Dataset {
            data: train_data,
            label: train_label,
        },
        test: Dataset {
            data: test_data,
            label: test_label,
        },
    }
}

predictディレクトリ

predictディレクトリの構造は以下のとおりです.

$ tree
.
├── Cargo.toml
├── Makefile # ベンチマークやビルド時のエイリアス等を記述
├── README.md
├── benches
│   └── criterion_bench.rs # 推論速度ベンチマーク用のコード
├── docs
│   ├── bench_reports # ベンチマークレポート結果比較用
│   └── report
│       └── index.html # ベンチマークレポート結果
└── src
    ├── lib.rs # runnerモジュールの公開設定
    ├── main.rs # メインロジックの呼び出しのみ実装
    └── runner
        ├── inference.rs # ニューラルネットワークの構造や推論処理のメインロジック等を実装
        └── mod.rs # メインロジックの公開設定

7 directories, 9 files

このディレクトリでは,sample_weight.pklに合わせたニューラルネットワークの推論処理を記述しました.

なお,benches以下には推論速度を計測するためのCriterionのベンチマークコードも置いています8. ただし本記事で扱うのは精度の算出までであるため,計測結果には触れません. また,main.rsとベンチマーク用のコードはいずれもinference.rsのメインロジックを呼び出しているだけであるため,詳細は割愛します.

inference.rsは以下のとおりです.

// inference.rs
#![allow(unused)]
use common::operation::functions::{sigmoid, softmax};
use common::operation::mnist_data::load_mnist;
use ndarray::{Array1, Array2};
use ndarray_npy::read_npy;
use std::path::Path;

// ネットワーク構造体を定義
struct Network {
    w1: Array2<f32>,
    w2: Array2<f32>,
    w3: Array2<f32>,
    b1: Array1<f32>,
    b2: Array1<f32>,
    b3: Array1<f32>,
}

impl Network {
    // 重み読み込み
    fn load_pretrained(npy_dir: &str) -> Result<Self, Box<dyn std::error::Error>> {
        let path = Path::new(npy_dir);
        Ok(Self {
            w1: read_npy(path.join("W1.npy"))?,
            w2: read_npy(path.join("W2.npy"))?,
            w3: read_npy(path.join("W3.npy"))?,
            b1: read_npy(path.join("b1.npy"))?,
            b2: read_npy(path.join("b2.npy"))?,
            b3: read_npy(path.join("b3.npy"))?,
        })
    }

    // 予測を行う
    fn predict(&self, x: &Array2<f32>) -> Array2<f32> {
        // 1層
        let a1 = x.dot(&self.w1) + &self.b1;
        let z1 = sigmoid(&a1);

        // 2層
        let a2 = z1.dot(&self.w2) + &self.b2;
        let z2 = sigmoid(&a2);

        // 3層
        let a3 = z2.dot(&self.w3) + &self.b3;
        softmax(&a3)
    }
}

pub fn main_logic() {
    // MNISTデータの読み込み
    let data_set_path = "../data";
    println!("MNISTデータの読み込み...");
    let dataset = load_mnist(data_set_path);

    // 学習済みの重みを読み込んでネットワークを構築
    println!("学習済みのパラメータを読み込み中...");
    let network =
        Network::load_pretrained("../weight/output_npy").expect("重みの読み込みに失敗しました");

    let x = &dataset.test.data;
    let t = &dataset.test.label;

    // 予測と精度の評価
    println!("推論を実行中...");
    let y = network.predict(x);

    let mut accuracy_count = 0;
    // yは(10000, 10)行列なので,一行ずつ取り出す
    for (i, row) in y.outer_iter().enumerate() {
        // 一番確率の高いインデックスを取得
        let (pred_label, _max_prob) = row
            .iter()
            .enumerate()
            .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
            .unwrap();

        // 正解ラベルと比較
        if pred_label == t[i] as usize {
            accuracy_count += 1;
        }
    }

    let accuracy = accuracy_count as f32 / t.len() as f32;
    println!("Accuracy: {:.4}", accuracy);
}

ここではネットワーク構造体を定義した上で,メソッドとして重み読み込みと推論実行処理を実装しました. 実装を構造体にまとめたのは,重みの名前空間をNetwork構造体の中に閉じ込めるためです.

実行結果

推論を行うには,predictディレクトリでcargo run --releaseを実行します. 以下は実行結果のログです.

$ cargo run --release
    Finished `release` profile [optimized] target(s) in 0.17s
     Running `/path/to/deep-learning-from-scratch-rs/target/release/predict`
MNISTデータの読み込み...
学習済みのパラメータを読み込み中...
推論を実行中...
Accuracy: 0.9352

学習済みモデルでの精度は93.52%となりました. これは書籍「ゼロつく」77ページ(3章6節)に記載された値と一致しており,実装した推論基盤が正しく動作していることを確認できました. なお,背景で述べた95%以上という目標は,自作モデルを学習させた場合のものです. 本記事で用いたのは書籍配布の学習済みモデルであり,その対象ではないため,ここでは93.52%を到達点としています.

まとめ

研究室の勉強会で,ゼロつくをRustで実装することになりました.

その中で私が最終的な目標に設定したのは,自作モデルでMNISTデータセットの手書き数字認識の推論精度を95%以上にすることです. 本記事では,その第一歩として,学習済みの重みを用いて推論を行う基盤の作成までをまとめました.

実験では,Rustのワークスペース機能を使い,学習フェーズと推論フェーズでプロジェクトを分離しました. また,学習と推論で同じように使用するコードはRustのライブラリクレート機能で管理しました.

一方で,課題として残ったのがパラメータの管理方法です. 今回は,pickle形式のデータをPythonでnpy形式に変換し,Rustから1ファイルずつ読み込みました. この方式では,層が増えるほど管理すべきnpyファイルも増えてしまいます.

そこで今後は,Dict[str, np.ndarray]形式のpickleファイルをRustから直接読み込めるようにしたいと考えています. 実装方針はまだ固まっていませんが,実現できればパラメータを1ファイルで扱えるようになり,データ読み出しコードの可読性も向上するはずです. 加えて,未実装の学習フェーズのコードを記述し,最終的な目的である推論精度95%の達成を目指します.