File size: 720 Bytes
ddaaeb2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 | from __future__ import annotations
import numpy as np
def generate_adding_problem(
samples: int, length: int, seed: int
) -> tuple[np.ndarray, np.ndarray]:
rng = np.random.default_rng(seed)
values = rng.random((samples, length), dtype=np.float32)
markers = np.zeros((samples, length), dtype=np.float32)
midpoint = length // 2
first = rng.integers(0, midpoint, size=samples)
second = rng.integers(midpoint, length, size=samples)
rows = np.arange(samples)
markers[rows, first] = 1
markers[rows, second] = 1
inputs = np.stack([values, markers], axis=-1)
targets = values[rows, first] + values[rows, second]
return inputs.astype(np.float32), targets.astype(np.float32)
|