✅Test: [Lab-2] 测试集下载测试

This commit is contained in:
2025-11-02 16:59:13 +08:00
parent f2aa426d7f
commit 1e1f9644fe
2 changed files with 46 additions and 1 deletions
+2 -1
View File
@@ -1,2 +1,3 @@
.idea .idea
*/img */img
*/data
+44
View File
@@ -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()