import jax.numpy as jnp
import matplotlib.pyplot as plt

from differt2d.geometry import Wall
from differt2d.scene import Scene
from differt2d.utils import received_power

ax = plt.gca()
scene = Scene.square_scene()
wall = Wall(xys=jnp.array([[.8, .2], [.8, .8]]))
scene = scene.add_objects(wall)
scene.plot(ax, receivers=True)

X, Y = scene.grid(300)
Z = scene.accumulate_on_receivers_grid_over_paths(
    X,
    Y,
    fun=received_power,
    reduce_all=True,
    max_order=2  # The default value was 1
)
ax.pcolormesh(X, Y, 10.0 * jnp.log10(Z), zorder=-1)
plt.show()