multipoly.py 5.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212
  1. class Variable:
  2. def __init__(self, name):
  3. self.name = name
  4. def __pow__(self, n):
  5. expr = MultiplyExpression()
  6. expr.set_symbol(self.name, n)
  7. return expr
  8. def __eq__(self, other):
  9. return self.name == other.name
  10. def termify(self):
  11. expr = MultiplyExpression()
  12. expr.set_symbol(self.name, 1)
  13. return expr
  14. class MultiplyExpression:
  15. def __init__(self):
  16. self.coeff = fp(1)
  17. self.symbols = {}
  18. def copy(self):
  19. result = MultiplyExpression()
  20. result.coeff = self.coeff
  21. result.symbols = self.symbols.copy()
  22. return result
  23. def clean(self):
  24. for symbol in list(self.symbols.keys()):
  25. if self.symbols[symbol] == 0:
  26. del self.symbols[symbol]
  27. def matches(self, other):
  28. return self.symbols == other.symbols
  29. def set_symbol(self, var_name, power):
  30. self.symbols[var_name] = power
  31. def __neg__(self):
  32. result = self.copy()
  33. result.coeff *= -1
  34. return result
  35. def __mul__(self, expr):
  36. result = MultiplyExpression()
  37. result.coeff = self.coeff
  38. result.symbols = self.symbols.copy()
  39. if hasattr(expr, "field"):
  40. result.coeff *= expr
  41. return result
  42. if isinstance(expr, Variable):
  43. expr = expr.termify()
  44. for var_name, power in expr.symbols.items():
  45. if var_name in result.symbols:
  46. result.symbols[var_name] += power
  47. else:
  48. result.symbols[var_name] = power
  49. # Remember to multiply the coefficients
  50. result.coeff *= expr.coeff
  51. return result
  52. def __add__(self, expr):
  53. if isinstance(expr, Variable):
  54. expr = expr.termify()
  55. if self.matches(expr):
  56. result = self.copy()
  57. result.coeff += expr.coeff
  58. return result
  59. return MultivariatePolynomial([self, expr])
  60. def __str__(self):
  61. repr = ""
  62. first = True
  63. if self.coeff != fp(1):
  64. repr += str(self.coeff)
  65. first = False
  66. for var_name, power in self.symbols.items():
  67. if first:
  68. first = False
  69. else:
  70. repr += " "
  71. if power == 1:
  72. repr += var_name
  73. else:
  74. repr += var_name + "^" + str(power)
  75. return repr
  76. class MultivariatePolynomial:
  77. def __init__(self, terms=[]):
  78. self.terms = terms
  79. def copy(self):
  80. terms = [term.copy() for term in self.terms]
  81. return MultivariatePolynomial(terms)
  82. # Operations can accept Variables and constants
  83. # so we make sure to convert them to MultiplyExpression types
  84. def _convert_term(self, term):
  85. if isinstance(term, Variable):
  86. term = term.termify()
  87. if hasattr(term, "field"):
  88. expr = MultiplyExpression()
  89. expr.coeff = term
  90. term = expr
  91. return term
  92. def __neg__(self):
  93. terms = [-term for term in self.terms]
  94. return MultivariatePolynomial(terms)
  95. def __add__(self, term):
  96. term = self._convert_term(term)
  97. if isinstance(term, MultivariatePolynomial):
  98. # Recursively apply addition operation
  99. result = self.copy()
  100. for other_term in term.terms:
  101. result += other_term
  102. return result
  103. assert isinstance(term, MultiplyExpression)
  104. # Delete ^0 variables
  105. term.clean()
  106. result = self.copy()
  107. result_term = result._find(term)
  108. if result_term is None:
  109. result.terms.append(term)
  110. else:
  111. result_term.coeff += term.coeff
  112. return result
  113. def __sub__(self, term):
  114. term = -term
  115. return self + term
  116. def __mul__(self, term):
  117. term = self._convert_term(term)
  118. if isinstance(term, MultivariatePolynomial):
  119. # Recursively apply addition operation
  120. result = MultivariatePolynomial()
  121. for other_term in term.terms:
  122. result += self * other_term
  123. return result
  124. assert isinstance(term, MultiplyExpression)
  125. # Delete ^0 variables
  126. term.clean()
  127. terms = [self_term * term for self_term in self.terms]
  128. result = MultivariatePolynomial(terms)
  129. return result
  130. def _find(self, other):
  131. for term in self.terms:
  132. if term.matches(other):
  133. return term
  134. return None
  135. def __str__(self):
  136. if not self.terms:
  137. return "0"
  138. repr = ""
  139. first = True
  140. for term in self.terms:
  141. if first:
  142. first = False
  143. else:
  144. repr += " + "
  145. repr += str(term)
  146. return repr
  147. if __name__ == "__main__":
  148. from finite_fields import finitefield
  149. p = 0x40000000000000000000000000000000224698fc094cf91b992d30ed00000001
  150. fp = finitefield.IntegersModP(p)
  151. x = Variable("X")
  152. y = Variable("Y")
  153. z = Variable("Z")
  154. print(y**2 + y**2)
  155. p = x**3 * y**2 * x**2 * fp(5) * fp(2) + x**3 * y + z + fp(6)
  156. q = x**3 * y * fp(3) + y
  157. print(p)
  158. print(q)
  159. print(p + q)
  160. print(p * q)
  161. print(-q)
  162. print(p - q)