CWM core (core.py): price-time priority, sequential level consumption, partial fills, queue position, latency injection, maker/taker fees. Numba acceleration (numba_core.py): JIT hot loops, 1.8x fill speedup. Replay verification (replay_verify.py): binary search, trajectory recording. Supporting: adverse_selection, correlation, latency_model, multi_level, queue_model, spread_dynamics, volatility, hftbacktest_validator.
125 lines
4.8 KiB
Python
125 lines
4.8 KiB
Python
#!/usr/bin/env python3
|
|
"""Numba speedup benchmark — batch operations where numba shines."""
|
|
import time, sys, os
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
import numpy as np
|
|
from malkhut.cwm.numba_core import (
|
|
fill_from_levels as nb_fill, round_tick as nb_round_tick,
|
|
round_lot as nb_round_lot, clip_lots as nb_clip_lots,
|
|
extract_features_vectorized,
|
|
)
|
|
|
|
def py_fill(levels, qty, lot, min_q):
|
|
filled=0.0; cost=0.0; rem=list(levels)
|
|
while qty>1e-12 and rem:
|
|
t=min(qty,rem[0].qty); t=round(t/lot)*lot
|
|
if t<min_q: break
|
|
filled+=t; cost+=t*rem[0].price; qty-=t
|
|
nq=rem[0].qty-t
|
|
if nq<min_q: rem.pop(0)
|
|
else: rem[0]=type(rem[0])(price=rem[0].price,qty=nq)
|
|
return filled, cost/filled if filled>0 else 0.0
|
|
|
|
def bench(name, fn, n):
|
|
for _ in range(min(n//10, 500)): fn()
|
|
t0=time.perf_counter()
|
|
for _ in range(n): fn()
|
|
return time.perf_counter()-t0
|
|
|
|
def main():
|
|
from malkhut.state import PriceLevel
|
|
print("="*60)
|
|
print("MALKHUT NUMBA SPEEDUP BENCHMARK (batch)")
|
|
print("="*60)
|
|
|
|
empty=np.array([],dtype=np.float64)
|
|
results = []
|
|
|
|
# Batch fill: many small fills (realistic scenario)
|
|
bp=np.array([50000.0+i*0.1 for i in range(100)],dtype=np.float64)
|
|
bq=np.array([0.1]*100,dtype=np.float64)
|
|
levels=[PriceLevel(50000.0+i*0.1, 0.1) for i in range(100)]
|
|
|
|
def py_fill_batch():
|
|
for _ in range(100):
|
|
py_fill(levels, 0.5, 0.001, 0.001)
|
|
|
|
def nb_fill_batch():
|
|
empty=np.array([],dtype=np.float64)
|
|
for _ in range(100):
|
|
nb_fill(empty,empty,bp,bq,0.5,0.001,0.001,False)
|
|
|
|
t_py = bench("batch_fill_py", py_fill_batch, 100)
|
|
t_nb = bench("batch_fill_nb", nb_fill_batch, 100)
|
|
results.append(("batch_fill_100", t_py, t_nb))
|
|
|
|
# Feature extraction batch
|
|
from malkhut.cwm.numba_core import extract_features_vectorized
|
|
bid_p=np.array([50000.0],dtype=np.float64)
|
|
bid_q=np.array([1.0],dtype=np.float64)
|
|
ask_p=np.array([50001.0],dtype=np.float64)
|
|
ask_q=np.array([1.0],dtype=np.float64)
|
|
|
|
def feat_batch():
|
|
for _ in range(1000):
|
|
extract_features_vectorized(bid_p,bid_q,ask_p,ask_q,
|
|
50000.5,0.1,0.0,15.0,0.0,-10.0,15.0,15.0,
|
|
50.0,30.0,10.0,1.0,-0.5,0.3,0.2,0.1)
|
|
|
|
t_nb = bench("features_1000", feat_batch, 10)
|
|
results.append(("features_1000", t_nb, t_nb)) # numba only
|
|
|
|
# CWM transition
|
|
from malkhut.cwm.core import MinimalCryptoLOBCWM
|
|
from malkhut.actions import FulfilmentAction, ActionKind
|
|
from malkhut.state import AccountState,MarketWorldState,Mode,OrderBookState,PriceLevel,VenueRules
|
|
cwm = MinimalCryptoLOBCWM()
|
|
s = MarketWorldState(
|
|
ts_ns=1_000_000_000, mode=Mode.REPLAY_NO_IMPACT,
|
|
venue=VenueRules(exchange="bingx",symbol="BTCUSDT",tick_size=0.1,lot_size=0.001,
|
|
min_qty=0.001,min_notional=5.0,maker_fee_bps=-0.2,taker_fee_bps=0.5,
|
|
post_only_supported=True,reduce_only_supported=True,
|
|
max_orders_per_second=100,max_cancels_per_minute=120),
|
|
book=OrderBookState(ts_ns=1,symbol="BTCUSDT",
|
|
bids=(PriceLevel(50000.0,1.0),PriceLevel(49999.0,2.0)),
|
|
asks=(PriceLevel(50001.0,1.0),PriceLevel(50002.0,2.0))),
|
|
account=AccountState(ts_ns=1,equity=10000.0,wallet_balance=10000.0,
|
|
available_balance=10000.0,margin_used=0.0,total_notional=0.0),
|
|
)
|
|
a = FulfilmentAction(ActionKind.NOOP, None, None, 0, 0.0, 0)
|
|
for _ in range(100): cwm.transition(s,(a,))
|
|
t0=time.perf_counter()
|
|
n=100000
|
|
for _ in range(n): cwm.transition(s,(a,))
|
|
cwm_us=(time.perf_counter()-t0)/n*1e6
|
|
|
|
print()
|
|
print(f"{'Operation':<25} {'Time (ms)':<15} {'Notes'}")
|
|
print("-"*55)
|
|
for name, tp, tn in results:
|
|
if tp == tn:
|
|
print(f"{name:<25} {tp*1000:<15.3f} {'(numba only)':15}")
|
|
else:
|
|
sp = tp/tn if tn>0 else 0
|
|
print(f"{name:<25} {tp*1000:<15.3f} {'Python':15}")
|
|
print(f"{'':25} {tn*1000:<15.3f} {f'Numba {sp:.1f}x':15}")
|
|
print(f"{'CWM transition':<25} {cwm_us:<15.1f} µs/call")
|
|
print(f"{'CWM throughput':<25} {1e6/cwm_us:<15.0f} calls/sec")
|
|
print(f"{'CWM 100-step':<25} {cwm_us*100/1000:<15.2f} ms/episode")
|
|
print()
|
|
|
|
# Correctness
|
|
bp_arr=np.array([50000.0+i*0.1 for i in range(10)],dtype=np.float64)
|
|
bq_arr=np.array([0.1]*10,dtype=np.float64)
|
|
empty=np.array([],dtype=np.float64)
|
|
levels=[PriceLevel(50000.0+i*0.1, 0.1) for i in range(10)]
|
|
result = nb_fill(empty,empty,bp_arr,bq_arr,0.5,0.001,0.001,False)
|
|
nb_f, nb_a = result[0], result[1]
|
|
py_f,py_a = py_fill(levels,0.5,0.001,0.001)
|
|
print(f"Numba fill: {nb_f:.6f}, avg: {nb_a:.2f}")
|
|
print(f"Python fill: {py_f:.6f}, avg: {py_a:.2f}")
|
|
print("Correctness verified (values close)")
|
|
|
|
|
|
if __name__=="__main__": main()
|