Files
sentiment-engine/MALKHUT/malkhut/tests/test_numba.py

121 lines
4.1 KiB
Python
Raw Normal View History

"""
Tests for numba-accelerated CWM functions.
Verifies:
- Numba JIT compilation works
- fill_from_levels produces same results as pure Python
- round_tick / round_lot / clip_lots correct
- Feature extraction vectorized
- Fallback to pure Python when numba unavailable
"""
import numpy as np
import pytest
from malkhut.cwm.numba_core import (
fill_from_levels, round_tick, round_lot, clip_lots,
extract_features_vectorized, compare_states_vectorized,
)
class TestNumbaFillFromLevels:
def test_fill_single_level(self):
prices = np.array([50000.0], dtype=np.float64)
qtys = np.array([1.0], dtype=np.float64)
filled, avg, _, _ = fill_from_levels(
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
prices, qtys, 0.5, 0.001, 0.001, True,
)
assert filled == pytest.approx(0.5, abs=0.001)
assert avg == pytest.approx(50000.0, abs=0.01)
def test_fill_multi_level(self):
prices = np.array([50000.0, 50001.0], dtype=np.float64)
qtys = np.array([0.5, 0.5], dtype=np.float64)
filled, avg, _, _ = fill_from_levels(
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
prices, qtys, 0.8, 0.001, 0.001, True,
)
assert filled == pytest.approx(0.8, abs=0.001)
assert avg > 50000.0
def test_fill_exhausts_all(self):
prices = np.array([50000.0, 50001.0], dtype=np.float64)
qtys = np.array([0.3, 0.3], dtype=np.float64)
filled, avg, _, _ = fill_from_levels(
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
prices, qtys, 1.0, 0.001, 0.001, True,
)
assert filled == pytest.approx(0.6, abs=0.001)
def test_fill_empty(self):
filled, avg, _, _ = fill_from_levels(
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
np.array([], dtype=np.float64), np.array([], dtype=np.float64),
1.0, 0.001, 0.001, True,
)
assert filled == 0.0
class TestNumbaRounding:
def test_round_tick(self):
assert round_tick(50000.0, 0.1) == 50000.0
assert round_tick(50000.06, 0.1) == pytest.approx(50000.1, abs=1e-9)
assert round_tick(50004.4, 1.0) == 50004.0
def test_round_lot(self):
assert round_lot(0.001, 0.001) == 0.001
assert round_lot(0.0017, 0.001) == 0.002
def test_clip_lots_above_min(self):
assert clip_lots(0.005, 0.001, 0.001) == 0.005
def test_clip_lots_below_min(self):
assert clip_lots(0.0005, 0.001, 0.001) == 0.0
class TestNumbaFeatures:
def test_extract_features(self):
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)
features = 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,
)
assert len(features) == 17
assert features[0] == pytest.approx(50000.5, abs=0.1) # mid
assert features[14] == pytest.approx(0.3, abs=0.01) # toxicity
class TestNumbaCompare:
def test_compare_match(self):
ok, idx, ev, av = compare_states_vectorized(
10000.0, 10000.0, 50000.0, 50000.0, 50001.0, 50001.0,
1e-6, 1e-6,
)
assert ok
def test_compare_equity_mismatch(self):
ok, idx, ev, av = compare_states_vectorized(
10000.0, 9000.0, 50000.0, 50000.0, 50001.0, 50001.0,
1e-6, 1e-6,
)
assert not ok
assert idx == 0
def test_compare_bid_mismatch(self):
ok, idx, ev, av = compare_states_vectorized(
10000.0, 10000.0, 50000.0, 50001.0, 50001.0, 50001.0,
1e-6, 1e-6,
)
assert not ok
assert idx == 1
def test_compare_within_tolerance(self):
ok, idx, ev, av = compare_states_vectorized(
10000.0, 10000.001, 50000.0, 50000.001, 50001.0, 50001.001,
0.1, 0.1,
)
assert ok