Views
No views yet
pip install mambavision
1from transformers import AutoModelForImageClassification
2from PIL import Image
3from timm.data.transforms_factory import create_transform
4import requests
5
6model = AutoModelForImageClassification.from_pretrained("nvidia/MambaVision-T2-1K", trust_remote_code=True)
7
8# eval mode for inference
9model.cuda().eval()
10
11# prepare image for the model
12url = 'http://images.cocodataset.org/val2017/000000020247.jpg'
13image = Image.open(requests.get(url, stream=True).raw)
14input_resolution = (3, 224, 224) # MambaVision supports any input resolutions
15
16transform = create_transform(input_size=input_resolution,
17 is_training=False,
18 mean=model.config.mean,
19 std=model.config.std,
20 crop_mode=model.config.crop_mode,
21 crop_pct=model.config.crop_pct)
22
23inputs = transform(image).unsqueeze(0).cuda()
24# model inference
25outputs = model(inputs)
26logits = outputs['logits']
27predicted_class_idx = logits.argmax(-1).item()
28print("Predicted class:", model.config.id2label[predicted_class_idx])brown bear, bruin, Ursus arctos.1from transformers import AutoModel
2from PIL import Image
3from timm.data.transforms_factory import create_transform
4import requests
5
6model = AutoModel.from_pretrained("nvidia/MambaVision-T2-1K", trust_remote_code=True)
7
8# eval mode for inference
9model.cuda().eval()
10
11# prepare image for the model
12url = 'http://images.cocodataset.org/val2017/000000020247.jpg'
13image = Image.open(requests.get(url, stream=True).raw)
14input_resolution = (3, 224, 224) # MambaVision supports any input resolutions
15
16transform = create_transform(input_size=input_resolution,
17 is_training=False,
18 mean=model.config.mean,
19 std=model.config.std,
20 crop_mode=model.config.crop_mode,
21 crop_pct=model.config.crop_pct)
22inputs = transform(image).unsqueeze(0).cuda()
23# model inference
24out_avg_pool, features = model(inputs)
25print("Size of the averaged pool features:", out_avg_pool.size()) # torch.Size([1, 640])
26print("Number of stages in extracted features:", len(features)) # 4 stages
27print("Size of extracted features in stage 1:", features[0].size()) # torch.Size([1, 80, 56, 56])
28print("Size of extracted features in stage 4:", features[3].size()) # torch.Size([1, 640, 7, 7])