7
2

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?

【全統計学博士号泣】ガウス過程、撃沈。高度で柔軟なモデルがボロ負けする惨劇をどうぞ

7
Posted at

はじめに

こんにちは、事業会社で働いているデータサイエンティストです。

本記事は、前回の記事の続編です。

前回の記事の内容を少しおさらいすると、データが急激に減った期間の平均を、どうすれば精度よく推定できるのか、というのがテーマでした。

具体的には、普段は各時点で500件のデータが得られているにもかかわらず、担当者が長期休暇に入ったことで、ある期間だけデータが1件しか得られなくなった、という極端な状況をシミュレーションしました。

そこで比較したのが、各時点の平均値を独立に推定するモデルと、

$$
\beta_t \sim Normal(\beta_{t-1}, \rho)
$$

のように、$t$期の平均値は$t-1$期の平均値の近くにあるはずだという時系列的な依存関係を利用するモデルです。

データ生成過程そのものにも同様の時系列的な依存関係を持たせているため、後者はいわば正しい構造を知っているモデルです。

だったら当然、こちらが毎回勝つ、、、、、、と思いたくなります。

しかし、1000回のシミュレーションを行った結果、依存関係を考慮したモデルの勝率は 73% でした。

つまり、モデルを正しく特定したからといって、有限標本で必ず高い推定精度が得られるわけではないという、なかなか面白い結果になりました。

では、ここで別の疑問が出てきます。

「もっと柔軟で高度な時系列モデルを使えば、さらに精度が上がるのでは?」

たしかに、前回使ったモデルはかなり素朴です。基本的には「今期は前期の近くにいる」という局所的な依存関係を仮定しているだけです。教科書の構成にもよりますが、これは入門から中級向けの教科書で登場するモデルで、そんなに高度なものではありません。

そこで今回は、もっと強そうな選手をリングに投入してみます。

ガウス過程(Gaussian Process; GP)です。

ガウス過程の概念は解説しますが、詳細には立ち入りません。興味があり、そもそもガウス過程って何って思った読者の方には、まずStan言語のドキュメントを読むことを強くお勧めします:

また、より踏み込んだ紹介に興味ある方は、下記の2本の論文を強くお勧めします:

  • Chen, Yehu, Roman Garnett, and Jacob M. Montgomery. "Polls, context, and time: A dynamic hierarchical Bayesian forecasting model for US senate elections." Political Analysis 31.1 (2023): 113-133. リンク
  • Dew, Ryan, and Asim Ansari. "Bayesian nonparametric customer base analysis with model-based visualizations." Marketing Science 37.2 (2018): 216-235. リンク

ガウス過程を使えば、「直前の時点だけ」を見るのではなく、時点間の距離に応じた共分散構造を通じて、時系列全体から情報を借りながら潜在的な平均の推移を推定できます。

今回は、Stan言語に用意されている指数二乗カーネルを使って、ガウス過程をかなり素直に実装します。

「そんな高度で柔軟なモデルを持ってきたら、前回の単純なモデルなんてボコボコにされるのでは?😎」

私もそう思いながら、とりあえず500回シミュレーションしてみました。

結果は、、、、、、

ガウス過程が3モデル中1位になったのは、推定に成功したケースのわずか7.02%でした。

😨😨😨😨😨

それだけではありません。

全シミュレーションの約6%では、ガウス過程の推定そのものが失敗しました。

どうしてこうなった。

本記事では、単に「ガウス過程が弱かった!おしまい!」とはしません。

まず、ガウス過程とはそもそも何を仮定しているモデルなのか、今回使用した指数二乗カーネルがどのような時系列を好むのかを、数式と直感の両方から確認します。

そのうえで500回のシミュレーション結果を比較し、さらにガウス過程が特にうまくいったシードと、推定そのものに失敗したシードを実際に可視化します。

なぜ、これほど高度で柔軟なモデルが、今回のような問題では単純なモデルにボロ負けしてしまったのでしょうか。

惨劇の現場を見ていきましょう。

シナリオ

今回も、前回の記事とまったく同じシナリオを使います。詳細については前回の記事をご確認ください:

ある指標について、通常は各時点500件のデータが得られているとします。しかし、担当者が長期休暇に入ったなどの事情により、途中の一定期間だけ各時点1件しかデータが得られなくなりました。

具体的には、500時点のデータを生成し、

  • 通常期間:各時点500件
  • 時点200〜250:各時点1件

とします。

また、真の平均値は完全にランダムに変化するのではなく、前の時点から少しずつ動いていくランダムウォークに従うものとします。

今回の目的は、このような極端にデータが少ない期間を含む状況で、次の3つのモデルのどれが真の平均値を最もうまく復元できるかを比較することです。

  1. 各時点を独立に推定するモデル
  2. 前時点との依存関係を利用する時系列モデル
  3. ガウス過程モデル

詳細なデータ生成過程については前回の記事で説明しているため、本記事ではガウス過程との比較に集中します。

モデル紹介:ガウス過程の簡単なおさらい

ここで、今回使うガウス過程について簡単に確認しておきます。

まず重要なのは、今回ガウス過程に従うと仮定しているのは観測されたデータそのものではなく、各時点における潜在的な平均値だということです。

今回のデータは、次のように考えています。

$$
y_i \sim Normal(\mu_{t_{i}}, \sigma)
$$

ここで、$y_i$ が実際に観測されるデータ、$\mu_t$ が時点 $t$ における真の平均値です。

そして、この平均値の時間的な変化にガウス過程を仮定します。

$$
\mu(t) \sim \mathrm{GP}(m(t), k(t,t'))
$$

つまり、

データがガウス過程に従うのではなく、データを生み出している平均値の推移がガウス過程に従う

という構造です。

今回使うのは、代表的な指数二乗カーネルです。

$$
k(t,t') = \alpha^2 \exp\left(-\frac{(t-t')^2}{2\rho^2}\right)
$$

数式がなんか怖いですが、言葉で説明すると、このカーネルでは、時間的に近い時点ほど平均値も似ていると考えます。極端な話をすると、2026年3月のデータは、1932年12月のデータより、2026年2月のデータに似ている、ということです。$\alpha$は平均値がどの程度大きく変動できるか、$\rho$は離れた時点同士の類似性がどれくらい速く弱くなるかを表します。

直感的には、各時点の平均値をバラバラに推定するのではなく、時系列全体の情報を使いながら、滑らかな平均値の推移を復元するモデルだと考えれば十分です。

一見すると、今回のように一部の時点だけデータが極端に少なくなる問題とは、非常に相性がよさそうに見えます。

では、本当にそうなのでしょうか?確認する前に、まずはStan言語によるモデル実装をご確認ください:

time_series_gp.stan
data {
  int time_type;
  array[time_type] real time_seq;
  
  int N;
  
  array[N] int time;
  array[N] real y;
}
parameters {
  real<lower=0> rho_quad;
  real<lower=0> alpha_quad;
  vector[time_type] eta_quad;
  
  real intercept;
  real<lower=0> sigma;
}
transformed parameters {
  vector[time_type] f_trend;
  {
    matrix[time_type, time_type] L_K;
    matrix[time_type, time_type] K = gp_exp_quad_cov(time_seq, alpha_quad, rho_quad);

    for (n in 1:time_type) {
      K[n, n] = K[n, n] + 1e-9;
    }

    L_K = cholesky_decompose(K);
    f_trend = L_K * eta_quad - mean(L_K * eta_quad);
  }
}
model {
  rho_quad ~ inv_gamma(5, 5);
  alpha_quad ~ normal(0, 1);
  eta_quad ~ normal(0, 1);

  sigma ~ inv_gamma(1, 1);
  intercept ~ normal(0, 5);
  y ~ normal(intercept + f_trend[time], sigma);
}

なお、ガウス過程の滑らかさはカーネルの種類だけではなく、長さ尺度 $\rho$ の事前分布にも影響されます。本記事ではガウス過程を最適化すること自体を目的とせず、上記の比較的素朴な仕様を一つの具体例として評価します。そのため、以下の結果を指数二乗カーネル一般の性能として解釈しないでください。

モデル推定結果

まず、シードが1の時の結果から確認しましょう:

set.seed(1)

coef_master <- 500 |>
  seq_len() |>
  tibble::tibble(
    time = _
  ) |>
  dplyr::mutate(
    beta_init = cumsum(rnorm(dplyr::n()))
  ) |>
  dplyr::mutate(
    beta = (beta_init - mean(beta_init))/sd(beta_init)
  ) |>
  dplyr::select(!beta_init)

sim_data <- 500 |>
  seq_len() |>
  purrr::map(
    \(i){
      if (dplyr::between(i, 200, 250)){
        rep(i, 1)
      } else {
        rep(i, 500)
      }
    }
  ) |>
  unlist() |>
  tibble::tibble(
    time = _
  ) |>
  dplyr::left_join(
    coef_master, by = "time"
  ) |>
  dplyr::mutate(
    e = rnorm(dplyr::n()),
    y = beta + e
  )

m_d_init <- cmdstanr::cmdstan_model("time_series_dependent.stan")

m_d_estimate <- m_d_init$variational(
  seed = 1,
  data = list(
    time_type = nrow(coef_master),

    N = nrow(sim_data),
    time = sim_data$time,
    y = sim_data$y
  )
)

m_d_summary <- m_d_estimate$summary()

m_id_init <- cmdstanr::cmdstan_model("time_series_independent.stan")

m_id_estimate <- m_id_init$variational(
  seed = 1,
  data = list(
    time_type = nrow(coef_master),

    N = nrow(sim_data),
    time = sim_data$time,
    y = sim_data$y
  )
)

m_id_summary <- m_id_estimate$summary()

m_gp_init <- cmdstanr::cmdstan_model("time_series_gp.stan")

m_gp_estimate <- m_gp_init$variational(
  seed = 1,
  data = list(
    time_type = nrow(coef_master),
    time_seq = (1:nrow(coef_master) - mean(1:nrow(coef_master)))/sd(1:nrow(coef_master)),

    N = nrow(sim_data),
    time = sim_data$time,
    y = sim_data$y
  )
)

m_gp_summary <- m_gp_estimate$summary()

変分推定のログはこちらです:

Model executable is up to date!
------------------------------------------------------------ 
EXPERIMENTAL ALGORITHM: 
  This procedure has not been thoroughly tested and may be unstable 
  or buggy. The interface is subject to change. 
------------------------------------------------------------ 
Gradient evaluation took 0.004463 seconds 
1000 transitions using 10 leapfrog steps per transition would take 44.63 seconds. 
Adjust your expectations accordingly! 
Begin eta adaptation. 
Iteration:   1 / 250 [  0%]  (Adaptation) 
Iteration:  50 / 250 [ 20%]  (Adaptation) 
Iteration: 100 / 250 [ 40%]  (Adaptation) 
Iteration: 150 / 250 [ 60%]  (Adaptation) 
Iteration: 200 / 250 [ 80%]  (Adaptation) 
Success! Found best value [eta = 1] earlier than expected. 
Begin stochastic gradient ascent. 
  iter             ELBO   delta_ELBO_mean   delta_ELBO_med   notes  
   100      -342478.219             1.000            1.000 
   200      -320802.404             0.534            1.000 
   300      -320433.505             0.356            0.068 
   400      -320109.078             0.267            0.068 
   500      -320376.887             0.214            0.001   MEDIAN ELBO CONVERGED 
Drawing a sample of size 1000 from the approximate posterior...  
COMPLETED. 
Finished in  4.3 seconds.
Model executable is up to date!
------------------------------------------------------------ 
EXPERIMENTAL ALGORITHM: 
  This procedure has not been thoroughly tested and may be unstable 
  or buggy. The interface is subject to change. 
------------------------------------------------------------ 
Gradient evaluation took 0.00552 seconds 
1000 transitions using 10 leapfrog steps per transition would take 55.2 seconds. 
Adjust your expectations accordingly! 
Begin eta adaptation. 
Iteration:   1 / 250 [  0%]  (Adaptation) 
Iteration:  50 / 250 [ 20%]  (Adaptation) 
Iteration: 100 / 250 [ 40%]  (Adaptation) 
Iteration: 150 / 250 [ 60%]  (Adaptation) 
Iteration: 200 / 250 [ 80%]  (Adaptation) 
Success! Found best value [eta = 1] earlier than expected. 
Begin stochastic gradient ascent. 
  iter             ELBO   delta_ELBO_mean   delta_ELBO_med   notes  
   100      -500065.773             1.000            1.000 
   200      -325870.618             0.767            1.000 
   300      -321878.813             0.516            0.535 
   400      -320651.662             0.388            0.535 
   500      -320870.667             0.310            0.012 
   600      -320653.306             0.259            0.012 
   700      -320566.340             0.222            0.004   MEDIAN ELBO CONVERGED 
Drawing a sample of size 1000 from the approximate posterior...  
COMPLETED. 
Finished in  5.2 seconds.
Model executable is up to date!
------------------------------------------------------------ 
EXPERIMENTAL ALGORITHM: 
  This procedure has not been thoroughly tested and may be unstable 
  or buggy. The interface is subject to change. 
------------------------------------------------------------ 
Gradient evaluation took 0.016151 seconds 
1000 transitions using 10 leapfrog steps per transition would take 161.51 seconds. 
Adjust your expectations accordingly! 
Begin eta adaptation. 
Iteration:   1 / 250 [  0%]  (Adaptation) 
Iteration:  50 / 250 [ 20%]  (Adaptation) 
Iteration: 100 / 250 [ 40%]  (Adaptation) 
Iteration: 150 / 250 [ 60%]  (Adaptation) 
Iteration: 200 / 250 [ 80%]  (Adaptation) 
Iteration: 250 / 250 [100%]  (Adaptation) 
Success! Found best value [eta = 0.1]. 
Begin stochastic gradient ascent. 
  iter             ELBO   delta_ELBO_mean   delta_ELBO_med   notes  
   100      -446507.853             1.000            1.000 
   200      -405757.927             0.550            1.000 
   300      -394067.079             0.377            0.100 
   400      -390720.258             0.285            0.100 
   500      -385864.345             0.230            0.030 
   600      -385261.220             0.192            0.030 
   700      -384043.288             0.165            0.013 
   800      -385168.812             0.145            0.013 
   900      -380673.118             0.130            0.012 
  1000      -381101.774             0.117            0.012 
  1100      -379895.323             0.018            0.009   MEDIAN ELBO CONVERGED 
Drawing a sample of size 1000 from the approximate posterior...  
COMPLETED. 
Finished in  26.3 seconds.

時点間の依存関係を考慮したモデル (dependent)、下段が各時点を独立に扱ったモデル (independent)が5秒ほどで推定できるのに対し、ガウス過程は26.3秒かかりました。

結果を可視化しましょう:

g_d_series <- m_d_summary |>
  dplyr::filter(stringr::str_detect(variable, "beta")) |>
  dplyr::mutate(
    time = dplyr::row_number()
  ) |>
  dplyr::bind_cols(
    answer = coef_master$beta
  ) |>
  dplyr::mutate(
    status = dplyr::case_when(
      dplyr::between(time, 200, 250) ~ "few data",
      TRUE ~ "normal"
    )
  ) |>
  ggplot2::ggplot() + 
  ggplot2::geom_line(ggplot2::aes(x = time, y = mean)) + 
  ggplot2::geom_ribbon(ggplot2::aes(x = time, ymin = q5, ymax = q95), fill = ggplot2::alpha("blue", 0.3)) + 
  ggplot2::geom_point(ggplot2::aes(x = time, y = answer, color = status), alpha = 0.3) + 
  ggplot2::geom_vline(xintercept = c(200, 250), linetype = "dashed", color = "red", linewidth = 1) +
  ggplot2::labs(
    title = "dependent"
  )

g_id_series <- m_id_summary |>
  dplyr::filter(stringr::str_detect(variable, "beta")) |>
  dplyr::mutate(
    time = dplyr::row_number()
  ) |>
  dplyr::bind_cols(
    answer = coef_master$beta
  ) |>
  dplyr::mutate(
    status = dplyr::case_when(
      dplyr::between(time, 200, 250) ~ "few data",
      TRUE ~ "normal"
    )
  ) |>
  ggplot2::ggplot() + 
  ggplot2::geom_line(ggplot2::aes(x = time, y = mean)) + 
  ggplot2::geom_ribbon(ggplot2::aes(x = time, ymin = q5, ymax = q95), fill = ggplot2::alpha("blue", 0.3)) + 
  ggplot2::geom_point(ggplot2::aes(x = time, y = answer, color = status), alpha = 0.3) + 
  ggplot2::geom_vline(xintercept = c(200, 250), linetype = "dashed", color = "red", linewidth = 1) +
  ggplot2::labs(
    title = "independent"
  )

g_gp_series <- m_gp_summary |>
  dplyr::filter(stringr::str_detect(variable, "f_trend")) |>
  dplyr::mutate(
    time = dplyr::row_number()
  ) |>
  dplyr::bind_cols(
    answer = coef_master$beta
  ) |>
  dplyr::mutate(
    status = dplyr::case_when(
      dplyr::between(time, 200, 250) ~ "few data",
      TRUE ~ "normal"
    )
  ) |>
  ggplot2::ggplot() + 
  ggplot2::geom_line(ggplot2::aes(x = time, y = mean)) + 
  ggplot2::geom_ribbon(ggplot2::aes(x = time, ymin = q5, ymax = q95), fill = ggplot2::alpha("blue", 0.3)) + 
  ggplot2::geom_point(ggplot2::aes(x = time, y = answer, color = status), alpha = 0.3) + 
  ggplot2::geom_vline(xintercept = c(200, 250), linetype = "dashed", color = "red", linewidth = 1) +
  ggplot2::labs(
    title = "gaussian process"
  )

gridExtra::grid.arrange(g_d_series, g_id_series, g_gp_series)

three_comparison_qiita.png

は?は?は?wwwwwwwwwwwwwww

"w" |> rep(1e+10000000) |> stringr::str_c(collapse = "w") |> print()

はあああああああああ!?!?!?

"w".join(["w" for i in range(1e+10000000)])

おい!ガウス過程くん!てめぇ舐めてんのか!直線推定しとるやないかい

図を見ると、時点間の依存関係を考慮したモデル(dependent)は、データが少ない期間でも前後の時点から情報を借りることで、真の平均値の動きをある程度追跡できています。

一方、各時点を独立に扱うモデル(independent)は、データが十分に存在する期間では問題ありませんが、データが1件しかない200〜250期では推定値が激しく暴れています。これは前回の記事で確認した通りです。

そして今回の主役であるガウス過程を見ると、、、、、、

めちゃくちゃ滑らかです。滑らかすぎます。

本当の平均値は上下に激しく動いているにもかかわらず、その細かな変動をほとんど無視して、時系列全体を貫く緩やかな曲線を推定しています。

ただし、ここで

ガウス過程はポンコツだ!

と結論づけるのは早すぎます。

今回使用している指数二乗カーネルは、単に「近い時点ほど似ている」というだけではなく、非常に滑らかな潜在関数を想定しています。

一方、今回のシミュレーションで真の平均値を生成しているのはランダムウォークです。

beta_init = cumsum(rnorm(dplyr::n()))

つまり、今回比較しているモデルはすべて「時点間に依存関係がある」という点では共通していても、その依存関係をどのようなものとして表現しているかが全く違います。

dependentモデルは、

$$
\mu_t \sim Normal(\mu_{t-1}, \rho)
$$

と、直前の時点から少しずつ変化していくことを直接モデル化しています。

対して今回のガウス過程は、指数二乗カーネルを通じて、時間的に近い時点の平均値は似ており、さらにその変化はかなり滑らかであるという構造を導入しています。

そのため、この結果からわかるのは「ガウス過程は弱い」ということではありません。

むしろ、

高度で柔軟なモデルを使えば、単純なモデルより高い精度が得られるとは限らない

ということです。

さらに重要なのは、「ガウス過程を使う」という選択だけではモデルは決まらないことです。どのカーネルを使うのかによって、モデルが想定する関数の形は大きく変わります。

とはいえ、この図だけを見せて、

「ガウス過程、撃沈!」

と喜ぶのも、それはそれで問題があります。

たまたまガウス過程にとって極端に相性の悪い乱数シードを引いただけかもしれません。

前回の記事でも、たった1回のシミュレーション結果だけを見ることの危険性を確認しました。あるシードではdependentモデルが圧勝しても、別のシードではindependentモデルが勝つことがあります。

ならば今回も同じです。

1回ガウス過程が盛大に撃沈したくらいで、勝利宣言してはいけません。

そこで次に、乱数シードを変えながらこの実験を大量に繰り返します。

大規模シミュレーション:本当にガウス過程は撃沈しているのか?

各シミュレーションでは新しいランダムウォークを生成し、まったく同じ条件のもとで、

  • 時点間の依存関係を直接モデル化したdependentモデル
  • 各時点を独立に推定するindependentモデル
  • 指数二乗カーネルを用いたガウス過程モデル

の3つを推定し、真の平均値に対するRMSEを比較します。

また、都合のいい結果だけをチェリーピッキングすることを避けるため、今回は前回の記事と同様に、データが少ない期間だけではなく、データが十分に存在する期間も含めた全期間のRMSEをメインの評価指標とします。

実は前回の記事を執筆する際、生成AIと壁打ちしていると、

「今回の問題設定なら、データが少ない期間だけでRMSEを評価すればいいのでは?」

と提案されました。

確かに、純粋に「データが少ない期間をどれだけうまく補えるか」を評価したいのであれば、それも合理的です。

しかし、今回想定している実務上の運用では、モデルにはデータが少ない時点だけではなく、すべての時点の平均値の推定を任せることになります。さらに現実のデータでは、「ここからここまでがデータ不足です」と事前に明確に線を引けるとは限りません。

そのため、「データが少ないところでは強いけれど、データが十分にあるところでは精度を落とす」というモデルを、単純に優秀とは評価したくありません。

そこで本記事では、全期間のRMSEを主たる評価指標として、ガウス過程が本当に「撃沈」したのかを判定します。

ただし、生成AIの指摘にも一理あります。全期間のRMSEだけでは、「どこで勝って、どこで負けたのか」が見えなくなってしまいます。

そこで今回は参考として、

  • データが少ない期間(200〜250期)
  • データが十分に存在する期間(それ以外)
  • 全期間

の3種類のRMSEをすべて記録し、後ほど可視化して比較します。

単に「誰が勝ったか」だけではなく、「どこで、なぜ勝ったのか」まで見てみましょう。

これなら、先ほどの衝撃的な結果が単なる偶然なのか、それとも今回のデータ生成過程においてガウス過程が系統的に苦戦しているのかを確認できます。

乱数シードを変えて何度も再戦すれば、最先端モデルは本来の実力を見せてくれるのでしょうか。

こちらがシミュレーションのコードなんですが、実行にものすごく時間がかかってしまうことに留意してください。また、ガウス過程モデルに推定の失敗が見られますので、例外処理も施しています:

future::plan(future::multisession(workers = 8))
seed_simulation_df <- 500 |>
  seq_len() |>
  furrr::future_map(
    \(simulation_id){
      print(stringr::str_c("now computing seed ", simulation_id))
      set.seed(simulation_id)

      coef_master <- 500 |>
        seq_len() |>
        tibble::tibble(
          time = _
        ) |>
        dplyr::mutate(
          beta_init = cumsum(rnorm(dplyr::n()))
        ) |>
        dplyr::mutate(
          beta = (beta_init - mean(beta_init))/sd(beta_init)
        ) |>
        dplyr::select(!beta_init)

      sim_data <- 500 |>
        seq_len() |>
        purrr::map(
          \(i){
            if (dplyr::between(i, 200, 250)){
              rep(i, 1)
            } else {
              rep(i, 500)
            }
          }
        ) |>
        unlist() |>
        tibble::tibble(
          time = _
        ) |>
        dplyr::left_join(
          coef_master, by = "time"
        ) |>
        dplyr::mutate(
          e = rnorm(dplyr::n()),
          y = beta + e
        )

      m_d_init <- cmdstanr::cmdstan_model("time_series_dependent.stan")

      m_d_estimate <- m_d_init$variational(
        seed = 1,
        data = list(
          time_type = nrow(coef_master),

          N = nrow(sim_data),
          time = sim_data$time,
          y = sim_data$y
        )
      )

      m_d_summary <- m_d_estimate$summary()

      m_id_init <- cmdstanr::cmdstan_model("time_series_independent.stan")

      m_id_estimate <- m_id_init$variational(
        seed = 1,
        data = list(
          time_type = nrow(coef_master),

          N = nrow(sim_data),
          time = sim_data$time,
          y = sim_data$y
        )
      )

      m_id_summary <- m_id_estimate$summary()

      m_gp_init <- cmdstanr::cmdstan_model("time_series_gp.stan")

      gp_result <- tryCatch(
        {
          m_gp_estimate <- m_gp_init$variational(
            seed = 1,
            data = list(
              time_type = nrow(coef_master),

              time_seq = (1:nrow(coef_master) - mean(1:nrow(coef_master)))/sd(1:nrow(coef_master)),

              N = nrow(sim_data),
              time = sim_data$time,
              y = sim_data$y
            )
          )
          m_gp_summary <- m_gp_estimate$summary()

          list(
            success = TRUE,
            summary = m_gp_summary,
            error = NA_character_
          )

        },
        error = function(e){

          message(
            "GP failed at seed ",
            simulation_id,
            ": ",
            conditionMessage(e)
          )

          list(
            success = FALSE,
            summary = NULL,
            error = conditionMessage(e)
          )
        }
      )

      if (gp_result$success){

        m_gp_summary <- gp_result$summary

        gp_rmse_all <- m_gp_summary |>
          dplyr::filter(stringr::str_detect(variable,"^f_trend")) |>
          dplyr::bind_cols(answer = coef_master$beta) |>
          dplyr::mutate(
            e2 = (answer - mean)^2
          ) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt()

        gp_rmse_fewdata <- m_gp_summary |>
          dplyr::filter(stringr::str_detect(variable, "^f_trend")) |>
          dplyr::bind_cols(
            answer = coef_master$beta,
            time = coef_master$time
          ) |>
          dplyr::filter(
            dplyr::between(time, 200, 250)
          ) |>
          dplyr::mutate(
            e2 = (answer - mean)^2
          ) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt()

        gp_rmse_fulldata <- m_gp_summary |>
          dplyr::filter(stringr::str_detect(variable,"^f_trend")) |>
          dplyr::bind_cols(
            answer = coef_master$beta,
            time = coef_master$time
          ) |>
          dplyr::filter(
            !dplyr::between(time, 200, 250)
          ) |>
          dplyr::mutate(
            e2 = (answer - mean)^2
          ) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt()

      } else {

        gp_rmse_all <- NA_real_
        gp_rmse_fewdata <- NA_real_
        gp_rmse_fulldata <- NA_real_

      }

      tibble::tibble(
        id = simulation_id,
        d_rmse_all = m_d_summary |>
          dplyr::filter(stringr::str_detect(variable, "beta")) |>
          dplyr::bind_cols(answer = coef_master$beta) |>
          dplyr::mutate(e2 = (answer - mean)^2) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt(),
        id_rmse_all = m_id_summary |>
          dplyr::filter(stringr::str_detect(variable, "beta")) |>
          dplyr::bind_cols(answer = coef_master$beta) |>
          dplyr::mutate(e2 = (answer - mean)^2) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt(),
        gp_rmse_all = gp_rmse_all,

        d_rmse_fewdata = m_d_summary |>
          dplyr::filter(stringr::str_detect(variable, "beta")) |>
          dplyr::bind_cols(answer = coef_master$beta, time = coef_master$time) |>
          dplyr::filter(dplyr::between(time, 200, 250)) |>
          dplyr::mutate(e2 = (answer - mean)^2) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt(),
        id_rmse_fewdata = m_id_summary |>
          dplyr::filter(stringr::str_detect(variable, "beta")) |>
          dplyr::bind_cols(answer = coef_master$beta, time = coef_master$time) |>
          dplyr::filter(dplyr::between(time, 200, 250)) |>
          dplyr::mutate(e2 = (answer - mean)^2) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt(),
        gp_rmse_fewdata = gp_rmse_fewdata,

        d_rmse_fulldata = m_d_summary |>
          dplyr::filter(stringr::str_detect(variable, "beta")) |>
          dplyr::bind_cols(answer = coef_master$beta, time = coef_master$time) |>
          dplyr::filter(!dplyr::between(time, 200, 250)) |>
          dplyr::mutate(e2 = (answer - mean)^2) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt(),
        id_rmse_fulldata = m_id_summary |>
          dplyr::filter(stringr::str_detect(variable, "beta")) |>
          dplyr::bind_cols(answer = coef_master$beta, time = coef_master$time) |>
          dplyr::filter(!dplyr::between(time, 200, 250)) |>
          dplyr::mutate(e2 = (answer - mean)^2) |>
          dplyr::pull(e2) |>
          mean() |>
          sqrt(),
        gp_rmse_fulldata = gp_rmse_fulldata
      )
    },
    .progress = TRUE,
    .options = furrr::furrr_options(seed = 1)
  ) |>
  dplyr::bind_rows()

ではまず、ガウス過程モデルの推定が成功する割合を確認しましょう:

> nrow(tidyr::drop_na(seed_simulation_df))/nrow(seed_simulation_df)
[1] 0.94

6%の割合で推定に失敗するんですね。これは運用に乗せることを考えると、なかなか怖い数字です。

では次に、ガウス過程が3モデルの中で最も全期間RMSEが低い、つまり最も高い推定精度を達成した割合を計算してみましょう。

ここでは、評価基準を2種類用意します。

  • gp_win_rate_success_XXX:ガウス過程の推定に成功したシードだけを対象に勝率を計算
  • gp_win_rate_XXX:ガウス過程の推定に失敗した場合も「負け」として勝率を計算

実際の運用を考えれば、「推定できなかったので今回はノーカウント!」とはいかないので、本記事では後者を主な評価基準とします。ただし、純粋に推定に成功した場合の性能推定失敗まで含めた運用上の性能を区別するため、両方を掲載します。

> seed_simulation_df |>
   dplyr::summarise(
     gp_win_rate_all =
       mean(
         !is.na(gp_rmse_all) &
         gp_rmse_all < id_rmse_all &
         gp_rmse_all < d_rmse_all
       ),
     gp_win_rate_success_all =
       mean(
          gp_rmse_all < id_rmse_all &
         gp_rmse_all < d_rmse_all,
         na.rm = TRUE
       ),
     
     gp_win_rate_fewdata =
       mean(
         !is.na(gp_rmse_fewdata) &
         gp_rmse_fewdata < id_rmse_fewdata &
         gp_rmse_fewdata < d_rmse_fewdata
       ),
     gp_win_rate_success_fewdata =
       mean(
         gp_rmse_fewdata < id_rmse_fewdata &
         gp_rmse_fewdata < d_rmse_fewdata,
         na.rm = TRUE
       ),
     
     gp_win_rate_fulldata =
       mean(
         !is.na(gp_rmse_fulldata) &
         gp_rmse_fulldata < id_rmse_fulldata &
         gp_rmse_fulldata < d_rmse_fulldata
       ),
     gp_win_rate_success_fulldata =
       mean(
         gp_rmse_fulldata < id_rmse_fulldata &
         gp_rmse_fulldata < d_rmse_fulldata,
         na.rm = TRUE
       )
   )
# A tibble: 1 × 6
  gp_win_rate_all gp_win_rate_success_all gp_win_rate_fewdata gp_win_rate_success_fewdata gp_win_rate_fulldata gp_win_rate_success_fulldata
            <dbl>                   <dbl>               <dbl>                       <dbl>                <dbl>                        <dbl>
1           0.066                  0.0702                0.38                       0.404                    0                            0

酷い。あまりにも酷すぎる。

全期間RMSEで比較すると、ガウス過程が3モデルの中で1位になったのはわずか6.6%です。 推定に失敗したケースを除外してガウス過程に少し有利な評価をしても、勝率はたったの7.02%。つまり、推定失敗だけがこの惨敗の原因というわけではありません。ちゃんと推定できた場合ですら、9割以上のシードで他のモデルに負けています。

ただし、ここで非常に面白い結果があります。

データが1時点あたり500件からわずか1件まで激減する期間だけを見ると、ガウス過程の勝率は38.0%、推定成功時に限定すれば40.4%まで上昇しています。

つまり、ガウス過程が何もできていないわけではありません。まさに今回救済したかった「データがほとんど存在しない期間」では、他時点の情報を滑らかにつなげるガウス過程の強みがかなり発揮されています。

ところが、さらに衝撃的なのがデータが十分に存在する期間です。

勝率0%

推定成功時だけに限定しても、

0%

一度も勝っていません。これこそがボロ負けです。

次に、ガウス過程の推定失敗を負けとして、各モデルの勝率をまとめました:

seed_simulation_df |>
  dplyr::mutate(
    winner_all = dplyr::case_when(
      is.na(gp_rmse_all) & d_rmse_all < id_rmse_all ~ "dependent",
      is.na(gp_rmse_all) & id_rmse_all < d_rmse_all ~ "independent",
      gp_rmse_all < d_rmse_all & gp_rmse_all < id_rmse_all ~ "GP",
      d_rmse_all < gp_rmse_all & d_rmse_all < id_rmse_all ~ "dependent",
      id_rmse_all < gp_rmse_all & id_rmse_all < d_rmse_all ~ "independent",
      TRUE ~ "tie"
    ),

    winner_fewdata = dplyr::case_when(
      is.na(gp_rmse_fewdata) & d_rmse_fewdata < id_rmse_fewdata ~ "dependent",
      is.na(gp_rmse_fewdata) & id_rmse_fewdata < d_rmse_fewdata ~ "independent",
      gp_rmse_fewdata < d_rmse_fewdata & gp_rmse_fewdata < id_rmse_fewdata ~ "GP",
      d_rmse_fewdata < gp_rmse_fewdata & d_rmse_fewdata < id_rmse_fewdata ~ "dependent",
      id_rmse_fewdata < gp_rmse_fewdata & id_rmse_fewdata < d_rmse_fewdata ~ "independent",
      TRUE ~ "tie"
    ),

    winner_fulldata = dplyr::case_when(
      is.na(gp_rmse_fulldata) & d_rmse_fulldata < id_rmse_fulldata ~ "dependent",
      is.na(gp_rmse_fulldata) & id_rmse_fulldata < d_rmse_fulldata ~ "independent",
      gp_rmse_fulldata < d_rmse_fulldata & gp_rmse_fulldata < id_rmse_fulldata ~ "GP",
      d_rmse_fulldata < gp_rmse_fulldata & d_rmse_fulldata < id_rmse_fulldata ~ "dependent",
      id_rmse_fulldata < gp_rmse_fulldata & id_rmse_fulldata < d_rmse_fulldata ~ "independent",
      TRUE ~ "tie"
    )
  ) |>
  dplyr::select(
    winner_all,
    winner_fewdata,
    winner_fulldata
  ) |>
  tidyr::pivot_longer(
    cols = dplyr::everything(),
    names_to = "situation",
    values_to = "model"
  ) |>
  dplyr::mutate(
    situation = dplyr::recode(
      situation,
      winner_all = "all",
      winner_fewdata = "fewdata",
      winner_fulldata = "fulldata"
    )
  ) |>
  dplyr::count(
    situation,
    model,
    name = "n"
  ) |>
  dplyr::group_by(situation) |>
  dplyr::mutate(
    win_rate = n / sum(n)
  ) |>
  dplyr::ungroup()

結果は:

# A tibble: 8 × 4
  situation model           n win_rate
  <chr>     <chr>       <int>    <dbl>
1 all       GP             33    0.066
2 all       dependent     363    0.726
3 all       independent   104    0.208
4 fewdata   GP            190    0.38 
5 fewdata   dependent     264    0.528
6 fewdata   independent    46    0.092
7 fulldata  dependent     212    0.424
8 fulldata  independent   288    0.576

ここから分かるように、データが少ない期間でガウス過程の勝率が38%まで高まったものの、時点間の依存関係を直接モデル化したdependentモデルの52.8%の勝率には及びません。

これは今回の結果を理解する上でかなり重要です。ガウス過程はデータの少ない期間では強力な補間能力を発揮する一方で、データが大量に存在し、各時点の平均をその時点のデータだけでもかなり正確に推定できる領域では、その「滑らかにつなぐ」という性質が逆に足かせになっている可能性があります。

今回の真の平均値はランダムウォークで生成されています。隣接時点には強い依存関係があるものの、その軌跡は決して滑らかな曲線ではなく、細かく上下に動き続けます。一方、今回使ったガウス過程の指数二乗カーネルは、時点間の距離が近ければ平均値も滑らかに変化するというかなり強い構造を持っています。

その結果、データが少ないところでは「周囲から情報を借りる」というメリットが勝つものの、データが十分にあるところでは、本当はデータから見えているランダムウォークの細かな凹凸まで勝手に滑らかにしてしまう。先ほど一部のシードで見た「おい!直線推定しとるやないかい!」は、単なる面白画像ではなく、このモデルの弱点をかなり象徴していたのかもしれません。

要するに今回の結果は、

「高度なモデルだから弱い」のではなく、「高度なモデルでも、データ生成過程に合わない構造を入れれば普通に負ける」

という、非常に当たり前でありながら重要な話を示しています。

もちろん、

「そもそも指数二乗カーネルが苦手なデータ生成過程をわざと選んでいるのだから不公平だ。統計学博士ならそんなカーネル選択はしない」

と指摘する方もいるでしょう。

それはもっともです。今回の真の平均値はランダムウォークで生成しており、非常に滑らかな関数を好む指数二乗カーネルにとって、決して有利な条件ではありません。

ただし、ここにも実務上かなり重要な論点があります。

特に、ガウス過程を潜在的な関数そのものを推定するために使う場合を考えてみましょう。

たとえば、求人の採用数を予測するときに「賃金と採用数の間にどのような非線形な関係があるのか」をガウス過程で表現したり、不動産価格を予測するときに緯度・経度から潜在的な地理的効果を推定したりするケースです。

こうした場面では、推定前に真の潜在関数を可視化して、

「あ!こいつランダムウォークだぞ!指数二乗カーネルはやめよう!」

と判断することは基本的にできません。

そもそも関数の形がわからないからこそ、柔軟に関数を推定できるガウス過程を使いたいわけです。

もし「真の関数がどれくらい滑らかなのかを事前に正確に知っていないと適切なカーネルを選べない」というのであれば、ガウス過程を使う実務上の難しさはかなり大きくなります。

もちろん、実際にはドメイン知識、事前予測チェック、複数カーネルの比較などを通じて、より妥当なモデルを選ぶことができます。しかし、それでも未知の潜在関数に対して、最初から正しいカーネルを選べるとは限りません。

そして指数二乗カーネルは、ガウス過程を学ぶと最初に登場することも多い、非常に代表的なカーネルの一つです。

だからこそ、代表的なカーネルを使ったガウス過程が、想定外に粗い潜在関数を前にしたとき、どれくらい壊れるのかを、ある程度悪意のあるデータ生成過程で確認しておくことには意味があります。

これはガウス過程をいじめるための実験というより、モデルが自分の得意な世界から外れたときに、どの程度ロバストなのかを確認するストレステストです。

実務では、データ生成過程の方がモデルに合わせてくれるわけではありません。

モデルの仮定にとって都合の悪い世界でも、それなりに動いてくれないと困るのです。

また、本記事で変分推論を使っているのは、実運用で扱う大規模データを念頭に置いているためです。もちろんMCMCを使った場合に結果が変わる可能性はありますが、本記事では「十分な計算時間をかければ推定できるか」ではなく、実務上現実的な計算方法と組み合わせたときにモデルがどう振る舞うかを評価します。

さてさて、ここまで来ると次に気になることがあります。

ガウス過程が勝ったシードでは、一体何が起きていたのでしょうか?

逆に、推定そのものが撃沈したシードや、とんでもなく悪いRMSEを叩き出したシードでは、真の時系列はどのような形をしていたのでしょうか?

そこで最後に、ガウス過程が特にうまくいったシードと、特に悲惨なことになったシードを実際に可視化して、何が勝敗を分けたのかを定性的に観察してみましょう。

成功したシードと失敗したシードを見比べてみる

ガウス過程のRMSEが最も低い(精度高い)シードと最も高いシードを出します:

> seed_simulation_df |> dplyr::arrange(gp_rmse_all) |> dplyr::select(id)
# A tibble: 500 × 1
      id
   <int>
 1    26
 2   299
 3    82
 4    41
 5   489
 6   500
 7   119
 8   327
 9   257
10    81
# ℹ 490 more rows
# ℹ Use `print(n = ...)` to see more rows
> seed_simulation_df |> dplyr::filter(is.na(gp_rmse_all)) |> dplyr::select(id)
# A tibble: 30 × 1
      id
   <int>
 1    23
 2    43
 3    65
 4    70
 5    75
 6    79
 7    85
 8    88
 9    95
10   107
# ℹ 20 more rows
# ℹ Use `print(n = ...)` to see more rows

大量にやってもあまり意味ないので、ここでは精度高い方26、299、82と精度低い方の23、43、65を確認します:

c(26, 299, 82) |>
  purrr::map(
    \(this_seed){
      set.seed(this_seed)

      500 |>
        seq_len() |>
        tibble::tibble(
          time = _
        ) |>
        dplyr::mutate(
          beta_init = cumsum(rnorm(dplyr::n()))
        ) |>
        dplyr::mutate(
          beta = (beta_init - mean(beta_init)) / sd(beta_init)
        ) |>
        dplyr::select(!beta_init) |>
        dplyr::mutate(
          seed = this_seed,
          status = "success"
        )
    }
  ) |>
  dplyr::bind_rows() |>
  dplyr::bind_rows(
    c(23, 43, 65) |>
      purrr::map(
        \(this_seed){
          set.seed(this_seed)

          500 |>
            seq_len() |>
            tibble::tibble(
              time = _
            ) |>
            dplyr::mutate(
              beta_init = cumsum(rnorm(dplyr::n()))
            ) |>
            dplyr::mutate(
              beta = (beta_init - mean(beta_init)) / sd(beta_init)
            ) |>
            dplyr::select(!beta_init) |>
            dplyr::mutate(
              seed = this_seed,
              status = "fail"
            )
        }
      ) |>
      dplyr::bind_rows()
  ) |>
  ggplot2::ggplot(ggplot2::aes(x = time, y = beta)) +
  ggplot2::geom_line() +
  ggplot2::facet_wrap(~ stringr::str_c(status, "_", seed),ncol = 3)

six_coefs.png

うーん、なんとも言えないですね。

少なくとも、「こういう形のランダムウォークならガウス過程は成功し、こういう形なら失敗する」という単純な規則が、この6例から明確に見えてくるわけではありません。成功例にもかなり激しく上下する系列がありますし、失敗例だけに一目で分かるような特殊な形状が存在するようにも見えません。

ここから考えられるのは、ガウス過程そのものが常に不安定というより、少なくとも今回の「ランダムウォークで生成された潜在平均 × 指数二乗カーネル × Stanの変分推論」という組み合わせでは、推定が不安定になる場合がある、ということです。

ここは重要な区別です。

今回確認された「撃沈」には、少なくとも2種類あります。

1つ目は、推定そのものは成功しているものの、指数二乗カーネルがランダムウォークの細かな変動を十分に捉えられず、過度に滑らかな関数を推定してしまうケースです。これは主として、データ生成過程とカーネルの仮定との相性の問題として考えることができます。

2つ目は、変分推論そのものがうまくいかず、利用可能な推定結果を得られないケースです。実際、今回のシミュレーションでは約6%の試行でガウス過程の推定結果を取得できませんでした。こちらについては、カーネルの性質だけでなく、共分散行列の数値的な条件や変分推論の最適化なども関係している可能性があります。

したがって、今回の結果だけから「ガウス過程は不安定なモデルである」と一般化するのは明らかに言い過ぎです。また、失敗した6%について原因を完全に特定するには、変分推論のシードや最適化設定を変えたり、MCMCと比較したり、カーネルのハイパーパラメータを詳しく調べたりする追加検証が必要です。

結論

いかがでしたか?かなり濃密なシミュレーション実験だったのではないかと思います。

記事の最後で恐縮ですが、一点強調させてください。本記事の目的はガウス過程の数値計算について徹底的にデバッグすることではありません

今回確認したかったのはもっと実務的な問題です。

「より高度で柔軟なモデルを使えば、単純なモデルより当然うまくいく」と考えてよいのか?

今回の実験については、その答えは明確にNOでした。

指数二乗カーネルを使ったガウス過程は、データが少ない期間だけを見ると約40%のシードで3モデル中最高の精度を達成しています。つまり、情報の少ない区間を周囲の情報から補間するというガウス過程の強みは、確かに確認できます。

しかし、全期間で評価すると勝率はわずか6〜7%程度でした。さらに、データが十分に存在する期間では一度も3モデル中最高にならず、約6%では推定結果そのものを取得できませんでした。

これは、ガウス過程が「弱いモデル」だからではありません。

むしろ逆です。非常に柔軟なモデルだからこそ、その性能はカーネルによって何を仮定するのか、そして実際のデータ生成過程とその仮定がどの程度一致しているのかに大きく左右されます

今回使った指数二乗カーネルは、滑らかな潜在関数を表現するには非常に強力です。しかし、今回のようなランダムウォーク的な変動に対しては、その滑らかさがむしろ弱点になり得ます。

そして、ここからが今回の記事で一番伝えたかったことです。

モデルが高度であることと、モデルが問題に適していることは、まったく別の話です。

「ガウス過程だから強い」「ベイズだから強い」「ノンパラメトリックだから柔軟」「最先端だから単純なモデルより高精度」。

そんな保証はどこにもありません。

必要なのはモデルの名前にビビることでも、逆にありがたがることでもなく、自分が解こうとしている問題に近い状況を作り、実際に殴り合わせることです

今回のリングでは、ガウス過程は撃沈しました。

ただし、これはガウス過程の葬式ではありません。

カーネルを変えれば、結果は大きく変わる可能性があります。

Matérnカーネルならどうなるのか。よりランダムウォークに近い構造を表現できるカーネルならどうなるのか。複数のカーネルを組み合わせたらどうなるのか。

その瞬間、今回ボコボコにされたガウス過程くんが、別人のような顔をしてリングに帰ってくる可能性があります。

しかし、それはまた別の記事で殴り合わせますのでお楽しみください。

最後に、私たちと一緒に働きたい方はぜひ下記のリンクもご確認ください:

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

Delete article

Deleted articles cannot be recovered.

Draft of this article would be also deleted.

Are you sure you want to delete this article?