# Vectorize the training step # params: axis 0 (each city has its own params) # x: axis 0 (each city has its own data) # y: axis 0 (each city has its own targets) # learning_rate: None (same for all) parallel_train_step = jax.jit( jax.vmap(train_step, in_axes=(0, 0, 0, None)) ) # Initialize parameters for all 3 cities # Shape: (3, 2) - 3 cities, 2 params each (slope, intercept) key, subkey = random.split(key) params = random.normal(subkey, (n_cities, 2)) print("Initial parameters:") for i in range(n_cities): print(f" City {i}: slope={params[i, 0]:.2f}, intercept={params[i, 1]:.2f}") __ __