Spaces:
Running on Zero
Running on Zero
Download src/models/caption_model.py from AdhamAshraf/image_caption_generator: direct link, hf CLI and curl.
- Browser
- Download file 1.86 kB
-
https://huggingface.co/spaces/AdhamAshraf/image_caption_generator/resolve/main/src/models/caption_model.py
- Command line
-
hf download hf://spaces/AdhamAshraf/image_caption_generator/src/models/caption_model.py
-
curl -L -o caption_model.py https://huggingface.co/spaces/AdhamAshraf/image_caption_generator/resolve/main/src/models/caption_model.py
1.86 kB
| """CaptionModel: wires encoder + decoder together. | |
| This class never needs to change when you swap architectures -- it only | |
| relies on the BaseEncoder/BaseDecoder contract. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| import torch.nn as nn | |
| from src.models.base import BaseDecoder, BaseEncoder | |
| from src.models.registry import build_decoder, build_encoder | |
| class CaptionModel(nn.Module): | |
| def __init__(self, encoder: BaseEncoder, decoder: BaseDecoder): | |
| super().__init__() | |
| self.encoder = encoder | |
| self.decoder = decoder | |
| def from_config(cls, config: dict, vocab_size: int) -> "CaptionModel": | |
| encoder = build_encoder(config) | |
| decoder = build_decoder(config, vocab_size) | |
| return cls(encoder, decoder) | |
| def forward(self, image_features: torch.Tensor, input_seq: torch.Tensor) -> torch.Tensor: | |
| image_embed = self.encoder(image_features) | |
| return self.decoder.forward(image_embed, input_seq) | |
| def generate( | |
| self, | |
| image_features: torch.Tensor, | |
| start_idx: int, | |
| end_idx: int, | |
| max_len: int, | |
| ) -> list[int]: | |
| self.eval() | |
| image_embed = self.encoder(image_features) | |
| return self.decoder.generate(image_embed, start_idx, end_idx, max_len) | |
| #for beam search | |
| def generate( | |
| self, | |
| image_features: torch.Tensor, | |
| start_idx: int, | |
| end_idx: int, | |
| max_len: int, | |
| decoding: str = "greedy", | |
| beam_width: int = 3, | |
| ) -> list[int]: | |
| self.eval() | |
| image_embed = self.encoder(image_features) | |
| if decoding == "beam": | |
| return self.decoder.generate_beam(image_embed, start_idx, end_idx, max_len, beam_width) | |
| return self.decoder.generate(image_embed, start_idx, end_idx, max_len) |