✅Test: [Lab-2] 测试集下载测试
This commit is contained in:
+2
-1
@@ -1,2 +1,3 @@
|
|||||||
.idea
|
.idea
|
||||||
*/img
|
*/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