mirror of
https://github.com/zebrajr/pytorch.git
synced 2025-12-06 12:20:52 +01:00
Test Plan: revert-hammer
Differential Revision:
D30279364 (b004307252)
Original commit changeset: c1ed77dfe43a
fbshipit-source-id: eab50857675c51e0088391af06ec0ecb14e2347e
17 lines
629 B
Python
17 lines
629 B
Python
import torch
|
|
import torchvision
|
|
|
|
print(torch.version.__version__)
|
|
|
|
resnet18 = torchvision.models.resnet18(pretrained=True)
|
|
resnet18.eval()
|
|
resnet18_traced = torch.jit.trace(resnet18, torch.rand(1, 3, 224, 224)).save("app/src/main/assets/resnet18.pt")
|
|
|
|
resnet50 = torchvision.models.resnet50(pretrained=True)
|
|
resnet50.eval()
|
|
torch.jit.trace(resnet50, torch.rand(1, 3, 224, 224)).save("app/src/main/assets/resnet50.pt")
|
|
|
|
mobilenet2q = torchvision.models.quantization.mobilenet_v2(pretrained=True, quantize=True)
|
|
mobilenet2q.eval()
|
|
torch.jit.trace(mobilenet2q, torch.rand(1, 3, 224, 224)).save("app/src/main/assets/mobilenet2q.pt")
|