tnh0527's picture
Update gliclass-std-base-v3-daecore-5facet-qint8-v2 documentation
6ab7c6d verified
Raw History Blame Contribute Delete
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()