import jax import jax.numpy as jnp from jax import random # Seed for reproducibility key = random.PRNGKey(42) # True parameters for 3 cities # City 0: price = 50 * size + 100 # City 1: price = 30 * size + 200 # City 2: price = 80 * size + 50 true_slopes = jnp.array([50.0, 30.0, 80.0]) true_intercepts = jnp.array([100.0, 200.0, 50.0]) n_cities = 3 n_samples = 100 # Generate features (house sizes) for each city key, subkey = random.split(key) X = random.uniform(subkey, (n_cities, n_samples, 1), minval=10, maxval=100) # Generate targets (prices) with some noise key, subkey = random.split(key) noise = random.normal(subkey, (n_cities, n_samples)) * 50 # Y[i] = true_slopes[i] * X[i] + true_intercepts[i] + noise[i] Y = (X[:, :, 0] * true_slopes[:, None] + true_intercepts[:, None] + noise) print(f"X shape: {X.shape}") # (3, 100, 1) print(f"Y shape: {Y.shape}") # (3, 100) __ __