mirror of
https://github.com/pytorch/pytorch.git
synced 2025-10-20 21:14:14 +08:00
See https://github.com/pytorch/pytorch/pull/129751#issue-2380881501. Most changes are auto-generated by linter. You can review these PRs via: ```bash git diff --ignore-all-space --ignore-blank-lines HEAD~1 ``` Pull Request resolved: https://github.com/pytorch/pytorch/pull/129752 Approved by: https://github.com/ezyang, https://github.com/malfet
25 lines
685 B
Python
25 lines
685 B
Python
from torchvision import models
|
|
|
|
import torch
|
|
|
|
|
|
print(torch.version.__version__)
|
|
|
|
resnet18 = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1)
|
|
resnet18.eval()
|
|
resnet18_traced = torch.jit.trace(resnet18, torch.rand(1, 3, 224, 224)).save(
|
|
"app/src/main/assets/resnet18.pt"
|
|
)
|
|
|
|
resnet50 = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V1)
|
|
resnet50.eval()
|
|
torch.jit.trace(resnet50, torch.rand(1, 3, 224, 224)).save(
|
|
"app/src/main/assets/resnet50.pt"
|
|
)
|
|
|
|
mobilenet2q = 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"
|
|
)
|