File size: 2,653 Bytes
ef81b32
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
from pathlib import Path
import sys

import numpy as np
import torch

ROOT = Path(__file__).resolve().parents[1]
if str(ROOT) not in sys.path:
    sys.path.insert(0, str(ROOT))

def build_parser() -> argparse.ArgumentParser:
    parser = argparse.ArgumentParser(description="Minimal Order-LM LOB-prefix inference example.")
    parser.add_argument("--model-size", default="base", choices=["tiny", "small", "base", "large"])
    parser.add_argument("--prompt-ids", default="examples/prompt_ids.npy")
    parser.add_argument("--conditioning", default="examples/conditioning.npz")
    parser.add_argument("--max-tokens", type=int, default=1024)
    parser.add_argument("--temperature", type=float, default=1.0)
    parser.add_argument("--top-p", type=float, default=1.0)
    parser.add_argument("--top-k", type=int, default=None)
    parser.add_argument("--device", default="cuda")
    parser.add_argument("--dtype", default="bf16")
    parser.add_argument("--tokenizer-device", default="cuda")
    return parser


def main() -> None:
    args = build_parser().parse_args()
    from vq_order_model.rollout_lob.cached_lob_rollout import (
        CachedLobRolloutConfig,
        TorchCachedLobRollout,
        load_lob_conditioning,
    )

    root = ROOT
    checkpoint = root / args.model_size / "best.pt"
    tokenizer_checkpoint = root / "tokenizer" / "base" / "best.pt"

    config = CachedLobRolloutConfig(
        checkpoint=str(checkpoint),
        tokenizer_checkpoint=str(tokenizer_checkpoint),
        device=args.device,
        dtype=args.dtype,
        tokenizer_device=args.tokenizer_device,
        max_tokens=int(args.max_tokens),
        temperature=float(args.temperature),
        top_p=float(args.top_p),
        top_k=args.top_k,
        include_prompt=False,
        denormalize=False,
        prepend_bos_for_model=True,
    )

    prompt_ids = np.load(root / args.prompt_ids).astype(np.int64)
    conditioning = load_lob_conditioning(root / args.conditioning, max_items=prompt_ids.shape[0])

    runner = TorchCachedLobRollout(config)
    generated_ids = runner.generate_ids_batched(
        prompt_ids=torch.as_tensor(prompt_ids, dtype=torch.long),
        lob=torch.as_tensor(conditioning.lob, dtype=torch.float32),
        initial_state_float=torch.as_tensor(conditioning.initial_state_float, dtype=torch.float32),
    )
    decoded = runner.decode_ids(generated_ids)

    print("generated ids:", tuple(generated_ids.shape))
    print("decoded features:", decoded["features"].shape)
    print("feature order:", decoded["feature_order"])


if __name__ == "__main__":
    main()