simit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- simit/__init__.py +39 -0
- simit/api.py +258 -0
- simit/backends/__init__.py +112 -0
- simit/backends/bagel.py +359 -0
- simit/backends/base.py +114 -0
- simit/backends/hf.py +336 -0
- simit/backends/imagegen.py +136 -0
- simit/backends/lance.py +338 -0
- simit/backends/mot/engine.py +844 -0
- simit/backends/mot/kernels.py +130 -0
- simit/backends/mot/modeling.py +514 -0
- simit/backends/mot/wan_vae.py +872 -0
- simit/backends/vllm.py +204 -0
- simit/config.py +112 -0
- simit/metrics.py +113 -0
- simit/pipeline/core.py +586 -0
- simit/pipeline/prompts.py +167 -0
- simit/skills/__init__.py +6 -0
- simit/skills/assets/MERMAID_LICENSE +21 -0
- simit/skills/assets/mermaid.min.js +3587 -0
- simit/skills/base.py +172 -0
- simit/skills/builtin.py +614 -0
- simit/skills/helpers.py +197 -0
- simit/skills/pool.py +262 -0
- simit/skills/prompts.py +2374 -0
- simit/skills/renderers.py +1292 -0
- simit/skills/routing.py +89 -0
- simit/skills/sandbox.py +188 -0
- simit/skills/web.py +253 -0
- simit/tune.py +323 -0
- simit/types.py +96 -0
- simit/utils.py +67 -0
- simit-0.1.0.dist-info/METADATA +388 -0
- simit-0.1.0.dist-info/RECORD +38 -0
- simit-0.1.0.dist-info/WHEEL +5 -0
- simit-0.1.0.dist-info/licenses/LICENSE +202 -0
- simit-0.1.0.dist-info/licenses/src/simit/skills/assets/MERMAID_LICENSE +21 -0
- simit-0.1.0.dist-info/top_level.txt +1 -0
simit/backends/vllm.py
ADDED
|
@@ -0,0 +1,204 @@
|
|
|
1
|
+
"""Fast engine for standard VLMs: vLLM (continuous batching, prefix caching,
|
|
2
|
+
CUDA graphs). Optional: ``pip install "simit[vllm]"``.
|
|
3
|
+
|
|
4
|
+
Prompts are built with the model's own chat template (as in the transformers
|
|
5
|
+
backend), so results match it up to numerics.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import itertools
|
|
11
|
+
import queue
|
|
12
|
+
import threading
|
|
13
|
+
import traceback
|
|
14
|
+
from concurrent.futures import Future
|
|
15
|
+
from typing import Optional
|
|
16
|
+
|
|
17
|
+
from PIL import Image
|
|
18
|
+
|
|
19
|
+
from ..utils import strip_thinking, to_pil
|
|
20
|
+
from .base import Backend, GenRequest, GenResult, ScoreRequest
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class VLLMBackend(Backend):
|
|
24
|
+
name = "vllm"
|
|
25
|
+
native_image_generation = False
|
|
26
|
+
verify_think = False
|
|
27
|
+
|
|
28
|
+
def __init__(self, model: str, *, gpu_memory_utilization: Optional[float] = None, max_model_len: int = 32768,
|
|
29
|
+
max_images_per_prompt: int = 8, max_image_pixels: Optional[int] = 1024 * 1024,
|
|
30
|
+
tensor_parallel_size: int = 1, think: bool = False, device: Optional[str] = None,
|
|
31
|
+
max_num_seqs: int = 64, **llm_kwargs):
|
|
32
|
+
"""``gpu_memory_utilization=None`` takes what is free on the GPU (minus
|
|
33
|
+
headroom), so an image generator loaded first keeps its share."""
|
|
34
|
+
import os
|
|
35
|
+
if device is not None and str(device) not in ("cuda", "cuda:0"):
|
|
36
|
+
raise ValueError("the vLLM engine runs on the first visible GPU; select GPUs with "
|
|
37
|
+
"CUDA_VISIBLE_DEVICES (and tensor_parallel_size for several)")
|
|
38
|
+
if gpu_memory_utilization is None:
|
|
39
|
+
import torch
|
|
40
|
+
free, total = torch.cuda.mem_get_info(0)
|
|
41
|
+
gpu_memory_utilization = max(0.3, min(0.9, (free - 6 * 2**30) / total))
|
|
42
|
+
# Keep vLLM's engine core in this process: a spawned core would re-import
|
|
43
|
+
# the user's __main__ (scripts without an `if __name__ == "__main__"` guard).
|
|
44
|
+
os.environ.setdefault("VLLM_ENABLE_V1_MULTIPROCESSING", "0")
|
|
45
|
+
# Portable defaults: no kernels that JIT-compile with the system nvcc
|
|
46
|
+
# (FlashInfer) or need a matching flash-attn build. Override via kwargs/env.
|
|
47
|
+
os.environ.setdefault("VLLM_USE_FLASHINFER_SAMPLER", "0")
|
|
48
|
+
os.environ.setdefault("VLLM_USE_DEEP_GEMM", "0") # FP8 checkpoints: DeepGEMM JIT needs nvcc >= 12.9
|
|
49
|
+
llm_kwargs.setdefault("attention_backend", "TRITON_ATTN")
|
|
50
|
+
llm_kwargs.setdefault("mm_encoder_attn_backend", "TORCH_SDPA")
|
|
51
|
+
llm_kwargs.setdefault("gdn_prefill_backend", "triton")
|
|
52
|
+
from transformers import AutoProcessor
|
|
53
|
+
from vllm import LLM
|
|
54
|
+
|
|
55
|
+
self.model_id = model
|
|
56
|
+
self.processor = AutoProcessor.from_pretrained(model)
|
|
57
|
+
self.tokenizer = getattr(self.processor, "tokenizer", self.processor)
|
|
58
|
+
self.llm = LLM(model=model, dtype="bfloat16", gpu_memory_utilization=gpu_memory_utilization,
|
|
59
|
+
max_model_len=max_model_len, limit_mm_per_prompt={"image": max_images_per_prompt},
|
|
60
|
+
enable_prefix_caching=True, tensor_parallel_size=tensor_parallel_size,
|
|
61
|
+
max_num_seqs=max_num_seqs, **llm_kwargs)
|
|
62
|
+
self.engine = self.llm.llm_engine
|
|
63
|
+
gc = self.llm.llm_engine.model_config.try_get_generation_config() or {}
|
|
64
|
+
eos = gc.get("eos_token_id")
|
|
65
|
+
eos = eos if isinstance(eos, list) else [eos] if eos is not None else []
|
|
66
|
+
self.eos_ids = {e for e in eos + [self.tokenizer.eos_token_id] if e is not None}
|
|
67
|
+
self.think = think
|
|
68
|
+
self.max_image_pixels = max_image_pixels
|
|
69
|
+
self._ids = itertools.count()
|
|
70
|
+
self._queue: "queue.Queue" = queue.Queue()
|
|
71
|
+
self._live: dict = {}
|
|
72
|
+
self._thread = threading.Thread(target=self._loop, daemon=True, name="simit-vllm")
|
|
73
|
+
self._thread.start()
|
|
74
|
+
|
|
75
|
+
# ------------------------------------------------------------- prompts
|
|
76
|
+
def _image(self, img):
|
|
77
|
+
img = to_pil(img)
|
|
78
|
+
if self.max_image_pixels and img.width * img.height > self.max_image_pixels:
|
|
79
|
+
s = (self.max_image_pixels / (img.width * img.height)) ** 0.5
|
|
80
|
+
img = img.resize((max(28, int(img.width * s)), max(28, int(img.height * s))), Image.BICUBIC)
|
|
81
|
+
return img
|
|
82
|
+
|
|
83
|
+
def _prompt(self, parts, think: bool):
|
|
84
|
+
content, images = [], []
|
|
85
|
+
for p in parts:
|
|
86
|
+
if isinstance(p, Image.Image):
|
|
87
|
+
content.append({"type": "image"})
|
|
88
|
+
images.append(self._image(p))
|
|
89
|
+
elif content and content[-1]["type"] == "text":
|
|
90
|
+
content[-1]["text"] += p
|
|
91
|
+
else:
|
|
92
|
+
content.append({"type": "text", "text": p})
|
|
93
|
+
messages = [{"role": "user", "content": content}]
|
|
94
|
+
try:
|
|
95
|
+
text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True,
|
|
96
|
+
enable_thinking=think)
|
|
97
|
+
except TypeError:
|
|
98
|
+
text = self.processor.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
|
99
|
+
return text, images
|
|
100
|
+
|
|
101
|
+
@staticmethod
|
|
102
|
+
def _inputs(text, images):
|
|
103
|
+
inp = {"prompt": text}
|
|
104
|
+
if images:
|
|
105
|
+
inp["multi_modal_data"] = {"image": images}
|
|
106
|
+
return inp
|
|
107
|
+
|
|
108
|
+
# ------------------------------------------------------------ requests
|
|
109
|
+
def submit_generate(self, req: GenRequest) -> Future:
|
|
110
|
+
from vllm import SamplingParams
|
|
111
|
+
think = req.think or self.think
|
|
112
|
+
text, images = self._prompt(req.parts, think)
|
|
113
|
+
sp = SamplingParams(max_tokens=req.max_new_tokens, temperature=req.temperature,
|
|
114
|
+
top_p=req.top_p if req.temperature > 0 else 1.0,
|
|
115
|
+
repetition_penalty=req.repetition_penalty or 1.0,
|
|
116
|
+
logprobs=0 if req.logprobs else None)
|
|
117
|
+
|
|
118
|
+
def post(out):
|
|
119
|
+
o = out.outputs[0]
|
|
120
|
+
raw = o.text
|
|
121
|
+
ids = list(o.token_ids)
|
|
122
|
+
if ids and ids[-1] in self.eos_ids: # like the other backends: the stop token is not part of the answer
|
|
123
|
+
ids = ids[:-1]
|
|
124
|
+
lps = None
|
|
125
|
+
if req.logprobs and o.logprobs is not None:
|
|
126
|
+
lps = [step[tid].logprob for tid, step in zip(ids, o.logprobs)]
|
|
127
|
+
text_ = strip_thinking(raw) if (think or "</think>" in raw) else raw.strip()
|
|
128
|
+
return GenResult(text=text_, raw_text=raw, token_logprobs=lps, num_tokens=len(ids))
|
|
129
|
+
|
|
130
|
+
return self._submit(self._inputs(text, images), sp, post)
|
|
131
|
+
|
|
132
|
+
def submit_score(self, req: ScoreRequest) -> Future:
|
|
133
|
+
from vllm import SamplingParams
|
|
134
|
+
text, images = self._prompt(req.parts, think=False)
|
|
135
|
+
n = len(self.tokenizer(req.target, add_special_tokens=False)["input_ids"])
|
|
136
|
+
if n == 0:
|
|
137
|
+
f: Future = Future()
|
|
138
|
+
f.set_result([])
|
|
139
|
+
return f
|
|
140
|
+
sp = SamplingParams(max_tokens=1, temperature=0.0, prompt_logprobs=0)
|
|
141
|
+
|
|
142
|
+
def post(out):
|
|
143
|
+
ids = out.prompt_token_ids[-n:]
|
|
144
|
+
plp = out.prompt_logprobs[-n:]
|
|
145
|
+
return [step[tid].logprob for tid, step in zip(ids, plp)]
|
|
146
|
+
|
|
147
|
+
return self._submit(self._inputs(text + req.target, images), sp, post)
|
|
148
|
+
|
|
149
|
+
def _submit(self, inputs, params, post) -> Future:
|
|
150
|
+
fut: Future = Future()
|
|
151
|
+
self._queue.put((inputs, params, post, fut))
|
|
152
|
+
return fut
|
|
153
|
+
|
|
154
|
+
def close(self):
|
|
155
|
+
self._queue.put(None)
|
|
156
|
+
|
|
157
|
+
# --------------------------------------------------------------- engine
|
|
158
|
+
def _loop(self):
|
|
159
|
+
while True:
|
|
160
|
+
block = not self._live
|
|
161
|
+
try:
|
|
162
|
+
item = self._queue.get(timeout=0.5) if block else self._queue.get_nowait()
|
|
163
|
+
while True:
|
|
164
|
+
if item is None:
|
|
165
|
+
return
|
|
166
|
+
inputs, params, post, fut = item
|
|
167
|
+
if not fut.cancelled():
|
|
168
|
+
rid = str(next(self._ids))
|
|
169
|
+
try:
|
|
170
|
+
self.engine.add_request(rid, inputs, params)
|
|
171
|
+
self._live[rid] = (post, fut)
|
|
172
|
+
except Exception as e:
|
|
173
|
+
fut.set_exception(e)
|
|
174
|
+
item = self._queue.get_nowait()
|
|
175
|
+
except queue.Empty:
|
|
176
|
+
pass
|
|
177
|
+
if not self._live:
|
|
178
|
+
continue
|
|
179
|
+
for rid, (post, fut) in list(self._live.items()):
|
|
180
|
+
if fut.cancelled():
|
|
181
|
+
try:
|
|
182
|
+
self.engine.abort_request([rid])
|
|
183
|
+
except Exception:
|
|
184
|
+
pass
|
|
185
|
+
self._live.pop(rid, None)
|
|
186
|
+
try:
|
|
187
|
+
outputs = self.engine.step()
|
|
188
|
+
except Exception as e:
|
|
189
|
+
tb = traceback.format_exc()
|
|
190
|
+
for post, fut in self._live.values():
|
|
191
|
+
if not fut.done():
|
|
192
|
+
fut.set_exception(RuntimeError(f"{e}\n{tb}"))
|
|
193
|
+
self._live.clear()
|
|
194
|
+
continue
|
|
195
|
+
for out in outputs:
|
|
196
|
+
if not out.finished or out.request_id not in self._live:
|
|
197
|
+
continue
|
|
198
|
+
post, fut = self._live.pop(out.request_id)
|
|
199
|
+
if fut.done():
|
|
200
|
+
continue
|
|
201
|
+
try:
|
|
202
|
+
fut.set_result(post(out))
|
|
203
|
+
except Exception as e:
|
|
204
|
+
fut.set_exception(e)
|
simit/config.py
ADDED
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
"""SIMIT hyperparameters. Defaults follow the paper; ``SIMIT.tune`` fits the
|
|
2
|
+
adaptive-budget (ABA) and difficulty-filter (DF) parameters on labeled
|
|
3
|
+
validation data."""
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
import json
|
|
8
|
+
from dataclasses import asdict, dataclass, fields
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
from typing import Optional
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@dataclass
|
|
14
|
+
class SIMITConfig:
|
|
15
|
+
# --- budget -----------------------------------------------------------
|
|
16
|
+
#: Maximum number of imagined demonstrations per query (K_max).
|
|
17
|
+
k_max: int = 4
|
|
18
|
+
#: Adaptive budget allocation: fewer demos for confidently answered queries.
|
|
19
|
+
use_aba: bool = True
|
|
20
|
+
#: ABA parameters (target error rate and log-accuracy curve A(K) = A1 + B ln K).
|
|
21
|
+
epsilon: float = 0.14
|
|
22
|
+
A0: float = 0.35
|
|
23
|
+
A1: float = 0.60
|
|
24
|
+
B: float = 0.2
|
|
25
|
+
#: Difficulty filtering: keep demos whose answer confidence lies in [t_low, t_high].
|
|
26
|
+
use_df: bool = True
|
|
27
|
+
t_low: float = 0.2
|
|
28
|
+
t_high: float = 0.9
|
|
29
|
+
|
|
30
|
+
# --- synthesis ----------------------------------------------------------
|
|
31
|
+
# ``None`` means "the backend's default" (e.g. Lance: decomposed synthesis,
|
|
32
|
+
# natural images only, no critic -- see ``Backend.pipeline_defaults``).
|
|
33
|
+
#: Triplet synthesis: "batch" (one call for K triplets, the paper's prompt)
|
|
34
|
+
#: or "decomposed" (question, answer, description in three short calls).
|
|
35
|
+
synthesis: Optional[str] = None
|
|
36
|
+
#: Diversity-encouraging triplet prompt (the paper's "DP").
|
|
37
|
+
diversity_prompt: bool = True
|
|
38
|
+
#: Use the structured skill library (False: realize everything as a natural image).
|
|
39
|
+
use_skills: Optional[bool] = None
|
|
40
|
+
#: Critic verification of every realized image.
|
|
41
|
+
verify: Optional[bool] = None
|
|
42
|
+
#: Minimum critic score (0-100) for an image to be accepted.
|
|
43
|
+
verify_threshold: int = 50
|
|
44
|
+
#: Run the critic in thinking mode (None: backend default; BAGEL: True).
|
|
45
|
+
verify_think: Optional[bool] = None
|
|
46
|
+
verify_max_new_tokens: int = 300
|
|
47
|
+
#: Realization attempts per demonstration slot (breadth-first slot filling).
|
|
48
|
+
attempts_per_slot: int = 4
|
|
49
|
+
#: Repair/regeneration rounds a failed candidate gets in the retry phase.
|
|
50
|
+
repair_retries: int = 2
|
|
51
|
+
verify_rounds: int = 5
|
|
52
|
+
#: Side length of natively generated images.
|
|
53
|
+
image_size: int = 400
|
|
54
|
+
#: Start every remaining candidate attempt at once (up to K_max) instead of only as many as the
|
|
55
|
+
#: budget still needs. Lower latency for a lone query on an otherwise idle GPU, wasted work under
|
|
56
|
+
#: load. None: on for ``imagine``, off for ``imagine_batch`` and ``tune``.
|
|
57
|
+
speculative: Optional[bool] = None
|
|
58
|
+
|
|
59
|
+
# --- answering ------------------------------------------------------------
|
|
60
|
+
answer_max_new_tokens: int = 128
|
|
61
|
+
#: Random seed for image generation (None: random).
|
|
62
|
+
seed: Optional[int] = None
|
|
63
|
+
|
|
64
|
+
def thresholds(self) -> dict[int, float]:
|
|
65
|
+
"""ABA confidence thresholds tau_K (Appendix D, Algorithm 1)."""
|
|
66
|
+
return compute_thresholds(self.epsilon, self.A0, self.A1, self.B, self.k_max)
|
|
67
|
+
|
|
68
|
+
def budget(self, p0: Optional[float]) -> int:
|
|
69
|
+
"""Number of demonstrations K*(q) for zero-shot confidence ``p0``."""
|
|
70
|
+
if not self.use_aba or p0 is None:
|
|
71
|
+
return self.k_max
|
|
72
|
+
tau = self.thresholds()
|
|
73
|
+
for k in range(self.k_max + 1):
|
|
74
|
+
if p0 >= tau[k]:
|
|
75
|
+
return k
|
|
76
|
+
return self.k_max
|
|
77
|
+
|
|
78
|
+
def in_band(self, c: Optional[float]) -> bool:
|
|
79
|
+
return (not self.use_df) or (c is not None and self.t_low <= c <= self.t_high)
|
|
80
|
+
|
|
81
|
+
def to_dict(self) -> dict:
|
|
82
|
+
return asdict(self)
|
|
83
|
+
|
|
84
|
+
def save(self, path) -> None:
|
|
85
|
+
Path(path).write_text(json.dumps(self.to_dict(), indent=2))
|
|
86
|
+
|
|
87
|
+
@classmethod
|
|
88
|
+
def load(cls, path) -> "SIMITConfig":
|
|
89
|
+
data = json.loads(Path(path).read_text())
|
|
90
|
+
names = {f.name for f in fields(cls)}
|
|
91
|
+
return cls(**{k: v for k, v in data.items() if k in names})
|
|
92
|
+
|
|
93
|
+
def replace(self, **kw) -> "SIMITConfig":
|
|
94
|
+
data = self.to_dict()
|
|
95
|
+
data.update(kw)
|
|
96
|
+
return SIMITConfig(**data)
|
|
97
|
+
|
|
98
|
+
|
|
99
|
+
def compute_thresholds(epsilon: float, A0: float, A1: float, B: float, k_max: int, eta: float = 1e-3) -> dict[int, float]:
|
|
100
|
+
import math
|
|
101
|
+
|
|
102
|
+
def clip(a):
|
|
103
|
+
return min(max(a, eta), 1 - eta)
|
|
104
|
+
|
|
105
|
+
a0 = clip(A0)
|
|
106
|
+
o0 = a0 / (1 - a0)
|
|
107
|
+
tau = {0: 1 - epsilon}
|
|
108
|
+
for k in range(1, k_max + 1):
|
|
109
|
+
ak = clip(A1 + B * math.log(k))
|
|
110
|
+
ok = ak / (1 - ak)
|
|
111
|
+
tau[k] = (1 - epsilon) * o0 / (epsilon * ok + (1 - epsilon) * o0)
|
|
112
|
+
return tau
|
simit/metrics.py
ADDED
|
@@ -0,0 +1,113 @@
|
|
|
1
|
+
"""Answer-scoring metrics for ``SIMIT.tune`` (higher is better, per sample in [0, 1]).
|
|
2
|
+
|
|
3
|
+
Every metric is ``metric(prediction: str, references: list[str]) -> float``;
|
|
4
|
+
pass any such callable to ``tune(metric=...)``.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
import re
|
|
10
|
+
from typing import Callable, Sequence, Union
|
|
11
|
+
|
|
12
|
+
_ARTICLES = {"a", "an", "the"}
|
|
13
|
+
_PUNCT = re.compile(r"[^\w\s]")
|
|
14
|
+
_NUMBER_WORDS = {"none": "0", "zero": "0", "one": "1", "two": "2", "three": "3", "four": "4", "five": "5",
|
|
15
|
+
"six": "6", "seven": "7", "eight": "8", "nine": "9", "ten": "10"}
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def normalize(text: str) -> str:
|
|
19
|
+
"""Lowercase, drop punctuation and articles, map number words to digits."""
|
|
20
|
+
text = _PUNCT.sub(" ", str(text).lower())
|
|
21
|
+
words = [_NUMBER_WORDS.get(w, w) for w in text.split() if w not in _ARTICLES]
|
|
22
|
+
return " ".join(words)
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
def _refs(references) -> list[str]:
|
|
26
|
+
if isinstance(references, str):
|
|
27
|
+
return [references]
|
|
28
|
+
return [str(r) for r in references]
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def exact_match(prediction: str, references) -> float:
|
|
32
|
+
p = normalize(prediction)
|
|
33
|
+
return float(any(p == normalize(r) for r in _refs(references)))
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def vqa_accuracy(prediction: str, references) -> float:
|
|
37
|
+
"""VQAv2-style soft accuracy: min(1, #matching annotators / 3)."""
|
|
38
|
+
p = normalize(prediction)
|
|
39
|
+
refs = _refs(references)
|
|
40
|
+
hits = sum(p == normalize(r) for r in refs)
|
|
41
|
+
return min(1.0, hits / 3.0) if len(refs) > 1 else float(hits > 0)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def contains(prediction: str, references) -> float:
|
|
45
|
+
p = normalize(prediction)
|
|
46
|
+
return float(any(normalize(r) and normalize(r) in p for r in _refs(references)))
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def _levenshtein(a: str, b: str) -> int:
|
|
50
|
+
if len(a) < len(b):
|
|
51
|
+
a, b = b, a
|
|
52
|
+
prev = list(range(len(b) + 1))
|
|
53
|
+
for i, ca in enumerate(a, 1):
|
|
54
|
+
cur = [i]
|
|
55
|
+
for j, cb in enumerate(b, 1):
|
|
56
|
+
cur.append(min(prev[j] + 1, cur[j - 1] + 1, prev[j - 1] + (ca != cb)))
|
|
57
|
+
prev = cur
|
|
58
|
+
return prev[-1]
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def anls(prediction: str, references, threshold: float = 0.5) -> float:
|
|
62
|
+
"""Average normalized Levenshtein similarity (DocVQA / InfographicVQA)."""
|
|
63
|
+
p = " ".join(str(prediction).lower().split())
|
|
64
|
+
best = 0.0
|
|
65
|
+
for r in _refs(references):
|
|
66
|
+
r = " ".join(r.lower().split())
|
|
67
|
+
d = _levenshtein(p, r) / max(len(p), len(r), 1)
|
|
68
|
+
best = max(best, 1 - d if d < threshold else 0.0)
|
|
69
|
+
return best
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
def relaxed_accuracy(prediction: str, references, tolerance: float = 0.05) -> float:
|
|
73
|
+
"""ChartQA: numbers within 5% relative error, otherwise exact match."""
|
|
74
|
+
def num(s):
|
|
75
|
+
s = str(s).strip().rstrip("%").replace(",", "")
|
|
76
|
+
try:
|
|
77
|
+
return float(s)
|
|
78
|
+
except ValueError:
|
|
79
|
+
return None
|
|
80
|
+
|
|
81
|
+
pv = num(prediction)
|
|
82
|
+
for r in _refs(references):
|
|
83
|
+
rv = num(r)
|
|
84
|
+
if pv is not None and rv is not None:
|
|
85
|
+
if (rv == 0 and pv == 0) or (rv != 0 and abs(pv - rv) / abs(rv) <= tolerance):
|
|
86
|
+
return 1.0
|
|
87
|
+
elif normalize(prediction) == normalize(r):
|
|
88
|
+
return 1.0
|
|
89
|
+
return 0.0
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
_LETTER = re.compile(r"\b([A-J])\b")
|
|
93
|
+
|
|
94
|
+
|
|
95
|
+
def multiple_choice(prediction: str, references) -> float:
|
|
96
|
+
"""Compare the first option letter in the prediction with the reference letter."""
|
|
97
|
+
m = _LETTER.search(str(prediction).strip().upper()) or re.match(r"\s*([A-J])", str(prediction).upper())
|
|
98
|
+
letter = m.group(1) if m else str(prediction).strip()[:1].upper()
|
|
99
|
+
return float(any(letter == str(r).strip().upper()[:1] for r in _refs(references)))
|
|
100
|
+
|
|
101
|
+
|
|
102
|
+
METRICS: dict[str, Callable] = {
|
|
103
|
+
"exact_match": exact_match, "vqa_accuracy": vqa_accuracy, "contains": contains, "anls": anls,
|
|
104
|
+
"relaxed_accuracy": relaxed_accuracy, "multiple_choice": multiple_choice,
|
|
105
|
+
}
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def get_metric(metric: Union[str, Callable]) -> Callable[[str, Sequence[str]], float]:
|
|
109
|
+
if callable(metric):
|
|
110
|
+
return metric
|
|
111
|
+
if metric not in METRICS:
|
|
112
|
+
raise ValueError(f"unknown metric {metric!r}; choose from {sorted(METRICS)} or pass a callable")
|
|
113
|
+
return METRICS[metric]
|