Numpyroで状態空間モデルをやってみる
概要
最近MMM(Marketing Mix Modeling)が流行っていて取り組む機会が多いです。 MMMを行うにもMMMに特化したライブラリであるRobynやlightweight mmm、pymc-marketingなどありますので、本来はそっちのほうが楽でしょう。
ただやはりより低レベルのstanやNumpyroで実質的にMMMとも言えるフレームワークにて状態空間モデル書けるようになっている方が、 トラブル時の対応や必要に駆られて色々とモデルをカスタムしたい時などは対応しやすいと思います。
そこで今回は状態空間モデルをやっていきたいのですが、stanに関しては世の中に状態空間モデルを実行する様々な例が公開されているので、 今回は特にNumpyroを使って状態空間モデルを行おうと思います。
参考にした記事と本記事の違い
Numpyroで状態空間モデルを実行するのに参考にしたのは下記のページです。
これらの記事と本記事の差異は説明変数を加えた状態空間モデルを行う点にあります。
MMMを行うには説明変数をモデルに入れる必要があるのでこの点が重要であり、 そのあたりの方法を詳しく書いているような記事は見つからなかったのでこの記事の価値がある点かと思います。
scan関数
Numpyroで時系列性を考慮するには、scan関数を使う必要があります((forループでもできるが非推奨とのこと))。
scan関数に関しては参考にした両記事にも触れている部分がありますが十分に説明されていないようなので説明を試みたいと思います。
scan関数はjaxのドキュメントページを見るに、
def scan(f, init, xs, length=None): if xs is None: xs = [None] * length carry = init ys = [] for x in xs: carry, y = f(carry, x) ys.append(y) return carry, np.stack(ys)
このような動作をするものとして説明されています。
ここからわかるのはscan関数の基本的な動作として、
xsの長さ分ループを回しながら引数に与えられた関数(f)に対してxsの中身を先頭から一つづつ取っていきxとして、
carry*1共に渡してその値を使って処理をしていくような関数だと言えるかと思います。
そして重要な部分はループ内で実行されている関数fの1つ目の返り値によってcarryの値が更新されていく点です。
これによって次のループに関数fの中で処理した結果を渡していくことができるので、時系列性を考慮したモデリングを行うことができるようになっています。
この関数fについてはどんなものを使うかまだ何も説明していないので分かりづらいとは思いますが実際にモデルを定義したコードを見ると理解できると思います。
実装
それではNumpyroで実際に実装してみたいと思います。 今回は3つのモデルを試したいと思います。具体的には以下の3つです。
- ローカル線形トレンドモデル
- 基本構造時系列モデル
- 基本構造時系列モデル+説明変数
各モデルがどのようなものかは詳しく説明しませんので下記書籍などを読むこと推奨します。
データについては上記書籍で利用されている5-6-1-sales-ts-4.csvを使います。
このようにトレンドと周期性を持った季節性が含まれているようなデータになっています。
本データは日付と売上のみのデータのため、説明変数に関しては基本構造時系列モデル+説明変数をやる際にランダムに生成した数値を用います。
事前準備
まずは必要なライブラリを読み込みます。
import numpy as np import pandas as pd import numpyro from numpyro.contrib.control_flow import scan from numpyro.infer import Predictive numpyro.set_host_device_count(3) #並列化 import jax import jax.numpy as jnp import arviz as az import matplotlib.pyplot as plt
Google Colaboratoryで実行する場合は最初に!pip install numpyroを実行してください。
データの読み込みはシンプルにcsvを読み込むだけです。
df = pd.read_csv('./5-6-1-sales-ts-4.csv')
ローカル線形トレンドモデル
ローカル線形トレンドモデルのモデル部分のコードは下記のようになります。
def model(l, y = None): s_w = numpyro.sample("s_w", numpyro.distributions.HalfNormal(2)) #過程誤差の標準偏差 s_v = numpyro.sample("s_v", numpyro.distributions.HalfNormal(2)) #観測誤差の標準偏差 mu_init = numpyro.sample("mu_init", numpyro.distributions.Normal(0, s_w)) #muの初期値サンプリング s_d = numpyro.sample("s_d", numpyro.distributions.HalfNormal(2)) #ドリフト誤差の標準偏差 delta_init = numpyro.sample("delta_init", numpyro.distributions.Normal(0, s_d)) timesteps = jnp.arange(l) def transition(carry, _): #水準+トレンド delta_prev = carry[1] delta = numpyro.sample("delta", numpyro.distributions.Normal(delta_prev, s_d)) delta_carry = delta mu_prev = carry[0] mu = numpyro.sample("mu", numpyro.distributions.Normal(mu_prev + delta_prev, s_w)) mu_carry = mu numpyro.sample("obs", numpyro.distributions.Normal(mu, s_v)) return [mu_carry, delta_carry], None with numpyro.handlers.condition(data={"obs": y}): scan(transition, [mu_init, delta_init], timesteps)
scan関数の中で使う関数fをtransitionとして定義しています。
transition関数では各パラメータをサンプリングしつつサンプリング結果を次のループに渡すために、
_carry変数として代入し渡したいものを配列に入れて返すようにしています。
このようにしてサンプリング結果を次のループに渡し、受けたっと数値でサンプリングを行うということを繰り返すでことで時系列性を考慮しています。
次にこの定義したモデルによってサンプリングを走らせます。
#型変換 y = jnp.array(df['sales']) rng_key= jax.random.PRNGKey(0) #乱数固定 kernel = numpyro.infer.NUTS(model) mcmc = numpyro.infer.MCMC(kernel, num_warmup=1000, num_samples=5000, num_chains=3) mcmc.run(rng_key = rng_key, l = len(y), y = y)
結果を確認するのは下記のように行って、
mcmc.print_summary()
r_hatを確認して収束していることがわかります。
mean std median 5.0% 95.0% n_eff r_hat
delta[0] 0.31 0.92 0.13 -0.97 1.52 530.74 1.01
delta[1] 0.37 1.05 0.17 -1.02 1.90 609.83 1.01
delta[2] 0.42 1.16 0.20 -1.07 2.28 600.28 1.01
delta[3] 0.49 1.26 0.25 -1.15 2.46 591.16 1.01
delta[4] 0.56 1.34 0.30 -1.24 2.62 547.05 1.01
delta[5] 0.61 1.38 0.33 -1.32 2.71 505.68 1.01
delta[6] 0.66 1.42 0.38 -1.16 3.05 462.66 1.01
delta[7] 0.70 1.47 0.40 -1.21 3.18 450.82 1.01
delta[8] 0.72 1.48 0.43 -1.41 3.10 505.09 1.01
delta[9] 0.74 1.51 0.45 -1.37 3.25 547.47 1.01
delta[10] 0.80 1.55 0.50 -1.44 3.25 483.91 1.01
delta[11] 0.87 1.60 0.55 -1.31 3.53 422.92 1.01
delta[12] 0.93 1.63 0.60 -1.22 3.77 382.25 1.01
delta[13] 0.97 1.66 0.63 -1.40 3.65 384.93 1.01
delta[14] 0.99 1.67 0.66 -1.42 3.66 385.07 1.01
delta[15] 1.00 1.69 0.67 -1.43 3.74 411.09 1.01
delta[16] 1.01 1.69 0.70 -1.47 3.76 436.66 1.01
delta[17] 1.06 1.71 0.74 -1.34 3.95 394.18 1.01
delta[18] 1.11 1.74 0.78 -1.30 4.01 344.82 1.01
delta[19] 1.15 1.75 0.80 -1.34 4.03 314.82 1.02
delta[20] 1.17 1.78 0.83 -1.31 4.04 313.14 1.02
delta[21] 1.18 1.77 0.84 -1.33 4.05 321.69 1.02
delta[22] 1.17 1.76 0.85 -1.36 4.03 355.50 1.01
delta[23] 1.15 1.75 0.86 -1.48 3.93 378.62 1.01
delta[24] 1.17 1.76 0.88 -1.39 4.00 403.30 1.01
delta[25] 1.21 1.79 0.91 -1.35 4.17 380.58 1.01
delta[26] 1.26 1.83 0.93 -1.25 4.33 341.18 1.02
delta[27] 1.28 1.84 0.96 -1.42 4.20 332.50 1.02
delta[28] 1.29 1.84 0.97 -1.41 4.22 327.29 1.02
delta[29] 1.28 1.83 0.96 -1.21 4.39 357.14 1.02
delta[30] 1.27 1.82 0.97 -1.37 4.27 371.41 1.02
delta[31] 1.30 1.82 1.01 -1.34 4.32 336.72 1.02
delta[32] 1.34 1.83 1.04 -1.24 4.40 301.93 1.02
delta[33] 1.39 1.85 1.07 -1.41 4.28 280.35 1.02
delta[34] 1.39 1.85 1.08 -1.19 4.48 273.80 1.02
delta[35] 1.38 1.84 1.10 -1.14 4.55 296.49 1.02
delta[36] 1.35 1.81 1.09 -1.23 4.35 333.59 1.02
delta[37] 1.33 1.79 1.07 -1.33 4.25 339.60 1.02
delta[38] 1.35 1.80 1.09 -1.12 4.45 344.76 1.02
delta[39] 1.37 1.80 1.12 -1.20 4.36 326.94 1.02
delta[40] 1.39 1.82 1.14 -1.12 4.53 326.39 1.02
delta[41] 1.40 1.81 1.16 -1.19 4.39 322.92 1.02
delta[42] 1.40 1.81 1.16 -1.25 4.35 318.89 1.02
delta[43] 1.37 1.81 1.14 -1.26 4.36 361.65 1.02
delta[44] 1.36 1.81 1.14 -1.33 4.29 386.44 1.02
delta[45] 1.37 1.80 1.14 -1.24 4.33 352.21 1.02
delta[46] 1.40 1.80 1.17 -1.16 4.44 330.35 1.02
delta[47] 1.41 1.82 1.18 -1.21 4.40 322.75 1.02
delta[48] 1.40 1.81 1.16 -1.35 4.31 340.24 1.02
delta[49] 1.38 1.80 1.17 -1.23 4.44 364.91 1.02
delta[50] 1.34 1.81 1.13 -1.34 4.34 434.12 1.01
delta[51] 1.30 1.80 1.11 -1.43 4.23 495.71 1.01
delta[52] 1.30 1.79 1.12 -1.32 4.30 483.82 1.01
delta[53] 1.33 1.80 1.13 -1.38 4.29 454.00 1.01
delta[54] 1.34 1.81 1.14 -1.25 4.37 455.27 1.01
delta[55] 1.33 1.82 1.14 -1.34 4.33 471.37 1.01
delta[56] 1.30 1.81 1.12 -1.33 4.33 511.51 1.01
delta[57] 1.25 1.80 1.08 -1.49 4.18 606.40 1.01
delta[58] 1.21 1.79 1.05 -1.46 4.17 686.17 1.01
delta[59] 1.22 1.80 1.06 -1.49 4.15 654.33 1.01
delta[60] 1.25 1.81 1.07 -1.56 4.12 609.14 1.01
delta[61] 1.27 1.81 1.08 -1.44 4.27 608.29 1.01
delta[62] 1.26 1.81 1.08 -1.60 4.10 622.47 1.01
delta[63] 1.25 1.82 1.06 -1.46 4.25 613.18 1.01
delta[64] 1.22 1.81 1.04 -1.48 4.26 717.03 1.01
delta[65] 1.19 1.81 1.03 -1.54 4.18 806.40 1.01
delta[66] 1.19 1.82 1.04 -1.44 4.33 775.47 1.01
delta[67] 1.22 1.82 1.05 -1.50 4.27 765.48 1.01
delta[68] 1.22 1.82 1.06 -1.62 4.17 742.59 1.01
delta[69] 1.21 1.82 1.06 -1.59 4.16 756.78 1.01
delta[70] 1.19 1.81 1.05 -1.45 4.31 801.49 1.01
delta[71] 1.16 1.80 1.02 -1.63 4.10 896.31 1.01
delta[72] 1.12 1.81 1.00 -1.59 4.14 1046.40 1.01
delta[73] 1.13 1.81 1.00 -1.46 4.28 1074.52 1.01
delta[74] 1.17 1.81 1.04 -1.59 4.15 873.71 1.01
delta[75] 1.20 1.83 1.04 -1.60 4.21 841.78 1.01
delta[76] 1.21 1.83 1.05 -1.51 4.28 794.68 1.01
delta[77] 1.20 1.83 1.04 -1.58 4.21 802.86 1.01
delta[78] 1.17 1.83 1.04 -1.69 4.15 923.18 1.01
delta[79] 1.13 1.84 0.99 -1.64 4.24 1007.10 1.00
delta[80] 1.14 1.85 1.01 -1.67 4.25 1068.28 1.01
delta[81] 1.17 1.87 1.02 -1.63 4.37 990.97 1.01
delta[82] 1.21 1.88 1.05 -1.66 4.35 831.53 1.01
delta[83] 1.20 1.90 1.03 -1.72 4.33 857.88 1.01
delta[84] 1.19 1.91 1.04 -1.63 4.53 948.67 1.01
delta[85] 1.16 1.92 1.01 -1.87 4.33 1031.67 1.00
delta[86] 1.13 1.94 0.99 -2.03 4.27 1199.48 1.00
delta[87] 1.14 1.97 1.00 -1.92 4.45 1264.95 1.00
delta[88] 1.18 1.99 1.02 -1.93 4.46 1179.48 1.01
delta[89] 1.21 2.02 1.04 -1.91 4.60 1046.26 1.01
delta[90] 1.22 2.05 1.05 -1.79 4.80 1059.62 1.01
delta[91] 1.23 2.06 1.06 -1.86 4.81 1074.92 1.01
delta[92] 1.20 2.10 1.03 -2.15 4.63 1329.93 1.01
delta[93] 1.17 2.15 1.00 -2.03 4.90 1519.29 1.00
delta[94] 1.19 2.20 1.02 -2.19 4.93 1506.21 1.00
delta[95] 1.24 2.26 1.03 -2.38 4.90 1498.44 1.00
delta[96] 1.30 2.34 1.07 -2.08 5.41 1348.59 1.01
delta[97] 1.33 2.42 1.10 -2.35 5.33 1282.19 1.01
delta[98] 1.36 2.51 1.10 -2.41 5.46 1283.86 1.01
delta[99] 1.36 2.59 1.09 -2.54 5.48 1388.16 1.01
delta_init 0.20 0.65 0.08 -0.68 1.13 592.60 1.01
mu[0] 60.45 12.09 60.16 41.01 80.69 4638.60 1.00
mu[1] 82.60 13.24 82.40 61.31 104.44 3923.83 1.00
mu[2] 87.02 12.18 86.89 66.20 105.82 5204.55 1.00
mu[3] 78.26 10.48 78.32 60.23 94.26 10327.70 1.00
mu[4] 79.63 10.46 79.54 62.56 96.92 15015.61 1.00
mu[5] 84.56 10.39 84.65 67.83 102.07 13829.38 1.00
mu[6] 89.68 10.56 89.63 72.43 107.01 8051.27 1.00
mu[7] 93.67 10.54 93.49 76.47 111.01 11807.77 1.00
mu[8] 104.26 11.03 104.32 86.46 122.57 8616.71 1.00
mu[9] 102.81 10.89 102.76 84.98 120.95 10544.23 1.00
mu[10] 91.80 10.55 91.72 74.75 109.55 11155.87 1.00
mu[11] 88.83 10.81 88.93 70.54 105.99 10104.64 1.00
mu[12] 94.62 10.55 94.54 76.66 111.61 11801.33 1.00
mu[13] 103.41 10.57 103.30 85.74 120.59 11971.95 1.00
mu[14] 109.62 10.56 109.45 92.69 127.11 11723.41 1.00
mu[15] 121.37 11.12 121.24 103.39 139.70 7849.96 1.00
mu[16] 122.33 11.08 122.34 103.80 140.03 7651.06 1.00
mu[17] 108.87 10.80 108.87 91.83 127.12 9960.20 1.00
mu[18] 104.49 11.32 104.54 85.43 122.82 6768.40 1.00
mu[19] 110.54 11.05 110.75 92.98 129.54 4509.79 1.00
mu[20] 123.48 10.67 123.55 106.57 141.57 9152.91 1.00
mu[21] 129.58 10.61 129.60 112.01 146.61 10950.28 1.00
mu[22] 142.72 11.08 142.49 123.78 160.20 7948.08 1.00
mu[23] 144.60 11.13 144.39 126.63 163.24 8281.30 1.00
mu[24] 133.60 10.65 133.58 115.33 150.06 12554.63 1.00
mu[25] 128.03 10.82 128.04 109.91 145.28 9342.99 1.00
mu[26] 128.48 10.88 128.58 110.70 146.50 8056.43 1.00
mu[27] 137.76 10.55 137.83 119.81 154.46 12119.42 1.00
mu[28] 147.20 10.83 147.08 129.10 164.95 3240.66 1.00
mu[29] 156.06 10.98 155.94 138.71 174.81 3464.94 1.00
mu[30] 156.62 10.94 156.51 138.65 174.58 6368.87 1.00
mu[31] 143.83 10.65 144.00 126.13 160.98 10026.00 1.00
mu[32] 140.85 10.96 140.95 123.99 159.86 7711.78 1.00
mu[33] 145.16 10.98 145.21 127.14 163.14 9350.36 1.00
mu[34] 159.51 10.71 159.55 141.81 177.03 12333.80 1.00
mu[35] 168.89 10.82 168.82 151.39 186.80 11308.06 1.00
mu[36] 177.32 11.12 177.30 159.03 195.40 8544.60 1.00
mu[37] 176.68 10.65 176.58 159.22 194.46 11288.45 1.00
mu[38] 165.15 10.60 165.05 147.68 182.46 10585.06 1.00
mu[39] 161.56 11.03 161.77 143.64 179.83 8183.15 1.00
mu[40] 168.93 10.70 169.10 150.76 186.09 11877.69 1.00
mu[41] 179.89 10.65 179.84 162.56 197.69 12726.42 1.00
mu[42] 183.69 10.60 183.80 165.63 200.66 13331.74 1.00
mu[43] 193.99 11.15 193.91 174.75 211.26 8282.14 1.00
mu[44] 192.84 10.77 192.94 175.48 210.55 9735.97 1.00
mu[45] 179.16 10.82 179.20 161.23 196.79 10265.60 1.00
mu[46] 180.14 10.80 180.21 162.70 197.87 9228.02 1.00
mu[47] 186.53 10.71 186.63 168.74 203.81 10027.11 1.00
mu[48] 198.80 10.63 198.76 181.42 216.49 12099.13 1.00
mu[49] 201.66 10.62 201.56 184.51 219.29 11087.92 1.00
mu[50] 214.99 11.29 214.85 195.83 232.98 7061.55 1.00
mu[51] 214.25 11.19 214.22 195.35 232.04 8490.54 1.00
mu[52] 198.89 10.82 198.97 181.84 217.19 10943.73 1.00
mu[53] 192.98 11.22 193.15 173.36 210.07 5831.96 1.00
mu[54] 199.31 10.86 199.36 180.56 216.21 5624.56 1.00
mu[55] 212.20 10.65 212.47 194.70 229.65 12980.73 1.00
mu[56] 219.48 10.52 219.56 202.84 237.05 12246.05 1.00
mu[57] 231.55 11.65 231.34 212.47 250.69 4035.75 1.00
mu[58] 228.32 11.45 228.18 209.30 246.56 6237.72 1.00
mu[59] 209.77 10.89 209.78 191.31 227.32 10659.89 1.00
mu[60] 201.80 11.36 201.96 183.58 220.84 7289.22 1.00
mu[61] 208.22 10.87 208.29 190.73 226.66 9735.40 1.00
mu[62] 221.04 10.71 221.00 202.78 238.26 11717.88 1.00
mu[63] 223.75 10.56 223.70 206.92 241.47 12244.14 1.00
mu[64] 234.49 11.01 234.50 216.84 252.83 6693.58 1.00
mu[65] 236.57 11.01 236.37 219.43 255.50 7891.75 1.00
mu[66] 220.98 10.85 220.99 202.83 238.53 8806.38 1.00
mu[67] 216.87 11.08 216.99 198.11 234.54 9373.84 1.00
mu[68] 221.70 10.77 221.78 203.70 239.15 9687.92 1.00
mu[69] 233.04 10.58 232.86 215.46 250.15 10803.44 1.00
mu[70] 235.67 10.36 235.68 218.16 252.19 13476.73 1.00
mu[71] 246.86 10.90 246.67 228.56 264.31 6599.89 1.00
mu[72] 250.10 11.29 249.89 232.52 269.65 7462.30 1.00
mu[73] 233.01 10.54 233.00 215.67 250.07 7963.59 1.00
mu[74] 223.64 11.15 223.81 205.07 241.62 6148.81 1.00
mu[75] 220.68 11.67 220.94 201.95 239.89 4692.11 1.00
mu[76] 236.69 10.52 236.74 219.36 253.99 6674.34 1.00
mu[77] 243.52 10.91 243.62 226.02 261.39 1856.73 1.00
mu[78] 256.25 11.37 256.24 237.55 274.94 3522.23 1.00
mu[79] 258.68 11.45 258.46 240.23 277.80 6883.68 1.00
mu[80] 242.18 10.60 242.29 224.20 259.23 11721.10 1.00
mu[81] 235.47 11.04 235.52 217.05 253.08 7931.17 1.00
mu[82] 235.35 11.27 235.57 217.43 254.52 7904.44 1.00
mu[83] 252.37 10.55 252.32 235.58 270.14 10789.05 1.00
mu[84] 258.58 10.70 258.68 241.42 276.48 11938.91 1.00
mu[85] 267.27 11.23 267.19 247.43 284.47 7754.68 1.00
mu[86] 266.53 11.18 266.38 248.45 285.24 8766.27 1.00
mu[87] 251.78 10.63 251.79 234.37 268.87 11280.49 1.00
mu[88] 245.01 10.95 245.02 226.86 262.93 10153.61 1.00
mu[89] 241.70 11.28 241.79 223.35 260.41 6773.23 1.00
mu[90] 255.64 10.63 255.64 237.44 272.28 12017.34 1.00
mu[91] 259.96 10.50 259.93 242.85 277.33 13100.31 1.00
mu[92] 271.47 11.07 271.49 252.93 289.39 8402.72 1.00
mu[93] 274.26 11.41 274.19 255.64 293.32 7426.24 1.00
mu[94] 256.39 10.60 256.33 239.28 273.90 12287.94 1.00
mu[95] 249.89 10.83 250.05 232.27 267.49 8260.93 1.00
mu[96] 243.68 11.40 243.65 224.83 262.07 7066.81 1.00
mu[97] 255.64 11.04 255.47 237.70 273.92 9897.58 1.00
mu[98] 263.62 11.60 263.51 245.45 283.40 9895.08 1.00
mu[99] 273.52 13.71 273.55 250.52 295.52 10089.69 1.00
mu_init 30.10 10.69 29.77 12.53 47.42 9389.52 1.00
s_d 0.46 0.41 0.34 0.02 0.98 48.85 1.06
s_v 19.26 1.37 19.30 17.02 21.51 3364.68 1.00
s_w 12.30 1.87 12.28 9.22 15.39 1903.30 1.00
Number of divergences: 14
トレースプロットは以下のように見ることができます。
az.plot_trace(mcmc_samples, var_names = ['delta', 'mu'])

Numpyroでは簡単に事前予測分布も見ることができます。
#予測 rng_key, rng_key_ = jax.random.split(rng_key) prior_predictive = Predictive(model, num_samples=100) prior_predictions = prior_predictive(rng_key_, l = len(y))['obs'] mean_prior_pred = jnp.mean(prior_predictions, axis=0) hpdi_prior_pred = numpyro.diagnostics.hpdi(prior_predictions, 0.9) #可視化 fig, ax = plt.subplots(figsize=(10, 4)) ax.plot(df["date"], df["sales"], "o") ax.plot(df["date"], mean_prior_pred) ax.fill_between( df["date"], hpdi_prior_pred[0], hpdi_prior_pred[1], alpha=0.3, interpolate=True ) fig.show()
データで学習していないのでデータに沿った予測になっておらず、また時間がたつほどに予測区間が広くなっています。

一方事後予測分布はというと、
#予測 rng_key, rng_key_ = jax.random.split(rng_key) predictive = Predictive(model, mcmc.get_samples()) predictions = predictive(rng_key_, l = len(y))['obs'] mean_pred = jnp.mean(predictions, axis=0) hpdi_pred = numpyro.diagnostics.hpdi(predictions, 0.9) #可視化 fig, ax = plt.subplots(figsize=(10, 4)) ax.plot(df["date"], df["sales"], "o") ax.plot(df["date"], mean_pred) ax.fill_between( df["date"], hpdi_pred[0], hpdi_pred[1], alpha=0.3, interpolate=True ) fig.show()

というようになっており、データに学習できている感じがあります。 ただトレンドしか考慮していないので精度としては心もとない感じでしょうか。
基本構造時系列モデル
次に季節性を考慮したいと思います。モデル部分のコードはこのようになります。
def model(l, s, y = None): s_v = numpyro.sample("s_v", numpyro.distributions.HalfNormal(2)) #観測誤差の標準偏差 s_w = numpyro.sample("s_w", numpyro.distributions.HalfNormal(2)) #過程誤差の標準偏差 mu_init = numpyro.sample("mu_init", numpyro.distributions.Normal(0, s_w)) #muの初期値サンプリング s_d = numpyro.sample("s_d", numpyro.distributions.HalfNormal(2)) #ドリフト誤差の標準偏差 delta_init = numpyro.sample("delta_init", numpyro.distributions.Normal(0, s_d)) s_s = numpyro.sample("s_s", numpyro.distributions.HalfNormal(2)) #季節性誤差の標準偏差 gamma_init = numpyro.sample("gamma_init", numpyro.distributions.Normal(0, s_s).expand([s-1])) #gammaの初期値サンプリング timesteps = jnp.arange(l) def transition(carry, _): #水準+トレンド delta_prev = carry[1] delta = numpyro.sample("delta", numpyro.distributions.Normal(delta_prev, s_d)) delta_carry = delta #トレンド mu_prev = carry[0] mu = numpyro.sample("mu", numpyro.distributions.Normal(mu_prev + delta_prev, s_w)) mu_carry = mu #季節性 gamma_prev = carry[2] gamma = numpyro.sample("gamma", numpyro.distributions.Normal(-gamma_prev.sum(), s_s)) gamma_prev = gamma_prev.at[:-1].set(gamma_prev[1:]) gamma_carry = gamma_prev.at[-1].set(gamma) #予測 numpyro.sample("obs", numpyro.distributions.Normal(mu + gamma, s_v)) return [mu_carry, delta_carry, gamma_carry], None with numpyro.handlers.condition(data={"obs": y}): scan(transition, [mu_init, delta_init, gamma_init], timesteps)
基本的な部分は同じでgammaにて季節性を表現してモデルに投入しています。
r_hat等は特に問題がないので割愛しますが予測に関しては可視化用のデータの軸など用意して12日先まで予測させてみます。
#可視化用軸 dates_new = pd.concat([df["date"], pd.DataFrame(['2010-04-11', '2010-04-12', '2010-04-13', '2010-04-14', '2010-04-15', '2010-04-16', '2010-04-17', '2010-04-18', '2010-04-19', '2010-04-20', '2010-04-21', '2010-04-22'])]) dates_new = dates_new.reset_index() dates_new.columns = ['idx', 'date'] #予測 rng_key, rng_key_ = jax.random.split(rng_key) predictive = Predictive(model, mcmc.get_samples()) predictions = predictive(rng_key_, l = len(y)+12, s=7)['obs'] mean_pred = jnp.mean(predictions, axis=0) hpdi_pred = numpyro.diagnostics.hpdi(predictions, 0.9) #可視化 fig, ax = plt.subplots(figsize=(10, 4)) ax.plot(df["date"], df["sales"], "o") ax.plot(dates_new['date'], mean_pred) ax.fill_between( dates_new["date"], hpdi_pred[0], hpdi_pred[1], alpha=0.3, interpolate=True ) fig.show()

トレンドのみのものよりもいい感じの予測になっています。 学習範囲の区間も狭いですし予測自体がデータにかなり近いです。学習範囲外も妥当な感じではないでしょうか。
基本構造時系列モデル+説明変数
最後に説明変数を加えてみます。2つの説明変数を加えたいのでデータを生成します。
#説明変数 rng_key, rng_key_ = jax.random.split(rng_key) x1 = jax.random.normal(rng_key_, (1, 112)) x2 = jax.random.normal(rng_key, (1, 112)) X = jnp.hstack([x1.reshape(-1, 1), x2.reshape(-1, 1)]) X_train = X[:100] X_test = X[100:] #係数 beta1 = 11.0 beta2 = -4.1 #目的変数 x1 = x1 * beta1 x2 = x2 * beta2 y = jnp.array(df['sales'] + X_train[:, 0]*beta1 + X_train[:, 1]*beta2 )
モデル部分はこんな感じです。
def model(l, s, X, y = None): s_v = numpyro.sample("s_v", numpyro.distributions.HalfNormal(2)) #観測誤差の標準偏差 s_w = numpyro.sample("s_w", numpyro.distributions.HalfNormal(2)) #過程誤差の標準偏差 mu_init = numpyro.sample("mu_init", numpyro.distributions.Normal(0, s_w)) #muの初期値サンプリング s_d = numpyro.sample("s_d", numpyro.distributions.HalfNormal(2)) #ドリフト誤差の標準偏差 delta_init = numpyro.sample("delta_init", numpyro.distributions.Normal(0, s_d)) s_s = numpyro.sample("s_s", numpyro.distributions.HalfNormal(2)) #季節性誤差の標準偏差 gamma_init = numpyro.sample("gamma_init", numpyro.distributions.Normal(0, s_s).expand([s-1])) #gammaの初期値サンプリング s_b = numpyro.sample("s_b", numpyro.distributions.HalfNormal(2)) beta = numpyro.sample("beta", numpyro.distributions.Normal(0, s_b).expand([X.shape[1]])) timesteps = jnp.arange(l) t = 0 def transition(carry, _): #水準+トレンド delta_prev = carry[1] delta = numpyro.sample("delta", numpyro.distributions.Normal(delta_prev, s_d)) delta_carry = delta #トレンド mu_prev = carry[0] mu = numpyro.sample("mu", numpyro.distributions.Normal(mu_prev + delta_prev, s_w)) mu_carry = mu #季節性 gamma_prev = carry[2] gamma = numpyro.sample("gamma", numpyro.distributions.Normal(-gamma_prev.sum(), s_s)) gamma_prev = gamma_prev.at[:-1].set(gamma_prev[1:]) gamma_carry = gamma_prev.at[-1].set(gamma) #時点処理 t = carry[3] t_carry = t + 1 #予測 numpyro.sample("obs", numpyro.distributions.Normal(mu + gamma + jnp.dot(X[t], beta), s_v)) return [mu_carry, delta_carry, gamma_carry, t_carry], None with numpyro.handlers.condition(data={"obs": y}): scan(transition, [mu_init, delta_init, gamma_init, t], timesteps)
説明変数部分の効果をbetaで表現してXを1レコードづつ取り出して積をとっていきます。
こちらも収束は問題ないですが、betaの値がデータ生成時と少しずれてしまっていますね。。。*2。
(2024/8/31 更新)データの作り方一部間違っていたので修正しました。concatではなくhstackでデータを結合したらbetaはいい感じになりました。
mean std median 5.0% 95.0% n_eff r_hat
beta[0] 10.53 1.18 10.43 8.65 12.56 10571.89 1.00
beta[1] -4.76 1.35 -4.72 -6.99 -2.64 379.84 1.01
トレースプロットでもbetaを可視化してみます。
az.plot_trace(mcmc_samples, var_names = ['beta'])

概ね問題はなさそうです。
betaの事後分布もみてみましょう。
az.plot_posterior(mcmc_samples, var_names = ['beta'])

真の値に近い部分が尖ってますね。
最後に事後予測分布です。
#予測 rng_key, rng_key_ = jax.random.split(rng_key) predictive = Predictive(model, mcmc.get_samples()) predictions = predictive(rng_key_, l = len(y)+12, s=7, X = X)['obs'] mean_pred = jnp.mean(predictions, axis=0) hpdi_pred = numpyro.diagnostics.hpdi(predictions, 0.9) #可視化 fig, ax = plt.subplots(figsize=(10, 4)) ax.plot(df["date"], df["sales"], "o") ax.plot(dates_new['date'], mean_pred) ax.fill_between( dates_new["date"], hpdi_pred[0], hpdi_pred[1], alpha=0.3, interpolate=True ) fig.show()

なかなかいい感じです。
おわりに
今回はNumpyroで状態空間モデルをやってみました。
本取り組みではbetaを正規分布にしていますが実際にMMMで媒体の効果のパラメータとして推定したいときは半正規分布を用いればよいです。
また飽和効果などを入れたい場合は追加で書く必要な部分もあるかと思いますが基本的な部分が理解できれば簡単だと思います。
本記事が参考になりましたら幸いです。間違いなど見つけましたらご指摘ください。