MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / forward

Method forward

inspiremusic/transformer/decoder.py:116–167  ·  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

114 self.tie_word_embedding = tie_word_embedding
115
116 def forward(
117 self,
118 memory: torch.Tensor,
119 memory_mask: torch.Tensor,
120 ys_in_pad: torch.Tensor,
121 ys_in_lens: torch.Tensor,
122 r_ys_in_pad: torch.Tensor = torch.empty(0),
123 reverse_weight: float = 0.0,
124 ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
125 """Forward decoder.
126 Args:
127 memory: encoded memory, float32 (batch, maxlen_in, feat)
128 memory_mask: encoder memory mask, (batch, 1, maxlen_in)
129 ys_in_pad: padded input token ids, int64 (batch, maxlen_out)
130 ys_in_lens: input lengths of this batch (batch)
131 r_ys_in_pad: not used in transformer decoder, in order to unify api
132 with bidirectional decoder
133 reverse_weight: not used in transformer decoder, in order to unify
134 api with bidirectional decode
135 Returns:
136 (tuple): tuple containing:
137 x: decoded token score before softmax (batch, maxlen_out,
138 vocab_size) if use_output_layer is True,
139 torch.tensor(0.0), in order to unify api with bidirectional decoder
140 olens: (batch, )
141 NOTE(xcsong):
142 We pass the `__call__` method of the modules instead of `forward` to the
143 checkpointing API because `__call__` attaches all the hooks of the module.
144 https://discuss.pytorch.org/t/any-different-between-model-input-and-model-forward-input/3690/2
145 """
146 tgt = ys_in_pad
147 maxlen = tgt.size(1)
148 # tgt_mask: (B, 1, L)
149 tgt_mask = ~make_pad_mask(ys_in_lens, maxlen).unsqueeze(1)
150 tgt_mask = tgt_mask.to(tgt.device)
151 # m: (1, L, L)
152 m = subsequent_mask(tgt_mask.size(-1),
153 device=tgt_mask.device).unsqueeze(0)
154 # tgt_mask: (B, L, L)
155 tgt_mask = tgt_mask & m
156 x, _ = self.embed(tgt)
157 if self.gradient_checkpointing and self.training:
158 x = self.forward_layers_checkpointed(x, tgt_mask, memory,
159 memory_mask)
160 else:
161 x = self.forward_layers(x, tgt_mask, memory, memory_mask)
162 if self.normalize_before:
163 x = self.after_norm(x)
164 if self.use_output_layer:
165 x = self.output_layer(x)
166 olens = tgt_mask.sum(1)
167 return x, torch.tensor(0.0), olens
168
169 def forward_layers(self, x: torch.Tensor, tgt_mask: torch.Tensor,
170 memory: torch.Tensor,

Callers

nothing calls this directly

Calls 5

forward_layersMethod · 0.95
make_pad_maskFunction · 0.90
subsequent_maskFunction · 0.90
embedMethod · 0.80

Tested by

no test coverage detected