import matplotlib.pyplot as plt # Reshape predictions to grid format for visualization Z_pred = model.predict(inputs) * outputs_std + outputs_mean Z_pred = Z_pred.reshape(X.shape) # Plot comparison of the true function and the model predictions fig, axes = plt.subplots(1, 2, figsize=(14, 6)) # Plot the true function axes[0].contourf(X, Y, Z, cmap='viridis') axes[0].set_title("True Function") axes[0].set_xlabel("X-axis") axes[0].set_ylabel("Y-axis") axes[0].colorbar = plt.colorbar(axes[0].contourf(X, Y, Z, cmap='viridis'), ax=axes[0], label="Function Value") # Plot the predicted function axes[1].contourf(X, Y, Z_pred, cmap='plasma') axes[1].set_title("NN Predicted Function") axes[1].set_xlabel("X-axis") axes[1].set_ylabel("Y-axis") axes[1].colorbar = plt.colorbar(axes[1].contourf(X, Y, Z_pred, cmap='plasma'), ax=axes[1], label="Function Value") plt.tight_layout() plt.show()