gipformer-extract / decoder.py
giangndm's picture
Upload decoder.py with huggingface_hub
0359e01 verified
Raw
History Blame Contribute Delete
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
)
@classmethod
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