multipoly.py 3.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124
  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 termify(self):
  9. expr = MultiplyExpression()
  10. expr.set_symbol(self.name, 1)
  11. return expr
  12. class MultiplyExpression:
  13. def __init__(self):
  14. self.coeff = None
  15. self.symbols = {}
  16. def clean(self):
  17. for symbol in list(self.symbols.keys()):
  18. if self.symbols[symbol] == 0:
  19. del self.symbols[symbol]
  20. def set_symbol(self, var_name, power):
  21. self.symbols[var_name] = power
  22. def __mul__(self, expr):
  23. result = MultiplyExpression()
  24. result.coeff = self.coeff
  25. result.symbols = self.symbols.copy()
  26. if hasattr(expr, "field"):
  27. if result.coeff is None:
  28. result.coeff = expr
  29. else:
  30. result.coeff *= expr
  31. return result
  32. if isinstance(expr, Variable):
  33. expr = expr.termify()
  34. for var_name, power in expr.symbols.items():
  35. if var_name in result.symbols:
  36. result.symbols[var_name] += power
  37. else:
  38. result.symbols[var_name] = power
  39. return result
  40. def __add__(self, expr):
  41. return MultivariatePolynomial([self, expr])
  42. def __str__(self):
  43. repr = ""
  44. first = True
  45. if self.coeff is not None:
  46. repr += str(self.coeff)
  47. first = False
  48. for var_name, power in self.symbols.items():
  49. if first:
  50. first = False
  51. else:
  52. repr += " "
  53. if power == 1:
  54. repr += var_name
  55. else:
  56. repr += var_name + "^" + str(power)
  57. return repr
  58. class MultivariatePolynomial:
  59. def __init__(self, terms=[]):
  60. self.terms = terms
  61. def __add__(self, term):
  62. if isinstance(term, Variable):
  63. term = term.termify()
  64. if hasattr(term, "field"):
  65. expr = MultiplyExpression()
  66. expr.coeff = term
  67. term = expr
  68. # Delete ^0 variables
  69. term.clean()
  70. # Skip zero terms
  71. if term.coeff is None or term.coeff == 0:
  72. return self
  73. result = MultivariatePolynomial(self.terms[:])
  74. result.terms.append(term)
  75. return result
  76. def __str__(self):
  77. if not self.terms:
  78. return "0"
  79. repr = ""
  80. first = True
  81. for term in self.terms:
  82. if first:
  83. first = False
  84. else:
  85. repr += " + "
  86. repr += str(term)
  87. return repr
  88. if __name__ == "__main__":
  89. from finite_fields import finitefield
  90. p = 0x40000000000000000000000000000000224698fc094cf91b992d30ed00000001
  91. fp = finitefield.IntegersModP(p)
  92. x = Variable("X")
  93. y = Variable("Y")
  94. z = Variable("Z")
  95. p = x**3 * y**2 * x**2 * fp(5) * fp(2) + x**3 * y + z + fp(6)
  96. print(p)