import%20marimo%0A%0A__generated_with%20%3D%20%220.24.0%22%0Aapp%20%3D%20marimo.App()%0A%0A%0A%40app.cell%0Adef%20_()%3A%0A%20%20%20%20import%20marimo%20as%20mo%0A%20%20%20%20import%20numpy%20as%20np%0A%20%20%20%20import%20pandas%20as%20pd%0A%20%20%20%20import%20plotly.graph_objects%20as%20go%0A%20%20%20%20import%20torch%0A%20%20%20%20from%20plotly.subplots%20import%20make_subplots%0A%0A%20%20%20%20return%20go%2C%20make_subplots%2C%20mo%2C%20np%2C%20pd%2C%20torch%0A%0A%0A%40app.cell%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%5B%E2%86%90%2041%20Pseudo%20R-squared%5D(41_pseudo_r2.py)%20%7C%20%5BIndex%5D(..%2Findex.html)%20%7C%20%5B43%20Energy%20Statistics%20%E2%86%92%5D(43_energy.py)%0A%0A%20%20%20%20%23%20Matrix%20Calculus%20of%20Multiclass%20Classification%3A%20Softmax%2C%20Cross-Entropy%2C%20and%20Vector-Jacobian%20Products%0A%0A%20%20%20%20%23%23%20%5Ba%5D%20Why%20do%20you%20need%20to%20know%20these%20concepts%3F%0A%0A%20%20%20%20Every%20modern%20classification%20neural%20network%E2%80%94from%20Vision%20Transformers%20(ViTs)%20and%20ResNets%20to%20Large%20Language%20Models%20(LLMs)%20predicting%20vocabulary%20distributions%20over%20hundreds%20of%20thousands%20of%20tokens%E2%80%94terminates%20with%20a%20**linear%20projection%20layer%20followed%20by%20Softmax%20and%20Cross-Entropy%20Loss**.%0A%0A%20%20%20%20%23%23%23%23%20The%20Elegant%20Cancellation%3A%20%24%5Cnabla_z%20L%20%3D%20p%20-%20y%24%0A%20%20%20%20When%20computing%20the%20gradient%20of%20Cross-Entropy%20loss%20through%20Softmax%2C%20a%20mathematical%20cancellation%20occurs%3A%0A%20%20%20%20-%20The%20gradient%20of%20the%20scalar%20loss%20with%20respect%20to%20probabilities%20is%20non-linear%20and%20divided%20by%20probabilities%3A%20%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20p_i%7D%20%3D%20-%5Cfrac%7By_i%7D%7Bp_i%7D%24.%0A%20%20%20%20-%20The%20Jacobian%20matrix%20of%20the%20Softmax%20function%20is%20dense%20with%20quadratic%20terms%3A%20%24J_%7Bi%2C%20j%7D%20%3D%20p_i(%5Cdelta_%7Bij%7D%20-%20p_j)%24.%0A%20%20%20%20-%20When%20multiplied%20via%20the%20chain%20rule%2C%20the%20%24p_i%24%20terms%20in%20the%20denominator%20cancel%20out%20completely%2C%20yielding%20the%20simple%20result%3A%0A%0A%20%20%20%20%24%24%5Cnabla_z%20L%20%3D%20%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z%7D%20%3D%20p%20-%20y%24%24%0A%0A%20%20%20%20The%20gradient%20with%20respect%20to%20the%20raw%20unnormalized%20logits%20is%20simply%20the%20**prediction%20error%20vector**%20(the%20predicted%20probability%20vector%20minus%20the%20one-hot%20target%20vector).%0A%0A%20%20%20%20%23%23%23%23%20Vector-Jacobian%20Products%20(VJPs)%20in%20Backpropagation%0A%20%20%20%20In%20reverse-mode%20automatic%20differentiation%20(backpropagation)%2C%20forming%20and%20materializing%20the%20full%20%24n%20%5Ctimes%20n%24%20Jacobian%20matrix%20%24%5Cfrac%7B%5Cpartial%20p%7D%7B%5Cpartial%20z%7D%24%20would%20consume%20%24O(n%5E2)%24%20memory%20and%20compute.%20For%20an%20LLM%20vocabulary%20size%20of%20%24n%20%3D%20128%2C000%24%2C%20an%20explicit%20Jacobian%20would%20require%20over%20%2465%24%20gigabytes%20of%20memory%20for%20a%20single%20token.%20By%20computing%20the%20Vector-Jacobian%20Product%20(VJP)%20analytically%2C%20automatic%20differentiation%20evaluates%20the%20backward%20pass%20in%20%24O(n)%24%20time%20and%20%24O(n)%24%20memory.%0A%0A%20%20%20%20%23%23%23%23%20Numerical%20Stability%20and%20the%20LogSumExp%20Trick%0A%20%20%20%20Directly%20evaluating%20%24e%5E%7Bz_i%7D%24%20leads%20to%20floating-point%20overflow%20for%20large%20logits%20(e.g.%2C%20%24e%5E%7B100%7D%20%5Capprox%202.68%20%5Ctimes%2010%5E%7B43%7D%24)%20and%20underflow%20for%20negative%20logits.%20In%20practice%2C%20production%20implementations%20combine%20Softmax%20and%20Cross-Entropy%20into%20a%20fused%2C%20numerically%20stable%20kernel%20using%20the%20LogSumExp%20identity%3A%0A%0A%20%20%20%20%24%24%5Cln%20%5Csum_%7Bj%3D1%7D%5En%20e%5E%7Bz_j%7D%20%3D%20M%20%2B%20%5Cln%20%5Csum_%7Bj%3D1%7D%5En%20e%5E%7Bz_j%20-%20M%7D%2C%20%5Cquad%20%5Ctext%7Bwhere%20%7D%20M%20%3D%20%5Cmax_%7Bj%7D%20z_j%24%24%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(mo)%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20%5Bb%5D%20Mathematical%20Foundations%20and%20Matrix%20Derivations%0A%0A%20%20%20%20%23%23%23%201.%20Model%20Formulation%0A%0A%20%20%20%20Let%20%24x%20%5Cin%20%5Cmathbb%7BR%7D%5Em%24%20be%20an%20input%20feature%20vector%2C%20%24W%20%5Cin%20%5Cmathbb%7BR%7D%5E%7Bn%20%5Ctimes%20m%7D%24%20the%20weight%20parameter%20matrix%2C%20and%20%24b%20%5Cin%20%5Cmathbb%7BR%7D%5En%24%20the%20bias%20vector%2C%20where%20%24n%24%20is%20the%20number%20of%20mutually%20exclusive%20classes.%0A%0A%20%20%20%20The%20linear%20pre-activation%20(logits)%20vector%20is%3A%0A%0A%20%20%20%20%24%24z%20%3D%20W%20x%20%2B%20b%20%5Cin%20%5Cmathbb%7BR%7D%5En%24%24%0A%0A%20%20%20%20The%20Softmax%20function%20normalizes%20logits%20into%20a%20valid%20probability%20distribution%20%24p%20%5Cin%20%5CDelta%5E%7Bn-1%7D%24%3A%0A%0A%20%20%20%20%24%24p_i%20%3D%20%5Coperatorname%7Bsoftmax%7D(z)_i%20%3D%20%5Cfrac%7Be%5E%7Bz_i%7D%7D%7B%5Csum_%7Bk%3D1%7D%5En%20e%5E%7Bz_k%7D%7D%2C%20%5Cquad%20%5Ctext%7Bfor%20%7D%20i%20%5Cin%20%5C%7B1%2C%20%5Cdots%2C%20n%5C%7D%24%24%0A%0A%20%20%20%20Let%20%24y%20%5Cin%20%5C%7B0%2C%201%5C%7D%5En%24%20be%20the%20one-hot%20encoded%20ground%20truth%20target%20vector%2C%20where%20%24y_c%20%3D%201%24%20for%20the%20true%20class%20%24c%24%20and%20%24y_j%20%3D%200%24%20for%20all%20%24j%20%5Cneq%20c%24.%20The%20Multiclass%20Cross-Entropy%20Loss%20is%20defined%20as%3A%0A%0A%20%20%20%20%24%24L%20%3D%20-%5Csum_%7Bi%3D1%7D%5En%20y_i%20%5Cln(p_i)%20%3D%20-y%5E%5Ctop%20%5Cln(p)%20%3D%20-%5Cln(p_c)%24%24%0A%0A%20%20%20%20%23%23%23%202.%20Jacobian%20Matrix%20of%20the%20Softmax%20Function%0A%0A%20%20%20%20The%20Softmax%20mapping%20%24%5Coperatorname%7Bsoftmax%7D%3A%20%5Cmathbb%7BR%7D%5En%20%5Cto%20%5Cmathbb%7BR%7D%5En%24%20produces%20an%20%24n%20%5Ctimes%20n%24%20Jacobian%20matrix%20%24J%20%5Cin%20%5Cmathbb%7BR%7D%5E%7Bn%20%5Ctimes%20n%7D%24%2C%20where%20%24J_%7Bi%2C%20j%7D%20%3D%20%5Cfrac%7B%5Cpartial%20p_i%7D%7B%5Cpartial%20z_j%7D%24.%0A%0A%20%20%20%20We%20evaluate%20two%20cases%20using%20the%20quotient%20rule%3A%0A%0A%20%20%20%20**Case%201%3A%20Diagonal%20Elements%20(%24i%20%3D%20j%24)**%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20p_i%7D%7B%5Cpartial%20z_i%7D%20%3D%20%5Cfrac%7B%5Cfrac%7B%5Cpartial%7D%7B%5Cpartial%20z_i%7D(e%5E%7Bz_i%7D)%20%5Ccdot%20%5Csum_%7Bk%7D%20e%5E%7Bz_k%7D%20-%20e%5E%7Bz_i%7D%20%5Ccdot%20%5Cfrac%7B%5Cpartial%7D%7B%5Cpartial%20z_i%7D(%5Csum_%7Bk%7D%20e%5E%7Bz_k%7D)%7D%7B%5Cleft(%5Csum_%7Bk%7D%20e%5E%7Bz_k%7D%5Cright)%5E2%7D%20%3D%20%5Cfrac%7Be%5E%7Bz_i%7D%7D%7B%5Csum%20e%5E%7Bz_k%7D%7D%20-%20%5Cleft(%5Cfrac%7Be%5E%7Bz_i%7D%7D%7B%5Csum%20e%5E%7Bz_k%7D%7D%5Cright)%5E2%20%3D%20p_i%20-%20p_i%5E2%20%3D%20p_i(1%20-%20p_i)%24%24%0A%0A%20%20%20%20**Case%202%3A%20Off-Diagonal%20Elements%20(%24i%20%5Cneq%20j%24)**%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20p_i%7D%7B%5Cpartial%20z_j%7D%20%3D%20%5Cfrac%7B0%20%5Ccdot%20%5Csum_%7Bk%7D%20e%5E%7Bz_k%7D%20-%20e%5E%7Bz_i%7D%20%5Ccdot%20e%5E%7Bz_j%7D%7D%7B%5Cleft(%5Csum_%7Bk%7D%20e%5E%7Bz_k%7D%5Cright)%5E2%7D%20%3D%20-%5Cfrac%7Be%5E%7Bz_i%7D%7D%7B%5Csum%20e%5E%7Bz_k%7D%7D%20%5Cfrac%7Be%5E%7Bz_j%7D%7D%7B%5Csum%20e%5E%7Bz_k%7D%7D%20%3D%20-p_i%20p_j%24%24%0A%0A%20%20%20%20Using%20the%20Kronecker%20delta%20%24%5Cdelta_%7Bij%7D%24%20(where%20%24%5Cdelta_%7Bij%7D%20%3D%201%24%20if%20%24i%3Dj%24%2C%20else%20%240%24)%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20p_i%7D%7B%5Cpartial%20z_j%7D%20%3D%20p_i(%5Cdelta_%7Bij%7D%20-%20p_j)%24%24%0A%0A%20%20%20%20In%20compact%20matrix%20form%3A%0A%0A%20%20%20%20%24%24J_%7B%5Ctext%7Bsoftmax%7D%7D%20%3D%20%5Cfrac%7B%5Cpartial%20p%7D%7B%5Cpartial%20z%7D%20%3D%20%5Coperatorname%7Bdiag%7D(p)%20-%20p%20p%5E%5Ctop%24%24%0A%0A%20%20%20%20%23%23%23%203.%20Derivation%20of%20the%20Logit%20Gradient%20%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z%7D%24%0A%0A%20%20%20%20The%20gradient%20of%20Cross-Entropy%20loss%20with%20respect%20to%20probability%20%24p_i%24%20is%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20p_i%7D%20%3D%20-%5Cfrac%7By_i%7D%7Bp_i%7D%24%24%0A%0A%20%20%20%20Applying%20the%20multivariate%20chain%20rule%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z_j%7D%20%3D%20%5Csum_%7Bi%3D1%7D%5En%20%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20p_i%7D%20%5Cfrac%7B%5Cpartial%20p_i%7D%7B%5Cpartial%20z_j%7D%20%3D%20%5Csum_%7Bi%3D1%7D%5En%20%5Cleft(-%5Cfrac%7By_i%7D%7Bp_i%7D%5Cright)%20%5Cleft%5B%20p_i(%5Cdelta_%7Bij%7D%20-%20p_j)%20%5Cright%5D%20%3D%20-%5Csum_%7Bi%3D1%7D%5En%20y_i%20(%5Cdelta_%7Bij%7D%20-%20p_j)%24%24%0A%0A%20%20%20%20Distributing%20the%20sum%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z_j%7D%20%3D%20-%5Csum_%7Bi%3D1%7D%5En%20y_i%20%5Cdelta_%7Bij%7D%20%2B%20p_j%20%5Csum_%7Bi%3D1%7D%5En%20y_i%24%24%0A%0A%20%20%20%20Because%20%24y%24%20is%20a%20one-hot%20distribution%2C%20%24%5Csum_%7Bi%3D1%7D%5En%20y_i%20%3D%201%24%2C%20and%20the%20sifting%20property%20gives%20%24%5Csum_%7Bi%3D1%7D%5En%20y_i%20%5Cdelta_%7Bij%7D%20%3D%20y_j%24.%20Therefore%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z_j%7D%20%3D%20-y_j%20%2B%20p_j%20%3D%20p_j%20-%20y_j%24%24%0A%0A%20%20%20%20In%20full%20vector%20notation%3A%0A%0A%20%20%20%20%24%24%5Cnabla_z%20L%20%3D%20%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z%7D%20%3D%20p%20-%20y%24%24%0A%0A%20%20%20%20%23%23%23%204.%20Gradient%20with%20Respect%20to%20Parameters%20%24W%24%20and%20%24b%24%0A%0A%20%20%20%20Since%20%24z%20%3D%20W%20x%20%2B%20b%24%2C%20the%20component-wise%20derivative%20is%20%24%5Cfrac%7B%5Cpartial%20z_j%7D%7B%5Cpartial%20W_%7Bj%2C%20k%7D%7D%20%3D%20x_k%24.%20Applying%20the%20chain%20rule%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20W_%7Bj%2C%20k%7D%7D%20%3D%20%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20z_j%7D%20%5Cfrac%7B%5Cpartial%20z_j%7D%7B%5Cpartial%20W_%7Bj%2C%20k%7D%7D%20%3D%20(p_j%20-%20y_j)%20x_k%24%24%0A%0A%20%20%20%20Expressing%20this%20as%20an%20outer%20product%20in%20matrix%20calculus%3A%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20W%7D%20%3D%20(%5Cnabla_z%20L)%20x%5E%5Ctop%20%3D%20(p%20-%20y)%20x%5E%5Ctop%20%5Cin%20%5Cmathbb%7BR%7D%5E%7Bn%20%5Ctimes%20m%7D%24%24%0A%0A%20%20%20%20%24%24%5Cfrac%7B%5Cpartial%20L%7D%7B%5Cpartial%20b%7D%20%3D%20%5Cnabla_z%20L%20%3D%20p%20-%20y%20%5Cin%20%5Cmathbb%7BR%7D%5En%24%24%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(np%2C%20torch)%3A%0A%20%20%20%20%23%20Fix%20seed%20for%20reproducible%20gradient%20audit%0A%20%20%20%20torch.manual_seed(42)%0A%20%20%20%20np.random.seed(42)%0A%0A%20%20%20%20%23%203-class%20classification%20with%202%20input%20features%0A%20%20%20%20m_features%20%3D%202%0A%20%20%20%20n_classes%20%3D%203%0A%0A%20%20%20%20%23%20Sample%20input%20x%20and%20weight%20matrix%20W%0A%20%20%20%20x_input%20%3D%20torch.tensor(%5B%5B2.5%5D%2C%20%5B-1.2%5D%5D%2C%20dtype%3Dtorch.float64)%0A%20%20%20%20w_matrix%20%3D%20torch.tensor(%0A%20%20%20%20%20%20%20%20%5B%5B0.8%2C%20-0.5%5D%2C%20%5B-0.3%2C%201.2%5D%2C%20%5B0.4%2C%200.6%5D%5D%2C%0A%20%20%20%20%20%20%20%20dtype%3Dtorch.float64%2C%0A%20%20%20%20%20%20%20%20requires_grad%3DTrue%2C%0A%20%20%20%20)%0A%20%20%20%20b_bias%20%3D%20torch.tensor(%5B%5B0.1%5D%2C%20%5B-0.2%5D%2C%20%5B0.3%5D%5D%2C%20dtype%3Dtorch.float64%2C%20requires_grad%3DTrue)%0A%0A%20%20%20%20%23%20Target%20class%3A%20Class%201%20(0-indexed%3A%20%5B0%2C%201%2C%200%5D%5ET)%0A%20%20%20%20target_idx%20%3D%201%0A%20%20%20%20y_onehot%20%3D%20torch.zeros(n_classes%2C%201%2C%20dtype%3Dtorch.float64)%0A%20%20%20%20y_onehot%5Btarget_idx%5D%20%3D%201.0%0A%0A%20%20%20%20%23%20Forward%20pass%0A%20%20%20%20z_logits%20%3D%20w_matrix%20%40%20x_input%20%2B%20b_bias%0A%20%20%20%20z_logits.retain_grad()%0A%0A%20%20%20%20p_probs%20%3D%20torch.softmax(z_logits%2C%20dim%3D0)%0A%20%20%20%20p_probs.retain_grad()%0A%0A%20%20%20%20loss_val%20%3D%20-torch.sum(y_onehot%20*%20torch.log(p_probs))%0A%20%20%20%20loss_val.backward(retain_graph%3DTrue)%0A%0A%20%20%20%20%23%20Analytical%20computations%0A%20%20%20%20p_np%20%3D%20p_probs.detach().numpy().flatten()%0A%20%20%20%20y_np%20%3D%20y_onehot.detach().numpy().flatten()%0A%20%20%20%20x_np%20%3D%20x_input.detach().numpy()%0A%0A%20%20%20%20%23%201.%20Softmax%20Jacobian%0A%20%20%20%20jacobian_analytical%20%3D%20np.diag(p_np)%20-%20np.outer(p_np%2C%20p_np)%0A%0A%20%20%20%20%23%202.%20Gradient%20w.r.t%20logits%3A%20p%20-%20y%0A%20%20%20%20grad_z_analytical%20%3D%20(p_np%20-%20y_np).reshape(-1%2C%201)%0A%0A%20%20%20%20%23%203.%20Gradient%20w.r.t%20weights%3A%20(p%20-%20y)%20x%5ET%0A%20%20%20%20grad_w_analytical%20%3D%20grad_z_analytical%20%40%20x_np.T%0A%20%20%20%20return%20(%0A%20%20%20%20%20%20%20%20b_bias%2C%0A%20%20%20%20%20%20%20%20grad_w_analytical%2C%0A%20%20%20%20%20%20%20%20grad_z_analytical%2C%0A%20%20%20%20%20%20%20%20jacobian_analytical%2C%0A%20%20%20%20%20%20%20%20w_matrix%2C%0A%20%20%20%20%20%20%20%20z_logits%2C%0A%20%20%20%20)%0A%0A%0A%40app.cell%0Adef%20_(go%2C%20grad_w_analytical%2C%20jacobian_analytical%2C%20make_subplots%2C%20mo%2C%20np)%3A%0A%20%20%20%20fig%20%3D%20make_subplots(%0A%20%20%20%20%20%20%20%20rows%3D1%2C%0A%20%20%20%20%20%20%20%20cols%3D2%2C%0A%20%20%20%20%20%20%20%20subplot_titles%3D%5B%0A%20%20%20%20%20%20%20%20%20%20%20%20%22%3Cb%3ESoftmax%20Jacobian%20Matrix%3A%20diag(p)%20-%20p%20p%5ET%3C%2Fb%3E%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%22%3Cb%3EWeight%20Gradient%20Outer%20Product%3A%20(p%20-%20y)%20x%5ET%3C%2Fb%3E%22%2C%0A%20%20%20%20%20%20%20%20%5D%2C%0A%20%20%20%20%20%20%20%20horizontal_spacing%3D0.14%2C%0A%20%20%20%20)%0A%0A%20%20%20%20class_names%20%3D%20%5B%22Class%200%22%2C%20%22Class%201%22%2C%20%22Class%202%22%5D%0A%20%20%20%20feat_names%20%3D%20%5B%22Feat%200%22%2C%20%22Feat%201%22%5D%0A%0A%20%20%20%20%23%20Left%3A%20Softmax%20Jacobian%20Heatmap%0A%20%20%20%20fig.add_trace(%0A%20%20%20%20%20%20%20%20go.Heatmap(%0A%20%20%20%20%20%20%20%20%20%20%20%20z%3Dnp.round(jacobian_analytical%2C%203)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20x%3Dclass_names%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20y%3Dclass_names%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20colorscale%3D%22Blues%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20text%3Dnp.round(jacobian_analytical%2C%203)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20texttemplate%3D%22%25%7Btext%7D%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20colorbar%3Ddict(title%3D%22dp_i%20%2F%20dz_j%22%2C%20x%3D0.44%2C%20len%3D0.8)%2C%0A%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20%20%20%20%20row%3D1%2C%0A%20%20%20%20%20%20%20%20col%3D1%2C%0A%20%20%20%20)%0A%0A%20%20%20%20%23%20Right%3A%20Gradient%20w.r.t%20W%20Heatmap%0A%20%20%20%20fig.add_trace(%0A%20%20%20%20%20%20%20%20go.Heatmap(%0A%20%20%20%20%20%20%20%20%20%20%20%20z%3Dnp.round(grad_w_analytical%2C%203)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20x%3Dfeat_names%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20y%3Dclass_names%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20colorscale%3D%22RdBu_r%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20text%3Dnp.round(grad_w_analytical%2C%203)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20texttemplate%3D%22%25%7Btext%7D%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20colorbar%3Ddict(title%3D%22dL%20%2F%20dW%22%2C%20x%3D1.02%2C%20len%3D0.8)%2C%0A%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20%20%20%20%20row%3D1%2C%0A%20%20%20%20%20%20%20%20col%3D2%2C%0A%20%20%20%20)%0A%0A%20%20%20%20fig.update_layout(%0A%20%20%20%20%20%20%20%20template%3D%22plotly_white%22%2C%0A%20%20%20%20%20%20%20%20height%3D480%2C%0A%20%20%20%20%20%20%20%20margin%3Ddict(l%3D40%2C%20r%3D40%2C%20t%3D70%2C%20b%3D40)%2C%0A%20%20%20%20)%0A%0A%20%20%20%20viz%20%3D%20mo.ui.plotly(fig)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_()%3A%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_(%0A%20%20%20%20b_bias%2C%0A%20%20%20%20grad_w_analytical%2C%0A%20%20%20%20grad_z_analytical%2C%0A%20%20%20%20jacobian_analytical%2C%0A%20%20%20%20mo%2C%0A%20%20%20%20np%2C%0A%20%20%20%20pd%2C%0A%20%20%20%20torch%2C%0A%20%20%20%20w_matrix%2C%0A%20%20%20%20z_logits%2C%0A)%3A%0A%20%20%20%20%23%20Example%201%3A%20Numerical%20Validation%3A%20Analytical%20Formulas%20vs%20PyTorch%20Autograd%0A%20%20%20%20autograd_grad_w%20%3D%20w_matrix.grad.detach().numpy()%0A%20%20%20%20autograd_grad_z%20%3D%20z_logits.grad.detach().numpy().flatten()%0A%20%20%20%20autograd_grad_b%20%3D%20b_bias.grad.detach().numpy().flatten()%0A%0A%20%20%20%20df_verification%20%3D%20pd.DataFrame(%0A%20%20%20%20%20%20%20%20%5B%0A%20%20%20%20%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Quantity%22%3A%20%22Logit%20Gradient%3A%20dL%20%2F%20dz%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Formula%22%3A%20%22p%20-%20y%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Values%22%3A%20str(grad_z_analytical.flatten().round(5))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22PyTorch_Autograd%22%3A%20str(autograd_grad_z.round(5))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Max_Absolute_Diff%22%3A%20float(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20np.max(np.abs(grad_z_analytical.flatten()%20-%20autograd_grad_z))%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Quantity%22%3A%20%22Weight%20Gradient%3A%20dL%20%2F%20dW%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Formula%22%3A%20%22(p%20-%20y)%20x%5ET%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Values%22%3A%20str(grad_w_analytical.flatten().round(5))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22PyTorch_Autograd%22%3A%20str(autograd_grad_w.flatten().round(5))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Max_Absolute_Diff%22%3A%20float(np.max(np.abs(grad_w_analytical%20-%20autograd_grad_w)))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Quantity%22%3A%20%22Bias%20Gradient%3A%20dL%20%2F%20db%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Formula%22%3A%20%22p%20-%20y%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Values%22%3A%20str(grad_z_analytical.flatten().round(5))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22PyTorch_Autograd%22%3A%20str(autograd_grad_b.round(5))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Max_Absolute_Diff%22%3A%20float(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20np.max(np.abs(grad_z_analytical.flatten()%20-%20autograd_grad_b))%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7D%2C%0A%20%20%20%20%20%20%20%20%5D%0A%20%20%20%20)%0A%0A%20%20%20%20%23%20Example%202%3A%20PyTorch%20Autograd%20Functional%20Jacobian%20vs%20Analytical%20Jacobian%0A%20%20%20%20def%20softmax_wrapper(z)%3A%0A%20%20%20%20%20%20%20%20return%20torch.softmax(z%2C%20dim%3D0)%0A%0A%20%20%20%20autograd_jacobian%20%3D%20(%0A%20%20%20%20%20%20%20%20torch.autograd.functional.jacobian(softmax_wrapper%2C%20z_logits).squeeze().detach().numpy()%0A%20%20%20%20)%0A%0A%20%20%20%20df_jacobian%20%3D%20pd.DataFrame(%0A%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%22Jacobian_Element%22%3A%20%5Bf%22J%5B%7Bi%7D%2C%7Bj%7D%5D%22%20for%20i%20in%20range(3)%20for%20j%20in%20range(3)%5D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%22Analytical_Value%22%3A%20jacobian_analytical.flatten().round(6)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%22Autograd_Functional%22%3A%20autograd_jacobian.flatten().round(6)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%22Difference%22%3A%20np.abs(jacobian_analytical.flatten()%20-%20autograd_jacobian.flatten()).round(%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%209%0A%20%20%20%20%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20%20%20%20%20%7D%0A%20%20%20%20)%0A%0A%20%20%20%20%23%20Example%203%3A%20Numerical%20Stability%20Benchmark%3A%20Naive%20Softmax%20vs%20LogSumExp%20Trick%0A%20%20%20%20extreme_logits%20%3D%20np.array(%5B1000.0%2C%201002.0%2C%20995.0%5D)%0A%0A%20%20%20%20%23%20Naive%20Softmax%3A%20Overflow%20occurs%20in%20exp(1000)%0A%20%20%20%20with%20np.errstate(over%3D%22ignore%22%2C%20invalid%3D%22ignore%22)%3A%0A%20%20%20%20%20%20%20%20exp_naive%20%3D%20np.exp(extreme_logits)%0A%20%20%20%20%20%20%20%20p_naive%20%3D%20exp_naive%20%2F%20np.sum(exp_naive)%0A%20%20%20%20%20%20%20%20loss_naive%20%3D%20-np.log(p_naive%5B1%5D)%0A%0A%20%20%20%20%23%20Numerically%20Stable%20Softmax%3A%20Subtract%20max(z)%0A%20%20%20%20max_z%20%3D%20np.max(extreme_logits)%0A%20%20%20%20exp_stable%20%3D%20np.exp(extreme_logits%20-%20max_z)%0A%20%20%20%20p_stable%20%3D%20exp_stable%20%2F%20np.sum(exp_stable)%0A%20%20%20%20loss_stable%20%3D%20-(extreme_logits%5B1%5D%20-%20max_z)%20%2B%20np.log(np.sum(exp_stable))%0A%0A%20%20%20%20df_stability%20%3D%20pd.DataFrame(%0A%20%20%20%20%20%20%20%20%5B%0A%20%20%20%20%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Implementation%22%3A%20%22Naive%20Softmax%20(exp(z)%20%2F%20sum(exp(z)))%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Max_Logit%22%3A%201002.0%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Probabilities%22%3A%20str(p_naive)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Calculated_Loss%22%3A%20str(loss_naive)%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Numerical_Status%22%3A%20%22Failed%20(NaN%20%2F%20Inf%20Overflow)%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7D%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7B%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Implementation%22%3A%20%22LogSumExp%20Stable%20(z%20-%20max(z))%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Max_Logit%22%3A%201002.0%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Probabilities%22%3A%20str(np.round(p_stable%2C%204))%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Calculated_Loss%22%3A%20f%22%7Bloss_stable%3A.4f%7D%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%22Numerical_Status%22%3A%20%22Exact%2C%20Mathematically%20Stable%22%2C%0A%20%20%20%20%20%20%20%20%20%20%20%20%7D%2C%0A%20%20%20%20%20%20%20%20%5D%0A%20%20%20%20)%0A%0A%20%20%20%20table_verif%20%3D%20mo.ui.table(df_verification)%0A%20%20%20%20table_jac%20%3D%20mo.ui.table(df_jacobian)%0A%20%20%20%20table_stab%20%3D%20mo.ui.table(df_stability)%0A%20%20%20%20return%0A%0A%0A%40app.cell%0Adef%20_()%3A%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
ba23ed1ffdaa4e307d190065c34e94b4