Buckets:
| #!/usr/bin/env python3 | |
| """1D signal datasets for Neural Thickets / RandOpt (numpy + tinygrad Tensors).""" | |
| from __future__ import annotations | |
| from typing import Callable | |
| import numpy as np | |
| from tinygrad import Tensor | |
| FREQ = 4.0 | |
| SCALE = 1.0 | |
| def generate_sigmoid(): | |
| phase = np.random.uniform(0, 2 * np.pi) - np.pi | |
| amp = np.random.uniform(0.5, 1.5) | |
| y_offset = np.random.uniform(-0.5, 0.5) | |
| def fn(x): | |
| x = np.asarray(x) | |
| return amp * np.tanh(0.1 * x + phase) + y_offset | |
| return fn | |
| def generate_line(): | |
| slope = np.random.uniform(-0.5, 0.5) | |
| intercept = np.random.uniform(-1.0, 1.0) | |
| return lambda x: slope * x + intercept | |
| def generate_one_line(): | |
| return lambda x: -0.25 * np.asarray(x) | |
| def generate_harmonic(): | |
| phase = np.random.uniform(0, 2 * np.pi) | |
| amp = np.random.uniform(0.8, 1.2) | |
| y_offset = np.random.uniform(-0.5, 0.5) | |
| def fn(x): | |
| x = np.asarray(x) | |
| return amp * (0.5 * np.sin(FREQ * x + phase) + 0.3 * np.sin(2 * FREQ * x)) + y_offset | |
| return fn | |
| def generate_sinusoid(): | |
| phase = np.random.uniform(0, 2 * np.pi) | |
| amp = np.random.uniform(0.8, 1.2) | |
| y_offset = np.random.uniform(-0.5, 0.5) | |
| def fn(x): | |
| x = np.asarray(x) | |
| return amp * np.sin(FREQ * x + phase) + y_offset | |
| return fn | |
| def generate_one_sinusoid(): | |
| def fn(x): | |
| x = np.asarray(x) | |
| return 0.5 * np.sin(FREQ * x) | |
| return fn | |
| def generate_squarewave(): | |
| phase = np.random.uniform(0, 2 * np.pi) | |
| amp = np.random.uniform(0.2, 0.4) | |
| y_offset = np.random.uniform(-0.5, 0.5) | |
| sharpness = np.random.uniform(4.0, 6.0) | |
| def fn(x): | |
| x = np.asarray(x) | |
| return amp * np.tanh(sharpness * np.sin(FREQ * x + phase)) + y_offset | |
| return fn | |
| def generate_one_squarewave(): | |
| def fn(x): | |
| x = np.asarray(x) | |
| return 0.3 * np.tanh(5.0 * np.sin(FREQ * x)) | |
| return fn | |
| def generate_sawtooth(): | |
| phase = np.random.uniform(0, 2 * np.pi) | |
| amp = np.random.uniform(0.8, 1.2) | |
| y_offset = np.random.uniform(-0.5, 0.5) | |
| def fn(x): | |
| x = np.asarray(x) | |
| t = FREQ * x + phase | |
| saw = np.sin(t) - 0.5 * np.sin(2 * t) + 0.33 * np.sin(3 * t) - 0.25 * np.sin(4 * t) | |
| return amp * saw * 0.5 + y_offset | |
| return fn | |
| def generate_mixed(): | |
| generators = [ | |
| generate_sinusoid, | |
| generate_squarewave, | |
| generate_sawtooth, | |
| generate_harmonic, | |
| generate_sigmoid, | |
| generate_line, | |
| ] | |
| return np.random.choice(generators)() | |
| DATASET_GENERATORS: dict[str, Callable] = { | |
| "line": generate_line, | |
| "one_line": generate_one_line, | |
| "sigmoid": generate_sigmoid, | |
| "harmonic": generate_harmonic, | |
| "sinusoid": generate_sinusoid, | |
| "one_sinusoid": generate_one_sinusoid, | |
| "squarewave": generate_squarewave, | |
| "one_squarewave": generate_one_squarewave, | |
| "sawtooth": generate_sawtooth, | |
| "mixed": generate_mixed, | |
| } | |
| def load_data(bsz: int, dataset: str, args) -> tuple[Tensor, Tensor, Tensor, Tensor]: | |
| if dataset not in DATASET_GENERATORS: | |
| raise ValueError(f"Dataset {dataset} not supported") | |
| generator = DATASET_GENERATORS[dataset] | |
| ctx_x_list, ctx_y_list, fut_x_list, fut_y_list = [], [], [], [] | |
| for _ in range(bsz): | |
| gt_fn = generator() | |
| start_x = -2.5 | |
| x_vals = start_x + np.arange(args.ctx_sz + args.fut_sz) * args.res_x | |
| y_vals = [float(gt_fn(x)) for x in x_vals] | |
| ctx_x_list.append(x_vals[: args.ctx_sz]) | |
| ctx_y_list.append(y_vals[: args.ctx_sz]) | |
| fut_x_list.append(x_vals[args.ctx_sz :]) | |
| fut_y_list.append(y_vals[args.ctx_sz :]) | |
| ctx_x = Tensor(np.asarray(ctx_x_list, dtype=np.float32) * SCALE) | |
| ctx_y = Tensor(np.asarray(ctx_y_list, dtype=np.float32) * SCALE) | |
| fut_x = Tensor(np.asarray(fut_x_list, dtype=np.float32) * SCALE) | |
| fut_y = Tensor(np.asarray(fut_y_list, dtype=np.float32) * SCALE) | |
| return ctx_x, ctx_y, fut_x, fut_y | |
Xet Storage Details
- Size:
- 3.96 kB
- Xet hash:
- 958e4a465d227d8340321f2763610a1cf74e719886c798c37cfe98bc7a64ec77
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.