import%20marimo%0A%0A__generated_with%20%3D%20%220.23.9%22%0Aapp%20%3D%20marimo.App(%0A%20%20%20%20width%3D%22medium%22%2C%0A%20%20%20%20app_title%3D%22nGPT%20scaling%3A%20flat%20across%20the%20width%20%C3%97%20depth%20grid%22%2C%0A%20%20%20%20css_file%3D%22..%2Freport.css%22%2C%0A%20%20%20%20auto_download%3D%5B%22html%22%5D%2C%0A)%0A%0Awith%20app.setup(hide_code%3DTrue)%3A%0A%20%20%20%20import%20json%0A%20%20%20%20import%20tempfile%0A%20%20%20%20from%20pathlib%20import%20Path%0A%0A%20%20%20%20import%20marimo%20as%20mo%20%20%23%20noqa%3A%20F401%0A%20%20%20%20import%20matplotlib.pyplot%20as%20plt%0A%20%20%20%20import%20numpy%20as%20np%0A%0A%20%20%20%20from%20mini.reports%20import%20report_bundle%2C%20use_publisher%0A%20%20%20%20from%20mini.store%20import%20project_store%0A%20%20%20%20from%20mini.vis%20import%20light_dark%2C%20themed%0A%0A%20%20%20%20use_publisher(report_bundle(__file__))%0A%0A%20%20%20%20%23%20Store%20ref%20published%20by%20experiment.py%20(kept%20in%20sync%20by%20hand).%0A%20%20%20%20CURVES_REF%20%3D%20%22reports%2Fngpt-scaling%2Fcurves%22%0A%20%20%20%20WIDTHS%20%3D%20%5B32%2C%2064%2C%20128%5D%0A%20%20%20%20DEPTHS%20%3D%20%5B4%2C%208%2C%2012%5D%0A%0A%20%20%20%20def%20load_curves()%20-%3E%20dict%5Bstr%2C%20np.ndarray%5D%20%7C%20None%3A%0A%20%20%20%20%20%20%20%20%22%22%22Resolve%20the%20val-loss%20curves%20from%20the%20store%2C%20or%20None%20if%20unpublished.%22%22%22%0A%20%20%20%20%20%20%20%20store%20%3D%20project_store()%0A%20%20%20%20%20%20%20%20art%20%3D%20store.get_ref(CURVES_REF)%0A%20%20%20%20%20%20%20%20if%20art%20is%20None%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20return%20None%0A%20%20%20%20%20%20%20%20with%20tempfile.TemporaryDirectory()%20as%20d%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20raw%20%3D%20json.loads(store.get(art%2C%20Path(d)%20%2F%20%22curves.json%22).read_text())%0A%20%20%20%20%20%20%20%20return%20%7Blabel%3A%20np.asarray(losses)%20for%20label%2C%20losses%20in%20raw.items()%7D%0A%0A%20%20%20%20def%20cell(curves%3A%20dict%5Bstr%2C%20np.ndarray%5D%2C%20w%3A%20int%2C%20d%3A%20int)%20-%3E%20np.ndarray%3A%0A%20%20%20%20%20%20%20%20return%20curves%5Bf%22d%7Bw%7D%7CL%7Bd%7D%22%5D%0A%0A%20%20%20%20def%20plateau(curves%3A%20dict%5Bstr%2C%20np.ndarray%5D%2C%20w%3A%20int%2C%20d%3A%20int)%20-%3E%20float%3A%0A%20%20%20%20%20%20%20%20%22%22%22Converged%20loss%3A%20mean%20of%20the%20last%2010%20epochs%20(per-epoch%20eval%20noise%20is%20~%C2%B10.08).%22%22%22%0A%20%20%20%20%20%20%20%20return%20float(cell(curves%2C%20w%2C%20d)%5B-10%3A%5D.mean())%0A%0A%20%20%20%20def%20width_shades()%20-%3E%20dict%5Bint%2C%20tuple%5D%3A%0A%20%20%20%20%20%20%20%20%22%22%22One%20ordered%20shade%20per%20width%2C%20picked%20with%20%60light_dark%60%20so%20the%20dark%20end%0A%20%20%20%20%20%20%20%20of%20the%20ramp%20stays%20legible%20on%20a%20dark%20background.%0A%20%20%20%20%20%20%20%20%22%22%22%0A%20%20%20%20%20%20%20%20stops%20%3D%20light_dark(%5B0.7%2C%200.45%2C%200.12%5D%2C%20%5B0.8%2C%200.55%2C%200.28%5D)%0A%20%20%20%20%20%20%20%20return%20dict(zip(WIDTHS%2C%20plt.cm.viridis(stops)%2C%20strict%3DTrue))%0A%0A%20%20%20%20def%20depth_shades()%20-%3E%20dict%5Bint%2C%20tuple%5D%3A%0A%20%20%20%20%20%20%20%20%22%22%22One%20ordered%20shade%20per%20depth%20(darker%20%3D%20deeper)%2C%20same%20convention.%22%22%22%0A%20%20%20%20%20%20%20%20stops%20%3D%20light_dark(%5B0.7%2C%200.45%2C%200.12%5D%2C%20%5B0.8%2C%200.55%2C%200.28%5D)%0A%20%20%20%20%20%20%20%20return%20dict(zip(DEPTHS%5B%3A%3A-1%5D%2C%20plt.cm.viridis(stops)%2C%20strict%3DTrue))%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_()%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%20nGPT%20scaling%3A%20flat%20across%20the%20width%20%C3%97%20depth%20grid%0A%0A%20%20%20%20Before%20we%20build%20the%20color-mixing%20experiments%20on%20top%20of%20this%20transformer%2C%20we%0A%20%20%20%20want%20to%20know%20that%20it%20holds%20its%20shape%20as%20it%20grows.%20This%20report%20trains%20the%20model%0A%20%20%20%20at%20a%20range%20of%20sizes%20and%20checks%20that%20none%20of%20them%20misbehave.%0A%0A%20%20%20%20The%20model%20is%20a%20simplified%20version%20of%20*nGPT*.%20The%20idea%20behind%20nGPT%20is%20to%20keep%0A%20%20%20%20the%20model's%20running%20state%20(the%20*residual%20stream*%2C%20the%20vector%20that%20each%20layer%0A%20%20%20%20reads%20from%20and%20writes%20back%20to)%20on%20the%20surface%20of%20a%20hypersphere%2C%20by%20normalizing%0A%20%20%20%20it%20after%20every%20step.%20We%20keep%20that%20residual%20update%2C%0A%20%20%20%20%60h%20%E2%86%90%20Norm(h%20%2B%20%CE%B1%C2%B7(Norm(sub(h))%20%E2%88%92%20h))%60%2C%20which%20moves%20the%20state%20%60h%60%20a%20fraction%20%60%CE%B1%60%0A%20%20%20%20of%20the%20way%20toward%20a%20sub-module's%20normalized%20output%20and%20then%20renormalizes.%20We%0A%20%20%20%20simplify%20two%20pieces%20of%20it.%20Each%20sub-module%20gets%20a%20single%20learned%20*gain*%20(one%0A%20%20%20%20number%20that%20scales%20its%20output)%20in%20place%20of%20nGPT's%20per-channel%20*eigen%20learning%0A%20%20%20%20rates*%2C%20and%20the%20residual%20step%20%60%CE%B1%60%20is%20fixed%20at%201%2Fn_layer%20instead%20of%20being%20learned%2C%0A%20%20%20%20which%20is%20about%20where%20the%20learned%20version%20settled%20anyway.%0A%0A%20%20%20%20For%20the%20claim%20that%20SCA%20(the%20concept-anchoring%20method%20this%20project%20studies)%0A%20%20%20%20carries%20over%20to%20language%20models%2C%20this%20pared-down%20architecture%20needs%20to%20stay%0A%20%20%20%20well-behaved%20as%20it%20scales.%20So%20the%20%5Bexperiment%5D(.%2Fexperiment.py)%20trains%20it%0A%20%20%20%20across%20a%20grid%3A%20three%20widths%20(how%20many%20numbers%20are%20in%20that%20state%20vector)%0A%20%20%20%20crossed%20with%20three%20depths%20(how%20many%20layers)%2C%20%7B32%2C%2064%2C%20128%7D%20%C3%97%20%7B4%2C%208%2C%2012%7D%2C%20with%0A%20%20%20%20everything%20else%20held%20fixed%3A%20batch%20size%2016%2C%20peak%20learning%20rate%2010%E2%81%BB%C2%B2%2C%20100%20epochs%2C%0A%20%20%20%20and%20*Pride%20and%20Prejudice*%20for%20the%20training%20text.%20Two%20outcomes%20would%20concern%20us.%0A%20%20%20%20One%20is%20a%20**depth%20penalty**%2C%20where%20adding%20layers%20at%20a%20fixed%20width%20makes%20the%0A%20%20%20%20model%20worse.%20The%20other%20is%20an%20instability%20that%20shows%20up%20only%20at%20large%20width%2C%0A%20%20%20%20where%20a%20run%20spikes%20or%20fails%20to%20train.%20We%20hope%20to%20see%20neither.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_()%3A%0A%20%20%20%20loaded%20%3D%20load_curves()%0A%20%20%20%20return%20(loaded%2C)%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(loaded)%3A%0A%20%20%20%20mo.stop(%0A%20%20%20%20%20%20%20%20loaded%20is%20None%2C%0A%20%20%20%20%20%20%20%20mo.md(%0A%20%20%20%20%20%20%20%20%20%20%20%20%22No%20results%20yet.%20Run%20the%20experiment%20first%3B%20it%20publishes%20the%20loss%20curves%20to%20the%20store%20when%20it%20finishes%3A%5Cn%5Cn%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22%60%60%60bash%5Cnbin%2Fmini%20run%20docs%2Fngpt-scaling%2Fexperiment.py%20--app%20modal%20--max-containers%209%5Cn%60%60%60%22%0A%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20)%0A%20%20%20%20curves%20%3D%20loaded%0A%20%20%20%20%23%20Largest%20within-a-width%20spread%20across%20depth%20%E2%80%94%20the%20%22depth%20penalty%22%20if%20there%20were%20one.%0A%20%20%20%20depth_spread%20%3D%20max(%0A%20%20%20%20%20%20%20%20max(plateau(curves%2C%20w%2C%20d)%20for%20d%20in%20DEPTHS)%20-%20min(plateau(curves%2C%20w%2C%20d)%20for%20d%20in%20DEPTHS)%20for%20w%20in%20WIDTHS%0A%20%20%20%20)%0A%20%20%20%20mo.md(%0A%20%20%20%20%20%20%20%20f%22%22%22%0A%20%20%20%20**The%20architecture%20scales%20cleanly.**%20We%20score%20each%20run%20by%20its%20converged%20loss%3A%20the%0A%20%20%20%20model's%20average%20error%20at%20predicting%20the%20next%20character%20once%20training%20has%0A%20%20%20%20settled%2C%20measured%20in%20*nats%20per%20character*%20(natural-log%20units%2C%20where%20lower%20is%0A%20%20%20%20better).%20That%20loss%20never%20rises%20as%20we%20add%20layers.%20At%20each%20width%2C%20the%20three%0A%20%20%20%20depths%20land%20within%20%7Bdepth_spread%3A.02f%7D%20nats%2Fchar%20of%20one%20another%2C%20well%20inside%0A%20%20%20%20the%20%C2%B10.08%20of%20noise%20we%20see%20between%20epochs%2C%20so%20the%20depth%20axis%20is%20flat.%20Width%0A%20%20%20%20behaves%20the%20way%20added%20capacity%20should%3A%20loss%20falls%20monotonically%20from%0A%20%20%20%20%7Bplateau(curves%2C%2032%2C%204)%3A.2f%7D%20at%20width%2032%20to%20%7Bplateau(curves%2C%20128%2C%2012)%3A.2f%7D%20at%0A%20%20%20%20width%20128%20with%2012%20layers.%20No%20cell%20spikes%20or%20stalls%2C%20and%20the%20same%20learning%20rate%0A%20%20%20%20(10%E2%81%BB%C2%B2)%20works%20everywhere.%20So%20fixing%20the%20residual%20step%20to%20a%20scalar%20constant%20seems%0A%20%20%20%20to%20be%20enough.%0A%20%20%20%20%22%22%22%0A%20%20%20%20)%0A%20%20%20%20return%20(curves%2C)%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_()%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Converged%20loss%20versus%20depth%0A%0A%20%20%20%20This%20chart%20shows%20converged%20validation%20loss%20(measured%20on%20held-out%20text%20and%0A%20%20%20%20averaged%20over%20the%20last%2010%20epochs)%20against%20depth%2C%20with%20one%20line%20per%20width.%20Read%0A%20%20%20%20each%20line%20from%20left%20to%20right%3A%20if%20adding%20layers%20hurt%2C%20the%20line%20would%20slope%20up.%0A%20%20%20%20Instead%20each%20one%20is%20nearly%20horizontal%2C%20so%20extra%20depth%20costs%20nothing%20at%20any%0A%20%20%20%20width.%20The%20lines%20also%20stack%20in%20width%20order%2C%20so%20a%20wider%20model%20is%20uniformly%0A%20%20%20%20better%2C%20and%20there%20is%20no%20far%20corner%2C%20wide%20and%20deep%20together%2C%20where%20the%20loss%0A%20%20%20%20turns%20back%20up.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(curves)%3A%0A%20%20%20%20%40themed(%0A%20%20%20%20%20%20%20%20name%3D%22plateau-vs-depth%22%2C%0A%20%20%20%20%20%20%20%20alt_text%3D(%0A%20%20%20%20%20%20%20%20%20%20%20%20%22Line%20chart%20of%20converged%20validation%20loss%20against%20depth%20(4%2C%208%2C%20and%2012%20layers)%2C%20with%20one%20line%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22per%20width.%20All%20three%20lines%20are%20close%20to%20flat%3A%20width%2032%20sits%20near%201.5%2C%20width%2064%20near%201.4%2C%20and%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22width%20128%20near%201.33%2C%20at%20every%20depth.%20Deeper%20models%20are%20no%20worse%20than%20shallow%20ones%20at%20any%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22width%2C%20and%20each%20wider%20model%20sits%20uniformly%20below%20the%20narrower%20ones.%22%0A%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20)%0A%20%20%20%20def%20_plot()%20-%3E%20plt.Figure%3A%0A%20%20%20%20%20%20%20%20fig%2C%20ax%20%3D%20plt.subplots(figsize%3D(6.2%2C%203.8))%0A%20%20%20%20%20%20%20%20shades%20%3D%20width_shades()%0A%20%20%20%20%20%20%20%20for%20w%20in%20WIDTHS%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20ys%20%3D%20%5Bplateau(curves%2C%20w%2C%20d)%20for%20d%20in%20DEPTHS%5D%0A%20%20%20%20%20%20%20%20%20%20%20%20ax.plot(DEPTHS%2C%20ys%2C%20%22o-%22%2C%20color%3Dshades%5Bw%5D%2C%20label%3Df%22width%20%7Bw%7D%22%2C%20lw%3D2.2)%0A%20%20%20%20%20%20%20%20ax.set(xlabel%3D%22depth%20(n_layer)%22%2C%20ylabel%3D%22converged%20validation%20loss%20(nats%2Fchar)%22)%0A%20%20%20%20%20%20%20%20ax.set_xticks(DEPTHS)%0A%20%20%20%20%20%20%20%20ax.grid(alpha%3D0.3)%0A%20%20%20%20%20%20%20%20ax.legend(fontsize%3D8)%0A%20%20%20%20%20%20%20%20return%20fig%0A%0A%20%20%20%20mo.Html(_plot())%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_()%3A%0A%20%20%20%20mo.md(r%22%22%22%0A%20%20%20%20%23%23%20Training%20curves%0A%0A%20%20%20%20The%20same%20runs%2C%20now%20shown%20across%20training%3A%20one%20panel%20per%20width%2C%20one%20line%20per%0A%20%20%20%20depth%2C%20loss%20on%20the%20vertical%20axis%20and%20epoch%20on%20the%20horizontal.%20Every%20run%0A%20%20%20%20descends%20through%20the%20learning-rate%20warmup%20(the%20opening%20stretch%20of%20training%2C%0A%20%20%20%20where%20the%20step%20size%20ramps%20up%20from%20small)%20and%20settles%20onto%20a%20plateau.%20Within%20a%0A%20%20%20%20panel%2C%20the%20depth%20lines%20sit%20on%20top%20of%20one%20another%20rather%20than%20fanning%20apart%2C%20and%0A%20%20%20%20the%20wider%20panels%20settle%20lower.%0A%20%20%20%20%22%22%22)%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(curves)%3A%0A%20%20%20%20%40themed(%0A%20%20%20%20%20%20%20%20name%3D%22convergence%22%2C%0A%20%20%20%20%20%20%20%20alt_text%3D(%0A%20%20%20%20%20%20%20%20%20%20%20%20%22Three%20line%20charts%20of%20validation%20loss%20against%20epoch%2C%20one%20per%20width%20(32%2C%2064%2C%20and%20128)%2C%20sharing%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22a%20y-axis%2C%20each%20with%20one%20line%20per%20depth%20(4%2C%208%2C%20and%2012%20layers).%20In%20every%20panel%2C%20all%20three%20depth%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22lines%20fall%20together%20through%20the%20first%20ten%20epochs%20of%20warmup%20and%20converge%20onto%20a%20single%20plateau%2C%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%22with%20no%20spikes%20or%20divergence.%20The%20plateau%20drops%20from%20about%201.5%20in%20the%20width-32%20panel%20to%20about%20%22%0A%20%20%20%20%20%20%20%20%20%20%20%20%221.4%20at%20width%2064%20to%20about%201.33%20at%20width%20128.%22%0A%20%20%20%20%20%20%20%20)%2C%0A%20%20%20%20)%0A%20%20%20%20def%20_plot()%20-%3E%20plt.Figure%3A%0A%20%20%20%20%20%20%20%20fig%2C%20axes%20%3D%20plt.subplots(1%2C%203%2C%20figsize%3D(10.5%2C%203.4)%2C%20sharey%3DTrue)%0A%20%20%20%20%20%20%20%20shades%20%3D%20depth_shades()%0A%20%20%20%20%20%20%20%20for%20ax%2C%20w%20in%20zip(axes%2C%20WIDTHS%2C%20strict%3DTrue)%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20ax.axvline(10%2C%20color%3D%22%238888%22%2C%20lw%3D1%2C%20ls%3D%22%3A%22%2C%20label%3D%22end%20of%20LR%20warmup%22)%0A%20%20%20%20%20%20%20%20%20%20%20%20for%20d%20in%20DEPTHS%3A%0A%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20%20ax.plot(cell(curves%2C%20w%2C%20d)%2C%20color%3Dshades%5Bd%5D%2C%20label%3Df%22%7Bd%7D%20layers%22)%0A%20%20%20%20%20%20%20%20%20%20%20%20ax.set(title%3Df%22width%20%7Bw%7D%22%2C%20xlabel%3D%22epoch%22)%0A%20%20%20%20%20%20%20%20%20%20%20%20ax.grid(alpha%3D0.3)%0A%20%20%20%20%20%20%20%20axes%5B0%5D.set_ylabel(%22validation%20loss%20(nats%2Fchar)%22)%0A%20%20%20%20%20%20%20%20axes%5B0%5D.legend(fontsize%3D8)%0A%20%20%20%20%20%20%20%20return%20fig%0A%0A%20%20%20%20mo.Html(_plot())%0A%20%20%20%20return%0A%0A%0A%40app.cell(hide_code%3DTrue)%0Adef%20_(curves)%3A%0A%20%20%20%20_best%20%3D%20plateau(curves%2C%20128%2C%2012)%0A%20%20%20%20mo.md(%0A%20%20%20%20%20%20%20%20f%22%22%22%0A%20%20%20%20%23%23%20Findings%0A%0A%20%20%20%20Across%20the%20whole%20grid%2C%20the%20simplified%20nGPT%20trains%20flat%20across%20depth%20and%20keeps%0A%20%20%20%20improving%20with%20width%2C%20and%20no%20cell%20destabilizes%20(%7B_best%3A.2f%7D%20nats%2Fchar%20at%20the%0A%20%20%20%20deepest%2C%20widest%20corner).%20That%20is%20what%20we%20were%20hoping%20for.%20The%20architecture%20that%20SCA%0A%20%20%20%20will%20anchor%20concepts%20in%20scales%20without%20a%20depth%20penalty%2C%20so%20if%20a%20later%0A%20%20%20%20experiment%20runs%20into%20trouble%2C%20the%20simplified%20architecture%20is%20unlikely%20to%20be%20the%0A%20%20%20%20reason.%0A%0A%20%20%20%20The%20grid%20tops%20out%20at%20width%20128%20with%2012%20layers%20on%20an%20L4%2C%20which%20is%20probably%0A%20%20%20%20fine%20for%20M2%20(this%20milestone).%20For%20M3%20(a%20future%0A%20%20%20%20milestone)%2C%20we%20should%20confirm%20that%20the%20fixed%20scalar%20residual%20step%20still%20holds%0A%20%20%20%20at%20a%20genuinely%20larger%20size%3A%20wider%20and%20deeper%2C%20on%20a%20bigger%20GPU%20with%20a%20bigger%0A%20%20%20%20batch.%0A%20%20%20%20%22%22%22%0A%20%20%20%20)%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
029ce2b74ee3e8ef4552ee908a8a3ef6