Views
No views yet

1import torch
2import torchvision.transforms as transforms
3from PIL import Image
4from safetensors.torch import load_model
5from huggingface_hub import hf_hub_download
6from timm import list_models, create_model
7import os
8import numpy as np
9
10# Download model from hub
11os.makedirs('/content/swin_s3_base_224', exist_ok=True)
12hf_hub_download(repo_id="LucyintheSky/lucy-feature-prediction", filename="model.safetensors", local_dir="/content/swin_s3_base_224")
13
14# Intialize the model
15model_name='swin_s3_base_224'
16model = create_model(
17 model_name,
18 num_classes=36
19)
20load_model(model,f'./{model_name}/model.safetensors')
21
22# Define class names
23class_names = ["3/4 Sleeve", "Accessory", "Babydoll", "Closed Back", "Corset", "Crochet", "Cutouts", "Draped", "Floral", "Gloves", "Halter", "Lace", "Long", "Long Sleeve", "Midi", "No Slit", "Off The Shoulder", "One Shoulder", "Open Back", "Pockets", "Print", "Puff Sleeve", "Ruched", "Satin", "Sequins", "Shimmer", "Short", "Short Sleeve", "Side Slit", "Square Neck", "Strapless", "Sweetheart Neck", "Tight", "V-Neck", "Velvet", "Wrap"]
24label2id = {c:idx for idx,c in enumerate(class_names)}
25id2label = {idx:c for idx,c in enumerate(class_names)}
26
27def predict_features(image_path):
28 # Load PIL image
29 pil_image = Image.open(image_path).convert('RGB')
30
31 # Define transformations to resize and convert image to tensor
32 transform = transforms.Compose([
33 transforms.Resize((224, 224)),
34 transforms.ToTensor()
35 ])
36 tensor_image = transform(pil_image)
37
38 inputs = tensor_image.unsqueeze(0)
39
40 with torch.no_grad():
41 logits = model(inputs)
42
43 # apply sigmoid activation to convert logits to probabilities
44 # getting labels with confidence threshold of 0.5
45 predictions = logits.sigmoid() > 0.5
46
47 # converting one-hot encoded predictions back to list of labels
48 predictions = predictions.float().numpy().flatten() # convert boolean predictions to float
49 pred_labels = np.where(predictions==1)[0] # find indices where prediction is 1
50 pred_labels = ([id2label[label] for label in pred_labels]) # converting integer labels to string
51
52 return pred_labels
53
54print(predict_features('image.jpg'))