from math import isclose


def nonlinear(x: float) -> float:
    return x * x


weights = [0.5, 0.5]
embeddings = [-1.0, 1.0]
mixed_input = sum(p * e for p, e in zip(weights, embeddings))
one_forward = nonlinear(mixed_input)
branch_average = sum(
    p * nonlinear(e) for p, e in zip(weights, embeddings)
)

assert isclose(one_forward, 0.0)
assert isclose(branch_average, 1.0)
print(one_forward, branch_average)  # 0.0 1.0
