AdhamAshraf's picture
Sync src/ with attention decoder support
7b4e05b
Raw History Blame Contribute Delete
12.4 kB
"""Decoder implementations."""
from __future__ import annotations
import torch
import torch.nn as nn
from src.models.base import BaseDecoder
class DecoderLSTM(BaseDecoder):
"""Image-as-first-token LSTM decoder (Show and Tell style baseline)."""
def __init__(
self,
vocab_size: int,
embed_dim: int = 256,
hidden_dim: int = 512,
num_layers: int = 1,
dropout: float = 0.0,
**kwargs,
):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
self.lstm = nn.LSTM(
embed_dim,
hidden_dim,
num_layers=num_layers,
batch_first=True,
dropout=dropout if num_layers > 1 else 0.0,
)
# NOTE: nn.LSTM's own `dropout` arg only applies BETWEEN stacked layers,
# so it has no effect with num_layers=1 (our baseline). This separate
# output_dropout applies to the LSTM's output before the final
# projection, giving a real regularization effect even for a
# single-layer LSTM.
self.output_dropout = nn.Dropout(dropout)
self.fc = nn.Linear(hidden_dim, vocab_size)
def forward(self, image_embed: torch.Tensor, input_seq: torch.Tensor) -> torch.Tensor:
# image_embed: (batch, embed_dim)
# input_seq: (batch, seq_len) token indices (teacher-forced)
word_embeds = self.embedding(input_seq) # (batch, seq_len, embed_dim)
image_step = image_embed.unsqueeze(1) # (batch, 1, embed_dim)
lstm_input = torch.cat([image_step, word_embeds], dim=1) # (batch, seq_len+1, embed_dim)
lstm_out, _ = self.lstm(lstm_input) # (batch, seq_len+1, hidden_dim)
lstm_out = self.output_dropout(lstm_out)
logits = self.fc(lstm_out) # (batch, seq_len+1, vocab_size)
# drop the output at the image step (index 0) -- it has no word target
return logits[:, 1:, :] # (batch, seq_len, vocab_size)
@torch.no_grad()
def generate(
self,
image_embed: torch.Tensor,
start_idx: int,
end_idx: int,
max_len: int,
) -> list[int]:
assert image_embed.size(0) == 1, "generate() supports batch_size=1"
device = image_embed.device
image_step = image_embed.unsqueeze(1)
_, hidden = self.lstm(image_step)
current_token = torch.tensor([[start_idx]], device=device)
generated: list[int] = []
for _ in range(max_len):
word_embed = self.embedding(current_token)
lstm_out, hidden = self.lstm(word_embed, hidden)
logits = self.fc(lstm_out.squeeze(1))
predicted_idx = int(logits.argmax(dim=-1).item())
if predicted_idx == end_idx:
break
generated.append(predicted_idx)
current_token = torch.tensor([[predicted_idx]], device=device)
return generated
@torch.no_grad()
def generate_beam(
self,
image_embed: torch.Tensor,
start_idx: int,
end_idx: int,
max_len: int,
beam_width: int = 3,
) -> list[int]:
"""Beam search decoding for a SINGLE image (batch=1).
Keeps the top `beam_width` candidate sequences at each step instead of
committing to a single best word (as generate()/greedy does), which
can recover from a locally-suboptimal early word choice. No retraining
required -- this only changes inference-time decoding.
"""
assert image_embed.size(0) == 1, "generate_beam() supports batch_size=1"
device = image_embed.device
image_step = image_embed.unsqueeze(1) # (1, 1, embed_dim)
_, init_hidden = self.lstm(image_step)
# each beam: (token_ids, cumulative_log_prob, hidden_state, finished)
beams = [([start_idx], 0.0, init_hidden, False)]
for _ in range(max_len):
all_candidates = []
for tokens, log_prob, hidden, finished in beams:
if finished:
all_candidates.append((tokens, log_prob, hidden, finished))
continue
last_token = torch.tensor([[tokens[-1]]], device=device)
word_embed = self.embedding(last_token)
lstm_out, new_hidden = self.lstm(word_embed, hidden)
logits = self.fc(lstm_out.squeeze(1))
log_probs = torch.log_softmax(logits, dim=-1).squeeze(0)
top_log_probs, top_indices = log_probs.topk(beam_width)
for lp, idx in zip(top_log_probs.tolist(), top_indices.tolist()):
new_tokens = tokens + [idx]
new_finished = idx == end_idx
all_candidates.append((new_tokens, log_prob + lp, new_hidden, new_finished))
all_candidates.sort(key=lambda c: c[1], reverse=True)
beams = all_candidates[:beam_width]
if all(finished for _, _, _, finished in beams):
break
best_tokens, _, _, _ = max(beams, key=lambda c: c[1])
result = best_tokens[1:]
if result and result[-1] == end_idx:
result = result[:-1]
return result
class BahdanauAttention(nn.Module):
"""Additive (Bahdanau-style) attention over spatial image regions.
At each decoder step, computes a weighted combination of the encoder's
spatial feature grid, where the weights depend on the current decoder
hidden state -- this is what lets the decoder "look at" different image
regions for different words, instead of relying on one static vector.
"""
def __init__(self, embed_dim: int, hidden_dim: int, attention_dim: int):
super().__init__()
self.encoder_att = nn.Linear(embed_dim, attention_dim)
self.decoder_att = nn.Linear(hidden_dim, attention_dim)
self.full_att = nn.Linear(attention_dim, 1)
self.relu = nn.ReLU()
self.softmax = nn.Softmax(dim=1)
def forward(self, encoder_out: torch.Tensor, decoder_hidden: torch.Tensor):
att1 = self.encoder_att(encoder_out)
att2 = self.decoder_att(decoder_hidden).unsqueeze(1)
att = self.full_att(self.relu(att1 + att2)).squeeze(2)
alpha = self.softmax(att)
context = (encoder_out * alpha.unsqueeze(2)).sum(dim=1)
return context, alpha
class DecoderAttentionLSTM(BaseDecoder):
"""LSTM decoder with Bahdanau-style attention over spatial image features.
Unlike DecoderLSTM (which injects one static image vector as the first
LSTM step), this decoder recomputes a fresh, weighted combination of
image regions at EVERY generation step, using an LSTMCell in an explicit
loop (required since each step's attention depends on that step's hidden
state -- unlike DecoderLSTM, the whole sequence can't be processed in one
nn.LSTM call).
"""
def __init__(
self,
vocab_size: int,
embed_dim: int = 256,
hidden_dim: int = 512,
attention_dim: int = 256,
dropout: float = 0.5,
**kwargs,
):
super().__init__()
self.embed_dim = embed_dim
self.hidden_dim = hidden_dim
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=0)
self.attention = BahdanauAttention(embed_dim, hidden_dim, attention_dim)
self.lstm_cell = nn.LSTMCell(embed_dim + embed_dim, hidden_dim)
self.init_h = nn.Linear(embed_dim, hidden_dim)
self.init_c = nn.Linear(embed_dim, hidden_dim)
self.f_beta = nn.Linear(hidden_dim, embed_dim)
self.sigmoid = nn.Sigmoid()
self.dropout = nn.Dropout(dropout)
self.fc = nn.Linear(hidden_dim, vocab_size)
def _init_hidden_state(self, encoder_out: torch.Tensor):
mean_encoder_out = encoder_out.mean(dim=1)
h = self.init_h(mean_encoder_out)
c = self.init_c(mean_encoder_out)
return h, c
def forward(self, image_embed: torch.Tensor, input_seq: torch.Tensor) -> torch.Tensor:
batch_size = image_embed.size(0)
seq_len = input_seq.size(1)
vocab_size = self.fc.out_features
device = image_embed.device
embeddings = self.embedding(input_seq)
h, c = self._init_hidden_state(image_embed)
outputs = torch.zeros(batch_size, seq_len, vocab_size, device=device)
for t in range(seq_len):
context, _ = self.attention(image_embed, h)
gate = self.sigmoid(self.f_beta(h))
context = gate * context
lstm_input = torch.cat([embeddings[:, t, :], context], dim=1)
h, c = self.lstm_cell(lstm_input, (h, c))
preds = self.fc(self.dropout(h))
outputs[:, t, :] = preds
return outputs
@torch.no_grad()
def generate(
self,
image_embed: torch.Tensor,
start_idx: int,
end_idx: int,
max_len: int,
) -> list[int]:
assert image_embed.size(0) == 1, "generate() supports batch_size=1"
device = image_embed.device
h, c = self._init_hidden_state(image_embed)
current_token = torch.tensor([start_idx], device=device)
generated: list[int] = []
for _ in range(max_len):
word_embed = self.embedding(current_token)
context, _ = self.attention(image_embed, h)
gate = self.sigmoid(self.f_beta(h))
context = gate * context
lstm_input = torch.cat([word_embed, context], dim=1)
h, c = self.lstm_cell(lstm_input, (h, c))
logits = self.fc(h)
predicted_idx = int(logits.argmax(dim=-1).item())
if predicted_idx == end_idx:
break
generated.append(predicted_idx)
current_token = torch.tensor([predicted_idx], device=device)
return generated
@torch.no_grad()
def generate_beam(
self,
image_embed: torch.Tensor,
start_idx: int,
end_idx: int,
max_len: int,
beam_width: int = 3,
) -> list[int]:
"""Beam search decoding for a SINGLE image (batch=1).
Same beam-search principle as DecoderLSTM.generate_beam(), adapted for
the LSTMCell + per-step attention loop: each beam carries its own
(h, c) hidden state, since attention (and therefore the next word
distribution) depends on each beam's own hidden state, not a shared one.
"""
assert image_embed.size(0) == 1, "generate_beam() supports batch_size=1"
device = image_embed.device
init_h, init_c = self._init_hidden_state(image_embed)
# each beam: (token_ids, cumulative_log_prob, h, c, finished)
beams = [([start_idx], 0.0, init_h, init_c, False)]
for _ in range(max_len):
all_candidates = []
for tokens, log_prob, h, c, finished in beams:
if finished:
all_candidates.append((tokens, log_prob, h, c, finished))
continue
last_token = torch.tensor([tokens[-1]], device=device)
word_embed = self.embedding(last_token)
context, _ = self.attention(image_embed, h)
gate = self.sigmoid(self.f_beta(h))
context = gate * context
lstm_input = torch.cat([word_embed, context], dim=1)
new_h, new_c = self.lstm_cell(lstm_input, (h, c))
logits = self.fc(new_h)
log_probs = torch.log_softmax(logits, dim=-1).squeeze(0)
top_log_probs, top_indices = log_probs.topk(beam_width)
for lp, idx in zip(top_log_probs.tolist(), top_indices.tolist()):
new_tokens = tokens + [idx]
new_finished = idx == end_idx
all_candidates.append((new_tokens, log_prob + lp, new_h, new_c, new_finished))
all_candidates.sort(key=lambda cand: cand[1], reverse=True)
beams = all_candidates[:beam_width]
if all(finished for _, _, _, _, finished in beams):
break
best_tokens, _, _, _, _ = max(beams, key=lambda cand: cand[1])
result = best_tokens[1:]
if result and result[-1] == end_idx:
result = result[:-1]
return result