Views
No views yet
bert-base-uncased)microsoft/resnet-34)1import torch
2from torch import nn
3from transformers import AutoModel
4from huggingface_hub import hf_hub_download
5from typing import Literal
6import json
7
8class MultimodalClassifier(nn.Module):
9 def __init__(
10 self,
11 text_encoder_id_or_path: str,
12 image_encoder_id_or_path: str,
13 projection_dim: int,
14 fusion_method: Literal["concat", "align", "cosine_similarity"] = "concat",
15 proj_dropout: float = 0.1,
16 fusion_dropout: float = 0.1,
17 num_classes: int = 1,
18 ) -> None:
19 super().__init__()
20
21 self.fusion_method = fusion_method
22 self.projection_dim = projection_dim
23 self.num_classes = num_classes
24
25 ##### Text Encoder
26 self.text_encoder = AutoModel.from_pretrained(text_encoder_id_or_path)
27 self.text_projection = nn.Sequential(
28 nn.Linear(self.text_encoder.config.hidden_size, self.projection_dim),
29 nn.Dropout(proj_dropout),
30 )
31
32 ##### Image Encoder (using ResNet34 from AutoModel with timm)
33 self.image_encoder = AutoModel.from_pretrained(image_encoder_id_or_path, trust_remote_code=True)
34 self.image_encoder.classifier = nn.Identity() # rm the classification head
35 self.image_projection = nn.Sequential(
36 nn.Linear(512, self.projection_dim),
37 nn.Dropout(proj_dropout),
38 )
39
40 ##### Fusion Layer
41 fusion_input_dim = self.projection_dim * 2 if fusion_method == "concat" else self.projection_dim
42 self.fusion_layer = nn.Sequential(
43 nn.Dropout(fusion_dropout),
44 nn.Linear(fusion_input_dim, self.projection_dim),
45 nn.GELU(),
46 nn.Dropout(fusion_dropout),
47 )
48
49 ##### Classification Layer
50 self.classifier = nn.Linear(self.projection_dim, self.num_classes)
51
52 def forward(self, pixel_values: torch.Tensor, input_ids: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor:
53 ##### Text Encoder Projection #####
54 full_text_features = self.text_encoder(input_ids=input_ids, attention_mask=attention_mask, return_dict=True).last_hidden_state
55 full_text_features = full_text_features[:, 0, :] # using cls token
56 full_text_features = self.text_projection(full_text_features)
57
58 ##### Image Encoder Projection #####
59 resnet_image_features = self.image_encoder(pixel_values=pixel_values).last_hidden_state
60
61 # global average pooling for resent image features (bad idea? dim problems)
62 resnet_image_features = resnet_image_features.mean(dim=[-2, -1])
63 resnet_image_features = self.image_projection(resnet_image_features)
64
65 ##### Fusion and Classification #####
66 if self.fusion_method == "concat":
67 fused_features = torch.cat([full_text_features, resnet_image_features], dim=-1)
68 else:
69 fused_features = full_text_features * resnet_image_features # don't think this works atm (should be dot prod)
70
71 # fusion and classifier layers
72 fused_features = self.fusion_layer(fused_features)
73 classification_output = self.classifier(fused_features)
74
75 return classification_output
76
77def load_model():
78 config_path = hf_hub_download(repo_id="maximuspowers/multimodal-bias-classifier", filename="config.json")
79 with open(config_path, "r") as f:
80 config = json.load(f)
81
82 model = MultimodalClassifier(
83 text_encoder_id_or_path=config["text_encoder_id_or_path"],
84 image_encoder_id_or_path="microsoft/resnet-34",
85 projection_dim=config["projection_dim"],
86 fusion_method=config["fusion_method"],
87 proj_dropout=config["proj_dropout"],
88 fusion_dropout=config["fusion_dropout"],
89 num_classes=config["num_classes"]
90 )
91
92 model_weights_path = hf_hub_download(repo_id="maximuspowers/multimodal-bias-classifier", filename="model_weights.pth")
93 checkpoint = torch.load(model_weights_path, map_location=torch.device('cpu'))
94 model.load_state_dict(checkpoint, strict=False)
95
96 return model1import torch
2from transformers import AutoTokenizer
3from PIL import Image
4import requests
5from torchvision import transforms
6
7model = load_model()
8model.eval()
9
10# text input
11text_tokenizer = AutoTokenizer.from_pretrained("bert-base-uncased")
12sample_text = "This is a sample sentence for bias classification."
13text_inputs = text_tokenizer(
14 sample_text,
15 return_tensors="pt",
16 padding="max_length",
17 truncation=True,
18 max_length=512
19)
20
21# image input
22image = Image.open("./random_image.jpg").convert("RGB")
23image_transform = transforms.Compose([
24 transforms.Resize((224, 224)),
25 transforms.ToTensor(),
26 transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
27])
28image_input = image_transform(image).unsqueeze(0) # add batch dim
29
30# run
31with torch.no_grad():
32 classification_output = model(
33 pixel_values=image_input,
34 input_ids=text_inputs["input_ids"],
35 attention_mask=text_inputs["attention_mask"]
36 )
37 predicted_class = torch.sigmoid(classification_output).round().item()
38print("Predicted class:", "Biased" if predicted_class == 1 else "Unbiased")