MCPcopy Create free account
hub / github.com/modelscope/FunASR / punc_forward

Method punc_forward

funasr/models/ct_transformer/model.py:113–125  ·  view source on GitHub ↗

Compute loss value from buffer sequences. Args: input (torch.Tensor): Input ids. (batch, len) hidden (torch.Tensor): Target ids. (batch, len)

(self, text: torch.Tensor, text_lengths: torch.Tensor, **kwargs)

Source from the content-addressed store, hash-verified

111 self.jieba_usr_dict = jieba
112
113 def punc_forward(self, text: torch.Tensor, text_lengths: torch.Tensor, **kwargs):
114 """Compute loss value from buffer sequences.
115
116 Args:
117 input (torch.Tensor): Input ids. (batch, len)
118 hidden (torch.Tensor): Target ids. (batch, len)
119
120 """
121 x = self.embed(text)
122 # mask = self._target_mask(input)
123 h, _, _ = self.encoder(x, text_lengths)
124 y = self.decoder(h)
125 return y, None
126
127 def with_vad(self):
128 """With vad."""

Callers 2

nllMethod · 0.95
inferenceMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected