File size: 2,332 Bytes
68e5880
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Run against a local merged HF snapshot; no project-local dependencies."""

import argparse
import json

from openjet_runtime import OpenJet


def decision_examples():
    binary = {
        "id": "example-binary",
        "group_id": "example-binary",
        "primitive": "choice",
        "state": "Order 731 has been delivered. The customer's message says thank you.",
        "instructions": "Select the appropriate next workflow action.",
        "criteria": [
            {"id": "close", "description": "Close the resolved support ticket."},
            {"id": "refund", "description": "Refund an undelivered order."},
        ],
    }
    browser = {
        "id": "example-browser",
        "group_id": "example-browser",
        "primitive": "choice",
        "state": "A settings page has 16 visible buttons labeled Page 1 through Page 16.",
        "instructions": "Navigate to Page 12 by choosing its matching button.",
        "criteria": [
            {"id": f"click-{i}", "description": f"Click the Page {i} button."}
            for i in range(1, 17)
        ],
    }
    return binary, browser


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("model", help="Local merged snapshot directory")
    parser.add_argument("--device", default="cuda:0")
    parser.add_argument("--dtype", choices=["float32", "bfloat16"], default="bfloat16")
    parser.add_argument("--effort", choices=["low", "high", "both"], default="both")
    parser.add_argument(
        "--text", action="store_true", help="Also run slow TYPE reference"
    )
    args = parser.parse_args()
    runtime = OpenJet.from_pretrained(args.model, args.device, args.dtype)
    efforts = ("low", "high") if args.effort == "both" else (args.effort,)
    for effort in efforts:
        for request in decision_examples():
            result = runtime.decide(request, effort)
            print(json.dumps({"example": request["id"], **result}, ensure_ascii=False))
        if args.text:
            result = runtime.generate_text(
                "Return only the literal text to type into a search box for 'red shoes'.",
                effort=effort,
                max_new_tokens=32,
            )
            print(json.dumps({"example": "browser-type", **result}, ensure_ascii=False))


if __name__ == "__main__":
    main()