0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

NotebookLMを使って、真剣にベイズ推定の問題を解いてみた【3.相互情報量を最大化する占い】

0
Last updated at Posted at 2026-03-13

はじめに

そもそものきっかけは、ベイズ推定とは何?という程度の知識レベルで、AtcoderのAH030を解こうとすると、適当なロジックで試行錯誤して、時間を浪費したところで大してスコアも伸びずギブアップするだけなので、ゼロから手を動かして知識を得たい。

一方で、従来型の独学での自習は限度があるので、NotebookLMを使うとどのようなメリットがあるのか試してみようかと思いました。

とっても参考になる題材はこちら→AHC典型解法シリーズ第3弾「ベイズ推定」

事前準備

コンパイルができるようにしたソースに、2.ランダムな占いと共通の修正を入れます。
特に乱数のブレに影響されるのには懲りた。

  • boolean DEBUG追加
  • int[][] DIJ追加(giveupのため)
  • Xorshiftクラスの中身を標準Randomに切り替え
  • Simクラスのquery、ansにログ追加
  • Simクラスのmine追加(giveupのため)
  • mainに共通の修正
  • now()、long startTime、Random rand追加

問題はあるが、あえて修正しない(byteの計算)

  • OilLayoutクラスのbyte[] volume;
  • Inputクラスのbyte[] getVolume(int[] topLefts)
  • OilStateクラスのList> topLeftQueryVolumes;
  • Stateクラスのbyte[] volumes;、moveTo、addQuery

seed0(N=15)でターン0の途中で4秒の強制終了(RuntimeException)。
seed20(N=10)でカクカクしながら、クエリ6回で正解。1.7秒かかる。

0029: pool.size()=360 total=41
0040:0 pool.get(0).pxIfR=0.002777777777777778
0350:1 pool.get(0).pxIfR=0.01980946762423068
0633:2 pool.get(0).pxIfR=0.0747073575523522
0915:3 pool.get(0).pxIfR=0.20875376533253384
1195:4 pool.get(0).pxIfR=0.2850241150288851
1468:5 pool.get(0).pxIfR=0.7048406672370957
1748:6 pool.get(0).pxIfR=0.9948737108880664
Score = 743523

比較のため、ランダムな占いでは、クエリ13回で正解。72ミリ秒(0.072秒)で終了。

0031: pool.size()=360 total=41
0043:0 pool.get(0).pxIfR=0.002777777777777778
0055:1 pool.get(0).pxIfR=0.011903008433828858
0057:2 pool.get(0).pxIfR=0.02185872979472808
0059:3 pool.get(0).pxIfR=0.05660424252495511
0061:4 pool.get(0).pxIfR=0.060522581772777244
0063:5 pool.get(0).pxIfR=0.08412290563641057
0064:6 pool.get(0).pxIfR=0.0937563643522413
0065:7 pool.get(0).pxIfR=0.10299773470508229
0067:8 pool.get(0).pxIfR=0.15715555916366344
0068:9 pool.get(0).pxIfR=0.19464025308765792
0069:10 pool.get(0).pxIfR=0.437790704628404
0070:11 pool.get(0).pxIfR=0.5942275205114096
0071:12 pool.get(0).pxIfR=0.6540323778263617
0072:13 pool.get(0).pxIfR=0.8312959909236446
Score = 1858524

つまり、相互情報量を最大化する山登りをすれば、性能はアップするが、とにかくどんくさい。
まずは、パフォーマンスに関するボトルネックを探して、速くすることを目的に改善する。

パフォーマンスの向上

ネックとなる部分を探す

メインループの主要部分を切り出せば、やはりflipとevalが怪しい。

        for (int t=0; true; t++) {
                for (int iter = 0; iter < 3; iter++) {
                    for (int ij : indices) {
                        q.flip(ij);
                        double nextEval = q.eval(sim, input);
                        else q.flip(ij);
                    }
                }
            }
        }

そこで、2つのメソッドに対して、カウントと実行時間の累計を取る。
本体をflip0とeval0に名前を変えて、計測用メソッドをはさむ。

        void flip(int ij) {
        	long mstart=System.currentTimeMillis();
        	try {
        		flip0(ij);
        	} finally {
        		flipcnt++;
        		fliptim+=(System.currentTimeMillis()-mstart);
        	}
        }
        double eval(Sim sim, Input input) {
        	long mstart=System.currentTimeMillis();
        	try {
        		return eval0(sim, input);
        	} finally {
        		evalcnt++;
        		evaltim+=(System.currentTimeMillis()-mstart);
        	}
        }

mainの最後に出力しているので、全体1746ミリ秒に対して、evaltimが1673ミリ秒なので、95.8%がevalにかかっている。
flipは何回実行されようが気にしなくてよい。

1746: evalcnt=1800
1748: evaltim=1673
1751: flipcnt=3097
1752: fliptim=5

likelihood計算のキャッシュ

eval0の中身を抜粋すると、sim.likelihoodが前半と後半の2重ループで呼ばれている。
まず、1度計算したパラメータをキャッシュ(メモ化)して、2度目はキャッシュを使うようにする。

        double eval0(Sim sim, Input input) {
            double[] pr_r = new double[sim.total + 2]; // 占い結果rの生起確率 p(r)
            for (int x = 0; x < pool.size(); x++) {
                for (int r = 0; r <= sim.total + 1; r++) {
                    pr_r[r] += sim.likelihood(mu, sigma, r) * pool.get(x).pxIfR;
                }
            }
            double info = 0.0;
            for (int x = 0; x < pool.size(); x++) {
                for (int r = 0; r <= sim.total + 1; r++) {
                    double p_r_x = sim.likelihood(mu, sigma, r);
                    if (p_r_x > SMALL_VALUE && pr_r[r] > SMALL_VALUE) {
                        info += p_r_x * pool.get(x).pxIfR * (Math.log(p_r_x) - Math.log(pr_r[r]));
                    }
                }
            }
            return info * Math.sqrt(k); // コスト 1/√k で割る = √k を掛ける
        }

muとsigmaの計算に、k,sと定数epsが必要なため、(k,s)をキーに、double[total+2]の値を管理する。

                double mu = (k - s) * input.eps + s * (1.0 - input.eps);
                double sigma = Math.sqrt(k * input.eps * (1.0 - input.eps));

kは1からn*n、sは0からtotal、totalの最大はn*n、この組み合わせがほぼ詰まっているので、Cacheクラスを作成し、2次元配列でキャッシュを管理する。mainにて、nとtotalが確定できた後、配列を確保する。

    static class Cache {
    	double[] pr_if_x;
    	Cache(double[] pr_r) {
    		this.pr_if_x = pr_r;
    	}
    }
    static Cache[][] cache;

//main()
        cache = new Cache[input.n2 + 1][input.total + 1];

getを追加し、キャッシュを参照し、なければ計算して登録する。
eval0の隣に作ったので、int kの情報はQueryクラスにあるが、Simクラスに移動した際に必要となるはずなので、明示しておく。

        Cache get(Sim sim, Input input, int k, int s) {
        	Cache c=cache[k][s];
        	if (c != null) return c;
            double[] pr_r = new double[sim.total + 2]; // 占い結果rの生起確率 p(r)
            double mu = (k - s) * input.eps + s * (1.0 - input.eps);
            double sigma = Math.sqrt(k * input.eps * (1.0 - input.eps));
            for (int r = 0; r <= sim.total + 1; r++) {
                pr_r[r] = sim.likelihood(mu, sigma, r);
            }
            c = new Cache(pr_r);
            cache[k][s] = c;
        	return c;
        }

eval0にて、getを呼び出すと、muとsigmaが不要となり、ちょっとすっきりする。

        double eval0(Sim sim, Input input) {
            if (k == 0) return -1e100;
            double[] pr_r = new double[sim.total + 2]; // 占い結果rの生起確率 p(r)
            for (int x = 0; x < pool.size(); x++) {
                int s = layoutVolumes[x];
                Cache c = get(sim, input, k, s);
                for (int r = 0; r <= sim.total + 1; r++) {
                    pr_r[r] += c.pr_if_x[r] * pool.get(x).pxIfR;
                }
            }
            double info = 0.0;
            for (int x = 0; x < pool.size(); x++) {
                int s = layoutVolumes[x];
                Cache c = get(sim, input, k, s);
                for (int r = 0; r <= sim.total + 1; r++) {
                    double p_r_x = c.pr_if_x[r];
                    if (p_r_x > SMALL_VALUE && pr_r[r] > SMALL_VALUE) {
                        info += p_r_x * pool.get(x).pxIfR * (Math.log(p_r_x) - Math.log(pr_r[r]));
                    }
                }
            }
            return info * Math.sqrt(k); // コスト 1/√k で割る = √k を掛ける
        }

seed20(N=10)で実行すると、カクカクだったのが、スパンと出る。
念のため、尤度の推移を比較し、一致することを確認する。
evaltimが1673ミリ秒から262ミリ秒に短縮した。
効果を実感するため、キャッシュにヒットした数(2回目以降)と、ミスした数(初回計算)もカウントして、表示したところ、1,970回の計算で済むところ、1,294,030回余計な計算をしていたらしい。

0030: pool.size()=360 total=41
0041:0 pool.get(0).pxIfR=0.002777777777777778
0144:1 pool.get(0).pxIfR=0.01980946762423068
0182:2 pool.get(0).pxIfR=0.0747073575523522
0222:3 pool.get(0).pxIfR=0.20875376533253384
0259:4 pool.get(0).pxIfR=0.2850241150288851
0295:5 pool.get(0).pxIfR=0.7048406672370957
0329:6 pool.get(0).pxIfR=0.9948737108880664
0330: evalcnt.hit=1294030
0333: evalcnt.mis=1970
0333: evalcnt=1800
0334: evaltim=262
0337: flipcnt=3097
0337: fliptim=1
Score = 743523

しかし、seed0(N=15)は3回クエリを出しただけで、giveupモードに移行している。
性能はまだまだ。

log計算の効率化

log計算がMath.log(p_r_x) - Math.log(pr_r[r])の2ヵ所がある。
前者は今回追加したCacheの中の値、後者はeval0の中で計算した値。

後者はループの前に移動するだけ。pr_r[r]は使わないので、対数値に上書き。
なお、後でlog(0)を計算すると-Infinityとなることが分かるが、今は何もしない。

            for (int r = 0; r <= sim.total + 1; r++) {
            	pr_r[r] = Math.log(pr_r[r]);
            }

このままMath.log(pr_r[r])pr_r[r]に置き換えただけでは、pr_r[r] > SMALL_VALUEの判定が意味をなさない(pr_r[r]は負の数ばかり)ので削除する。

前者はCacheの中に対数も追加し、コンストラクタで計算してしまう。

    static class Cache {
    	double[] pr_if_x;
    	double[] ln_pr_if_x;
    	Cache(double[] pr_r) {
    		this.pr_if_x = pr_r;
    		this.ln_pr_if_x = new double[pr_if_x.length];
            for (int r = 0; r < pr_if_x.length; r++) {
            	ln_pr_if_x[r] = Math.log(pr_if_x[r]);
            }
    	}
    }

実行すると、evaltimが262ミリ秒から92ミリ秒に短縮した。
なぜかキャッシュヒット、ミス数が前回から5個変わっているが、実は意味不明。
尤度の推移の推移は変わっていないので、気にしない。

0030: pool.size()=360 total=41
0042:0 pool.get(0).pxIfR=0.002777777777777778
0120:1 pool.get(0).pxIfR=0.01980946762423068
0130:2 pool.get(0).pxIfR=0.0747073575523522
0138:3 pool.get(0).pxIfR=0.20875376533253384
0147:4 pool.get(0).pxIfR=0.2850241150288851
0156:5 pool.get(0).pxIfR=0.7048406672370957
0164:6 pool.get(0).pxIfR=0.9948737108880664
0165: evalcnt.hit=1294025
0168: evalcnt.mis=1975
0168: evalcnt=1800
0168: evaltim=92
0172: flipcnt=3099
0172: fliptim=9
Score = 743523

seed0(N=15)では、正解までたどり着いた。
元々は5000万回も余計にlikelihoodを呼んでいたのか。

0030: pool.size()=7623 total=38
0047:0 pool.get(0).pxIfR=1.3118194936376755E-4
0458:1 pool.get(0).pxIfR=0.0012722450518718035
0806:2 pool.get(0).pxIfR=0.01607010382583862
1156:3 pool.get(0).pxIfR=0.09382865777370121
1504:4 pool.get(0).pxIfR=0.3717922675106972
1862:5 pool.get(0).pxIfR=0.8519443643710758
1863: evalcnt.hit=51450317
1867: evalcnt.mis=4933
1867: evalcnt=3375
1868: evaltim=1589
1872: flipcnt=5811
1872: fliptim=184
Score = 387269

自分も含めて、普通の人はここまでじゃないかな。
後は正しい計算をやっているに過ぎない。

正規分布の値の散らばり

正規分布は、平均をピークにして、左右になだらかに減衰する形になってます。
知らなくても、とりあえず計算したキャッシュの中身をダンプしてみればいいです。

k=10 s=10
[0.0, 0.0, 2.886579864025407E-15, 1.533295712619065E-11, 2.1488135271141573E-8,
8.043096030863062E-6, 8.162022359360366E-4, 0.02293920551220724, 0.1835032096941212, 0.42944731154816285,
0.29844032788088315, 0.06117497205864908, 0.0036107961642023456, 5.964083818565946E-5, 2.6914192585714147E-7,
3.261197978332575E-10, 1.0480505352461478E-13, 0.0, 0.0, 0.0,
0.0, 0.0, 0.0, 0.0, 0.0,
0.0, 0.0, 0.0, 0.0, 0.0,
0.0, 0.0, 0.0, 0.0, 0.0,
0.0, 0.0, 0.0, 0.0, 0.0,
0.0, 0.0, 0.0]

total=41だったので、配列のサイズは43。中に0.0が28個あります。
何気にif (p_r_x > SMALL_VALUE) {で計算をスキップしているようで、for (int r = 0; r <= sim.total + 1; r++)のループの空回り自体がコストになっています。

なお、この中の最大値はr=9の0.42944で、eps=0.08のため、(10-10)*0.08+10*(1-0.08)=9.2からmu=9と一致します。sigma^2=10*0.08*(1-0.08)=0.736で、1e-6の閾値で絞ると前後4個の合計9個あればよいことが分かります。

ちなみにAIチャットで、この配列に必要なサイズを見積もってもらったら、n=20、eps=0.2のときにk=400、sigma^2=64、sigma=8で、mu=320。eps=0.01のとき、mu=396。なぜか、mu=396とsigma=8で、5*sigma程度を確保すると、400に対して40くらい余計に確保すればと言われたが、epsが異なるからどこか間違ってますな。

さて、今のtotal + 2のままの大きさの配列を使い、muから小さい側とmuから大きい側の2回ループし、閾値SMALL_VALUEで打ち切ります。
cpp版では打ち切った位置がlb(おそらくLower Boundの意味)としているので、大きい方はub(Upper Bound)にします。

            int lb = 0;
            for (int r = (int)Math.round(mu); r >= 0; r--) {
            	double v = sim.likelihood(mu, sigma, r);
            	if (v < SMALL_VALUE) {
            		lb = r + 1;
            		break;
            	}
                pr_r[r] = v;
            }
            int ub = 0;
            for (int r = (int)Math.round(mu) + 1; true; r++) {
            	double v = sim.likelihood(mu, sigma, r);
            	if (v < SMALL_VALUE) {
            		ub = r;
            		break;
            	}
                pr_r[r] = v;
            }
            c = new Cache(pr_r, lb, ub);

Cache側でlbからubまでを切り取って保管し、lb値も保持します。1

    	Cache(double[] pr_r, int lb, int ub) {
    		this.lb = lb;
    		this.pr_if_x = Arrays.copyOfRange(pr_r, lb, ub);
    		this.ln_pr_if_x = new double[pr_if_x.length];
            for (int r = 0; r < pr_if_x.length; r++) {
            	ln_pr_if_x[r] = Math.log(pr_if_x[r]);
            }
    	}

使う側としては、
pr_r[r]pr_r[r + c.lb]の下駄を履かせます。
もうif (p_r_x > SMALL_VALUE) {は要らないです。
pool.get(x).pxIfRもループ前に出せるけど、まあそのまま。

                Cache c = get(sim, input, k, s);
                for (int r = 0; r < c.pr_if_x.length; r++) {
                    pr_r[r + c.lb] += c.pr_if_x[r] * pool.get(x).pxIfR;
                }

                Cache c = get(sim, input, k, s);
                for (int r = 0; r < c.pr_if_x.length; r++) {
                    double p_r_x = c.pr_if_x[r];
                    double ln_p_r_x = c.ln_pr_if_x[r];
                    info += p_r_x * pool.get(x).pxIfR * (ln_p_r_x - pr_r[r + c.lb]);
                }

これで実行すると、実はtotal+2では足りない。total+8でも足りない。total+9で動いた。
ここが固定長配列の難しいところ。2

Exception in thread "main" java.lang.ArrayIndexOutOfBoundsException: Index 43 out of bounds for length 43

AIの回答はtotal+40くらいで十分そうだけど、total+100あれば足りなくなることはないんじゃないという見解に従う。

ダンプにlbを追加して、同じk=10 s=10を出力すると、有用な値のみに圧縮されている。

k=10 s=10 lb=5
[8.043096030863062E-6, 8.162022359360366E-4, 0.02293920551220724, 0.1835032096941212, 0.42944731154816285,
0.29844032788088315, 0.06117497205864908, 0.0036107961642023456, 5.964083818565946E-5]

seed20(N=10)とseed0(N=15)で実行すると、seed20のクエリが1つ減っているのはただの運。
evaltim/evalcntが、92ミリ秒/1800回と87ミリ秒/1500回はあまり変わっていない。
seed0はevaltimが1589ミリ秒から1108ミリ秒、全体でも500ミリ秒弱短縮している。

execute 0020
0029: pool.size()=360 total=41
0041:0 pool.get(0).pxIfR=0.002777777777777778
0098:1 pool.get(0).pxIfR=0.01980946762423068
0136:2 pool.get(0).pxIfR=0.09390202362604827
0144:3 pool.get(0).pxIfR=0.25813033301337246
0151:4 pool.get(0).pxIfR=0.6505518140209513
0157:5 pool.get(0).pxIfR=0.9612801235559135
0158: evalcnt.hit=1078152
0160: evalcnt.mis=1848
0160: evalcnt=1500
0161: evaltim=87
0164: flipcnt=2653
0165: fliptim=6
Score = 682701

execute 0000
0032: pool.size()=7623 total=38
0049:0 pool.get(0).pxIfR=1.3118194936376755E-4
0368:1 pool.get(0).pxIfR=0.00121782200745726
0623:2 pool.get(0).pxIfR=0.01833915888421839
0872:3 pool.get(0).pxIfR=0.16474092737979626
1121:4 pool.get(0).pxIfR=0.5364608401147308
1388:5 pool.get(0).pxIfR=0.9960255129583347
1389: evalcnt.hit=51450246
1392: evalcnt.mis=5004
1392: evalcnt=3375
1392: evaltim=1108
1396: flipcnt=5830
1397: fliptim=184
Score = 389358

相互情報量の確認、山登りの修正

きっとcpp版と同程度の速度は出るようになったはずなので、本題の相互情報量の確認と、山登りの修正(必ずしも改善とは言えない)を行う。

まずmainからクエリ生成部分を、List<Integer> getDivinationQuery(Input input, List<OilLayout> pool, Sim sim, int t)に切り出しておく。
先に山登りの修正をしないといけないのが、今はshuffleしたリストの順にクエリを選んでいるだけで、そもそも各マスの評価値を計算していない。

        double[] evals = new double[input.n2];
        List<Integer> indices = new ArrayList<>();
        for (int ij = 0; ij < input.n2; ij++) {
        	q.flip(ij);
        	evals[ij] = q.eval(sim, input);
        	q.flip(ij);
        	indices.add(ij);
        }
        Collections.sort(indices, (a, b) -> Double.compare(evals[b], evals[a]));

ここでseed20で、evalsをダンプしてみる。

t=0
-7.14e-15    0.0908     0.165     0.271     0.271     0.271     0.271     0.209 -7.14e-15 -7.14e-15 
   0.0908     0.325     0.411     0.530     0.572     0.603     0.544     0.474     0.240     0.137 
    0.209     0.499     0.616     0.683     0.691     0.691     0.674     0.603     0.355     0.209 
    0.302     0.578     0.683     0.668     0.605     0.641     0.695     0.661     0.463     0.209 
    0.302     0.511     0.672     0.666     0.642     0.685     0.691     0.616     0.325    0.0908 
    0.209     0.401     0.643     0.695     0.682     0.694     0.699     0.646     0.411     0.165 
    0.165     0.282     0.599     0.682     0.684     0.684     0.677     0.555     0.282    0.0908 
    0.165     0.282     0.580     0.646     0.656     0.656     0.633     0.446     0.228    0.0908 
    0.165     0.282     0.507     0.555     0.572     0.572     0.533     0.327     0.228    0.0908 
   0.0908     0.165     0.325     0.371     0.371     0.371     0.325     0.165    0.0908 -7.14e-15 

t=1
-3.46e-15     0.159     0.260     0.336     0.359     0.303     0.136       NaN -3.46e-15 -3.46e-15 
    0.159     0.402     0.510     0.629     0.650     0.566     0.418     0.323     0.147    0.0252 
    0.265     0.469     0.600     0.660     0.712     0.591     0.518     0.361     0.103       NaN 
    0.334     0.450     0.574     0.500     0.579     0.533     0.511     0.416     0.123       NaN 
    0.323     0.488     0.581     0.548     0.642     0.465     0.469     0.192       NaN  0.000324 
    0.265     0.438     0.573     0.544     0.615     0.430     0.463     0.247       NaN  0.000645 
    0.273     0.396     0.418     0.542     0.646     0.441     0.408     0.177    0.0112  0.000324 
    0.266     0.383     0.479     0.662     0.682     0.507     0.357    0.0989   0.00445  0.000324 
    0.245     0.360     0.535     0.624     0.628     0.423     0.258    0.0506  0.000965  0.000324 
    0.131     0.212     0.393     0.438     0.400     0.250       NaN  0.000645  0.000324 -3.46e-15 

NaNって何だ。
NaNと実数を比較すると、Double.compareでは、NaNが最大だとAIが返したが、ホントだった。
cppだとおそらく>演算子で比較しているだろうから、比較結果は常にfalseで順序はバラバラ。
doubleとDoubleのまとめ

java
7:NaN 29:NaN 39:NaN 48:NaN 58:NaN 96:NaN 24:0.712 74:0.682 73:0.662 23:0.660 14:0.650 64:0.646 44:0.642 
13:0.629 84:0.628 83:0.624 54:0.615 22:0.600 25:0.591 42:0.581 34:0.579 32:0.574 52:0.573 15:0.566 
43:0.548 53:0.544 63:0.542 82:0.535 35:0.533 26:0.518 36:0.511 12:0.510 75:0.507 33:0.500 41:0.488 
72:0.479 46:0.469 21:0.469 45:0.465 56:0.463 31:0.450 65:0.441 51:0.438 93:0.438 55:0.430 85:0.423 
62:0.418 16:0.418 37:0.416 66:0.408 11:0.402 94:0.400 61:0.396 92:0.393 71:0.383 27:0.361 81:0.360 
4:0.359 76:0.357 3:0.336 30:0.334 40:0.323 17:0.323 5:0.303 60:0.273 70:0.266 50:0.265 20:0.265 
2:0.260 86:0.258 95:0.250 57:0.247 80:0.245 91:0.212 47:0.192 67:0.177 1:0.159 10:0.159 18:0.147 
6:0.136 90:0.131 38:0.123 28:0.103 77:0.0989 87:0.0506 19:0.0252 68:0.0112 78:0.00445 88:0.000965 
59:0.000645 97:0.000645 49:0.000324 79:0.000324 69:0.000324 89:0.000324 98:0.000324 0:-3.46e-15 8:-3.46e-15 9:-3.46e-15 99:-3.46e-15 

cpp
24:0.729434 22:0.685362 23:0.679181 74:0.66918 84:0.658234 13:0.648738 73:0.644198 64:0.635868 52:0.632751 83:0.621286 14:0.618386 25:0.598762 
34:0.57657 42:0.569816 62:0.568456 63:0.566425 32:0.56322 12:0.557356 35:0.55622 72:0.554622 45:0.542977 75:0.539409 55:0.532425 82:0.525989 
53:0.523331 36:0.522689 65:0.522619 21:0.518534 43:0.501979 26:0.496878 41:0.486018 33:0.481686 44:0.478949 15:0.478279 54:0.476765 56:0.475503 
93:0.467479 94:0.463314 51:0.461527 31:0.457673 46:0.450217 85:0.447103 92:0.405041 11:0.394959 61:0.393152 16:0.391319 3:0.387665 81:0.376622 
71:0.372189 4:0.351646 66:0.340899 37:0.337103 30:0.311627 40:0.30844 76:0.305143 27:0.294382 95:0.274377 2:0.271191 17:0.269029 86:0.267469 
5:0.25479 50:0.25281 60:0.246778 91:0.242284 20:0.230506 70:0.220442 80:0.220442 58:-nan 57:0.175115 48:-nan 39:-nan 38:-nan 29:-nan 28:-nan 
10:0.158152 7:-nan 6:-nan 1:0.158152 47:0.131336 90:0.124678 77:0.0731114 96:0.0712212 87:0.0630948 67:0.0563739 18:0.0312636 19:0.00425073 
78:0.00227881 68:0.000759251 88:0.000724218 97:0.000689172 59:7.1715e-05 79:3.60186e-05 69:3.60186e-05 49:3.60186e-05 98:3.60186e-05 
89:3.60186e-05 99:-1.26989e-15 9:-1.26989e-15 8:-1.26989e-15 0:-1.26989e-15 

何でNaNになったのか追うと、プールの中にpxIfR=0があり、evalでのMath.log(pr_r[r])で-Infinityになり、p_r_x * pool.get(x).pxIfR * (ln_p_r_x - pr_r[r + c.lb])で-Infinityの符号が反転し0をかけるとNaNとなる。
そもそもプールの中のpxIfR=0は、対数lnPRifX=-Infinityで、exp(-Infinity)=0で作られたもの。
プールの対数はqueryで計算したlikelihood(mu, sig, res)が0で、このときのlogで-Infinityになっていた。

log(0)の回避は、log(max(x, 1e-300))で底上げするか、電卓でln(10^-300))=-690程度なのでこの定数にするか、-999にしてもよいかと。
ただしexp(-999)は確実に0なので、またlog(0)となる場所が増える可能性がある。

検証としてquery部分だけlog(0)の回避を入れて、ダンプするとNaNが消えて、実数となる。
最終的には直接Math.logを呼ぶのを止める。(今のソースで3か所query,eval,Cache)

t=1
-3.46e-15     0.159     0.260     0.336     0.359     0.303     0.136    0.0255 -3.46e-15 -3.46e-15 
    0.159     0.402     0.510     0.629     0.650     0.566     0.418     0.323     0.147    0.0252 
    0.265     0.469     0.600     0.660     0.712     0.591     0.518     0.361     0.103    0.0111 
    0.334     0.450     0.574     0.500     0.579     0.533     0.511     0.416     0.123    0.0111 
    0.323     0.488     0.581     0.548     0.642     0.465     0.469     0.192    0.0151  0.000324 
    0.265     0.438     0.573     0.544     0.615     0.430     0.463     0.247    0.0220  0.000645 
    0.273     0.396     0.418     0.542     0.646     0.441     0.408     0.177    0.0112  0.000324 
    0.266     0.383     0.479     0.662     0.682     0.507     0.357    0.0989   0.00445  0.000324 
    0.245     0.360     0.535     0.624     0.628     0.423     0.258    0.0506  0.000965  0.000324 
    0.131     0.212     0.393     0.438     0.400     0.250    0.0565  0.000645  0.000324 -3.46e-15 

seed20(N=10)とseed0(N=15)で実行すると、seed20のクエリが1つ減っているのはただの運な気がする。
seed0はevalcntが増えているので、evaltimが伸びているが、なぜだろうか。スコアは伸びている。

execute 0020
0036: pool.size()=360 total=41
0050:0 pool.get(0).pxIfR=0.002777777777777778
0120:1 pool.get(0).pxIfR=0.01866200286940894
0136:2 pool.get(0).pxIfR=0.04719578996251157
0146:3 pool.get(0).pxIfR=0.2734146071596286
0156:4 pool.get(0).pxIfR=0.9320931598025174
0157: evalcnt.hit=1149904
0159: evalcnt.mis=2096
0159: evalcnt=1600
0160: evaltim=74
0162: flipcnt=2889
0162: fliptim=12
Score = 532335

execute 0000
0030: pool.size()=7623 total=38
0048:0 pool.get(0).pxIfR=1.3118194936376755E-4
0495:1 pool.get(0).pxIfR=0.0016099461940363303
0867:2 pool.get(0).pxIfR=0.012840956884832791
1226:3 pool.get(0).pxIfR=0.2809390102645636
1595:4 pool.get(0).pxIfR=0.7793096168878012
1946:5 pool.get(0).pxIfR=0.9989717619209165
1947: evalcnt.hit=68599933
1951: evalcnt.mis=7067
1951: evalcnt=4500
1951: evaltim=1585
1956: flipcnt=7887
1956: fliptim=264
Score = 373896

細かい修正

giveup,mine,query,ansにexitを入れているので、せっかくのevalcnt等が出力しない。
全部AppExceptionに変えて、mainでcatchし、最後に必ず出るようにする。
mainの処理が大きいので、solveに移動する。
byteの計算バグは、ランダム占いのソースを手作業でマージする。
その他、ソース間の揺れを最小になるように手作業でマージする。

おまけ

今まではcppファイルのみをNotebookLMのソースに追加していたが、今度は自分が作ったjavaファイルもソースに追加して、いろいろ指示を足した。

03_all_pool_hill_climb.cppをjavaに翻訳して。ただし、整数はbyteやByteを使わずに、intやIntegerを使ってください。クラス名やメソッド名、変数名は04_all_pool_random_divination.javaと合わせてください。

整数型はすべて int または Integer を使用し、変数名やクラス構成は前回の 04_all_pool_random_divination.java に合わせています。また、議論にあった 0 * log(0) による NaN 回避 や、負の数を返さない Xorshift の実装を盛り込んでいます。


中身をみたら、確かにbyteは無くなった(そもそもintだったのが、javaのGC時間がどうのとチャットしていたら、きっと値の範囲も考慮せずにbyteに変えてきた)。
logはまさにqueryとevalの問題の起きる位置に、1e-300の手当てをつけている。
Xorshiftは新たにLong.toUnsignedDoubleという使ったことないものを入れてきた。3

今回の自分で作ったCache部分を、cppと同じようにint[][] prIfXLb;ProbPair[][][] prIfX;を持つ、Simクラスに持たせて、コンストラクタで事前計算してきた。

山登りの評価値でソートする部分が、List<int[]> evalsint[]{ij, (int)(ev * 1000000)}を入れるとか、int[2]大好きなのか。

相互情報量とランダムの比較

1st: java相互情報量
2nd: javaランダム
3rd: cpp相互情報量
4th: cppランダム
先頭に!のついているセルはgiveup状態

seed N M eps total 1st 2nd 3rd 4th
0 15 2 0.01 38 373,896 1,127,426 !136,155,980 866,481
2 13 2 0.07 68 1,363,505 5,070,120 553,241 2,170,685
3 19 2 0.08 114 !162,200,181 3,196,577 !162,064,685 2,237,383
8 11 2 0.13 46 1,172,376 5,406,208 1,290,767 11,219,032
9 14 2 0.17 41 1,303,602 12,744,716 !72,218,526 17,157,175
20 10 2 0.08 41 532,335 1,858,524 745,690 2,684,369
21 13 2 0.15 50 1,391,050 14,009,250 !84,212,600 10,730,681
25 13 2 0.2 27 !91,208,916 15,233,752 !90,117,851 104,569,930
27 12 2 0.2 56 3,172,628 9,278,628 !53,995,803 21,683,423
38 13 2 0.06 95 451,205 1,434,019 281,400 1,761,736
42 19 2 0.04 148 248,635 1,194,401 !159,065,372 1,334,046
45 15 2 0.07 83 330,335 3,470,842 !124,250,612 2,540,834
61 15 2 0.12 52 !66,053,660 5,292,366 !65,093,659 10,925,874
65 14 2 0.03 62 342,619 1,521,813 !140,179,886 1,300,246
68 20 2 0.19 190 !181,658,220 10,247,067 !181,062,869 10,099,656
85 10 2 0.09 50 545,268 3,715,208 427,115 4,904,346
88 15 2 0.04 68 471,945 1,874,357 !97,164,123 2,790,980
90 12 2 0.15 49 1,293,744 14,811,889 !78,932,263 7,001,686
93 10 2 0.05 23 2,761,135 9,131,199 2,505,022 4,608,564
94 16 2 0.03 105 363,839 2,982,593 367,023 1,234,386

作業用バッファサイズ

よくよく考えたら、Cache[][]の2次元配列に入れているし、cppとの誤差を調べるために、cppとjavaのどちらもeps=0.01とeps=0.2について、400*401=160,404件をファイルに落として、比較したくらいなのだから、バッファサイズの最大値なんて全件で調べられるな。4
というわけで、epsのすべてについてmaxを求めると、当然k=400,s=400で、意外にもeps=0.01の方がバッファが必要だった。

なんでjavaとcppに差があるのだろうか(単なる運)と思っていたが、seed42(N=19,total=148,eps=0.04)では19*19*149=53,789件をメモリに持つため、eps=0.04の202ミリ秒の30%くらいが事前計算の時間がかかっているんじゃないかな。さすがにアドホックなキャッシュで全部使うとは思えないし。キャッシュミス数を調べたら、27,660件でした。

192ms eps=0.01 k=400 s=400 mu=396 sigma=1.99 lb=387 max=406
387: [8.814248554311899E-6, 7.229444280432151E-5, 4.627208401298266E-4, 
390: 0.0023114776527381475, 0.0090130094651934, 0.027435103125845017, 0.0651998502963278, 0.12098691756703794, 
395: 0.17531572391729944, 0.1983863679632028, 0.17531572391729955, 0.12098691756703794, 0.0651998502963278, 
400: 0.027435103125844962, 0.009013009465193345, 0.002311477652738203, 4.627208401297711E-4, 7.229444280443253E-5, 
405: 8.814248554256388E-6]
219ms eps=0.02 k=400 s=400 mu=392 sigma=2.80 lb=379 max=406
282ms eps=0.03 k=400 s=400 mu=388 sigma=3.41 lb=372 max=405
202ms eps=0.04 k=400 s=400 mu=384 sigma=3.92 lb=366 max=403
218ms eps=0.05 k=400 s=400 mu=380 sigma=4.36 lb=360 max=401
241ms eps=0.06 k=400 s=400 mu=376 sigma=4.75 lb=354 max=399
239ms eps=0.07 k=400 s=400 mu=372 sigma=5.10 lb=348 max=397
308ms eps=0.08 k=400 s=400 mu=368 sigma=5.43 lb=343 max=394
261ms eps=0.09 k=400 s=400 mu=364 sigma=5.72 lb=337 max=392
260ms eps=0.10 k=400 s=400 mu=360 sigma=6.00 lb=332 max=389
277ms eps=0.11 k=400 s=400 mu=356 sigma=6.26 lb=327 max=386
274ms eps=0.12 k=400 s=400 mu=352 sigma=6.50 lb=322 max=383
291ms eps=0.13 k=400 s=400 mu=348 sigma=6.73 lb=317 max=380
291ms eps=0.14 k=400 s=400 mu=344 sigma=6.94 lb=312 max=377
308ms eps=0.15 k=400 s=400 mu=340 sigma=7.14 lb=307 max=374
313ms eps=0.16 k=400 s=400 mu=336 sigma=7.33 lb=302 max=371
323ms eps=0.17 k=400 s=400 mu=332 sigma=7.51 lb=297 max=368
324ms eps=0.18 k=400 s=400 mu=328 sigma=7.68 lb=293 max=364
324ms eps=0.19 k=400 s=400 mu=324 sigma=7.85 lb=288 max=361
329ms eps=0.20 k=400 s=400 mu=320 sigma=8.00 lb=283 max=358

ソース置き場

  • 04 相互情報量を最大化する占い
    • 03_all_pool_hill_climb_cache1.java 基準ファイル
    • 03_all_pool_hill_climb_cache2.java 調査ログ
    • 03_all_pool_hill_climb_cache3.java pr_if_xのキャッシュ
    • 03_all_pool_hill_climb_cache4.java ln_pr_if_xのキャッシュ
    • 03_all_pool_hill_climb_cache5.java SMALL_VALUE以上に限定
    • 03_all_pool_hill_climb_climb1.java 山登り調査用

  1. 実はここで手を抜いて、pr_rとpr_if_xを同じ名前にしてたら、logの方が切り取らない方で実行していて、大量のlog(0)を返すことになる。恐ろしいほど山を登らなくなった。

  2. 後から作り直したら、List<Double>を使ってた。まあ動くけど、個数分のDoubleインスタンスを作るんか。

  3. 使ったことがないのではなく、そもそもそんなメソッドは存在しない。存在するのはtoUnsignedString。

  4. eps=0.01で24MB、eps=0.2で88MBくらい。

0
0
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
0
0

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?