MCPcopy Create free account
hub / github.com/ASLP-lab/OSUM / forward

Method forward

OSUM/wenet/transformer/decoder.py:146–201  ·  view source on GitHub ↗

Forward decoder. Args: memory: encoded memory, float32 (batch, maxlen_in, feat) memory_mask: encoder memory mask, (batch, 1, maxlen_in) ys_in_pad: padded input token ids, int64 (batch, maxlen_out) ys_in_lens: input lengths of this batch (batch

(
        self,
        memory: torch.Tensor,
        memory_mask: torch.Tensor,
        ys_in_pad: torch.Tensor,
        ys_in_lens: torch.Tensor,
        r_ys_in_pad: torch.Tensor = torch.empty(0),
        reverse_weight: float = 0.0,
    )

Source from the content-addressed store, hash-verified

144 self.use_sdpa = use_sdpa
145
146 def forward(
147 self,
148 memory: torch.Tensor,
149 memory_mask: torch.Tensor,
150 ys_in_pad: torch.Tensor,
151 ys_in_lens: torch.Tensor,
152 r_ys_in_pad: torch.Tensor = torch.empty(0),
153 reverse_weight: float = 0.0,
154 ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
155 """Forward decoder.
156 Args:
157 memory: encoded memory, float32 (batch, maxlen_in, feat)
158 memory_mask: encoder memory mask, (batch, 1, maxlen_in)
159 ys_in_pad: padded input token ids, int64 (batch, maxlen_out)
160 ys_in_lens: input lengths of this batch (batch)
161 r_ys_in_pad: not used in transformer decoder, in order to unify api
162 with bidirectional decoder
163 reverse_weight: not used in transformer decoder, in order to unify
164 api with bidirectional decode
165 Returns:
166 (tuple): tuple containing:
167 x: decoded token score before softmax (batch, maxlen_out,
168 vocab_size) if use_output_layer is True,
169 torch.tensor(0.0), in order to unify api with bidirectional decoder
170 olens: (batch, )
171 NOTE(xcsong):
172 We pass the `__call__` method of the modules instead of `forward` to the
173 checkpointing API because `__call__` attaches all the hooks of the module.
174 https://discuss.pytorch.org/t/any-different-between-model-input-and-model-forward-input/3690/2
175 """
176 tgt = ys_in_pad
177 maxlen = tgt.size(1)
178 # tgt_mask: (B, 1, L)
179 tgt_mask = ~make_pad_mask(ys_in_lens, maxlen).unsqueeze(1)
180 tgt_mask = tgt_mask.to(tgt.device)
181 # m: (1, L, L)
182 m = subsequent_mask(tgt_mask.size(-1),
183 device=tgt_mask.device).unsqueeze(0)
184 # tgt_mask: (B, L, L)
185 tgt_mask = tgt_mask & m
186 if self.use_sdpa:
187 tgt_mask = mask_to_bias(tgt_mask, memory.dtype)
188 memory_mask = mask_to_bias(memory_mask, memory.dtype)
189
190 x, _ = self.embed(tgt)
191 if self.gradient_checkpointing and self.training:
192 x = self.forward_layers_checkpointed(x, tgt_mask, memory,
193 memory_mask)
194 else:
195 x = self.forward_layers(x, tgt_mask, memory, memory_mask)
196 if self.normalize_before:
197 x = self.after_norm(x)
198 if self.use_output_layer:
199 x = self.output_layer(x)
200 olens = tgt_mask.sum(1)
201 return x, torch.tensor(0.0), olens
202
203 def forward_layers(self, x: torch.Tensor, tgt_mask: torch.Tensor,

Callers

nothing calls this directly

Calls 5

forward_layersMethod · 0.95
make_pad_maskFunction · 0.90
subsequent_maskFunction · 0.90
mask_to_biasFunction · 0.90

Tested by

no test coverage detected