Browse Source

calculate sonic arithmetization equations

narodnik 5 years ago
parent
commit
0b865a69c8
2 changed files with 101 additions and 2 deletions
  1. 54 0
      scripts/halo/multipoly.py
  2. 47 2
      scripts/halo/sonic.py

+ 54 - 0
scripts/halo/multipoly.py

@@ -1,3 +1,8 @@
+import numpy as np
+from finite_fields import finitefield
+p = 0x40000000000000000000000000000000224698fc094cf91b992d30ed00000001
+fp = finitefield.IntegersModP(p)
+
 class Variable:
 
     def __init__(self, name):
@@ -11,6 +16,9 @@ class Variable:
     def __eq__(self, other):
         return self.name == other.name
 
+    def __hash__(self):
+        return hash(self.name)
+
     def termify(self):
         expr = MultiplyExpression()
         expr.set_symbol(self.name, 1)
@@ -39,6 +47,10 @@ class MultiplyExpression:
     def set_symbol(self, var_name, power):
         self.symbols[var_name] = power
 
+    def __eq__(self, other):
+        return (self.coeff == other.coeff and
+                self.symbols == other.symbols)
+
     def __neg__(self):
         result = self.copy()
         result.coeff *= -1
@@ -49,6 +61,9 @@ class MultiplyExpression:
         result.coeff = self.coeff
         result.symbols = self.symbols.copy()
 
+        if isinstance(expr, np.int64) or isinstance(expr, int):
+            expr = fp(int(expr))
+
         if hasattr(expr, "field"):
             result.coeff *= expr
             return result
@@ -77,6 +92,20 @@ class MultiplyExpression:
 
         return MultivariatePolynomial([self, expr])
 
+    def __sub__(self, expr):
+        expr = -expr
+        return self + expr
+
+    def evaluate(self, symbol_map):
+        result = MultiplyExpression()
+        for symbol, power in self.symbols.items():
+            if symbol in symbol_map:
+                value = symbol_map[symbol]
+                result *= value**power
+            else:
+                result *= Variable(symbol)**power
+        return result
+
     def __str__(self):
         repr = ""
         first = True
@@ -118,6 +147,12 @@ class MultivariatePolynomial:
 
         return term
 
+    def __bool__(self):
+        return bool(self.terms)
+
+    def __eq__(self, other):
+        return self.terms == other.terms
+
     def __neg__(self):
         terms = [-term for term in self.terms]
         return MultivariatePolynomial(terms)
@@ -136,6 +171,10 @@ class MultivariatePolynomial:
         # Delete ^0 variables
         term.clean()
 
+        # Skip terms where the coeff is 0
+        if term.coeff == fp(0):
+            return self
+
         result = self.copy()
         result_term = result._find(term)
 
@@ -164,17 +203,32 @@ class MultivariatePolynomial:
         # Delete ^0 variables
         term.clean()
 
+        # Skip terms where the coeff is 0
+        if term.coeff == fp(0):
+            return self
+
         terms = [self_term * term for self_term in self.terms]
         result = MultivariatePolynomial(terms)
 
         return result
 
+    def divmod(self, poly):
+        assert isinstance(poly, MultivariatePolynomial)
+        # https://www.win.tue.nl/~aeb/2WF02/groebner.pdf
+
     def _find(self, other):
         for term in self.terms:
             if term.matches(other):
                 return term
         return None
 
+    def evaluate(self, variable_map):
+        p = MultivariatePolynomial()
+        for term in self.terms:
+            assert isinstance(term, MultiplyExpression)
+            p += term.evaluate(variable_map)
+        return p
+
     def __str__(self):
         if not self.terms:
             return "0"

+ 47 - 2
scripts/halo/sonic.py

@@ -106,15 +106,60 @@ assert u.shape == w.shape
 
 k = np.array((k1, k2, k3, k4, k5, k6, k7))
 
+x = Variable("X")
 y = Variable("Y")
 p = MultivariatePolynomial()
 for i, (a_i, b_i, c_i) in enumerate(zip(a, b, c), 1):
     #print(a_i, "\t", b_i, "\t", c_i)
     p += y**i * (a_i * b_i - c_i)
-print("Polynomial:", p)
+assert not p
 
 p = MultivariatePolynomial()
 for q, (u_q, v_q, w_q, k_q) in enumerate(zip(u, v, w, k)):
     p += y**q * (a.dot(u_q) + b.dot(v_q) + c.dot(w_q) - k_q)
-print("Polynomial:", p)
+assert not p
+
+n = len(a)
+assert len(b) == n
+assert len(c) == n
+
+assert u.shape == (7, n)
+assert v.shape == u.shape
+assert w.shape == u.shape
+assert k.shape == (7,)
+
+r_x_y = MultivariatePolynomial()
+s_x_y = MultivariatePolynomial()
+for i, (a_i, b_i, c_i) in enumerate(zip(a, b, c), 1):
+    assert 1 <= i <= n
+
+    r_x_y += x**i * y**i * a_i
+    r_x_y += x**-i * y**-i * b_i
+    r_x_y += x**(-i - n) * y**(-i - n) * c_i
+
+    u_i = u.T[i - 1]
+    v_i = v.T[i - 1]
+    w_i = w.T[i - 1]
+    u_i_Y = MultivariatePolynomial()
+    v_i_Y = MultivariatePolynomial()
+    w_i_Y = MultivariatePolynomial()
+    for q, (u_q_i, v_q_i, w_q_i) in enumerate(zip(u_i, v_i, w_i), 1):
+        assert 1 <= q <= 7
+
+        u_i_Y += y**(q + n) * u_q_i
+        v_i_Y += y**(q + n) * v_q_i
+        w_i_Y += -y**i - y**(-i) + y**(q + n) * v_q_i
+
+    s_x_y += u_i_Y * x**-i + v_i_Y * x**i + w_i_Y * x**(i + n)
+
+k_y = MultivariatePolynomial()
+for q, k_q in enumerate(k, 1):
+    assert 1 <= q <= 7
+    k_y += y**(q + n) * k_q
+
+r_prime_x_y = r_x_y + s_x_y
+r_x_1 = r_x_y.evaluate({y.name: fp(1)})
+t_x_y = r_x_1 * r_prime_x_y - k_y
+print()
+print(t_x_y)