File size: 6,878 Bytes
d6e1c8a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
#!/usr/bin/env python3
"""Sync 32B benchmark results into doc/nips.tex (Table tab:main_results).

Idempotent. Run after any new *results.json* lands. Updates only the
Qwen3-VL-32B (dense) section's Vanilla / SD-RPN / GroundFlow rows.

Usage:  python3 scripts/sync_to_nips.py
"""
from __future__ import annotations
import json, re, sys
from pathlib import Path
from glob import glob

REPO = Path("/opt/tiger/thothvl_pretrain")
RESULTS = Path("/mnt/bn/leonworkspace/terry/results")
TEX = REPO / "doc" / "nips.tex"

# (tag → row label in 32B section) — note: GroundFlow row has \textbf{} prefix
TAG_TO_ROW = {
    "32b-vanilla-576": "Vanilla",
    "32b-sdrpn-576": "SD-RPN",
    "publish-32b-ours-576": "GroundFlow",
}

# (task → column index 0..9, metric_field, scale)
# Column order in .tex: OCR, Info, Doc, Chart, Text, V*, HR4K, HR8K, RWQA, MME-RW
# Scale: multiplier from raw lmms-eval value to .tex display value.
TASK_TO_COL = {
    "ocrbench":      (0, "ocrbench_accuracy,none", 1000.0),     # raw 0-1 → 0-1000 (e.g. 0.862 → 862)
    "infovqa_val":   (1, "anls,none", 100.0),                   # 0-1 → 0-100
    "docvqa_val":    (2, "anls,none", 100.0),
    "chartqa":       (3, "relaxed_overall,none", 100.0),
    "textvqa_val":   (4, "exact_match,none", 100.0),
    "vstar_bench":   (5, "vstar_overall_acc,none", 1.0),        # already 0-100
    "hrbench4k":     (6, "average,none", 100.0),
    "hrbench8k":     (7, "average,none", 100.0),
    "realworldqa":   (8, "exact_match,flexible-extract", 100.0),
    "mmerealworld":  (9, "mme_realworld_score,none", 100.0),
}

def extract_score(task: str, results: dict) -> float | None:
    """Extract the canonical metric. Falls back to alternates for older eval runs."""
    col, field, scale = TASK_TO_COL[task]
    fallbacks = {
        "ocrbench": ("ocrbench_accuracy,none", "exact_match,none", "score,none"),
        "vstar_bench": ("vstar_overall_acc,none", "vstar_bench_accuracy,none", "accuracy,none"),
        "realworldqa": ("exact_match,flexible-extract", "exact_match,none"),
        "mmerealworld": ("mme_realworld_score,none", "lite_overall,none", "overall,none", "score,none"),
    }
    keys_to_try = fallbacks.get(task, (field,))
    for _, m in results.items():
        if not isinstance(m, dict):
            continue
        for k in keys_to_try:
            if k in m and isinstance(m[k], (int, float)):
                return m[k] * scale
    return None

def latest_results_json(tag: str, task: str) -> Path | None:
    cands = sorted(glob(str(RESULTS / tag / task / "**/*results.json"), recursive=True))
    return Path(cands[-1]) if cands else None

def collect_scores() -> dict[str, dict[str, float]]:
    """Returns {row: {col_idx: value_or_None}} for the 3 32B rows."""
    out: dict[str, dict[int, float]] = {row: {} for row in TAG_TO_ROW.values()}
    for tag, row in TAG_TO_ROW.items():
        for task, (col, _, _) in TASK_TO_COL.items():
            f = latest_results_json(tag, task)
            if not f:
                continue
            try:
                data = json.load(open(f))
            except Exception:
                continue
            results = data.get("results", {})
            score = extract_score(task, results)
            if score is not None:
                out[row][col] = score
    return out

def fmt(col: int, v: float | None) -> str:
    if v is None:
        return "---"
    if col == 0:  # OCR native 0-1000
        return f"{round(v)}"
    return f"{v:.1f}"

def patch_tex(scores: dict[str, dict[int, float]]) -> bool:
    text = TEX.read_text()
    # The 32B section is bounded by the `Qwen3-VL-32B` header and the next \bottomrule
    # Each row is a multi-line LaTeX block. We rewrite the value lines.
    # The pattern: <ROW LABEL>\n& v1 & v2 & v3 & v4 & v5\n& v6 & v7 & v8 & v9 & v10\n& AVG \\
    section_re = re.compile(
        r"(\\multicolumn\{12\}\{@\{\}l\}\{\\emph\{Qwen3-VL-32B \(dense\)\}.*?)"
        r"(\\bottomrule)", re.DOTALL)
    m = section_re.search(text)
    if not m:
        print("ERROR: 32B section not found in nips.tex", file=sys.stderr)
        return False
    section = m.group(1)
    end = m.group(2)
    new_section = section
    for row_label, cells in scores.items():
        # Match the row block: label line + 2 cell lines + avg line, all & --- or numbers
        # Allow \textbf{} wrap on label
        label_pat = re.escape(row_label)
        # Match 10 cells then Avg
        row_re = re.compile(
            r"(\\textbf\{)?" + label_pat + r"(\})?\n"
            r"(?:\\rowcolor\{[^}]+\}\n)?"  # tolerate optional rowcolor (precedes label normally)
            r"& [^\n]+\n"
            r"& [^\n]+\n"
            r"& [^\n]+ \\\\")
        # Build replacement values — reuse existing if cell missing in `scores`
        # Parse current cells from the row to preserve them.
        existing = re.search(
            r"(\\textbf\{)?" + label_pat + r"(\})?\n"
            r"((?:\\rowcolor\{[^}]+\}\n)?)"
            r"& ([^\n]+)\n"
            r"& ([^\n]+)\n"
            r"& ([^\n]+) \\\\",
            new_section)
        if not existing:
            print(f"WARN: row {row_label!r} block not found in 32B section", file=sys.stderr)
            continue
        bf_open, bf_close, _, line1, line2, avg_line = existing.groups()
        cur1 = [c.strip() for c in line1.split("&")]
        cur2 = [c.strip() for c in line2.split("&")]
        cur_cells = cur1 + cur2  # 10 cells
        # Apply updates from scores
        for col, v in cells.items():
            cur_cells[col] = fmt(col, v)
        # Re-emit (preserve any \textbf{} on individual cells? — drop them, keep simple)
        new_line1 = "& " + " & ".join(cur_cells[0:5])
        new_line2 = "& " + " & ".join(cur_cells[5:10])
        new_avg = avg_line.strip()
        label_line = (bf_open or "") + row_label + (bf_close or "")
        new_block = f"{label_line}\n{new_line1}\n{new_line2}\n& {new_avg} \\\\"
        # Replace
        old_block = existing.group(0)
        new_section = new_section.replace(old_block, new_block, 1)
    if new_section == section:
        return False
    new_text = text[:m.start(1)] + new_section + end + text[m.end(2):]
    TEX.write_text(new_text)
    return True

def main():
    scores = collect_scores()
    # Print summary
    print("== Collected 32B scores ==")
    col_names = ["OCR","Info","Doc","Chart","Text","V*","HR4K","HR8K","RWQA","MME-RW"]
    for row in ("Vanilla","SD-RPN","GroundFlow"):
        cells = scores.get(row, {})
        if not cells:
            print(f"  {row}: <none>")
            continue
        s = "  ".join(f"{col_names[c]}={fmt(c,v)}" for c,v in sorted(cells.items()))
        print(f"  {row}: {s}")
    changed = patch_tex(scores)
    print("[PATCH]", "updated nips.tex" if changed else "no changes (already in sync)")

if __name__ == "__main__":
    main()