Views
No views yet
pip install mlx-image1import mlx.core as mx
2from mlxim.model import create_model
3from mlxim.io import read_rgb
4from mlxim.transform import ImageNetTransform
5from mlxim.utils.imagenet import IMAGENET2012_CLASSES
6
7transform = ImageNetTransform(train=False, img_size=380)
8x = transform(read_rgb("cat.jpg"))
9x = mx.array(x)
10x = mx.expand_dims(x, 0)
11
12model = create_model("efficientnet_b4")
13model.eval()
14
15logits = model(x)
16predicted_idx = mx.argmax(logits, axis=-1).item()
17predicted_class = list(IMAGENET2012_CLASSES.values())[predicted_idx]
18
19print(f"Predicted class: {predicted_class}")1import mlx.core as mx
2from mlxim.model import create_model
3from mlxim.io import read_rgb
4from mlxim.transform import ImageNetTransform
5
6transform = ImageNetTransform(train=False, img_size=380)
7x = transform(read_rgb("cat.jpg"))
8x = mx.array(x)
9x = mx.expand_dims(x, 0)
10
11# first option
12model = create_model("efficientnet_b4", num_classes=0)
13model.eval()
14
15embeds = model(x)
16
17# second option
18model = create_model("efficientnet_b4")
19model.eval()
20
21embeds = model.get_features(x)