def predict(params, x): """Predict price for one city's houses.""" slope, intercept = params[0], params[1] return x[:, 0] * slope + intercept def loss_fn(params, x, y): """MSE loss for one city.""" predictions = predict(params, x) return jnp.mean((predictions - y) ** 2) def train_step(params, x, y, learning_rate): """One gradient descent step for one city.""" loss, grads = jax.value_and_grad(loss_fn)(params, x, y) new_params = params - learning_rate * grads return new_params, loss __ __