Views
No views yet
GraphPropagationLayer and ImageCaptioningModel (or import them) before loading the weights.1pip install torch torchvision transformers pillow huggingface_hub
2
3### Inference Script
4python
5import torch
6import torch.nn as nn
7from torchvision import transforms, models
8from torchvision.models.feature_extraction import create_feature_extractor
9from transformers import BertTokenizer
10from huggingface_hub import hf_hub_download
11from PIL import Image
12
13# 1. Define the Model Architecture (Must match training)
14class GraphPropagationLayer(nn.Module):
15def __init__(self, embed_dim, window_size=3, max_steps=3):
16super().__init__()
17self.window_size = window_size
18self.max_steps = max_steps
19self.msg_linear = nn.Linear(embed_dim, embed_dim)
20self.activation = nn.Hardtanh()
21
22def forward(self, x):
23B, L, D = x.shape
24pad = self.window_size // 2
25for _ in range(self.max_steps):
26neighbors = []
27for i in range(-pad, pad + 1):
28neighbors.append(torch.roll(x, shifts=i, dims=1))
29
30stacked = torch.stack(neighbors, dim=2)
31aggregated = stacked.mean(dim=2)
32messages = self.msg_linear(aggregated)
33x = self.activation(x + messages)
34return x
35
36class ImageCaptioningModel(nn.Module):
37def __init__(self, vocab_size, embed_dim=512, num_layers=5, num_heads=8):
38super().__init__()
39
40# Encoder
41mobilenet = models.mobilenet_v3_small(weights=models.MobileNet_V3_Small_Weights.IMAGENET1K_V1)
42self.encoder = create_feature_extractor(mobilenet, return_nodes={'features': 'features'})
43self.vis_proj = nn.Linear(576, embed_dim)
44
45# Graph
46self.graph_layer = GraphPropagationLayer(embed_dim, window_size=3, max_steps=3)
47
48# Decoder
49self.embedding = nn.Embedding(vocab_size, embed_dim)
50self.pos_encoder = nn.Parameter(torch.randn(1, 2000, embed_dim))
51
52decoder_layer = nn.TransformerDecoderLayer(d_model=embed_dim, nhead=num_heads, dim_feedforward=1024, batch_first=True)
53self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_layers)
54self.fc_out = nn.Linear(embed_dim, vocab_size)
55
56def encode_images(self, images):
57visual_features = self.encoder(images)['features']
58B, C, H, W = visual_features.shape
59visual_tokens = visual_features.view(B, C, -1).permute(0, 2, 1)
60visual_tokens = self.vis_proj(visual_tokens)
61memory = self.graph_layer(visual_tokens)
62return memory
63
64@torch.no_grad()
65def generate(self, images, max_len=80, tokenizer=None):
66self.eval()
67memory = self.encode_images(images)
68batch_size = images.shape[0]
69
70generated = torch.full((batch_size, 1), tokenizer.cls_token_id, dtype=torch.long, device=images.device)
71
72
73
74## Usage
75
76This model uses a custom PyTorch architecture. Therefore, you need to instantiate the `GraphPropagationLayer` and `ImageCaptioningModel` classes to load the weights.
77
78### Installation
79
80First, install the required Python packages:
81
82
83
84```bash
85pip install torch torchvision transformers pillow huggingface_hub
86
87### Inference in Python
88
89Here is a complete script to load the model and generate a caption for a medical image.
90
91
92
93
94python
95import torch
96import torch.nn as nn
97from torchvision import transforms, models
98from torchvision.models.feature_extraction import create_feature_extractor
99from transformers import BertTokenizer
100from huggingface_hub import hf_hub_download
101from PIL import Image
102
103# --- 1. Define the Model Architecture ---
104# These classes must match the training script exactly.
105
106class GraphPropagationLayer(nn.Module):
107def __init__(self, embed_dim, window_size=3, max_steps=3):
108super().__init__()
109self.window_size = window_size
110self.max_steps = max_steps
111self.msg_linear = nn.Linear(embed_dim, embed_dim)
112self.activation = nn.Hardtanh()
113
114def forward(self, x):
115B, L, D = x.shape
116pad = self.window_size // 2
117for _ in range(self.max_steps):
118neighbors = []
119for i in range(-pad, pad + 1):
120neighbors.append(torch.roll(x, shifts=i, dims=1))
121stacked = torch.stack(neighbors, dim=2)
122aggregated = stacked.mean(dim=2)
123messages = self.msg_linear(aggregated)
124x = self.activation(x + messages)
125return x
126
127class ImageCaptioningModel(nn.Module):
128def __init__(self, vocab_size, embed_dim=512, num_layers=5, num_heads=8):
129super().__init__()
130# Encoder
131mobilenet = models.mobilenet_v3_small(weights=models.MobileNet_V3_Small_Weights.IMAGENET1K_V1)
132self.encoder = create_feature_extractor(mobilenet, return_nodes={'features': 'features'})
133self.vis_proj = nn.Linear(576, embed_dim)
134
135# Graph
136self.graph_layer = GraphPropagationLayer(embed_dim, window_size=3, max_steps=3)
137
138# Decoder
139self.embedding = nn.Embedding(vocab_size, embed_dim)
140self.pos_encoder = nn.Parameter(torch.randn(1, 2000, embed_dim))
141decoder_layer = nn.TransformerDecoderLayer(d_model=embed_dim, nhead=num_heads, dim_feedforward=1024, batch_first=True)
142self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_layers)
143self.fc_out = nn.Linear(embed_dim, vocab_size)
144
145def encode_images(self, images):
146visual_features = self.encoder(images)['features']
147B, C, H, W = visual_features.shape
148visual_tokens = visual_features.view(B, C, -1).permute(0, 2, 1)
149visual_tokens = self.vis_proj(visual_tokens)
150return self.graph_layer(visual_tokens)
151
152@torch.no_grad()
153def generate(self, images, max_len=80, tokenizer=None):
154self.eval()
155memory = self.encode_images(images)
156batch_size = images.shape[0]
157
158generated = torch.full((batch_size, 1), tokenizer.cls_token_id, dtype=torch.long, device=images.device)
159
160for _ in range(max_len):
161tgt_emb = self.embedding(generated)
162seq_len = tgt_emb.size(1)
163tgt_emb = tgt_emb + self.pos_encoder[:, :seq_len, :]
164
165# Causal Mask
166tgt_mask = torch.triu(torch.ones((seq_len, seq_len), device=images.device) * float('-inf'), diagonal=1)
167# Padding Mask
168tgt_key_padding_mask = (generated == tokenizer.pad_token_id).bool().to(images.device)
169
170output = self.decoder(
171tgt=tgt_emb, memory=memory,
172tgt_mask=tgt_mask,
173tgt_key_padding_mask=tgt_key_padding_mask
174)
175
176next_token = self.fc_out(output[:, -1, :]).argmax(dim=-1).unsqueeze(1)
177if (next_token == tokenizer.sep_token_id).all():
178break
179generated = torch.cat([generated, next_token], dim=1)
180return generated
181
182# --- 2. Load Model & Tokenizer ---
183
184device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
185tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
186
187# Initialize model architecture
188model = ImageCaptioningModel(tokenizer.vocab_size).to(device)
189
190# Download weights from Hub
191model_path = hf_hub_download(repo_id="erfanasghariyan/txttyt", filename="pytorch_model.bin")
192model.load_state_dict(torch.load(model_path, map_location=device))
193model.eval()
194
195print(f"Model loaded on {device}")
196
197# --- 3. Prepare Data ---
198
199transform = transforms.Compose([
200transforms.Resize((256, 256)),
201transforms.CenterCrop(224),
202transforms.ToTensor(),
203transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
204])
205
206# Load an image
207image_path = "path_to_your_medical_image.jpg" # Replace with actual path
208image = Image.open(image_path).convert("RGB")
209image_tensor = transform(image).unsqueeze(0).to(device)
210
211# --- 4. Run Inference ---
212
213output_ids = model.generate(image_tensor, max_len=80, tokenizer=tokenizer)
214caption = tokenizer.decode(output_ids.squeeze(0).cpu().numpy(), skip_special_tokens=True)
215
216print(f"Generated Caption: {caption}")
217
218### Hugging Face Space Demo
219
220You can also try the model without writing code by visiting the **Hugging Face Space**:
221
222[](https://huggingface.co/spaces/erfansghariyan/rad-vqa-demo)
223*(Note: Update the link above once you have created and deployed your Space)*