Download evaluation/figures.py from Daecore/gliclass-knowledge-classifier-v2: direct link, hf CLI and curl.
- Browser
- Download file 5.49 kB
-
https://huggingface.co/Daecore/gliclass-knowledge-classifier-v2/resolve/main/evaluation/figures.py
- Command line
-
hf download hf://Daecore/gliclass-knowledge-classifier-v2/evaluation/figures.py
-
curl -L -o figures.py https://huggingface.co/Daecore/gliclass-knowledge-classifier-v2/resolve/main/evaluation/figures.py
5.49 kB
| """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() | |