| 155 | return self.linear(x) |
| 156 | |
| 157 | class SonarLossWrapper(nn.Module): |
| 158 | def __init__(self, llama_model, forward_proj, reverse_proj, sonar_decoder): |
| 159 | super().__init__() |
| 160 | self.llama_model = llama_model |
| 161 | self.forward_proj = forward_proj |
| 162 | self.reverse_proj = reverse_proj |
| 163 | self.sonar_decoder = sonar_decoder |
| 164 | |
| 165 | for p in self.sonar_decoder.parameters(): |
| 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 |