Views
No views yet
1from urllib.request import urlopen
2from PIL import Image
3import jax
4
5import jaxnn
6
7img = Image.open(urlopen(
8 'https://huggingface.co/datasets/huggingface/cats-image/resolve/main/cats_image.jpeg'
9))
10
11model = jaxnn.create_model('resnet50.fb_swsl_ig1b_ft_in1k', pretrained=True)
12model.eval()
13
14# get model specific transforms (normalization, resize)
15data_config = jaxnn.data.resolve_model_data_config(model)
16transforms = jaxnn.data.create_transform(**data_config, is_training=False)
17
18output = model(jax.numpy.expand_dims(transforms(img), 0)) # unsqueeze single image into batch of 1
19
20top5_probabilities, top5_class_indices = jax.lax.top_k(jax.nn.softmax(output, axis=-1) * 100, k=5)
211from urllib.request import urlopen
2from PIL import Image
3import jax
4
5import jaxnn
6
7img = Image.open(urlopen(
8 'https://huggingface.co/datasets/huggingface/cats-image/resolve/main/cats_image.jpeg'
9))
10
11model = jaxnn.create_model(
12 'resnet50.fb_swsl_ig1b_ft_in1k',
13 pretrained=True,
14 features_only=True,
15)
16model.eval()
17
18# get model specific transforms (normalization, resize)
19data_config = jaxnn.data.resolve_model_data_config(model)
20transforms = jaxnn.data.create_transform(**data_config, is_training=False)
21
22output = model(jax.numpy.expand_dims(transforms(img), 0)) # jax.numpy.expand_dims single image into batch of 1
23
24for o in output:
25 # print shape of each feature map in output in format [Batch, Height, Width, Channels]
26 # e.g.:
27 # (1, 112, 112, 64)
28 # (1, 56, 56, 64)
29 # (1, 28, 28, 128)
30 # (1, 14, 14, 256)
31 # (1, 7, 7, 512)
32
33 print(o.shape)1from urllib.request import urlopen
2from PIL import Image
3import jax
4
5import jaxnn
6
7img = Image.open(urlopen(
8 'https://huggingface.co/datasets/huggingface/cats-image/resolve/main/cats_image.jpeg'
9))
10
11model = jaxnn.create_model(
12 'resnet50.fb_swsl_ig1b_ft_in1k',
13 pretrained=True,
14 num_classes=0, # remove classifier nn.Linear
15)
16model.eval()
17
18# get model specific transforms (normalization, resize)
19data_config = jaxnn.data.resolve_model_data_config(model)
20transforms = jaxnn.data.create_transform(**data_config, is_training=False)
21
22output = model(jax.numpy.expand_dims(transforms(img), 0)) # output is (batch_size, num_features) shaped Array
23
24# or equivalently (without needing to set num_classes=0)
25
26output = model.forward_features(jax.numpy.expand_dims(transforms(img), 0))
27# output is unpooled, a (1, 7, 7, 512) shaped tensor
28