Download decoder.py from giangndm/gipformer-extract: direct link, hf CLI and curl.
- Browser
- Download file 7.25 kB
-
https://huggingface.co/giangndm/gipformer-extract/resolve/main/decoder.py
- Command line
-
hf download hf://giangndm/gipformer-extract/decoder.py
-
curl -L -o decoder.py https://huggingface.co/giangndm/gipformer-extract/resolve/main/decoder.py
7.25 kB
| import json | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from typing import List, Tuple, Optional, Union, Dict, Any | |
| from huggingface_hub import hf_hub_download | |
| from safetensors.torch import load_file | |
| # --- HELPER CLASSES --- | |
| class Balancer(nn.Module): | |
| def __init__(self, *args, **kwargs): | |
| super().__init__() | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return x | |
| def ScaledLinear(*args, initial_scale: float = 1.0, **kwargs) -> nn.Linear: | |
| return nn.Linear(*args, **kwargs) | |
| # --- DECODER & JOINER --- | |
| class Decoder(nn.Module): | |
| def __init__( | |
| self, | |
| vocab_size: int, | |
| decoder_dim: int, | |
| blank_id: int, | |
| context_size: int, | |
| ): | |
| super().__init__() | |
| self.embedding = nn.Embedding( | |
| num_embeddings=vocab_size, | |
| embedding_dim=decoder_dim, | |
| ) | |
| self.balancer = Balancer( | |
| decoder_dim, | |
| channel_dim=-1, | |
| min_positive=0.0, | |
| max_positive=1.0, | |
| min_abs=0.5, | |
| max_abs=1.0, | |
| prob=0.05, | |
| ) | |
| self.blank_id = blank_id | |
| assert context_size >= 1, context_size | |
| self.context_size = context_size | |
| self.vocab_size = vocab_size | |
| if context_size > 1: | |
| self.conv = nn.Conv1d( | |
| in_channels=decoder_dim, | |
| out_channels=decoder_dim, | |
| kernel_size=context_size, | |
| padding=0, | |
| groups=decoder_dim // 4, | |
| bias=False, | |
| ) | |
| self.balancer2 = Balancer( | |
| decoder_dim, | |
| channel_dim=-1, | |
| min_positive=0.0, | |
| max_positive=1.0, | |
| min_abs=0.5, | |
| max_abs=1.0, | |
| prob=0.05, | |
| ) | |
| else: | |
| self.conv = nn.Identity() | |
| self.balancer2 = nn.Identity() | |
| def forward(self, y: torch.Tensor, need_pad: bool = True) -> torch.Tensor: | |
| y = y.to(torch.int64) | |
| embedding_out = self.embedding(y.clamp(min=0)) * (y >= 0).unsqueeze(-1) | |
| embedding_out = self.balancer(embedding_out) | |
| if self.context_size > 1: | |
| embedding_out = embedding_out.permute(0, 2, 1) | |
| if need_pad is True: | |
| embedding_out = F.pad(embedding_out, pad=(self.context_size - 1, 0)) | |
| else: | |
| assert embedding_out.size(-1) == self.context_size | |
| embedding_out = self.conv(embedding_out) | |
| embedding_out = embedding_out.permute(0, 2, 1) | |
| embedding_out = F.relu(embedding_out) | |
| embedding_out = self.balancer2(embedding_out) | |
| return embedding_out | |
| class Joiner(nn.Module): | |
| def __init__( | |
| self, | |
| encoder_dim: int, | |
| decoder_dim: int, | |
| joiner_dim: int, | |
| vocab_size: int, | |
| ): | |
| super().__init__() | |
| self.encoder_proj = ScaledLinear(encoder_dim, joiner_dim, initial_scale=0.25) | |
| self.decoder_proj = ScaledLinear(decoder_dim, joiner_dim, initial_scale=0.25) | |
| self.output_linear = nn.Linear(joiner_dim, vocab_size) | |
| def forward( | |
| self, | |
| encoder_out: torch.Tensor, | |
| decoder_out: torch.Tensor, | |
| project_input: bool = True, | |
| ) -> torch.Tensor: | |
| assert encoder_out.ndim == decoder_out.ndim, ( | |
| encoder_out.shape, | |
| decoder_out.shape, | |
| ) | |
| if project_input: | |
| logit = self.encoder_proj(encoder_out) + self.decoder_proj(decoder_out) | |
| else: | |
| logit = encoder_out + decoder_out | |
| logit = self.output_linear(torch.tanh(logit)) | |
| return logit | |
| # --- DECODING HELPER --- | |
| def greedy_search( | |
| model: nn.Module, | |
| encoder_out: torch.Tensor, | |
| max_sym_per_frame: int = 1, | |
| blank_penalty: float = 0.0, | |
| ) -> List[int]: | |
| assert encoder_out.ndim == 3 | |
| assert encoder_out.size(0) == 1, encoder_out.size(0) | |
| blank_id = model.decoder.blank_id | |
| context_size = model.decoder.context_size | |
| unk_id = getattr(model, "unk_id", blank_id) | |
| device = encoder_out.device | |
| decoder_input = torch.tensor( | |
| [-1] * (context_size - 1) + [blank_id], device=device, dtype=torch.int64 | |
| ).reshape(1, context_size) | |
| decoder_out = model.decoder(decoder_input, need_pad=False) | |
| decoder_out = model.joiner.decoder_proj(decoder_out) | |
| encoder_out = model.joiner.encoder_proj(encoder_out) | |
| T = encoder_out.size(1) | |
| t = 0 | |
| hyp = [blank_id] * context_size | |
| max_sym_per_utt = 1000 | |
| sym_per_frame = 0 | |
| sym_per_utt = 0 | |
| while t < T and sym_per_utt < max_sym_per_utt: | |
| if sym_per_frame >= max_sym_per_frame: | |
| sym_per_frame = 0 | |
| t += 1 | |
| continue | |
| current_encoder_out = encoder_out[:, t:t+1, :].unsqueeze(2) | |
| logits = model.joiner( | |
| current_encoder_out, decoder_out.unsqueeze(1), project_input=False | |
| ) | |
| if blank_penalty != 0: | |
| logits[:, :, :, 0] -= blank_penalty | |
| y = logits.argmax().item() | |
| if y not in (blank_id, unk_id): | |
| hyp.append(y) | |
| decoder_input = torch.tensor([hyp[-context_size:]], device=device).reshape( | |
| 1, context_size | |
| ) | |
| decoder_out = model.decoder(decoder_input, need_pad=False) | |
| decoder_out = model.joiner.decoder_proj(decoder_out) | |
| sym_per_utt += 1 | |
| sym_per_frame += 1 | |
| else: | |
| sym_per_frame = 0 | |
| t += 1 | |
| hyp = hyp[context_size:] | |
| return hyp | |
| # --- WRAPPER CLASSES --- | |
| class PurePyTorchDecoder(nn.Module): | |
| """ | |
| Decoupled Decoder containing stateless predictor (decoder) | |
| and joint network (joiner). | |
| """ | |
| def __init__(self, config: dict): | |
| super().__init__() | |
| self.config = config | |
| vocab_size = config.get("vocab_size", 2000) | |
| decoder_dim = config.get("decoder_dim", 512) | |
| joiner_dim = config.get("joiner_dim", 512) | |
| blank_id = config.get("blank_id", 0) | |
| context_size = config.get("context_size", 2) | |
| self.decoder = Decoder( | |
| vocab_size=vocab_size, | |
| decoder_dim=decoder_dim, | |
| blank_id=blank_id, | |
| context_size=context_size | |
| ) | |
| self.joiner = Joiner( | |
| encoder_dim=decoder_dim, | |
| decoder_dim=decoder_dim, | |
| joiner_dim=joiner_dim, | |
| vocab_size=vocab_size | |
| ) | |
| def from_pretrained(cls, repo_id="giangndm/gipformer-extract", device="cpu") -> "PurePyTorchDecoder": | |
| config_path = hf_hub_download(repo_id=repo_id, filename="decoder.json") | |
| with open(config_path, "r") as f: | |
| config = json.load(f) | |
| model = cls(config) | |
| weights_path = hf_hub_download(repo_id=repo_id, filename="gipformer_decoder.safetensors") | |
| state_dict = load_file(weights_path) | |
| model.load_state_dict(state_dict, strict=True) | |
| model.to(device) | |
| return model | |
| class ModelContainer(nn.Module): | |
| def __init__(self, encoder, decoder_joiner): | |
| super().__init__() | |
| self.encoder = encoder | |
| self.decoder = decoder_joiner.decoder | |
| self.joiner = decoder_joiner.joiner | |