(self, embeddings_1024, texts, seq_lens)
| 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 |
nothing calls this directly
no outgoing calls
no test coverage detected