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)