From 1bf14f57f7e5cba48a00a6ae18ce6a6464fad7a4 Mon Sep 17 00:00:00 2001 From: Didictateur Date: Mon, 16 Feb 2026 15:26:31 +0100 Subject: [PATCH] optimisation evaluate --- legrandchien.py | 90 ++++++++++++++++++++++++++++++++----------------- 1 file changed, 59 insertions(+), 31 deletions(-) diff --git a/legrandchien.py b/legrandchien.py index b116d70..d355322 100644 --- a/legrandchien.py +++ b/legrandchien.py @@ -414,41 +414,59 @@ class Equation: if self.compiled_function is not None: return self.compiled_function - if self.op == Operation.CONST: - self.compiled_function = lambda dico : self.value.value - if self.op == Operation.VAR: - self.compiled_function = lambda dico : dico[self.value.name] - if self.op == Operation.ADD: - left_func = self.left.compile() - right_func = self.right.compile() - self.compiled_function = lambda dico : left_func(dico) + right_func(dico) - if self.op == Operation.SUB: - left_func = self.left.compile() - right_func = self.right.compile() - self.compiled_function = lambda dico : left_func(dico) - right_func(dico) - if self.op == Operation.MULT: - left_func = self.left.compile() - right_func = self.right.compile() - self.compiled_function = lambda dico : left_func(dico) * right_func(dico) - if self.op == Operation.DIV: - left_func = self.left.compile() - right_func = self.right.compile() - self.compiled_function = lambda dico : left_func(dico) / right_func(dico) - if self.op == Operation.SIN: - value_func = self.value.compile() - self.compiled_function = lambda dico : np.sin(value_func(dico)) - if self.op == Operation.COS: - value_func = self.value.compile() - self.compiled_function = lambda dico : np.cos(value_func(dico)) - + code = self.generate_code() + self.compiled_function = eval(code) return self.compiled_function + + def generate_code(self): + code, var_count = self._generate_code_recursive(0) + return f"lambda dico: {code}" + + def _generate_code_recursive(self, var_count): + if self.op == Operation.CONST: + val = self.value.value + return (str(val), var_count) + + elif self.op == Operation.VAR: + val = self.value.name + return (f"dico['{val}']", var_count) + + elif self.op == Operation.ADD: + left_code, var_count = self.left._generate_code_recursive(var_count) + right_code, var_count = self.right._generate_code_recursive(var_count) + return (f"({left_code}) + ({right_code})", var_count) + + elif self.op == Operation.SUB: + left_code, var_count = self.left._generate_code_recursive(var_count) + right_code, var_count = self.right._generate_code_recursive(var_count) + return (f"({left_code}) - ({right_code})", var_count) + + elif self.op == Operation.MULT: + left_code, var_count = self.left._generate_code_recursive(var_count) + right_code, var_count = self.right._generate_code_recursive(var_count) + return (f"({left_code}) * ({right_code})", var_count) + + elif self.op == Operation.DIV: + left_code, var_count = self.left._generate_code_recursive(var_count) + right_code, var_count = self.right._generate_code_recursive(var_count) + return (f"({left_code}) / ({right_code})", var_count) + + elif self.op == Operation.SIN: + val_code, var_count = self.value._generate_code_recursive(var_count) + return (f"np.sin({val_code})", var_count) + + elif self.op == Operation.COS: + val_code, var_count = self.value._generate_code_recursive(var_count) + return (f"np.cos({val_code})", var_count) + + return ("0", var_count) def evaluate(self, dico): if self.compiled_function is None: self.compile() return self.compiled_function(dico) - def solve(self, dico, tmax=10, dt=0.01): + def solve(self, dico, tmax=10, dt=0.01, progress_bar=True): self.simplify() variables = self.getAllVar() @@ -478,11 +496,21 @@ class Equation: equations = list(equations) J_template = np.empty((n, n), dtype=object) + F_template = np.array((n, n), dtype=object) for i in range(n): + f = equations[i] + f.compile() + F_template[i] = f for j in range(n): - J_template[i][j] = equations[i].partial(unknown[j].name) + eq = equations[i].partial(unknown[j].name) + eq.compile() + J_template[i][j] = eq - for t in tqdm([n * dt for n in range(int(tmax/dt))]): + if progress_bar: + bar = tqdm([n * dt for n in range(int(tmax/dt))]) + else: + bar = [n * dt for n in range(int(tmax/dt))] + for t in bar: err = 1 iterations = 0 max_iteration = 50 @@ -493,7 +521,7 @@ class Equation: J = np.zeros((n, n)) F = np.zeros(n) for i in range(n): - F[i] = equations[i].evaluate(dico) + F[i] = F_template[i].evaluate(dico) for j in range(n): # print(equations[0]) J[i][j] = J_template[i][j].evaluate(dico)