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

Class SonarLossWrapper

train_sonarllm.py:157–226  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

155 return self.linear(x)
156
157class 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

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected