import matplotlib.pyplot as plt
import jax.numpy as jnp
from differt2d.geometry import Wall

ax = plt.gca()
wall = Wall(xys=jnp.array([[0., 0.], [1., 0.]]))
_ = wall.plot(ax)
for vertex in wall.get_vertices():
    _ = vertex.plot(ax)
plt.show()  # doctest: +SKIP