"""Regenerate model comparison SVGs from the published summaries. Run with Python 3.11+: python figures.py --output ../figures The repository shares its renderer from eval/lib; HF packages include that same renderer next to this script. """ from __future__ import annotations import argparse import importlib.util import json import sys from pathlib import Path import metrics HERE = Path(__file__).resolve().parent renderer = HERE / 'svg_figures.py' if not renderer.is_file(): renderer = HERE.parents[1] / 'lib/svg_figures.py' spec = importlib.util.spec_from_file_location('model_card_svg', renderer) svg = importlib.util.module_from_spec(spec) sys.modules[spec.name] = svg spec.loader.exec_module(svg) UPSTREAM, FIT, REFERENCE = '#0072B2', '#D55E00', '#777777' def classifier(summary: dict) -> str: facets = [*metrics.FACETS, 'macro'] def values(model): result = summary['models'][model] return [result['facets'][f]['average_precision'] for f in metrics.FACETS] + [result['macro']['average_precision']] return svg.render_bars( facets, [('Upstream', values('upstream'), UPSTREAM), ('TF-IDF + logistic regression', values('linear'), REFERENCE), ('Daecore fine-tune', values('finetuned'), FIT)], title='Classifier: matched before and after fine-tuning', subtitle='1,300 generated passages · 3 held-out families · resolved labels only', ylabel='average precision', ylim=(0, 1.05), separator_before=5, width=840, height=360, notes=['Model-generated labels; these results do not establish accuracy on unrelated real documents.'], ) def reranker(summary: dict) -> str: cutoffs = list(metrics.CUTOFFS) definitions = [('hit', 'At least one useful passage'), ('precision', 'Useful passages / returned passages'), ('ndcg', 'Graded ranking quality')] panels = [] legend = [('Upstream', UPSTREAM, None), ('Daecore fine-tune', FIT, None), ('Random pool order', REFERENCE, '5 3')] for field, title in definitions: series = [] for model, (label, color, dash) in zip(('upstream', 'finetuned', 'random'), legend, strict=True): values = summary['models'][model]['answerable'] series.append(svg.Series(label, list(range(len(cutoffs))), [values[str(k)][field] for k in cutoffs], color, dash=dash)) panels.append(svg.Panel(title, series, xlabel='rank cutoff', ylabel={'hit': 'Hit@k', 'precision': 'Precision@k', 'ndcg': 'nDCG@k'}[field], xlim=(0, len(cutoffs) - 1), ylim=(0, 1.02), xticks=list(enumerate(map(str, cutoffs))))) return svg.render_grid(panels, columns=3, title='Ettin: identical candidates, different ordering', subtitle='873 answerable pools · 3–20 returned in Daecore · 50 is the full candidate pool', legend=legend, panel_width=300, panel_height=270) def retriever(summary: dict) -> str: corpus_passages = summary['corpus_passages'] summary = summary['answerable'] cutoffs = sorted(map(int, summary['models']['upstream'])) x = list(range(len(cutoffs))) panels = [] legend = [('Upstream Gemma', UPSTREAM, None), ('Daecore v2', FIT, None)] for field, title in [('hit', 'Find at least one useful passage'), ('precision', 'Useful passages / retained passages'), ('ndcg', 'Graded ranking quality')]: series = [svg.Series(label, x, [summary['models'][name][str(k)][field] for k in cutoffs], color) for name, (label, color, _) in zip(('upstream', 'finetuned'), legend, strict=True)] panels.append(svg.Panel(title, series, xlabel='rank cutoff', ylabel={'hit': 'Hit@k', 'precision': 'Precision@k', 'ndcg': 'nDCG@k'}[field], xlim=(0, len(cutoffs) - 1), ylim=(0, 1.02), xticks=list(enumerate(map(str, cutoffs))))) return svg.render_grid(panels, columns=3, title='Gemma v2: matched dense retrieval on Daecore data', subtitle=f"{summary['queries']:,} known-answerable queries · {corpus_passages:,} passages · abstentions excluded", legend=legend, panel_width=300, panel_height=270) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument('--directory', type=Path, default=HERE) parser.add_argument('--output', type=Path, required=True) parser.add_argument('--model', choices=('classifier', 'reranker', 'retriever')) args = parser.parse_args() args.output.mkdir(parents=True, exist_ok=True) definitions = {'classifier': ('classifier', 'classifier-comparison.svg'), 'reranker': ('reranker', 'ettin-comparison.svg'), 'retriever': ('retriever', 'gemma-comparison.svg')} names = [args.model] if args.model else [ name for name, (record, _) in definitions.items() if (args.directory / f'{record}-summary.json').is_file() ] if not names: parser.error('No model comparison summaries found in the selected directory') for name in names: record, filename = definitions[name] data = json.loads((args.directory / f'{record}-summary.json').read_text(encoding='utf-8')) (args.output / filename).write_text(globals()[name](data), encoding='utf-8', newline='\n') if __name__ == '__main__': main()