Spaces:
Running on Zero
Running on Zero
File size: 2,134 Bytes
143710c 7b4e05b 143710c 7b4e05b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 | """Abstract base classes.
Any encoder/decoder you add later (attention, transformer, EfficientNet...)
MUST implement these methods with these exact signatures. That's what lets
Trainer, evaluate.py, and predict.py stay architecture-agnostic.
"""
from __future__ import annotations
from abc import ABC, abstractmethod
import torch
import torch.nn as nn
class BaseEncoder(nn.Module, ABC):
"""Takes precomputed CNN feature vectors and projects them to embed_dim."""
@abstractmethod
def forward(self, image_features: torch.Tensor) -> torch.Tensor:
"""image_features: (batch, feature_dim) -> returns (batch, embed_dim)."""
raise NotImplementedError
class BaseDecoder(nn.Module, ABC):
"""Consumes image embedding + token sequence, predicts next-word logits."""
@abstractmethod
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 input
returns: (batch, seq_len, vocab_size) logits aligned with target_seq
"""
raise NotImplementedError
#for greedy search
@abstractmethod
def generate(
self,
image_embed: torch.Tensor,
start_idx: int,
end_idx: int,
max_len: int,
) -> list[int]:
"""Autoregressive greedy generation for a SINGLE image (batch=1).
Returns a list of generated token indices, excluding <start> and <end>.
"""
raise NotImplementedError
#for beam search
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 generation for a SINGLE image (batch=1).
Optional: not marked @abstractmethod so existing decoders that only
implement greedy generate() remain valid; override in subclasses
that support beam search (see DecoderLSTM).
"""
raise NotImplementedError("This decoder does not implement beam search.") |