Downloads · 30 days
0
pgg3/resnet18_mnist
resnet18_mnist is a machine learning model from pgg3. Use it for the machine learning task on the model card, and read the license before you ship it in a product. It is set up for timm.
Downloads · 30 days
0
Access
Public
Updated Jul 16, 2023
Repo size
44.8 MB
Likes
0
Public
Click a slice to open those files.
.pth44.8 MB · 100%
From the Hugging Face model README
import timm
import torchvision
MNIST_PATH = './datasets/mnist'
net = timm.create_model("resnet18", pretrained=False, num_classes=10)
net.conv1 = torch.nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)
net.load_state_dict(
torch.hub.load_state_dict_from_url(
"https://huggingface.co/gpcarl123/resnet18_mnist/resolve/main/resnet18_mnist.pth",
map_location="cpu",
file_name="resnet18_mnist.pth",
)
)
preprocessor = torchvision.transforms.Normalize((0.1307,), (0.3081,))
transform = transforms.Compose([transforms.ToTensor()])
test_set = datasets.MNIST(root=MNIST_PATH, train=False, download=True, transform=transform)
test_loader = data.DataLoader(test_set, batch_size=5, shuffle=False, num_workers=2)
for data, target in test_loader:
print(net(preprocessor(data)))
print(target)
break