(image_size=224, normalization="imagenet")
| 71 | |
| 72 | |
| 73 | def get_dog_image_tensor(image_size=224, normalization="imagenet"): |
| 74 | url, filename = ( |
| 75 | "https://github.com/pytorch/hub/raw/master/images/dog.jpg", |
| 76 | "dog.jpg", |
| 77 | ) |
| 78 | try: |
| 79 | urllib.URLopener().retrieve(url, filename) |
| 80 | except: |
| 81 | urllib.request.urlretrieve(url, filename) |
| 82 | |
| 83 | from PIL import Image |
| 84 | from torchvision import transforms |
| 85 | |
| 86 | input_image = Image.open(filename).convert("RGB") |
| 87 | |
| 88 | transforms_list = [ |
| 89 | transforms.Resize((image_size, image_size)), |
| 90 | transforms.ToTensor(), |
| 91 | ] |
| 92 | if normalization == "imagenet": |
| 93 | transforms_list.append( |
| 94 | transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]), |
| 95 | ) |
| 96 | |
| 97 | preprocess = transforms.Compose(transforms_list) |
| 98 | |
| 99 | input_tensor = preprocess(input_image) |
| 100 | input_batch = input_tensor.unsqueeze(0) |
| 101 | input_batch = (input_batch,) |
| 102 | return input_batch |
| 103 | |
| 104 | |
| 105 | def init_model(model_name): |
no test coverage detected