multipoly.py 4.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197
  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 __mul__(self, expr):
  32. result = MultiplyExpression()
  33. result.coeff = self.coeff
  34. result.symbols = self.symbols.copy()
  35. if hasattr(expr, "field"):
  36. result.coeff *= expr
  37. return result
  38. if isinstance(expr, Variable):
  39. expr = expr.termify()
  40. for var_name, power in expr.symbols.items():
  41. if var_name in result.symbols:
  42. result.symbols[var_name] += power
  43. else:
  44. result.symbols[var_name] = power
  45. # Remember to multiply the coefficients
  46. result.coeff *= expr.coeff
  47. return result
  48. def __add__(self, expr):
  49. if isinstance(expr, Variable):
  50. expr = expr.termify()
  51. if self.matches(expr):
  52. result = self.copy()
  53. result.coeff += expr.coeff
  54. return result
  55. return MultivariatePolynomial([self, expr])
  56. def __str__(self):
  57. repr = ""
  58. first = True
  59. if self.coeff != fp(1):
  60. repr += str(self.coeff)
  61. first = False
  62. for var_name, power in self.symbols.items():
  63. if first:
  64. first = False
  65. else:
  66. repr += " "
  67. if power == 1:
  68. repr += var_name
  69. else:
  70. repr += var_name + "^" + str(power)
  71. return repr
  72. class MultivariatePolynomial:
  73. def __init__(self, terms=[]):
  74. self.terms = terms
  75. def copy(self):
  76. terms = [term.copy() for term in self.terms]
  77. return MultivariatePolynomial(terms)
  78. # Operations can accept Variables and constants
  79. # so we make sure to convert them to MultiplyExpression types
  80. def _convert_term(self, term):
  81. if isinstance(term, Variable):
  82. term = term.termify()
  83. if hasattr(term, "field"):
  84. expr = MultiplyExpression()
  85. expr.coeff = term
  86. term = expr
  87. return term
  88. def __add__(self, term):
  89. term = self._convert_term(term)
  90. if isinstance(term, MultivariatePolynomial):
  91. # Recursively apply addition operation
  92. result = self.copy()
  93. for other_term in term.terms:
  94. result += other_term
  95. return result
  96. assert isinstance(term, MultiplyExpression)
  97. # Delete ^0 variables
  98. term.clean()
  99. result = self.copy()
  100. result_term = result.find(term)
  101. if result_term is None:
  102. result.terms.append(term)
  103. else:
  104. result_term.coeff += term.coeff
  105. return result
  106. def __mul__(self, term):
  107. term = self._convert_term(term)
  108. if isinstance(term, MultivariatePolynomial):
  109. # Recursively apply addition operation
  110. result = MultivariatePolynomial()
  111. for other_term in term.terms:
  112. result += self * other_term
  113. return result
  114. assert isinstance(term, MultiplyExpression)
  115. # Delete ^0 variables
  116. term.clean()
  117. terms = [self_term * term for self_term in self.terms]
  118. result = MultivariatePolynomial(terms)
  119. return result
  120. def find(self, other):
  121. for term in self.terms:
  122. if term.matches(other):
  123. return term
  124. return None
  125. def __str__(self):
  126. if not self.terms:
  127. return "0"
  128. repr = ""
  129. first = True
  130. for term in self.terms:
  131. if first:
  132. first = False
  133. else:
  134. repr += " + "
  135. repr += str(term)
  136. return repr
  137. if __name__ == "__main__":
  138. from finite_fields import finitefield
  139. p = 0x40000000000000000000000000000000224698fc094cf91b992d30ed00000001
  140. fp = finitefield.IntegersModP(p)
  141. x = Variable("X")
  142. y = Variable("Y")
  143. z = Variable("Z")
  144. print(y**2 + y**2)
  145. p = x**3 * y**2 * x**2 * fp(5) * fp(2) + x**3 * y + z + fp(6)
  146. q = x**3 * y * fp(3) + y
  147. print(p)
  148. print(q)
  149. print(p + q)
  150. print(p * q)