Skip to content

Commit 13d43a0

Browse files
committed
refactoring
1 parent d9872d6 commit 13d43a0

2 files changed

Lines changed: 26 additions & 39 deletions

File tree

hydradx/model/solver/amm_constraints.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -115,9 +115,19 @@ def get_profit_coefs(self, global_asset_list, buffer_fee: float = 0.0, fee_match
115115
return np.hstack((X_coefs, L_coefs, a_coefs))
116116

117117
@abstractmethod
118-
def get_amm_bounds(self, amm: Exchange, amm_i: AmmIndexObject, approx: str, scaling: dict) -> tuple:
118+
def get_amm_bounds(self, approx: str, scaling: dict) -> tuple:
119119
pass
120120

121+
def get_amm_constraint_matrix(self, approx: str, scaling: dict, amm_directions: list, last_amm_deltas: list,
122+
trading_tkns: list):
123+
A1, b1, cones1, cone_sizes1 = self.get_amm_limits_A(amm_directions, last_amm_deltas, trading_tkns)
124+
A2, b2, cones2, cone_sizes2 = self.get_amm_bounds(approx, scaling)
125+
A = np.vstack([A1, A2])
126+
b = np.concatenate([b1, b2])
127+
cones = cones1 + cones2
128+
cone_sizes = cone_sizes1 + cone_sizes2
129+
return A, b, cones, cone_sizes
130+
121131

122132
class XykConstraints(AmmConstraints):
123133
def __init__(self, amm: ConstantProductPoolState):

hydradx/model/solver/omnix_solver.py

Lines changed: 15 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99

1010
from hydradx.model.amm.exchange import Exchange
1111
from hydradx.model.amm.omnipool_amm import OmnipoolState
12-
from hydradx.model.solver.amm_constraints import XykConstraints, StableswapConstraints, AmmConstraints
12+
from hydradx.model.solver.amm_constraints import XykConstraints, StableswapConstraints
1313
from hydradx.model.solver.amm_constraints import AmmIndexObject
1414
from hydradx.model.solver.omnix import validate_and_execute_solution
1515
from hydradx.model.amm.stableswap_amm import StableSwapPoolState
@@ -686,27 +686,6 @@ def _expand_submatrix(A, k: int, start: int):
686686
return A_limits_i
687687

688688

689-
def _get_amm_bounds(p, indices_to_keep=None):
690-
# CFMM invariants must be respected
691-
n, u, m, N = p.n, p.u, p.m, p.N
692-
k = 4 * n + p.amm_vars + m
693-
694-
A5 = np.zeros((0, k))
695-
b5 = np.array([])
696-
cones5 = []
697-
cones_count5 = []
698-
for j, amm in enumerate(p.amm_list):
699-
amm_constraints = p.amm_constraints[j]
700-
approx = p.get_amm_approx(j)
701-
A5j_small, b5j, cones5j, cones_count5j = amm_constraints.get_amm_bounds(approx, p._scaling)
702-
A5j = _expand_submatrix(A5j_small, k, p.amm_i[j].shares_net)
703-
A5 = np.vstack([A5, A5j])
704-
b5 = np.concatenate([b5, b5j])
705-
cones5 = cones5 + cones5j
706-
cones_count5 = cones_count5 + cones_count5j
707-
return A5[:, indices_to_keep], b5, cones5, cones_count5
708-
709-
710689
def _find_solution_unrounded(
711690
p: ICEProblem,
712691
allow_loss: bool = False
@@ -897,9 +876,6 @@ def _find_solution_unrounded(
897876
b4 = np.append(b4, b4i)
898877
A4_trimmed = A4[:, indices_to_keep]
899878

900-
# CFMM invariants must be respected
901-
A5_trimmed, b5, cones5, cone_sizes5 = _get_amm_bounds(p, indices_to_keep)
902-
903879
# inequality constraints on comparison of lrna_lambda to yi, lambda to xi
904880
A6 = np.zeros((0, k))
905881
for i in range(n):
@@ -916,28 +892,29 @@ def _find_solution_unrounded(
916892
cone6 = cb.NonnegativeConeT(A6.shape[0])
917893
cone_sizes6 = [A6.shape[0]]
918894

919-
A_limits_amms = np.zeros((0, k))
920-
b_limits_amms = np.zeros(0)
921-
cones_limits_amms = []
895+
A_amms = np.zeros((0, k))
896+
b_amms = np.zeros(0)
897+
cones_amms = []
922898
cone_sizes_amms = []
923899
for i, amm in enumerate(p.amm_list):
924900
amm_constraints = p.amm_constraints[i]
901+
approx = p.get_amm_approx(i)
925902
directions = [] if len(amm_directions) <= i else amm_directions[i]
926903
last_amm_deltas = [] if p._last_amm_deltas is None else p._last_amm_deltas[i]
927-
x = amm_constraints.get_amm_limits_A(directions, last_amm_deltas, p.trading_tkns)
928-
A_limits_amm_i, b_limits_amm_i, cones_limits_amm_i, cones_sizes_amm_i = x
929-
A_limits_i = _expand_submatrix(A_limits_amm_i, k, p.amm_i[i].shares_net)
904+
x = amm_constraints.get_amm_constraint_matrix(approx, p._scaling, directions, last_amm_deltas, p.trading_tkns)
905+
A_amm_i_small, b_amm_i, cones_amm_i, cones_sizes_amm_i = x
906+
A_amm_i = _expand_submatrix(A_amm_i_small, k, p.amm_i[i].shares_net)
930907

931-
A_limits_amms = np.vstack([A_limits_amms, A_limits_i])
932-
b_limits_amms = np.concatenate([b_limits_amms, b_limits_amm_i])
933-
cones_limits_amms.extend(cones_limits_amm_i)
908+
A_amms = np.vstack([A_amms, A_amm_i])
909+
b_amms = np.concatenate([b_amms, b_amm_i])
910+
cones_amms.extend(cones_amm_i)
934911
cone_sizes_amms.extend(cones_sizes_amm_i)
935912

936-
A = np.vstack([A1_trimmed, A2_trimmed, A3_trimmed, A4_trimmed, A5_trimmed, A6_trimmed, A_limits_amms])
913+
A = np.vstack([A1_trimmed, A2_trimmed, A3_trimmed, A4_trimmed, A6_trimmed, A_amms])
937914
A_sparse = sparse.csc_matrix(A)
938-
b = np.concatenate([b1, b2, b3, b4, b5, b6, b_limits_amms])
939-
cones = cones1 + [cone2, cone3] + cones4 + cones5 + [cone6] + cones_limits_amms
940-
cone_sizes = cone_sizes1 + cone_sizes2 + cone_sizes3 + cone_sizes4 + cone_sizes5 + cone_sizes6 + cone_sizes_amms
915+
b = np.concatenate([b1, b2, b3, b4, b6, b_amms])
916+
cones = cones1 + [cone2, cone3] + cones4 + [cone6] + cones_amms
917+
cone_sizes = cone_sizes1 + cone_sizes2 + cone_sizes3 + cone_sizes4 + cone_sizes6 + cone_sizes_amms
941918

942919
# solve
943920
settings = clarabel.DefaultSettings()

0 commit comments

Comments
 (0)