Views
No views yet


1from PIL import Image
2import torch
3import torch.nn as nn
4import torch.nn.functional as F
5import torchvision.transforms as v2
6from transformers import AutoImageProcessor, SwinModel, SwinConfig
7from huggingface_hub import PyTorchModelHubMixin
8
9device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu')
10
11ckpt = "yainage90/fashion-image-feature-extractor"
12encoder_config = SwinConfig.from_pretrained(ckpt)
13encoder_image_processor = AutoImageProcessor.from_pretrained(ckpt)
14
15class ImageEncoder(nn.Module, PyTorchModelHubMixin):
16 def __init__(self):
17 super(ImageEncoder, self).__init__()
18 self.swin = SwinModel(config=encoder_config)
19 self.embedding_layer = nn.Linear(encoder_config.hidden_size, 128)
20
21 def forward(self, image_tensor):
22 features = self.swin(image_tensor).pooler_output
23 embeddings = self.embedding_layer(features)
24 embeddings = F.normalize(embeddings, p=2, dim=1)
25
26 return embeddings
27
28encoder = ImageEncoder().from_pretrained('yainage90/fashion-image-feature-extractor').to(device)
29
30transform = v2.Compose([
31 v2.Resize((encoder_config.image_size, encoder_config.image_size)),
32 v2.ToTensor(),
33 v2.Normalize(mean=encoder_image_processor.image_mean, std=encoder_image_processor.image_std),
34])
35
36image = Image.open('<path/to/image>').convert('RGB')
37image = transform(image)
38with torch.no_grad():
39 embedding = encoder(image.unsqueeze(0).to(device)).cpu().numpy()


