1# !pip install transformers sentencepiece torch -q
2import torch
3from transformers import AutoTokenizer, AutoModel
4from transformers.models.m2m_100.modeling_m2m_100 import M2M100Encoder
5
6# 1. Load SONAR encoder
7sonar_model_name = "cointegrated/SONAR_200_text_encoder"
8encoder = M2M100Encoder.from_pretrained(sonar_model_name)
9tokenizer = AutoTokenizer.from_pretrained(sonar_model_name)
10
11def encode_mean_pool(texts, tokenizer, encoder, lang='eng_Latn', norm=False):
12 tokenizer.src_lang = lang
13 with torch.inference_mode():
14 batch = tokenizer(texts, return_tensors='pt', padding=True)
15 seq_embs = encoder(**batch).last_hidden_state
16 mask = batch.attention_mask
17 mean_emb = (seq_embs * mask.unsqueeze(-1)).sum(1) / mask.unsqueeze(-1).sum(1)
18 if norm:
19 mean_emb = torch.nn.functional.normalize(mean_emb)
20 return mean_emb
21
22# Example sentences
23src_sentences = ["Le chat s'assit sur le tapis."]
24mt_sentences = ["The cat sat down on the carpet."] # Example MT output
25ref_sentences = ["The cat sat on the mat."] # Example reference translation
26
27# Encode source and MT sentences
28src_embs = encode_mean_pool(src_sentences, tokenizer, encoder, lang="fra_Latn")
29mt_embs = encode_mean_pool(mt_sentences, tokenizer, encoder, lang="eng_Latn")
30ref_embs = encode_mean_pool(ref_sentences, tokenizer, encoder, lang="eng_Latn")
31
32# 2. Load BLASER Ref model (ported)
33ref_model_name = "oist/blaser_2_0_ref_ported"
34ref_model = AutoModel.from_pretrained(ref_model_name, trust_remote_code=True)
35ref_model.eval() # set to evaluation mode
36
37# 3. Compute Ref scores
38with torch.inference_mode():
39 ref_scores = ref_model(src_embs, mt_embs, ref_embs) # expects source and MT embeddings, and ref embeddings
40 print("Blaser score shape:", ref_scores.shape)
41 print("Blaser scores:", ref_scores[0])