Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 0 additions & 4 deletions amigo/unary_operations.py
Original file line number Diff line number Diff line change
Expand Up @@ -73,10 +73,6 @@ def log(expr):
return Expr(UnaryNode("log", expr))


def log10(expr):
return Expr(UnaryNode("log10", expr))


def atan2(a, b):
if isinstance(a, (int, float)) and isinstance(b, (int, float)):
raise ValueError("Neither argument is active")
Expand Down
18 changes: 9 additions & 9 deletions include/component_group.h
Original file line number Diff line number Diff line change
Expand Up @@ -111,17 +111,17 @@ class SerialGroupBackend {
Vector<T>& res) const {
if constexpr (ncomp > 0) {
int num_elems = layout.get_num_elements();
for (int i = 0; i < num_elems; i++) {
for (int elem = 0; elem < num_elems; elem++) {
Data data;
Input input, gradient, direction, result;
data_layout.get_values(i, data_vec, data);
layout.get_values(i, vec, input);
layout.get_values(i, dir, direction);
data_layout.get_values(elem, data_vec, data);
layout.get_values(elem, vec, input);
layout.get_values(elem, dir, direction);
gradient.zero();
result.zero();
compute_hessian<T, Data, Input, Components...>(
alpha, data, input, direction, gradient, result);
layout.add_values(i, result, res);
layout.add_values(elem, result, res);
}
}
}
Expand All @@ -138,13 +138,13 @@ class SerialGroupBackend {
int num_elems = layout.get_num_elements();
Data data;
Input input, gradient, direction, result;
for (int i = 0; i < num_elems; i++) {
for (int elem = 0; elem < num_elems; elem++) {
int index[ncomp], index_global[ncomp];
layout.get_indices(i, index);
layout.get_indices(elem, index);
owners.local_to_global(ncomp, index, index_global);

data_layout.get_values(i, data_vec, data);
layout.get_values(i, vec, input);
data_layout.get_values(elem, data_vec, data);
layout.get_values(elem, vec, input);

for (int j = 0; j < ncomp; j++) {
direction.zero();
Expand Down
97 changes: 97 additions & 0 deletions tests/functional/expressions/test_expressions.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
import numpy as np
import amigo as am

unary_expressions = [
"sqrt",
"exp",
"log",
"sin",
"cos",
"tan",
"asin",
"acos",
"atan",
"sinh",
"cosh",
"tanh",
"asinh",
"acosh",
"atanh",
]


class Expressions(am.Component):
def __init__(self):
super().__init__()

value = 0.34
self.add_input("x", value=value, lower=-am.inf, upper=am.inf)
self.add_objective("obj")

delta = 0.1
for expr in unary_expressions:
if expr == "acosh":
fval = getattr(np, expr)(value + 1)
else:
fval = getattr(np, expr)(value)

lower = fval - delta
upper = fval + delta
self.add_constraint(f"{expr}_value", lower=lower, upper=upper)

for expr in unary_expressions:
fval = getattr(np, expr)(value)
self.add_output(f"{expr}_output")

def compute(self):
x = self.inputs["x"]
self.objective["obj"] = x**2

for expr in unary_expressions:
if expr == "acosh":
self.constraints[f"{expr}_value"] = getattr(am, expr)(x + 1)
else:
self.constraints[f"{expr}_value"] = getattr(am, expr)(x)

def compute_output(self):
x = self.inputs["x"]
for expr in unary_expressions:
if expr == "acosh":
self.outputs[f"{expr}_output"] = getattr(am, expr)(x + 1)
else:
self.outputs[f"{expr}_output"] = getattr(am, expr)(x)


def test_expressions():

model = am.Model("expr_test")
expr = Expressions()
model.add_component("expr", 1, expr)

model.build_module()
model.initialize()

x = model.create_vector()
opt = am.Optimizer(model, x)
opt.optimize()

g = model.create_vector()
model.eval_gradient(x, g)

output = model.create_output_vector()
model.compute_output(x, output)

tol = 1e-8
xval = x["expr.x"]
for expr in unary_expressions:
am_val = output[f"expr.{expr}_output"]
if expr == "acosh":
np_val = getattr(np, expr)(xval + 1)
else:
np_val = getattr(np, expr)(xval)

assert abs(am_val - np_val) < tol


if __name__ == "__main__":
test_expressions()
2 changes: 1 addition & 1 deletion tests/functional/interp/test_smt_quadratic.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ def compute(self):
"solver": solver,
"max_iterations": 200,
"initial_barrier_param": 1.0,
"convergence_tolerance": 1e-8,
"convergence_tolerance": 1e-5,
"max_line_search_iterations": 10,
"init_affine_step_multipliers": False,
"init_least_squares_multipliers": False,
Expand Down
Loading