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

ax = plt.gca()
ray = Ray(xys=jnp.array([[0., 0.], [1., 1.]]))
_ = ray.plot(ax)
plt.show()  # doctest: +SKIP