# Quick visualisation (run after training)
def plot_decision_boundary(model, xlim=(-2,2), ylim=(-2,2), res=100):
    x = torch.linspace(xlim[0], xlim[1], res)
    y = torch.linspace(ylim[0], ylim[1], res)
    xx, yy = torch.meshgrid(x, y)
    grid = torch.stack([xx.flatten(), yy.flatten()], dim=1).to(device)
    with torch.no_grad():
        # Integrate from each grid point
        zT = integrate(grid, model.attractors, model.steps, model.dt)
        diff = zT[:, None, :] - model.attractors[None, :, :]
        dist_sq = torch.sum(diff**2, dim=-1)
        pred = torch.argmin(dist_sq, dim=1)   # closest attractor
    pred = pred.cpu().numpy().reshape(res, res)
    plt.imshow(pred, extent=(xlim[0], xlim[1], ylim[0], ylim[1]), origin='lower', cmap='tab10')
    plt.scatter(attractors[:,0], attractors[:,1], c='black', marker='*', s=100)
    plt.title('Decision boundary (basins)')
    plt.show()

plot_decision_boundary(model)