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