1import type { WorkspaceFileSystem } from 'minwebide';
2import {
3 effective_sample_size,
4 mean,
5 percentile,
6 split_potential_scale_reduction,
7 std_deviation,
8} from 'mcmc-stats';
10// Writes a completed run into the .sample file's output directory:
11//
12// <output_dir>/chain_1.csv ... one CSV per chain, header = parameter names
13// <output_dir>/summary.csv mean, MCSE, sd, percentiles, ESS, Rhat
14// <output_dir>/sampling_opts.json the exact configuration used
15// <output_dir>/console.txt sampler console output
16//
17// (the per-chain CSV layout matches stan-playground's "download multiple
18// CSVs" export)
20export interface RunOutputs {
21 /** draws[param][draw], chains concatenated along the draw axis (tinystan). */
22 draws: number[][];
23 paramNames: string[];
24 numChains: number;
25 consoleText: string;
26 samplingOpts: Record<string, unknown>;
27 computeTimeSec: number;
28}
30export async function writeRunOutputs(fs: WorkspaceFileSystem, outputDir: string, run: RunOutputs): Promise<string[]> {
31 const written: string[] = [];
32 const write = async (name: string, contents: string) => {
33 const path = `${outputDir}/${name}`;
34 await fs.writeFile(path, contents);
35 written.push(path);
36 };
38 // clear previous results so the directory holds exactly this run
39 await fs.deleteFile(outputDir);
41 const numDraws = run.draws[0]?.length ?? 0;
42 const perChain = Math.floor(numDraws / run.numChains);
44 for (let chain = 0; chain < run.numChains; chain++) {
45 const lines = [run.paramNames.join(',')];
46 for (let draw = chain * perChain; draw < (chain + 1) * perChain; draw++) {
47 lines.push(run.draws.map(paramDraws => String(paramDraws[draw])).join(','));
48 }
49 await write(`chain_${chain + 1}.csv`, lines.join('\n') + '\n');
50 }
52 await write('summary.csv', summaryCsv(run));
53 await write('sampling_opts.json', JSON.stringify(run.samplingOpts, null, 2) + '\n');
54 await write('console.txt', run.consoleText);
56 return written;
57}
59function summaryCsv(run: RunOutputs): string {
60 const numDraws = run.draws[0]?.length ?? 0;
61 const perChain = Math.floor(numDraws / run.numChains);
63 // model parameters first, sampler diagnostics (lp__, divergent__, ...) last
64 const order = [...run.paramNames.keys()].sort((a, b) =>
65 Number(run.paramNames[a].endsWith('__')) - Number(run.paramNames[b].endsWith('__')));
67 const lines = ['parameter,mean,mcse,sd,p5,median,p95,ess,ess_per_sec,rhat'];
68 for (const index of order) {
69 const flat = run.draws[index];
70 const byChain = Array.from({ length: run.numChains }, (_, chain) =>
71 flat.slice(chain * perChain, (chain + 1) * perChain));
72 const sorted = [...flat].sort((a, b) => a - b);
74 const ess = safe(() => effective_sample_size(byChain));
75 const sd = safe(() => std_deviation(sorted));
76 const row = [
77 safe(() => mean(sorted)),
78 sd / Math.sqrt(ess),
79 sd,
80 safe(() => percentile(sorted, 0.05)),
81 safe(() => percentile(sorted, 0.5)),
82 safe(() => percentile(sorted, 0.95)),
83 ess,
84 run.computeTimeSec > 0 ? ess / run.computeTimeSec : NaN,
85 safe(() => split_potential_scale_reduction(byChain)),
86 ];
87 lines.push([run.paramNames[index], ...row.map(formatStat)].join(','));
88 }
89 return lines.join('\n') + '\n';
90}
92function safe(compute: () => number): number {
93 try {
94 return compute();
95 } catch {
96 return NaN;
97 }
98}
100function formatStat(value: number): string {
101 return Number.isFinite(value) ? String(Number(value.toPrecision(6))) : 'NaN';
102}