MCPcopy Create free account
hub / github.com/FusionBrainLab/SONAR-LLM / forward

Method forward

train_sonarllm.py:168–226  ·  view source on GitHub ↗
(self, embeddings_1024, texts, seq_lens)

Source from the content-addressed store, hash-verified

166 p.requires_grad = False
167
168 def forward(self, embeddings_1024, texts, seq_lens):
169 device = embeddings_1024.device
170 B, T, _ = embeddings_1024.shape
171
172 embs_projected = self.forward_proj(embeddings_1024)
173 llama_out = self.llama_model(
174 inputs_embeds=embs_projected,
175 output_hidden_states=True
176 )
177 last_hidden = llama_out.hidden_states[-1]
178
179 pred_hidden_list = []
180 ref_texts_list = []
181 for b in range(B):
182 seqlen = seq_lens[b]
183 for k in range(1, seqlen):
184 pred_hidden_list.append(last_hidden[b, k - 1, :])
185 ref_texts_list.append(texts[b][k])
186
187 if len(pred_hidden_list) == 0:
188 return torch.tensor(0.0, device=device)
189
190 pred_hidden_batch = torch.stack(pred_hidden_list, dim=0)
191 pred_emb_1024 = self.reverse_proj(pred_hidden_batch)
192
193 with torch.no_grad():
194 target_text_encoder = self.sonar_decoder.tokenizer.create_encoder(
195 task="translation", lang="eng_Latn", mode="target", device=device
196 )
197 encoded_texts = [target_text_encoder(t) for t in ref_texts_list]
198 lengths = [et.size(0) for et in encoded_texts]
199 max_len = min(max(lengths), 256)
200
201 pad_idx = self.sonar_decoder.tokenizer.vocab_info.pad_idx
202 dec_ids = torch.full((len(encoded_texts), max_len), pad_idx, dtype=torch.long, device=device)
203 labels = torch.full((len(encoded_texts), max_len), pad_idx, dtype=torch.long, device=device)
204 for i, et in enumerate(encoded_texts):
205 dec_ids[i, : min(len(et), max_len)] = et[:max_len]
206 et = torch.cat([et[1:], torch.tensor([3]).to(device)])
207 labels[i, : min(len(et), max_len)] = et[:max_len]
208
209 enc_output = pred_emb_1024.unsqueeze(1)
210 dec_out, dec_pad_mask = self.sonar_decoder.model.decode(
211 seqs=dec_ids,
212 padding_mask=None,
213 encoder_output=enc_output,
214 encoder_padding_mask=None
215 )
216 final_out = self.sonar_decoder.model.project(dec_out, dec_pad_mask)
217 logits = final_out.logits
218
219 vocab_size = logits.size(-1)
220 logits_2d = logits.view(-1, vocab_size)
221 labels_1d = labels.view(-1)
222
223 ce_fn = nn.CrossEntropyLoss(ignore_index=pad_idx, reduction="mean")
224 total_ce = ce_fn(logits_2d, labels_1d)
225

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected