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()
|