はじめに
そもそものきっかけは、ベイズ推定とは何?という程度の知識レベルで、AtcoderのAH030を解こうとすると、適当なロジックで試行錯誤して、時間を浪費したところで大してスコアも伸びずギブアップするだけなので、ゼロから手を動かして知識を得たい。
一方で、従来型の独学での自習は限度があるので、NotebookLMを使うとどのようなメリットがあるのか試してみようかと思いました。
- NotebookLMを初めて使っての気づき 【NotebookLM編】
- AHC030に関する気づき
- 1.機械的にjavaへ移植
- 2.ランダムな占い
- 3.相互情報量を最大化する占い
- 4.事後確率とgiveup改良【この記事】
- 5.数を絞ったランダムなプール
- 6.なんちゃってテキスト変換ツールでcppソースをjavaソースに変換
- 7.java完成版
とっても参考になる題材はこちら→AHC典型解法シリーズ第3弾「ベイズ推定」
事前確率と事後確率
今回の問題にベイズの定理を適用したとき、事前確率と事後確率とは何なのか。
M=2において、プールに全候補を入れた状態は、まだどれが正しいか分からないので、すべて同じ確率としましょうというのが、対数尤度lnPRifX=0で初期化している状態。pxIfR=0も入れているが、こちらはまだ未計算の状態。
まずt=0ターン目に、l.lnPRifX = sim.getLnPRifX(state.oilStates, l.volume, l.topLefts);を実行しても、まだクエリを実行していないので、変化しない。
その後、shuffleし、降順にsortし、最大値maxLn(トップの対数尤度)を得て、l.pxIfR = Math.exp(l.lnPRifX - maxLn);で対数尤度を尤度に戻して、for (OilLayout l : pool) l.pxIfR /= sumPx;にて、合計が1.0となるように正規化する。
この状態のpxIfRが事前確率で、プール内のすべて同じ値を正規化するので、同じ値1/pool.size()となる。
次にt=1以降、クエリを出して、その結果から対数尤度lnPRifXを更新する。実際には、減点主義で、ちょっと間違ったら-1pt、全然違ったら-100ptが対数尤度のようなものだと思えばよい。対数を取らないと、0.1*0.1*0.1と0.01*0.01*0.01では小さくなるスピードが違うため、アンダーフローを起こして比較ができなくなる。底10の対数を取ると-1-1-1と-2-2-2のスケールで大小比較することになる。(ソースでは自然対数なので底eで計算している)
事後確率の推移
ランダムな占いにて、プール内を正規化した後のトップの尤度の推移は以前に出した。
0051: pool.size()=130321
0088:0 pool.get(0).pxIfR=7.673360394717659E-6
0193:1 pool.get(0).pxIfR=4.9632872855075336E-5
0238:2 pool.get(0).pxIfR=1.7614041235043389E-4
:
2616:32 pool.get(0).pxIfR=0.09734594967055978
2725:33 pool.get(0).pxIfR=0.1039300549575216
2833:34 pool.get(0).pxIfR=0.08912393859176292
ここで、プールの中のトップ10の尤度の推移を表示したところで、あまり面白くもないので、尤度の重みを使って、n*nのマス目ごとに、油田が存在する、存在しないの予想を投票させる。
AIが吐き出したコードは、ここでtotalを計算し、後で割っていたが、呼び出し元ですでに正規化しているので、負荷を軽くするためtotalをコメントアウトした。
static double[] getPosteriorProbability(Input input, List<OilLayout> pool) {
// 盤面の事後確率を求める
// 前提としてpxIfRが正規化されている(全poolの合計が1)
// 正規化されていない場合、ここで累計double totalを計算し、もう一つループしてprob[ij]/=totalで最大1の確率にする
double[] prob = new double[input.n2];
// double total = 0.0;
for (OilLayout l : pool) {
for (int ij = 0; ij < input.n2; ij++) {
if (l.volume[ij] > 0) prob[ij] += l.pxIfR;
}
// total += l.pxIfR;
}
// for (int ij = 0; ij < input.n2; ij++) {
// prob[ij] /= total;
// }
return prob;
}
ランダムな占いのソースの尤度を表示しているところに、このダンプを表示させる。
なお、idを表示すると便利なので、OilLayoutクラスにint idを追加し、4重ループで総当たりの候補を作るときに、シーケンス番号を入れておく。
if (DEBUG) {
String msg=String.format(" 1st(%d)=%.2f 2nd(%d)=%.2f",
pool.get(0).id, pool.get(0).pxIfR*100,
pool.get(1).id, pool.get(1).pxIfR*100);
System.err.println(now()+t+msg);
double[] prob = getPosteriorProbability(input, pool);
for (int ij = 0; ij < prob.length; ij++) {
System.err.printf("%3.0f ", prob[ij]*100);
if (ij % input.n == input.n - 1) System.err.println();
}
}
まずseed20にて、事前確率(均等な尤度1/360=0.277)のときの盤面の確率。
ここで、0となっているマスは、総当たりの投票で0なので、確実に存在しない。
左上の2つの4のマスは、油田を左上に置いたとき、油田Aに存在せず、油田Bに存在する。したがって、油田Bを左上に固定し、油田Aを自由に配置したときの組み合わせは(10-7)*(10-5)=15通りあり、全候補360通りに対する割合は15/360=4.16の数字となっている。
0029: pool.size()=360 total=41 max[0]=(7,5) max[1]=(6,4)
0047:0 1st(87)=0.28 2nd(219)=0.28
0 4 8 14 14 14 14 11 0 0
4 18 26 37 42 47 39 33 13 7
11 33 48 64 71 71 60 47 21 11
17 44 64 81 89 86 73 57 30 11
17 36 62 79 84 77 67 48 18 4
11 24 53 73 80 77 71 54 26 8
8 17 47 64 67 67 61 42 17 4
8 17 43 54 57 57 51 30 13 4
8 17 35 42 46 46 39 21 13 4
4 8 18 22 22 22 18 8 4 0
ランダムなクエリ(q 48個)を出し、対数尤度を更新した後のt=1の盤面。
マスごと、上がっている部分、下がっている部分、変わらない部分に分かれる。
1stと2ndのidが変わっているのは、まだ同率がたくさんいるのか。
0071:1 1st(262)=1.19 2nd(22)=1.19
0 1 1 13 7 4 6 4 0 0
1 21 16 20 38 28 22 11 7 3
20 36 50 50 52 41 45 32 15 4
23 60 67 85 90 75 59 44 23 9
17 48 79 94 87 78 61 38 8 2
14 21 58 84 85 77 73 54 31 8
9 10 43 60 59 54 55 35 9 0
16 25 51 65 50 52 43 23 5 2
9 28 55 68 68 73 57 35 21 8
1 12 36 40 50 49 37 19 8 0
飛ばして、トップが50%を超えたt=11の盤面。
このときの1st.id=164、2nd.id=156は最後まで変わらない。
0130:11 1st(164)=59.42 2nd(156)=12.37
0 1 1 2 2 0 0 0 0 0
1 3 4 3 90 4 2 0 0 0
2 19 92 12 95 89 88 0 0 0
15 95 97 100 100 9 7 3 2 0
3 95 100 100 99 92 4 0 0 0
2 7 35 100 100 85 68 2 0 0
1 1 18 97 99 68 2 0 0 0
13 16 22 96 86 8 4 0 0 0
13 29 98 100 86 83 67 0 0 0
0 14 81 82 80 68 0 0 0 0
最後、トップシェアが80%を超えて、回答をしたt=13の盤面。
0140:13 1st(164)=83.13 2nd(163)=6.13
0 0 0 1 1 0 0 0 0 0
0 1 2 1 94 2 1 0 0 0
1 4 95 6 97 94 93 0 0 0
2 96 98 100 100 6 4 2 1 0
2 96 100 100 100 95 2 0 0 0
1 4 13 100 100 98 90 1 0 0
0 1 4 98 100 90 1 0 0 0
1 2 5 98 93 4 2 0 0 0
1 8 99 100 99 98 89 0 0 0
0 7 96 97 96 90 0 0 0 0
極端な例
以前手で作ったtotalが最大となるテストデータのN=10版。
※javaはScannerが改行も区別せずに解析するので、本来1行のところが6行に分かれてます
10 2 0.01
50
0 0 0 1 0 2 0 3 0 4 0 5 0 6 0 7 0 8 0 9
1 0 1 1 1 2 1 3 1 4 1 5 1 6 1 7 1 8 1 9
2 0 2 1 2 2 2 3 2 4 2 5 2 6 2 7 2 8 2 9
3 0 3 1 3 2 3 3 3 4 3 5 3 6 3 7 3 8 3 9
4 0 4 1 4 2 4 3 4 4 4 5 4 6 4 7 4 8 4 9
50
0 0 0 1 0 2 0 3 0 4 0 5 0 6 0 7 0 8 0 9
1 0 1 1 1 2 1 3 1 4 1 5 1 6 1 7 1 8 1 9
2 0 2 1 2 2 2 3 2 4 2 5 2 6 2 7 2 8 2 9
3 0 3 1 3 2 3 3 3 4 3 5 3 6 3 7 3 8 3 9
4 0 4 1 4 2 4 3 4 4 4 5 4 6 4 7 4 8 4 9
0 0
5 0
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
1 1 1 1 1 1 1 1 1 1
初期状態t=0。
0042: pool.size()=36 total=100 max[0]=(4,9) max[1]=(4,9)
0059:0 1st(9)=2.78 2nd(1)=2.78
31 31 31 31 31 31 31 31 31 31
56 56 56 56 56 56 56 56 56 56
75 75 75 75 75 75 75 75 75 75
89 89 89 89 89 89 89 89 89 89
97 97 97 97 97 97 97 97 97 97
97 97 97 97 97 97 97 97 97 97
89 89 89 89 89 89 89 89 89 89
75 75 75 75 75 75 75 75 75 75
56 56 56 56 56 56 56 56 56 56
31 31 31 31 31 31 31 31 31 31
t=5にて、もう1stと2ndで独占してるが、1stのシェア80%を条件としているので、最大200クエリ投げてさようなら。
0096:5 1st(5)=50.00 2nd(30)=50.00
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
100 100 100 100 100 100 100 100 100 100
※これは極端過ぎるので、油田の形が同じならば同じハッシュ(zhash)を振って、順番が違うだけのものは総当たりの時にスキップすればよい。
バイナリーエントロピー
エントロピーの説明ではよくコインの表と裏の確率50%がよく見かけるが、確率pが決まれば、反対の確率(1-p)が計算できるときのエントロピーを、バイナリーエントロピーと呼ぶらしい。Wikipediaでは二値エントロピーとなっているな。(なんかダサい)
計算式:$-(p*log_2(p)+(1-p)*log_2(1-p))$
要は、事後確率によって、3分類できる。
- 1だった(確実にある):二値エントロピーは0
- 0.5だった(意見が半々に割れてよく分からない):二値エントロピーは1
- 0だった(確実にない):二値エントロピーは0
埋もれていたソースから計算する関数を持ってきた。1
ここでは、pが0と1を事前チェックで弾いている。
static double calculateBinaryEntropy(double p) {
if (p<=0 || p>=1) return 0; // 確信している状態
return -(p * Math.log(p) + (1 - p) * Math.log(1 - p)) / Math.log(2);
}
左に確率、右に二値エントロピーを並べてみる。
seed20のt=0
二値エントロピー100のマスは、47から53辺りのよく分かっていない部分。
確率4のところ、正確にはp=15/360=1/24で、-p*log(p)/log(2)-(1-p)*log(1-p)/log(2)に入れて、式を変形すると、log(24)/log(2)-23/24*(log(23)/log(2))からの24.9882となるらしい。
0031: pool.size()=360 total=41 max[0]=(7,5) max[1]=(6,4)
0048:0 1st(87)=0.28 2nd(219)=0.28
0: 0 4: 25 8: 41 14: 60 14: 60 14: 60 14: 60 11: 49 0: 0 0: 0
4: 25 18: 69 26: 83 37: 95 42: 98 47:100 39: 96 33: 91 13: 57 7: 35
11: 49 33: 92 48:100 64: 94 71: 87 71: 87 60: 97 47:100 21: 73 11: 49
17: 66 44: 99 64: 94 81: 71 89: 49 86: 60 73: 84 57: 98 30: 88 11: 49
17: 66 36: 94 62: 96 79: 74 84: 62 77: 78 67: 92 48:100 18: 69 4: 25
11: 49 24: 80 53:100 73: 84 80: 72 77: 78 71: 87 54: 99 26: 83 8: 41
8: 41 17: 65 47:100 64: 94 67: 92 67: 92 61: 96 42: 98 17: 65 4: 25
8: 41 17: 65 43: 99 54: 99 57: 98 57: 98 51:100 30: 88 13: 54 4: 25
8: 41 17: 65 35: 93 42: 98 46: 99 46: 99 39: 96 21: 74 13: 54 4: 25
4: 25 8: 41 18: 69 22: 76 22: 76 22: 76 18: 69 8: 41 4: 25 0: 0
t=1
0071:1 1st(262)=1.19 2nd(22)=1.19
0: 0 1: 9 1: 11 13: 57 7: 37 4: 25 6: 33 4: 23 0: 0 0: 0
1: 9 21: 74 16: 63 20: 72 38: 96 28: 85 22: 76 11: 50 7: 35 3: 21
20: 72 36: 94 50:100 50:100 52:100 41: 98 45: 99 32: 90 15: 61 4: 23
23: 77 60: 97 67: 91 85: 61 90: 47 75: 81 59: 98 44: 99 23: 77 9: 45
17: 66 48:100 79: 74 94: 32 87: 57 78: 76 61: 97 38: 95 8: 41 2: 17
14: 58 21: 74 58: 98 84: 63 85: 62 77: 79 73: 84 54: 99 31: 90 8: 41
9: 43 10: 48 43: 99 60: 97 59: 97 54:100 55: 99 35: 93 9: 42 0: 0
16: 63 25: 80 51:100 65: 94 50:100 52:100 43: 99 23: 79 5: 28 2: 17
9: 43 28: 85 55: 99 68: 90 68: 91 73: 84 57: 99 35: 93 21: 75 8: 40
1: 9 12: 54 36: 95 40: 97 50:100 49:100 37: 95 19: 70 8: 40 0: 0
t=11
0133:11 1st(164)=59.42 2nd(156)=12.37
0: 0 1: 6 1: 8 2: 12 2: 16 0: 0 0: 0 0: 0 0: 0 0: 0
1: 6 3: 18 4: 23 3: 20 90: 46 4: 23 2: 16 0: 0 0: 0 0: 0
2: 13 19: 70 92: 40 12: 53 95: 27 89: 50 88: 52 0: 1 0: 0 0: 0
15: 62 95: 31 97: 18 100: 3 100: 1 9: 44 7: 37 3: 18 2: 12 0: 0
3: 18 95: 28 100: 1 100: 0 99: 6 92: 41 4: 25 0: 0 0: 0 0: 0
2: 13 7: 36 35: 93 100: 0 100: 0 85: 62 68: 91 2: 13 0: 0 0: 0
1: 9 1: 11 18: 68 97: 20 99: 7 68: 91 2: 13 0: 0 0: 0 0: 0
13: 56 16: 63 22: 76 96: 23 86: 58 8: 39 4: 25 0: 0 0: 0 0: 0
13: 56 29: 87 98: 12 100: 1 86: 59 83: 65 67: 91 0: 1 0: 0 0: 0
0: 3 14: 57 81: 71 82: 67 80: 72 68: 91 0: 1 0: 0 0: 0 0: 0
t=13
0142:13 1st(164)=83.13 2nd(163)=6.13
0: 0 0: 2 0: 3 1: 6 1: 11 0: 0 0: 0 0: 0 0: 0 0: 0
0: 2 1: 8 2: 15 1: 10 94: 33 2: 15 1: 11 0: 0 0: 0 0: 0
1: 6 4: 23 95: 31 6: 34 97: 21 94: 34 93: 37 0: 2 0: 0 0: 0
2: 13 96: 24 98: 16 100: 0 100: 0 6: 34 4: 25 2: 13 1: 6 0: 0
2: 12 96: 22 100: 0 100: 0 100: 2 95: 30 2: 16 0: 0 0: 0 0: 0
1: 10 4: 23 13: 56 100: 0 100: 0 98: 12 90: 47 1: 7 0: 0 0: 0
0: 4 1: 5 4: 22 98: 14 100: 3 90: 47 1: 7 0: 0 0: 0 0: 0
1: 9 2: 14 5: 31 98: 16 93: 37 4: 23 2: 16 0: 0 0: 0 0: 0
1: 7 8: 42 99: 6 100: 0 99: 11 98: 14 89: 49 0: 0 0: 0 0: 0
0: 0 7: 36 96: 23 97: 21 96: 23 90: 48 0: 0 0: 0 0: 0 0: 0
あまり使い物にならずお蔵入りとなった理由は、確率の閾値でも判断できる。二値エントロピーの最大値を探しても、確率13%の56が、ほぼ無いかもしれないけど、ひょっとしたら有るかもしれない、はあまり価値がない。
giveupの効率化
二値エントロピーはイマイチで、その手前の盤面の事後確率が役に立ったこと。
時間切れでgiveupモードになった際、今の実装では中央の座標から土地が見つかれば上下左右をキューの先頭に入れ、海ならキューの末尾に入れて、total個見つかるまでbfsする。
このキューの初期状態を事後確率で降順ソートした順に登録すると、土地があると高確率で予想している座標から探していく。
その後は、今まで通り、土地だったら上下左右をキューの先頭に入れる。海だったら何もしない。
なお、キューの順番のみを使うと、上手くいく場合は格段にスコアが良くなるが、本来は存在するのに、隅っこに極端に小さい確率のマスがあるだけで、最後まで外れを引き続けるので、当たりの次は隣接する上下左右を確率に関係なく優先する方がよさげ。
void giveup(Input input, List<OilLayout> pool) {
System.err.println(now()+"giveup start");
double[] prob = getPosteriorProbability(input, pool);
List<Integer> indices = new ArrayList<>();
for (int ij = 0; ij < input.n2; ij++) indices.add(ij);
Collections.sort(indices, (a, b) -> Double.compare(prob[b], prob[a]));
Deque<int[]> que = new ArrayDeque<>();
for (int ij : indices) que.add(new int[] {ij / input.n, ij % input.n});
//ret == 0のときのaddLastは要らない
for (int[] d : DIJ) {
int nr = r + d[0], nc = c + d[1];
if (nr >= 0 && nr < n && nc >= 0 && nc < n && !used[nr][nc]) {
if (ret == 1) que.addFirst(new int[]{nr, nc});
}
}
意図的にt=10で強制的にgiveupした場合。左の島は確率50くらい。
0136:10 1st(1417)=43.38 2nd(7498)=19.34
0137:giveup start
0 0 0 0 0 0 4 1 0 5 0 0 0 0 0
0 0 0 0 4 1 5 6 0 9 1 0 5 0 0
0 0 0 48 6 5 11 10 37 11 50 6 5 0 0
0 43 0 44 6 6 50 11 39 8 71 6 0 0 1
43 43 44 44 49 50 46 11 35 2 71 1 1 1 1
0 43 44 44 45 45 1 1 28 28 98 72 45 45 0
0 43 43 0 44 0 0 0 28 30 72 66 22 22 0
0 0 0 0 0 0 0 28 28 73 99 70 48 4 0
0 0 0 0 0 0 0 1 28 51 99 71 68 2 1
0 0 0 0 0 0 0 0 29 7 100 36 70 1 2
0 0 0 0 0 0 0 0 3 26 56 29 71 3 23
0 0 0 0 0 0 0 1 43 26 50 26 35 43 21
0 0 0 0 0 0 0 19 20 44 46 43 48 41 0
0 0 0 0 0 0 0 0 20 43 41 21 42 0 0
0 0 0 0 0 0 0 0 19 19 0 20 0 0 0
※(11,10)の海2つ続いているのはバグのように見えて、左の島(3,6)の50よりも大きいため、初期状態のキューにて前にいる。
相互情報量との比較
ランダムな占いは一つもgiveupしていなかったので、giveupしていた相互情報量と結果を比較する。
何かあまり変わらないと思ったら、3/4が中央に固まっていた。それじゃ効果はないよ。
| seed | N | M | eps | total | 相互情報量+giveup | 相互情報量 | コメント |
|---|---|---|---|---|---|---|---|
| 3 | 19 | 2 | 0.08 | 114 | 162,267,448 | 162,200,181 | 島が1つ |
| 25 | 13 | 2 | 0.2 | 27 | 72,034,838 | 91,208,916 | 島が2つ |
| 61 | 15 | 2 | 0.12 | 52 | 68,053,660 | 66,053,660 | 島が1つ |
| 68 | 20 | 2 | 0.19 | 190 | 177,556,685 | 181,658,220 | 島が1つ |
seed25だけ何度も実行すると、内部の乱数のseedは同じなのに、さらに良い57,208,916になったり、結果が異なる。あれれ。
クエリが1つ違うだけで、尤度の違いがここまで出るとは。
2579:7 pool.get(0).pxIfR=0.011112465663302186
Score = 72034838
2459:7 pool.get(0).pxIfR=0.011112465663302186
2809:8 pool.get(0).pxIfR=0.04006837741377529
Score = 57208916
AIチャットとパフォーマンスについて話していて、javaはメソッド内で一時的な配列が短命のヒープメモリ確保となり、メモリ不足にはならないが、YoungGCによる非同期に割込みがかかり、思うようにパフォーマンスがでないかも(という一般論)を聞いたが、件数によってはあながち無視できないのかもしれない。対策はstatic変数に固定バッファを用意し、使いまわす。シングルスレッドならば、固定バッファが競合するのは、メソッドの子供や孫が同じバッファを使うときくらいで、静的にチェックできる。
ソース置き場
- 05 事後確率とgiveup改良
- 04_all_pool_random_divination_debug.java 事後確率とバイナリーエントロピー調査用
-
自力で解いていたとき、AIチャットが盤面のバイナリーエントロピーと、クエリのエントロピーの2通りを出してきて、後者はn=20だと
2^400通り(から、全く選択しない1と、1つしか選択しない400を引く)の中からエントロピーの減少が最大となる400ビットのパターンを見つけるのかとか思って、簡単な盤面のエントロピーを見てたら、あまりスコアは伸びなかった。 ↩



