ONNX
English
vons
research
candidate-selection
File size: 2,017 Bytes
49ad2ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import importlib.util
import sys
from pathlib import Path

import pytest

from vons.data import Example

_SPEC = importlib.util.spec_from_file_location("benchmark_onnx_shapes", Path(__file__).parents[1] / "tools/benchmark_onnx_shapes.py")
assert _SPEC and _SPEC.loader
_MODULE = importlib.util.module_from_spec(_SPEC)
sys.modules[_SPEC.name] = _MODULE
_SPEC.loader.exec_module(_MODULE)
_options_for = _MODULE._options_for
shape_cells = _MODULE.shape_cells


def _row(options: tuple[str, ...]) -> Example:
    return Example(
        id="fixture",
        task_group="shape",
        state="ready",
        question="Choose",
        options=options,
        label=options[0],
        answerable=True,
        split="smoke",
        provenance={},
        metadata={},
    )


def test_shape_cells_keep_the_four_padding_comparisons_explicit() -> None:
    cells = shape_cells(live_candidates=4, short_sequence=64, full_sequence=512)
    assert [cell.cell_id for cell in cells] == [
        "direct_live_short",
        "direct_padded_short",
        "direct_padded_full",
        "diffusion_padded_full",
    ]
    assert [(cell.backend, cell.slots, cell.sequence_length) for cell in cells] == [
        ("direct", 4, 64),
        ("direct", 32, 64),
        ("direct", 32, 512),
        ("diffusion", 32, 512),
    ]


def test_shape_fixture_expands_candidates_without_changing_live_prefix() -> None:
    options = _options_for(_row(("a", "b")), 4)
    assert options[:2] == ("a", "b")
    assert len(options) == 4
    assert len(set(options)) == 4


@pytest.mark.parametrize(
    "kwargs",
    [
        {"live_candidates": 1, "short_sequence": 64, "full_sequence": 512},
        {"live_candidates": 33, "short_sequence": 64, "full_sequence": 512},
        {"live_candidates": 4, "short_sequence": 512, "full_sequence": 64},
    ],
)
def test_shape_cells_reject_invalid_dimensions(kwargs: dict[str, int]) -> None:
    with pytest.raises(ValueError):
        shape_cells(**kwargs)