import%20marimo%0A%0A__generated_with%20%3D%20%220.24.0%22%0Aapp%20%3D%20marimo.App(width%3D%22medium%22)%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%20Fitting%20Input-output%20Hidden%20Markov%20Models%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_()%3A%0A%20%20%20%20import%20warnings%0A%20%20%20%20from%20pathlib%20import%20Path%0A%0A%20%20%20%20warnings.filterwarnings(%22ignore%22)%0A%0A%20%20%20%20import%20cvxpy%20as%20cp%0A%20%20%20%20import%20marimo%20as%20mo%0A%20%20%20%20import%20matplotlib.pyplot%20as%20plt%0A%20%20%20%20import%20numpy%20as%20np%0A%0A%20%20%20%20from%20dbcp%20import%20BiconvexProblem%0A%0A%20%20%20%20_example_directory%20%3D%20Path(__file__).resolve().parent%0A%20%20%20%20plt.style.use(_example_directory%20%2F%20%22zhlatex.mplstyle%22)%0A%20%20%20%20figure_directory%20%3D%20_example_directory%20%2F%20%22figures%22%0A%20%20%20%20figure_directory.mkdir(parents%3DTrue%2C%20exist_ok%3DTrue)%0A%0A%20%20%20%20np.random.seed(1)%0A%20%20%20%20return%20BiconvexProblem%2C%20cp%2C%20figure_directory%2C%20mo%2C%20np%2C%20plt%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Introduction%0A%0A%20%20%20%20We%20consider%20the%20fitting%20problem%20of%20a%20logistic%20input-output%20hidden%20Markov%20model%20(IO-HMM)%20to%20some%20dataset.%0A%20%20%20%20Suppose%20we%20are%20given%20a%20dataset%20%24(x(t)%2C%20y(t))%24%2C%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24%2C%20where%20each%20sample%20consists%20of%20an%20input%0A%20%20%20%20feature%20vector%20%24x(t)%20%5Cin%20%5Cmathbf%7BR%7D%5En%24%20and%20an%20output%20label%20%24y(t)%20%5Cin%20%5C%7B0%2C%201%5C%7D%24%2C%20generated%20from%20a%20%24K%24-state%0A%20%20%20%20IO-HMM%2C%20according%20to%20the%20following%20procedure%3A%0A%20%20%20%20Let%20%24%5Chat%7Bz%7D(t)%20%5Cin%20%5C%7B1%2C%20%5Cldots%2C%20K%5C%7D%24%2C%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24%2C%20be%20the%20state%20label%20of%20the%20IO-HMM%20with%20initial%0A%20%20%20%20state%20distribution%20%24p_%7B%5Crm%20init%7D%20%5Cin%20%5Cmathbf%7BR%7D%5EK%24%20with%20%24%5Cmathbf%7B1%7D%5ET%20p_%7B%5Crm%20init%7D%20%3D%201%24%20and%20transition%20matrix%0A%20%20%20%20%24P_%7B%5Crm%20tr%7D%20%5Cin%20%5Cmathbf%7BR%7D%5E%7BK%20%5Ctimes%20K%7D%24%20with%20%24P_%7B%5Crm%20tr%7D%20%5Cmathbf%7B1%7D%20%3D%20%5Cmathbf%7B1%7D%24.%0A%20%20%20%20At%20the%20time%20step%20%24t%24%2C%20the%20state%20label%20%24%5Chat%7Bz%7D(t)%24%20is%20sampled%20according%20to%0A%0A%20%20%20%20%5C%5B%0A%20%20%20%20%20%20%20%20%5Chat%7Bz%7D(t)%20%5Csim%20%5Cleft%5C%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%5Cbegin%7Barray%7D%7Bll%7D%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%7B%5Crm%20Cat%7D(p_%7B%5Crm%20init%7D)%20%26%20t%20%3D%201%5C%5C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%7B%5Crm%20Cat%7D(p_%7B%5Chat%7Bz%7D(t%20-%201)%7D)%20%26%20t%20%3E%201%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%5Cend%7Barray%7D%5Cright.%0A%20%20%20%20%5C%5D%0A%0A%20%20%20%20where%20the%20vector%20%24p_%7B%5Chat%7Bz%7D(t-1)%7D%20%5Cin%20%5Cmathbf%7BR%7D%5EK%24%20denotes%20the%20%24%5Chat%7Bz%7D(t-1)%24th%20row%20of%20the%20matrix%0A%20%20%20%20%24P_%7B%5Crm%20tr%7D%24%2C%20and%20%24%7B%5Crm%20Cat%7D(p)%24%20denotes%20the%20categorical%20distribution%20with%20%24p%24%20being%20the%20vector%20of%20category%0A%20%20%20%20probabilities.%0A%20%20%20%20Then%2C%20given%20the%20feature%20vector%20%24x(t)%20%5Cin%20%5Cmathbf%7BR%7D%5En%24%2C%20the%20output%20%24y(t)%20%5Cin%20%5C%7B0%2C%201%5C%7D%24%20of%20this%20IO-HMM%20at%0A%20%20%20%20time%20step%20%24t%24%20is%20then%20generated%20from%20a%20logistic%20model%2C%20i.e.%2C%0A%0A%20%20%20%20%5C%5B%0A%20%20%20%20%20%20%20%20%5Cmathop%7B%5Cbf%20prob%7D(y(t)%20%3D%201)%20%3D%20%5Cfrac%7B1%7D%7B1%20%2B%20%5Cexp(-%7Bx(t)%7D%5ET%20%5Ctheta_%7B%5Chat%7Bz%7D(t)%7D)%7D%2C%0A%20%20%20%20%5C%5D%0A%0A%20%20%20%20where%20%24%5Ctheta_%7B%5Chat%7Bz%7D(t)%7D%20%5Cin%20%5C%7B%5Ctheta_1%2C%20%5Cldots%2C%20%5Ctheta_K%5C%7D%20%5Csubseteq%20%5Cmathbf%7BR%7D%5En%24%20is%20the%20coefficient.%0A%0A%20%20%20%20We%20are%20interested%20in%20recovering%20the%20transition%20matrix%20%24P_%7B%5Crm%20tr%7D%24%2C%20the%20model%20parameters%0A%20%20%20%20%24%5Ctheta_1%2C%20%5Cldots%2C%20%5Ctheta_K%24%2C%20and%20the%20unobserved%20state%20labels%20%24%5Chat%7Bz%7D(1)%2C%20%5Cldots%2C%20%5Chat%7Bz%7D(m)%24%2C%20given%20the%0A%20%20%20%20dataset%20%24(x(t)%2C%20y(t))%24%2C%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24.%0A%20%20%20%20Noticing%20that%20the%20transition%20matrix%20%24P_%7B%5Crm%20tr%7D%24%20can%20be%20easily%20estimated%20from%20the%20state%20labels%20%24%5Chat%7Bz%7D(t)%24%2C%0A%20%20%20%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24%2C%20we%20consider%20the%20following%20biconvex%20optimization%20problem%20for%20fitting%20the%20IO-HMM%3A%0A%0A%20%20%20%20%5C%5B%0A%20%20%20%20%20%20%20%20%5Cbegin%7Barray%7D%7Bll%7D%0A%20%20%20%20%20%20%20%20%20%20%20%20%5Ctext%7Bminimize%7D%20%26%20-%5Csum_%7Bt%20%3D%201%7D%5E%7Bm%7D%20%7Bz(t)%7D%5ET%20%7B%5Cleft(y(t)%7Bx(t)%7D%5ET%20%5Ctheta_k%20-%0A%20%20%20%20%5Clog(1%20%2B%20%5Cexp(%7Bx(t)%7D%5ET%20%5Ctheta_k))%5Cright)%7D_%7Bk%20%3D%201%7D%5EK%5C%5C%0A%20%20%20%20%20%20%20%20%20%20%20%20%26%5Cqquad%20%2B%20%5Calpha_%5Ctheta%20%5Csum_%7Bk%20%3D%201%7D%5E%7BK%7D%20%7B%5C%7C%5Ctheta_k%5C%7C%7D%5E2_2%20%2B%0A%20%20%20%20%5Calpha_z%20%5Csum_%7Bt%20%3D%201%7D%5E%7Bm%20-%201%7D%20D_%7B%5Crm%20kl%7D(z(t)%2C%20z(t%20%2B%201))%5C%5C%0A%20%20%20%20%20%20%20%20%20%20%20%20%5Ctext%7Bsubject%20to%7D%20%26%200%20%5Cpreceq%20z(t)%20%5Cpreceq%20%5Cmathbf%7B1%7D%2C%5Cquad%20%5Cmathbf%7B1%7D%5ET%20z(t)%20%3D%201%2C%5Cquad%20t%20%3D%201%2C%20%5Cldots%2C%20m%5C%5C%0A%20%20%20%20%20%20%20%20%20%20%20%20%26%20%5Ctheta_k%20%5Cin%20%7B%5Ccal%20C%7D_k%2C%5Cquad%20k%20%3D%201%2C%20%5Cldots%2C%20K%2C%0A%20%20%20%20%20%20%20%20%5Cend%7Barray%7D%0A%20%20%20%20%5C%5D%0A%0A%20%20%20%20where%20the%20optimization%20variables%20are%20%24%5Ctheta_k%20%5Cin%20%5Cmathbf%7BR%7D%5En%24%2C%20%24k%20%3D%201%2C%20%5Cldots%2C%20K%24%2C%20and%0A%20%20%20%20%24z(t)%20%5Cin%20%5Cmathbf%7BR%7D%5EK%24%2C%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24.%0A%20%20%20%20Note%20that%20the%20variable%20%24z(t)%24%20is%20a%20soft%20assignment%20vector%20for%20the%20hidden%20state%20label%20%24%5Chat%7Bz%7D(t)%24%2C%20where%20the%0A%20%20%20%20%24k%24th%20entry%20of%20%24z(t)%24%20indicates%20the%20probability%20of%20the%20state%20being%20%24k%24%20at%20time%20step%20%24t%24%2C%20and%20%24%5Chat%7Bz%7D(t)%24%20can%0A%20%20%20%20be%20estimated%20as%20the%20index%20of%20the%20largest%20entry%20of%20%24z(t)%24%20after%20solving%20the%20problem%20above.%0A%0A%20%20%20%20Each%20component%20of%20this%20problem%20can%20be%20interpreted%20as%20follows%3A%0A%20%20%20%20The%20first%20term%20in%20the%20objective%20function%20is%20the%20negative%20log-likelihood%20of%20the%20observed%20data%20under%20the%20IO-HMM%0A%20%20%20%20model%2C%20given%20the%20state%20assignment%20probabilities%20%24z(t)%24%2C%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24%2C%20and%20the%20model%20parameters%0A%20%20%20%20%24%5Ctheta_k%24%2C%20%24k%20%3D%201%2C%20%5Cldots%2C%20K%24.%0A%20%20%20%20The%20second%20term%20is%20a%20Tikhonov%20regularization%20on%20the%20model%20parameters%20%24%5Ctheta_k%24%2C%20with%20regularization%20parameter%0A%20%20%20%20%24%5Calpha_%5Ctheta%20%3E%200%24.%0A%20%20%20%20The%20third%20term%20is%20a%20temporal%20smoothness%20regularization%20on%20the%20state%20assignment%20probabilities%2C%20where%0A%20%20%20%20%24D_%7B%5Crm%20kl%7D(p%2C%20q)%24%20denotes%20the%20Kullback-Leibler%20divergence%20between%20two%20probability%20distributions%20%24p%24%20and%20%24q%24%2C%0A%20%20%20%20and%20%24%5Calpha_z%20%3E%200%24%20is%20the%20corresponding%20regularization%20parameter.%0A%20%20%20%20The%20constraints%20on%20the%20variables%20%24z(t)%24%2C%20%24t%20%3D%201%2C%20%5Cldots%2C%20m%24%2C%20ensure%20that%20they%20are%20valid%20probability%20distributions.%0A%20%20%20%20The%20sets%20%24%7B%5Ccal%20C%7D_k%20%5Csubseteq%20%5Cmathbf%7BR%7D%5En%24%2C%20%24k%20%3D%201%2C%20%5Cldots%2C%20K%24%2C%20are%20nonempty%20closed%20convex%20sets%20that%20encode%0A%20%20%20%20potential%20prior%20knowledge%20about%20the%20model%20parameters%20%24%5Ctheta_k%24.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Generate%20problem%20data%0A%0A%20%20%20%20We%20consider%20the%20case%20of%20%24n%20%3D%202%24%2C%20and%20the%20feature%20vector%20for%20each%20sample%20is%20generated%20according%20to%0A%0A%20%20%20%20%5C%5B%0A%20%20%20%20%20%20%20%20x(t)%20%5Csim%20(%7B%5Ccal%20U%7D(-5%2C%205)%2C%5C%201)%2C%0A%20%20%20%20%5C%5D%0A%0A%20%20%20%20where%20%24%7B%5Ccal%20U%7D(a%2C%20b)%24%20denotes%20a%20uniform%20distribution%20over%20the%20interval%20%24%5Ba%2C%20b%5D%24%2C%20and%20the%20second%20entry%20of%0A%20%20%20%20%24x(t)%24%20is%20always%20%241%24%20to%20account%20for%20the%20bias%20term.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(np)%3A%0A%20%20%20%20m%20%3D%201800%0A%20%20%20%20n%20%3D%202%0A%20%20%20%20K%20%3D%203%0A%20%20%20%20coefs%20%3D%20np.array(%5B%5B-1%2C%200%5D%2C%20%5B2%2C%206%5D%2C%20%5B2%2C%20-6%5D%5D)%0A%20%20%20%20p_tr%20%3D%20np.array(%5B%5B0.95%2C%200.025%2C%200.025%5D%2C%20%5B0.025%2C%200.95%2C%200.025%5D%2C%20%5B0.025%2C%200.025%2C%200.95%5D%5D)%0A%0A%20%20%20%20xs%20%3D%20np.random.uniform(-5%2C%205%2C%20m)%0A%20%20%20%20xs%20%3D%20np.vstack(%5Bxs%2C%20np.ones(m)%5D).T%0A%0A%20%20%20%20ys%20%3D%20np.zeros(m)%0A%20%20%20%20labels%20%3D%20np.zeros(m%2C%20dtype%3Dint)%0A%0A%20%20%20%20_s%20%3D%200%0A%20%20%20%20for%20_i%2C%20_feat%20in%20enumerate(xs)%3A%0A%20%20%20%20%20%20%20%20ys%5B_i%5D%20%3D%201%20if%20np.random.uniform()%20%3C%201%20%2F%20(1%20%2B%20np.exp(-_feat%20%40%20coefs%5B_s%5D))%20else%200%0A%20%20%20%20%20%20%20%20labels%5B_i%5D%20%3D%20_s%0A%20%20%20%20%20%20%20%20_s%20%3D%20np.random.choice(K%2C%20p%3Dp_tr%5B_s%5D)%0A%20%20%20%20return%20K%2C%20coefs%2C%20labels%2C%20m%2C%20n%2C%20p_tr%2C%20xs%2C%20ys%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Specify%20and%20solve%20the%20problem%0A%0A%20%20%20%20To%20fully%20specify%20the%20biconvex%20problem%2C%20it%20is%20assumed%20that%20we%20are%20given%20the%20following%20prior%20knowledge%20about%20the%0A%20%20%20%20coefficients%3A%0A%0A%20%20%20%20%5C%5B%0A%20%20%20%20%20%20%20%20%5Ctheta_%7B1%2C1%7D%20%5Cleq%200%2C%5Cquad%20%5Ctheta_%7B2%2C%201%7D%20%5Cgeq%200%2C%5Cquad%20%5Ctheta_%7B3%2C%201%7D%20%5Cgeq%200%2C%0A%20%20%20%20%20%20%20%20%5Cquad%20%5Ctheta_%7B2%2C%202%7D%20%5Cgeq%20%5Ctheta_%7B3%2C%202%7D%2C%0A%20%20%20%20%5C%5D%0A%0A%20%20%20%20where%20%24%5Ctheta_%7Bi%2C%20j%7D%24%20denotes%20the%20%24j%24th%20entry%20of%20the%20vector%20%24%5Ctheta_i%24.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(BiconvexProblem%2C%20K%2C%20cp%2C%20m%2C%20n%2C%20xs%2C%20ys)%3A%0A%20%20%20%20thetas%20%3D%20cp.Variable((K%2C%20n))%0A%20%20%20%20zs%20%3D%20cp.Variable((m%2C%20K)%2C%20nonneg%3DTrue)%0A%0A%20%20%20%20alpha_theta%20%3D%200.1%0A%20%20%20%20alpha_z%20%3D%202%0A%0A%20%20%20%20rs%20%3D%20%5B-cp.multiply(ys%2C%20xs%20%40%20thetas%5Bk%5D)%20%2B%20cp.logistic(xs%20%40%20thetas%5Bk%5D)%20for%20k%20in%20range(K)%5D%0A%20%20%20%20obj%20%3D%20cp.Minimize(%0A%20%20%20%20%20%20%20%20cp.sum(cp.multiply(zs%2C%20cp.vstack(rs).T))%0A%20%20%20%20%20%20%20%20%2B%20alpha_theta%20*%20cp.sum_squares(thetas)%0A%20%20%20%20%20%20%20%20%2B%20alpha_z%20*%20cp.sum(cp.kl_div(zs%5B%3A-1%5D%2C%20zs%5B1%3A%5D))%0A%20%20%20%20)%0A%20%20%20%20constr%20%3D%20%5B%0A%20%20%20%20%20%20%20%20thetas%5B0%5D%5B0%5D%20%3C%3D%200%2C%0A%20%20%20%20%20%20%20%20thetas%5B1%5D%5B0%5D%20%3E%3D%200%2C%0A%20%20%20%20%20%20%20%20thetas%5B2%5D%5B0%5D%20%3E%3D%200%2C%0A%20%20%20%20%20%20%20%20thetas%5B1%5D%5B1%5D%20%3E%3D%20thetas%5B2%5D%5B1%5D%2C%0A%20%20%20%20%20%20%20%20zs%20%3C%3D%201%2C%0A%20%20%20%20%20%20%20%20cp.sum(zs%2C%20axis%3D1)%20%3D%3D%201%2C%0A%20%20%20%20%5D%0A%0A%20%20%20%20prob%20%3D%20BiconvexProblem(obj%2C%20%5Bzs%5D%2C%20%5Bthetas%5D%2C%20constr)%0A%20%20%20%20prob.solve(solver%3Dcp.CLARABEL%2C%20mode%3D%22penalty%22%2C%20nu%3D1e2%2C%20lbd%3D0.1%2C%20abs_tol%3D1e-3)%0A%20%20%20%20return%20thetas%2C%20zs%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Plot%20the%20results%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(K%2C%20coefs%2C%20figure_directory%2C%20labels%2C%20m%2C%20np%2C%20plt%2C%20thetas%2C%20zs)%3A%0A%20%20%20%20fig%2C%20axs%20%3D%20plt.subplots(2%2C%201%2C%20figsize%3D(6.5%2C%207))%0A%0A%20%20%20%20axs%5B0%5D.plot(labels%2C%20linestyle%3D%22dashed%22%2C%20color%3D%22k%22%2C%20linewidth%3D1%2C%20zorder%3D10)%0A%20%20%20%20axs%5B0%5D.plot(np.argmax(zs.value%2C%20axis%3D-1)%2C%20color%3D%22r%22%2C%20linewidth%3D2)%0A%0A%20%20%20%20inputs%20%3D%20np.linspace(-5%2C%205%2C%20m)%0A%20%20%20%20inputs%20%3D%20np.vstack(%5Binputs%2C%20np.ones(m)%5D).T%0A%20%20%20%20for%20_i%20in%20range(K)%3A%0A%20%20%20%20%20%20%20%20axs%5B1%5D.plot(inputs%5B%3A%2C%200%5D%2C%201%20%2F%20(1%20%2B%20np.exp(-inputs%20%40%20coefs%5B_i%5D))%2C%20linestyle%3D%22dashed%22%2C%20color%3D%22k%22%2C%20zorder%3D10)%0A%20%20%20%20%20%20%20%20axs%5B1%5D.plot(inputs%5B%3A%2C%200%5D%2C%201%20%2F%20(1%20%2B%20np.exp(-inputs%20%40%20thetas%5B_i%5D.value)))%0A%0A%20%20%20%20axs%5B0%5D.set_xlabel(%22%24t%24%22)%0A%20%20%20%20axs%5B0%5D.set_ylabel(r%22%24%5Chat%7Bz%7D(t)%24%22)%0A%20%20%20%20axs%5B0%5D.set_yticks(%5B0%2C%201%2C%202%5D)%0A%20%20%20%20axs%5B0%5D.set_yticklabels(%5B1%2C%202%2C%203%5D)%0A%0A%20%20%20%20axs%5B1%5D.set_xlabel(r%22%24x_1%24%22)%0A%20%20%20%20axs%5B1%5D.set_ylabel(r%22%241%2F(1%20%2B%20%5Cexp(-x%5ET%20%5Ctheta))%24%22)%0A%0A%20%20%20%20fig.tight_layout()%0A%20%20%20%20fig.savefig(figure_directory%20%2F%20%22iohmm.pdf%22%2C%20bbox_inches%3D%22tight%22)%0A%20%20%20%20plt.show()%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(K%2C%20m%2C%20np%2C%20p_tr%2C%20zs)%3A%0A%20%20%20%20p_tr_hat%20%3D%20np.zeros_like(p_tr)%0A%20%20%20%20z_hat%20%3D%20np.argmax(zs.value%2C%20axis%3D-1)%0A%20%20%20%20for%20zi%20in%20range(K)%3A%0A%20%20%20%20%20%20%20%20z_idx%20%3D%20np.where(z_hat%20%3D%3D%20zi)%5B0%5D%0A%20%20%20%20%20%20%20%20z_idx%20%3D%20np.delete(z_idx%2C%20np.where(z_idx%20%3D%3D%20m%20-%201)%5B0%5D)%0A%20%20%20%20%20%20%20%20_%2C%20nz_num%20%3D%20np.unique(z_hat%5Bz_idx%20%2B%201%5D%2C%20return_counts%3DTrue)%0A%20%20%20%20%20%20%20%20p_tr_hat%5Bzi%5D%20%3D%20nz_num%20%2F%20len(z_idx)%0A%0A%20%20%20%20print(p_tr_hat)%0A%20%20%20%20return%0A%0A%0Aif%20__name__%20%3D%3D%20%22__main__%22%3A%0A%20%20%20%20app.run()%0A
ce75a4df00eaad515ec0f27bd48929e3