multipoly.py 4.2 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170
  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 clean(self):
  19. for symbol in list(self.symbols.keys()):
  20. if self.symbols[symbol] == 0:
  21. del self.symbols[symbol]
  22. def matches(self, other):
  23. return self.symbols == other.symbols
  24. def set_symbol(self, var_name, power):
  25. self.symbols[var_name] = power
  26. def __mul__(self, expr):
  27. result = MultiplyExpression()
  28. result.coeff = self.coeff
  29. result.symbols = self.symbols.copy()
  30. if hasattr(expr, "field"):
  31. result.coeff *= expr
  32. return result
  33. if isinstance(expr, Variable):
  34. expr = expr.termify()
  35. for var_name, power in expr.symbols.items():
  36. if var_name in result.symbols:
  37. result.symbols[var_name] += power
  38. else:
  39. result.symbols[var_name] = power
  40. return result
  41. def __add__(self, expr):
  42. if isinstance(expr, Variable):
  43. expr = expr.termify()
  44. return MultivariatePolynomial([self, expr])
  45. def __str__(self):
  46. repr = ""
  47. first = True
  48. if self.coeff != fp(1):
  49. repr += str(self.coeff)
  50. first = False
  51. for var_name, power in self.symbols.items():
  52. if first:
  53. first = False
  54. else:
  55. repr += " "
  56. if power == 1:
  57. repr += var_name
  58. else:
  59. repr += var_name + "^" + str(power)
  60. return repr
  61. class MultivariatePolynomial:
  62. def __init__(self, terms=[]):
  63. self.terms = terms
  64. def copy(self):
  65. return MultivariatePolynomial(self.terms[:])
  66. def __add__(self, term):
  67. if isinstance(term, Variable):
  68. term = term.termify()
  69. if hasattr(term, "field"):
  70. expr = MultiplyExpression()
  71. expr.coeff = term
  72. term = expr
  73. if isinstance(term, MultiplyExpression):
  74. # Delete ^0 variables
  75. term.clean()
  76. # Skip zero terms
  77. #if term.coeff is None or term.coeff == 0:
  78. # return self
  79. result = self.copy()
  80. result_term = result.find(term)
  81. if result_term is None:
  82. result.terms.append(term)
  83. else:
  84. result_term.coeff += term.coeff
  85. return result
  86. else:
  87. assert isinstance(term, MultivariatePolynomial)
  88. result = self.copy()
  89. for other_term in term.terms:
  90. result += other_term
  91. return result
  92. def __mul__(self, other):
  93. if isinstance(term, Variable):
  94. term = term.termify()
  95. if hasattr(term, "field"):
  96. expr = MultiplyExpression()
  97. expr.coeff = term
  98. term = expr
  99. if isinstance(term, MultiplyExpression):
  100. # Delete ^0 variables
  101. term.clean()
  102. return None
  103. else:
  104. assert isinstance(term, MultivariatePolynomial)
  105. return None
  106. def find(self, other):
  107. for term in self.terms:
  108. if term.matches(other):
  109. return term
  110. return None
  111. def __str__(self):
  112. if not self.terms:
  113. return "0"
  114. repr = ""
  115. first = True
  116. for term in self.terms:
  117. if first:
  118. first = False
  119. else:
  120. repr += " + "
  121. repr += str(term)
  122. return repr
  123. if __name__ == "__main__":
  124. from finite_fields import finitefield
  125. p = 0x40000000000000000000000000000000224698fc094cf91b992d30ed00000001
  126. fp = finitefield.IntegersModP(p)
  127. x = Variable("X")
  128. y = Variable("Y")
  129. z = Variable("Z")
  130. p = x**3 * y**2 * x**2 * fp(5) * fp(2) + x**3 * y + z + fp(6)
  131. q = x**3 * y * fp(3) + y
  132. print(p)
  133. print(q)
  134. print(p + q)