はじめに
なんかベイズ推定の問題を解いてみると、他人のソースからいろいろ見えてくることもあるなということで、さくっとモンテカルロ法の問題を解いてみた。
とっても参考になる題材はこちら→AHC典型解法シリーズ第1弾「モンテカルロ法」
モンテカルロ法とは
30年くらい前の大昔に、モンテカルロ法の例題として、πの近似値を計算するために、ランダムに(x, y)の点を打って、それが円の中に入っている点の割合からπを出すとかやってた。
その頃のパソコン(i386くらい?)だと、何万個シミュレーションしたか覚えていないが、誤差が凄かったような。
というわけで、試してみた。
static void test(long max) {
long start = System.currentTimeMillis();
Random rand = new Random(0);
double x;
double y;
long cnt = 0;
for (int i = 0; i < max; i++) {
x = rand.nextDouble();
y = rand.nextDouble();
if (x * x + y * y <= 1)
cnt++;
}
double pi = (double) cnt / max * 4.0;
long end = System.currentTimeMillis();
System.out.println("" + (end - start) + "ms cnt=" + cnt + " max=" + max + " pi=" + pi);
}
public static void main(String[] args) {
long base = 1000 * 1000;
test(1 * base);
test(10 * base);
test(100 * base);
test(1000 * base);
}
何か良くなっていないような。
38ms cnt=785539 max=1000000 pi=3.142156
331ms cnt=7853518 max=10000000 pi=3.1414072
2999ms cnt=78540271 max=100000000 pi=3.14161084
30198ms cnt=785401280 max=1000000000 pi=3.14160512
AHC015
とりあえずcppソースをまたjavaソースに機械的に変換して、スコアを出してみた。
TimeKeeperの設定が、1950ミリ秒だと、AtcoderのサーバーではすべてTLEとなってしまうので、1900ミリ秒に変更する。
auto state = base_state;がcppではコピーだが、javaは参照値(ポインタ)が代入されるだけで、シミュレーション結果が本体のStateに上書きしてしまうので、コピーコンストラクタを使って複製した。1
cpp版でもきっと問題なんだろうけど、update(p_it[simulate_cnt][this.t_]);にて、this.t_==100でも呼ばれて例外が出るので、条件分岐した。cpp版はおそらく隣のメモリを参照して、update()はすべて埋まっているので、実害はなく動いている。なお、「問題特性を利用しない、シンプルなモンテカルロ法」ではreverse_t != 0の条件分岐をして0で割って落ちるのを回避している。2
| 手法 | java | cpp |
|---|---|---|
| 問題特性に応じた工夫を加えたモンテカルロ法 | 159077184 | 161276478 |
| 問題特性に応じたルールベース手法 | 134104748 | 134104748 |
| 問題特性を利用しない、探索部分のみに工夫を加えたモンテカルロ法 | 132037572 | 133207822 |
| 問題特性を利用しない、シンプルなモンテカルロ法 | 115859105 | 123178860 |
高速化
「問題特性に応じた工夫を加えたモンテカルロ法」を高速化してみる。
といっても、なかなかネックとなるところが見当たらないが、getScore()を何度も呼んで評価に使っているので、ここをピンポイントに手を加える。
まず、ArrayDeque<PairInt>をArrayDeque<Integer>にしてみる。実際はあまり変わらず。
そこで、自前のMyDequeに置き換えて、中身をint[]の実装にしてみる。あまり変わらず。
AIチャットで、pos(y, x)=y*w+x gety(pos)=y/w getx(pos)=y%wはlookup tableを用意するといいですと言っていたな。なんかisdigit()みたいだ。というやり取りを思い出して、int[][] pos_tbl int[] gety_tbl int[] getx_tblを用意すると、これがなかなかの効果あり。
デバッグで、各ターンにタイムアウトまでに何回シミュレーションしたのか最後に出力するようにして、比較した。
修正前
time=19 turn=0 simulate_cnt=32 simulate_sum=32
time=38 turn=1 simulate_cnt=42 simulate_sum=74
time=57 turn=2 simulate_cnt=104 simulate_sum=178
time=76 turn=3 simulate_cnt=109 simulate_sum=287
time=95 turn=4 simulate_cnt=115 simulate_sum=402
:
time=1824 turn=95 simulate_cnt=1499 simulate_sum=29682
time=1843 turn=96 simulate_cnt=1416 simulate_sum=31098
time=1862 turn=97 simulate_cnt=1732 simulate_sum=32830
time=1881 turn=98 simulate_cnt=1690 simulate_sum=34520
time=1900 turn=99 simulate_cnt=1969 simulate_sum=36489
修正後
time=19 turn=0 simulate_cnt=34 simulate_sum=34
time=38 turn=1 simulate_cnt=111 simulate_sum=145
time=57 turn=2 simulate_cnt=135 simulate_sum=280
time=76 turn=3 simulate_cnt=141 simulate_sum=421
time=95 turn=4 simulate_cnt=128 simulate_sum=549
:
time=1824 turn=95 simulate_cnt=1667 simulate_sum=36361
time=1843 turn=96 simulate_cnt=1760 simulate_sum=38121
time=1862 turn=97 simulate_cnt=2004 simulate_sum=40125
time=1882 turn=98 simulate_cnt=2122 simulate_sum=42247
time=1900 turn=99 simulate_cnt=2004 simulate_sum=44251
ちょっと期待して、スコアを出してみると、まあ良くはなった。
| ソース | 問題特性に応じた工夫を加えたモンテカルロ法 |
|---|---|
| java高速化 | 160166676 |
| java修正前 | 159077184 |
| cpp | 161276478 |
高速化2
updateにもi % W, i / Wがあったので、getx,getyに置き換えたが、修正前より落ちてしまったか。
| ソース | 問題特性に応じた工夫を加えたモンテカルロ法 |
|---|---|
| java高速化2 | 158351353 |
| java高速化 | 160166676 |
| java修正前 | 159077184 |
| cpp | 161276478 |
ローカルで実行すると効果はあるぽいのだけど。
time=19 turn=0 simulate_cnt=35 simulate_sum=35
time=38 turn=1 simulate_cnt=93 simulate_sum=128
time=57 turn=2 simulate_cnt=138 simulate_sum=266
time=76 turn=3 simulate_cnt=159 simulate_sum=425
time=95 turn=4 simulate_cnt=159 simulate_sum=584
:
time=1824 turn=95 simulate_cnt=1584 simulate_sum=39018
time=1843 turn=96 simulate_cnt=1764 simulate_sum=40782
time=1862 turn=97 simulate_cnt=1988 simulate_sum=42770
time=1881 turn=98 simulate_cnt=2038 simulate_sum=44808
time=1900 turn=99 simulate_cnt=2095 simulate_sum=46903
高速化3
NotebookLMにjavaソースを追加して、高速化が可能か聞いてみた。
- getScoreの中のインスタンス生成を抑えろ
- MyDequeの使いまわし
- boolean[][] checkedをint[] visitedIdに変更し、常に異なるcurrentIdでチェック(getScoreごとに初期化しない)
- int[][] board_を1次元配列int[] board_に変更
というわけで、board_に合わせて、swapを(y, x)で交換するのではなく、posで交換するようにして、destとかいじった。
time=19 turn=0 simulate_cnt=42 simulate_sum=42
time=38 turn=1 simulate_cnt=131 simulate_sum=173
time=57 turn=2 simulate_cnt=156 simulate_sum=329
time=76 turn=3 simulate_cnt=165 simulate_sum=494
time=95 turn=4 simulate_cnt=165 simulate_sum=659
:
time=1824 turn=95 simulate_cnt=2386 simulate_sum=45633
time=1843 turn=96 simulate_cnt=2542 simulate_sum=48175
time=1862 turn=97 simulate_cnt=2825 simulate_sum=51000
time=1881 turn=98 simulate_cnt=3027 simulate_sum=54027
time=1900 turn=99 simulate_cnt=3069 simulate_sum=57096
simulate_cntが増えたが、スコアはモンテカルロ(運次第)。
| ソース | 問題特性に応じた工夫を加えたモンテカルロ法 |
|---|---|
| java高速化3 | 158969679 |
| java高速化2 | 158351353 |
| java高速化 | 160166676 |
| java修正前 | 159077184 |
| cpp | 161276478 |
調整
Scannerは遅いと指摘されていたので、BufferedReaderに変更した。まあ、件数が少ないから、特に変化なし。
ログを見ると、明らかに最終ターンにこんなにシミュレーションする必要はないだろうということで、ターン数が小さいときに多めに時間を割り振りたいという要望に、こんなコードを出してくれた。
boolean isTimeOver() {
var now = System.currentTimeMillis();
// var whole_diff = now - this.start_time_;
var whole_diff = this.before_time_ - this.start_time_;
var last_diff = now - this.before_time_;
var remaining_time = time_threshold_ - whole_diff;
// 均等配分(残り時間 / 残りターン数)
double average_remaining = (double) remaining_time / (end_turn_ - this.turn_);
// 重みを 2.0 から 1.0 へ線形に減少させる
double weight = 2.0 - (double) this.turn_ / this.end_turn_;
long now_threshold = (long) (average_remaining * weight);
return last_diff >= now_threshold;
}
今まで均等配分だった時間にウェイトをかけるだけなのだが、毎ターンごとにremaining_timeを計算しているから、そのまま2倍とかかけてもよいのね。
time=38 turn=0 simulate_cnt=66 simulate_sum=66
time=75 turn=1 simulate_cnt=297 simulate_sum=363
time=111 turn=2 simulate_cnt=328 simulate_sum=691
time=147 turn=3 simulate_cnt=331 simulate_sum=1022
time=182 turn=4 simulate_cnt=316 simulate_sum=1338
:
time=1866 turn=95 simulate_cnt=861 simulate_sum=33034
time=1874 turn=96 simulate_cnt=927 simulate_sum=33961
time=1882 turn=97 simulate_cnt=911 simulate_sum=34872
time=1891 turn=98 simulate_cnt=1414 simulate_sum=36286
time=1900 turn=99 simulate_cnt=1441 simulate_sum=37727
最新版のseed0で、832243をよく見るようになった。
| ソース | 問題特性に応じた工夫を加えたモンテカルロ法 | seed0 |
|---|---|---|
| java調整 | 160317137 | 775729,832243,817371 |
| java高速化3 | 158969679 | 753123,753123,800714 |
| java高速化2 | 158351353 | 804283,847115,765021 |
| java高速化 | 160166676 | 803093,823914,803093 |
| java修正前 | 159077184 | 659131,737656,719215 |
| cpp | 161276478 | 723974,680547,680547 |
cppにも各ターンのシミュレーション回数を出力したら、あらら、こんなものか。
time=19 turn=0 simulate_cnt=59 simulate_sum=59
time=38 turn=1 simulate_cnt=60 simulate_sum=119
time=57 turn=2 simulate_cnt=61 simulate_sum=180
time=77 turn=3 simulate_cnt=62 simulate_sum=242
time=96 turn=4 simulate_cnt=63 simulate_sum=305
:
time=1851 turn=95 simulate_cnt=384 simulate_sum=11621
time=1871 turn=96 simulate_cnt=389 simulate_sum=12010
time=1891 turn=97 simulate_cnt=441 simulate_sum=12451
time=1911 turn=98 simulate_cnt=482 simulate_sum=12933
time=1931 turn=99 simulate_cnt=519 simulate_sum=13452
あまりにもAIがYoungGCとか言うものだから、出してみたら、2回、3.5ミリ秒だよ。
[0.010s] Using G1
[0.010s] ConcGCThreads: 3 offset 22
[0.010s] ParallelGCThreads: 10
[0.010s] Initialize mark stack with 4096 chunks, maximum 524288
[0.890s] GC(0) Pause Young (Normal) (G1 Evacuation Pause) 22M->1M(254M) 2.062ms
[1.438s] GC(1) Pause Young (Normal) (G1 Evacuation Pause) 33M->1M(254M) 1.524ms
ソース置き場
- MC1.java : 問題特性を利用しない、シンプルなモンテカルロ法
- MC2.java : 問題特性を利用しない、探索部分のみに工夫を加えたモンテカルロ法
- MC3.java : 問題特性に応じたルールベース手法
- MC4.java : 問題特性に応じた工夫を加えたモンテカルロ法
- MC4mkII.java : MC4.javaの高速化
- MC4mkIII.java : MC4.javaの高速化2
- MC4mkIV.java : MC4.javaの高速化3
- MC4mkV.java : MC4IV.javaの調整
- PI.java : πの近似値