はじめに
PRISMはプログラム内に学習機能を組み込んだ論理プログラミング言語であり、ルールベースの複雑な統計モデルを構築できます。一方で、人工知能(AI)や機械学習の研究が進む中、ニューラルネットワークはその中核技術として広く利用されています。こうした技術とルールを統合的に実装するために提案されているのが、PRISMのテンソル拡張であるT-PRISMです。
本記事では、T-PRISMの概要を解説し、どのようにしてルールなどの知識をニューラルネットワークなどの学習に組み込めるのかを紹介します。
本記事の実装(colab notebook)
T-PRISMの準備
まずは、T-PRISMをインストールしましょう。colabワークスペースなどでは、以下のコマンド(T-PRISM チュートリアル参照)を実行してインストールできます。
# GitHubから、事前ビルド済みのPRISMバイナリをダウンロードする
!wget "https://github.com/prismplp/prism/releases/download/v2.4.2a(T-PRISM)-prerelease/prism_linux_dev4colab.auto.zip"
# ダウンロードしたzipファイルを解凍する
!unzip -q -o prism_linux_dev4colab.auto.zip
# Protocol BuffersのPython実装を、C++実装ではなくPython実装に設定する
%env PROTOCOL_BUFFERS_PYTHON_IMPLEMENTATION=python
# PyPRISMをGitHub上のPyPRISMリポジトリからインストールする
!pip install -I "git+https://github.com/prismplp/pyprism.git"
# T-PRISMをGitHub上のPRISMリポジトリからインストールする
!pip install -I "git+https://github.com/prismplp/prism.git#egg=t-prism&subdirectory=bin"
これでT-PRISMの環境が整いました。次に、T-PRISMの基本概念について学びましょう。
ルールとTensor atom
Tensor atomは、T-PRISMにおける基本的なデータ構造です。
以下のコードで、Tensor atomの作成と、ルールの適用を行えます。
この例では、rel1は7種類の対象に関する相互作用的な関係に対応しています。
tensor_atom(rel1,[7,7]).
以下のような記法によってrel1に関して成立している論理ルールを定義することができます。この例ではrel1はrという述語から導かれています。
rel1(X,X):- member(X,[a,b,c,d,e,f,g]).
rel1(A,B):- r(A,B).
rel1(C,A):- r(A,B), r(B,C).
さらに以下のようにして追加で事前知識を入れることができます。
r(a,b).
r(b,c).
r(d,e).
r(e,f).
r(c,d).
r(d,g).
これをT-PRISMで実行すると、以下の推論結果が得られます。
array([[1., 1., 0., 0., 0., 0., 0.],
[0., 1., 1., 0., 0., 0., 0.],
[1., 0., 1., 1., 0., 0., 0.],
[0., 1., 0., 1., 1., 0., 1.],
[0., 0., 1., 0., 1., 1., 0.],
[0., 0., 0., 1., 0., 1., 0.],
[0., 0., 1., 0., 0., 0., 1.]])
Tensor atomはT-PRISMのルールにおける最小単位であり、同時にTensorでもあるため、行列分解やテンソル分解の対象にもなります。(詳しくはT-PRISM チュートリアル参照)
T-PRISMでニューラルネットワークを学習する
T-PRISMでは、ニューラルネットワークの構造を「論理ルール」として宣言的に書き、学習すべき重みをtensor_atomとして定義します。T-PRISM チュートリアルの多層パーセプトロンの例では、MNIST画像を入力として、以下のように2層のネットワークを記述しています。
tensor_atom(w(0),[10,256]). % 出力層の重み
tensor_atom(w(1),[256,784]). % 中間層の重み
tensor_atom(in(_),[784]). % 入力画像
output(Y,X):-layer0(X,Y).
layer0(X,Y):-operator(softmax), tensor(w(0),[i,j]),layer1(X,Y).
layer1(X,Y):-operator(sigmoid), tensor(w(1),[j,k]),layer2(X,Y).
layer2(X,Y):-tensor(in(X),[k]).
ここで重要なのは、PyTorchのようにforward()を逐次実装する代わりに、「どのテンソルをどう結合するか」をルールで書ける点です。tensor(w(1),[j,k])は中間層の重み行列、operator(sigmoid)は活性化関数を表し、これらをつないだ結果としてoutput(Y,X)が各クラスの予測値になります。つまり、ネットワークの計算グラフを論理プログラムとしてそのまま表現しているわけです。
学習の流れも比較的明快です。まず、教師データをoutput(正解ラベル, サンプル番号).というゴール集合として用意します。次にsave_expl_graph(...)で、そのゴールを導くための説明グラフを生成します。これは「どのルールとテンソル演算を通って出力に到達するか」を表した計算グラフです。その後、Python側でload_explanation_graph(...)とTprismModel(...)を使ってモデルを構築し、fit()を呼ぶと、tensor_atomとして定義した重みが勾配法で更新されます。
graph, tensor_shapes, flags = load_explanation_graph(expl_file, flag_file)
flags.embedding = ['data09/mnist.h5']
flags.vocab = "data09/vocab.pkl"
flags.sgd_learning_rate = 0.001
flags.max_iterate = 50
model = TprismModel(flags, tensor_shapes, graph, loss_cls=CE)
model.build(input_data=None, load_vocab=False, embedding_key="train")
evaluator = model.fit(verbose=True)
このように、T-PRISMでのニューラルネットワーク学習は、tensor_atomでパラメータを置き、operator(...)で非線形変換を指定し、最後にfit()で学習する、という3段階で理解できます。通常の深層学習フレームワークとの違いは、このネットワーク記述の中にそのままルールや記号的な制約を混ぜ込めることです。次の節では、この特徴を使って、ニューラルネットワークに医療知識を組み込む例を見ていきます。
ニューラルネットワークにルールを組み込む
T-PRISMで作成したニューラルネットワークをベースに、以下のような医療知識を組み込むことができます。
- 「高血糖」かつ「代謝関連リスク高」なら「糖尿病リスク高」
- 「高齢」かつ「代謝関連リスク高」なら「糖尿病リスク高」
- 「肥満」かつ「インスリン関連リスク高」なら「糖尿病リスク高」
ここでのポイントは、GlucoseやBMIのような観測できる特徴量をそのままルールに使うのではなく、まず「高血糖」「高齢」「肥満」のような分かりやすい条件をフラグとして定義し、それをニューラルネットワークが学習する潜在概念と組み合わせることです。今回のコードでは、代謝関連リスクとインスリン関連リスクという2つの潜在概念を別々のネットワークとして学習し、それらをルールで最終予測につなげています。
まず、各患者についてルールに使うフラグを作ります。pima_indians_diabetes_tprism_classification_0515_vscode_2.pyでは、以下の3つを明示的に定義しています。
# Explicit rule flags
rule_df["high_glucose_flag"] = (rule_df["Glucose"] >= 140).astype(int)
rule_df["obese_flag"] = (rule_df["BMI"] >= 30).astype(int)
rule_df["older_age_flag"] = (rule_df["Age"] >= 40).astype(int)
次に、これらのフラグをT-PRISMの事実として埋め込み、潜在概念を表す2つのニューラルネットワークと接続します。実際のコードではこれを1つの関数でまとめて定義していますが、記事としては分けて考えたほうが理解しやすいです。
まず、各患者の入力は9次元の表形式データとして与えます。
tensor_atom(get(tab,_), [9]).
続いて、先ほどPythonで作ったフラグをT-PRISM側の述語として読み込めるようにします。ここでの rule_high_glucose(X) などは、各サンプルについて事前に生成した事実です。
high_glucose(X) :- rule_high_glucose(X).
obese(X) :- rule_obese(X).
older_age(X) :- rule_older_age(X).
そのうえで、2つの潜在概念をニューラルネットワークで推定します。metabolic_nn は代謝関連リスク、insulin_nn はインスリン関連リスクに対応します。
metabolic_concept(Y, X) :-
operator(metabolic_nn),
vector(get(tab, X), [k]).
insulin_concept(Y, X) :-
operator(insulin_nn),
vector(get(tab, X), [k]).
最後に、潜在概念の出力と明示的なルール条件を組み合わせて最終判定を作ります。以下のルールは、「代謝関連リスクが高く、かつ高血糖なら糖尿病」「インスリン関連リスクが高く、かつ肥満なら糖尿病」といった形で読めます。
output(1, X) :-
metabolic_concept(1, X),
high_glucose(X).
output(1, X) :-
insulin_concept(1, X),
obese(X).
output(1, X) :-
metabolic_concept(1, X),
older_age(X).
つまり、このモデルでは「数値特徴量から潜在概念を学習する部分」と「その潜在概念を使って糖尿病を判定する部分」が分かれています。前者はニューラルネットワーク、後者はルールです。実装上はこれらを1つのT-PRISMプログラムとしてまとめて説明グラフに変換し、全体を通して学習します。
この構成では、通常の分類器のように1つのネットワークが直接「糖尿病あり・なし」を出力するのではありません。まず各患者の9次元入力特徴量から2つの潜在概念を推定し、その後で「潜在概念の出力」と「明示的なルール条件」を論理的に結合して最終判定を作ります。言い換えると、ニューラルネットワークはルールの中で使う中間概念を学習し、最終的な意思決定はルールが担っています。
学習時には、各サンプルの特徴量をdiabetes.h5に保存し、サンプルIDと教師ラベルをtrain.dat / test.datとして用意します。そのうえで、rule_high_glucose(ex123).のようなルール事実を生成し、save_expl_graph(...)で学習用・評価用の説明グラフを作ります。Python側では metabolic_nn と insulin_nn をそれぞれ BaseOperator として実装し、T-PRISMに登録して学習します。これにより、「数値データから潜在概念を学習する部分」と「専門知識に基づいて判定する部分」を1つの枠組みで扱えます。
モデルの評価
以下は、PyTorchで構築した通常のニューラルネットワークと、T-PRISMでルールを組み込んだモデルのテスト結果比較です。今回の設定では、ルールを組み込んだモデルのほうがAccuracy、Precision、F1-scoreでわずかに上回りました。一方でRecallは通常のNNのほうが高く、ルールを入れることで予測がやや保守的になっていることも分かります。
| Model | Accuracy | Precision | Recall | F1-score |
|---|---|---|---|---|
| NN (Pytorch) | 0.7273 | 0.5968 | 0.6852 | 0.6379 |
| NN + Rule (T-PRISM) | 0.7403 | 0.6207 | 0.6667 | 0.6429 |
今回の結果では、ルールを組み込んだことで「それらしい陽性判定」を出しやすくなる一方、ルールに合わない症例は拾いにくくなる可能性があります。つまり、精度の改善だけでなく、「どの条件で陽性と判定したのか」を説明しやすくなる点が大きな利点です。医療のように判断根拠が重要な場面では、このような解釈可能性の向上は、純粋な予測性能と同じくらい重要です。
まとめ
本記事では、T-PRISMの使い方と応用例ついて解説しました。Tensorの操作からニューラルネットワークの構築、トレーニング、評価までの流れを理解することで、T-PRISMでは専門知識を組み込んだモデルを構築できます。T-PRISMは柔軟で強力なツールであり、その活用方法を学ぶことで、PyTorchでは表現できない複雑な推論を実現することができます。
さらに詳細を学びたい方は、T-PRISMの公式ドキュメントを参照し、より高度なトピックにも挑戦してみてください!