From 7b6b1f32b104a4923013fa983787ec0247af5a8c Mon Sep 17 00:00:00 2001 From: AllentDan <41138331+AllentDan@users.noreply.github.com> Date: Mon, 30 Jan 2023 13:30:47 +0800 Subject: [PATCH] [Fix] fix typo: test_mnist -> datasets (#492) * test_mnist -> datasets * fix mnist_gan --- examples/mnist_gan.py | 5 +++-- extra/augment.py | 2 +- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/examples/mnist_gan.py b/examples/mnist_gan.py index 7e55f84aeb..a58d103bdc 100644 --- a/examples/mnist_gan.py +++ b/examples/mnist_gan.py @@ -6,10 +6,10 @@ from tqdm import tqdm sys.path.append(os.getcwd()) sys.path.append(os.path.join(os.getcwd(), 'test')) -from tinygrad.tensor import Tensor, Function, register +from tinygrad.tensor import Tensor from extra.utils import get_parameters import tinygrad.nn.optim as optim -from test_mnist import X_train +from datasets import fetch_mnist from torchvision.utils import make_grid, save_image import torch GPU = os.getenv("GPU") is not None @@ -61,6 +61,7 @@ if __name__ == "__main__": disc_loss = [] output_folder = "outputs" os.makedirs(output_folder, exist_ok=True) + X_train = fetch_mnist()[0] train_data_size = len(X_train) ds_noise = Tensor(np.random.randn(64,128).astype(np.float32), requires_grad=False) n_steps = int(train_data_size/batch_size) diff --git a/extra/augment.py b/extra/augment.py index 02f3e55852..1be046d38d 100644 --- a/extra/augment.py +++ b/extra/augment.py @@ -4,7 +4,7 @@ import os import sys sys.path.append(os.getcwd()) sys.path.append(os.path.join(os.getcwd(), 'test')) -from test_mnist import fetch_mnist +from datasets import fetch_mnist from tqdm import trange def augment_img(X, rotate=10, px=3):