Spaces:
Running on Zero
Running on Zero
File size: 4,372 Bytes
1a1ea81 | 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 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 | """Transducer decoder interface module."""
from dataclasses import dataclass
from typing import Any
from typing import Dict
from typing import List
from typing import Optional
from typing import Tuple
from typing import Union
import torch
@dataclass
class Hypothesis:
"""Default hypothesis definition for transducer search algorithms."""
score: float
yseq: List[int]
dec_state: Union[
Tuple[torch.Tensor, Optional[torch.Tensor]],
List[Optional[torch.Tensor]],
torch.Tensor,
]
lm_state: Union[Dict[str, Any], List[Any]] = None
@dataclass
class ExtendedHypothesis(Hypothesis):
"""Extended hypothesis definition for NSC beam search and mAES."""
dec_out: List[torch.Tensor] = None
lm_scores: torch.Tensor = None
class TransducerDecoderInterface:
"""Decoder interface for transducer models."""
def init_state(
self,
batch_size: int,
) -> Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
]:
"""Initialize decoder states.
Args:
batch_size: Batch size.
Returns:
state: Initial decoder hidden states.
"""
raise NotImplementedError("init_state(...) is not implemented")
def score(
self,
hyp: Hypothesis,
cache: Dict[str, Any],
) -> Tuple[
torch.Tensor,
Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
],
torch.Tensor,
]:
"""One-step forward hypothesis.
Args:
hyp: Hypothesis.
cache: Pairs of (dec_out, dec_state) for each token sequence. (key)
Returns:
dec_out: Decoder output sequence.
new_state: Decoder hidden states.
lm_tokens: Label ID for LM.
"""
raise NotImplementedError("score(...) is not implemented")
def batch_score(
self,
hyps: Union[List[Hypothesis], List[ExtendedHypothesis]],
dec_states: Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
],
cache: Dict[str, Any],
use_lm: bool,
) -> Tuple[
torch.Tensor,
Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
],
torch.Tensor,
]:
"""One-step forward hypotheses.
Args:
hyps: Hypotheses.
dec_states: Decoder hidden states.
cache: Pairs of (dec_out, dec_states) for each label sequence. (key)
use_lm: Whether to compute label ID sequences for LM.
Returns:
dec_out: Decoder output sequences.
dec_states: Decoder hidden states.
lm_labels: Label ID sequences for LM.
"""
raise NotImplementedError("batch_score(...) is not implemented")
def select_state(
self,
batch_states: Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[torch.Tensor]
],
idx: int,
) -> Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
]:
"""Get specified ID state from decoder hidden states.
Args:
batch_states: Decoder hidden states.
idx: State ID to extract.
Returns:
state_idx: Decoder hidden state for given ID.
"""
raise NotImplementedError("select_state(...) is not implemented")
def create_batch_states(
self,
states: Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
],
new_states: List[
Union[
Tuple[torch.Tensor, Optional[torch.Tensor]],
List[Optional[torch.Tensor]],
]
],
l_tokens: List[List[int]],
) -> Union[
Tuple[torch.Tensor, Optional[torch.Tensor]], List[Optional[torch.Tensor]]
]:
"""Create decoder hidden states.
Args:
batch_states: Batch of decoder states
l_states: List of decoder states
l_tokens: List of token sequences for input batch
Returns:
batch_states: Batch of decoder states
"""
raise NotImplementedError("create_batch_states(...) is not implemented")
|