class SyntheticRegressionData(d2l.DataModule):
"""Synthetic data for linear regression."""
def __init__(self, w, b, noise=0.01, num_train=1000, num_val=1000,
batch_size=32, key=None):
super().__init__()
self.save_hyperparameters()
# Resolve the key at call time rather than reusing a key in the signature.
key = jax.random.key(0) if key is None else key
n = num_train + num_val
key1, key2 = jax.random.split(key)
self.X = jax.random.normal(key1, (n, w.shape[0]))
eps = jax.random.normal(key2, (n, 1)) * noise
self.y = d2l.matmul(self.X, d2l.reshape(w, (-1, 1))) + b + eps