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()