File size: 2,451 Bytes
c69aaec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""One decision request -> probability distributions, shared by the benchmark engine and the server.

Every question of a request is scored against the same state in token-budgeted, length-sorted
batches (one forward pass per batch). Nothing is truncated: a request that does not fit raises
CapacityError, whose message carries the standard capacity markers.
"""

from collections.abc import Mapping, Sequence
from typing import Any, cast

import torch

from kev.model import MAX_OPTIONS, DecisionModel, options
from kev.train import length_estimate, microbatches
from kev.types import Example, ImageInput


class CapacityError(ValueError):
    pass


def decide(model: DecisionModel, state: object, questions: Mapping[str, Mapping[str, Any]], *, temperature: float,
           max_tokens: int, token_budget: int, batch_size: int, images: Sequence[ImageInput] = ()) -> tuple[dict[str, list[float]], int]:
    """Normalized probabilities per question key (in option order) and the input tokens used."""
    rows: list[Example] = []
    for key, question in questions.items():
        count = len(options(cast(Any, question))[0])
        if count > MAX_OPTIONS:
            raise CapacityError(f"at most {MAX_OPTIONS} options per choice question are supported ({count} given)")
        row = {"state": state, "question": dict(question), "id": key, "suite": "", "family": key, "label": "", "target": "",
               "source": {}}
        if images:
            row["images"] = list(images)
        rows.append(cast(Example, row))
    distributions: dict[str, list[float]] = {}
    input_tokens = 0
    for batch in microbatches(sorted(rows, key=length_estimate), batch_size, token_budget):
        try:
            prepared = model.prepare(batch, max_length=max_tokens)
        except ValueError as error:
            if "token limit" in str(error):
                raise CapacityError(f"request exceeds the maximum context length of {max_tokens} tokens") from error
            raise
        input_tokens += prepared.input_tokens
        with torch.inference_mode():
            probabilities = (model(prepared) / temperature).softmax(-1).float().cpu().tolist()
        for item, values, count in zip(batch, probabilities, prepared.counts, strict=True):
            total = sum(values[:count])
            distributions[item["id"]] = [value / total for value in values[:count]]
    return {key: distributions[key] for key in questions}, input_tokens