✅Test: [Lab-2] 测试集下载测试
This commit is contained in:
@@ -1,2 +1,3 @@
|
||||
.idea
|
||||
*/img
|
||||
*/data
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user