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

ax = plt.gca()
p1 = Point(xy=jnp.array([0., 0.]))
_ = p1.plot(ax)
p2 = Point(xy=jnp.array([1., 1.]))
_ = p2.plot(ax, color="b", annotate="$p_2$")
plt.show()  # doctest: +SKIP