File size: 5,490 Bytes
e56b151
4289dff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e56b151
4289dff
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
e56b151
 
4289dff
 
 
 
 
 
 
e56b151
4289dff
 
e56b151
4289dff
e56b151
 
 
 
 
6ab7c6d
 
 
e56b151
 
 
 
 
 
 
 
 
 
 
 
 
6ab7c6d
4289dff
 
 
 
 
 
 
e56b151
4289dff
 
 
e56b151
 
4289dff
 
e56b151
4289dff
 
e56b151
4289dff
 
e56b151
4289dff
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
"""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()