From 1e1f9644fe809c95eb032aef42d42146dac55bd0 Mon Sep 17 00:00:00 2001 From: wonder Date: Sun, 2 Nov 2025 16:59:13 +0800 Subject: [PATCH] =?UTF-8?q?=20=E2=9C=85Test:=20[Lab-2]=20=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=E9=9B=86=E4=B8=8B=E8=BD=BD=E6=B5=8B=E8=AF=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .gitignore | 3 ++- Lab-2/test/verify_env.py | 44 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 46 insertions(+), 1 deletion(-) create mode 100644 Lab-2/test/verify_env.py diff --git a/.gitignore b/.gitignore index 1719c51..06703ae 100644 --- a/.gitignore +++ b/.gitignore @@ -1,2 +1,3 @@ .idea -*/img \ No newline at end of file +*/img +*/data \ No newline at end of file diff --git a/Lab-2/test/verify_env.py b/Lab-2/test/verify_env.py new file mode 100644 index 0000000..d60086b --- /dev/null +++ b/Lab-2/test/verify_env.py @@ -0,0 +1,44 @@ +import unittest + +import torch +import torchvision +from torchvision import transforms +import matplotlib.pyplot as plt + + +class MyTestCase(unittest.TestCase): + def test_verify_env(self): + self.assertIsNotNone(torch.__version__) + print("== Env test pass ==\n" + f"[torch] {torch.__version__}\n") + def test_using_dataset(selfs): + transform = transforms.Compose([ + transforms.ToTensor(), + transforms.Normalize((0.5,), (0.5,)) + ]) + test_dataset = torchvision.datasets.MNIST( + root='../data', + train=False, + download=True, + transform=transform + ) + test_loader = torch.utils.data.DataLoader( + test_dataset, + batch_size=64, + shuffle=False + ) + images, labels = next(iter(test_loader)) + fig, axes = plt.subplots(2, 5, figsize=(10, 4)) + for i, ax in enumerate(axes.flat): + ax.imshow(images[i][0], cmap='gray') + ax.set_title(f'Label: {labels[i]}') + ax.axis('off') + + plt.tight_layout + plt.show() + + print(f"Successfully load MNIST dataset and show 10 images") + return images, labels + +if __name__ == '__main__': + unittest.main()