model_083084092_perceiver_tiny / model_083084092_perceiver_tiny.py
HarutoKondo's picture
Upload 2 files
df2a9ba verified
Raw
History Blame Contribute Delete
4.65 kB
import torch
import torch.nn as nn
import math
class PerceiverModel(nn.Module):
'''
perceiver model with dilated attention.
Scale: tiny (dim=128, layers=3, heads=2)
Fusion: cross_attention, Task: retrieval
'''
def __init__(self, embed_dim=128, num_layers=3, num_heads=2, num_classes=10):
super().__init__()
self.embed_dim = embed_dim
self.num_layers = num_layers
self.num_heads = num_heads
# image patch embedding
self.patch_proj = nn.Conv2d(3, embed_dim, kernel_size=16, stride=16)
self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim) * 0.02)
self.pos_embed = nn.Parameter(torch.randn(1, 197, embed_dim) * 0.02)
self.dropout = nn.Dropout(0.1)
# image transformer blocks
self.image_blocks = nn.ModuleList([
nn.TransformerEncoderLayer(
embed_dim, num_heads, embed_dim * 4, 0.1,
activation='gelu', batch_first=True, norm_first=True
)
for _ in range(num_layers)
])
self.image_norm = nn.LayerNorm(embed_dim)
# text embedding
self.text_embed = nn.Embedding(30522, embed_dim, padding_idx=0)
self.text_pos = nn.Parameter(torch.randn(1, 128, embed_dim) * 0.02)
self.text_blocks = nn.ModuleList([
nn.TransformerEncoderLayer(
embed_dim, num_heads, embed_dim * 4, 0.1,
activation='gelu', batch_first=True, norm_first=True
)
for _ in range(num_layers)
])
self.text_norm = nn.LayerNorm(embed_dim)
# fusion
self.fusion_blocks = nn.ModuleList([
nn.TransformerEncoderLayer(
embed_dim, num_heads, embed_dim * 4, 0.1,
activation='gelu', batch_first=True, norm_first=True
)
for _ in range(2)
])
self.fusion_norm = nn.LayerNorm(embed_dim)
# task head
self.classifier = nn.Sequential(
nn.Linear(embed_dim, embed_dim),
nn.GELU(approximate='tanh'),
nn.Dropout(0.1),
nn.Linear(embed_dim, num_classes),
)
self._initialize_weights()
def _initialize_weights(self):
for m in self.modules():
if isinstance(m, nn.Linear):
nn.init.kaiming_normal_(m.weight, mode='fan_out')
if m.bias is not None:
nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.trunc_normal_(m.weight, std=0.02)
if m.padding_idx is not None:
m.weight[m.padding_idx].zero_()
elif isinstance(m, nn.LayerNorm):
nn.init.ones_(m.weight)
nn.init.zeros_(m.bias)
def encode_image(self, images):
x = self.patch_proj(images)
x = x.flatten(2).transpose(1, 2)
cls = self.cls_token.expand(x.size(0), -1, -1)
x = torch.cat([cls, x], dim=1)
x = x + self.pos_embed
x = self.dropout(x)
for block in self.image_blocks:
x = block(x)
return self.image_norm(x)
def encode_text(self, input_ids, attention_mask=None):
x = self.text_embed(input_ids)
x = x + self.text_pos[:, :input_ids.size(1)]
x = self.dropout(x)
padding = (attention_mask == 0) if attention_mask is not None else None
for block in self.text_blocks:
x = block(x, src_key_padding_mask=padding)
return self.text_norm(x)
def forward(self, images, input_ids, attention_mask=None, labels=None):
image_features = self.encode_image(images)
text_features = self.encode_text(input_ids, attention_mask)
fused = text_features
for block in self.fusion_blocks:
fused = block(fused)
fused = self.fusion_norm(fused[:, 0])
logits = self.classifier(fused)
loss = None
if labels is not None:
loss = nn.functional.cross_entropy(logits, labels)
return {'logits': logits, 'loss': loss}
if __name__ == '__main__':
model = PerceiverModel()
total = sum(p.numel() for p in model.parameters())
print(f'PerceiverModel: {total:,} params ({total/1e6:.2f}M)')
img = torch.randn(2, 3, 224, 224)
ids = torch.randint(0, 30522, (2, 128))
mask = torch.ones(2, 128)
out = model(img, ids, mask, torch.tensor([0, 1]))
print(f'Output: {out["logits"].shape}, Loss: {out["loss"].item():.4f}')