Skip to content

Commit 5a71877

Browse files
committed
Fixed a bug with HY-MT
1 parent bd8bb41 commit 5a71877

2 files changed

Lines changed: 32 additions & 16 deletions

File tree

app/src/main/java/nie/translator/rtranslator/voice_translation/neural_networks/translation/Tokenizer.java

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -51,10 +51,10 @@ public Tokenizer(String vocab_path, int mode) throws IOException {
5151
}
5252

5353
public TokenizerResult tokenize(String text, String srcLanguage, String tgtLanguage){
54-
return tokenize(text, srcLanguage, tgtLanguage, null, null);
54+
return tokenize(text, srcLanguage, tgtLanguage, null, null, false);
5555
}
5656

57-
public TokenizerResult tokenize(String text, String srcLanguage, String tgtLanguage, @Nullable Translator.HyLanguageInfo srcLangInfo, @Nullable Translator.HyLanguageInfo tgtLangInfo) {
57+
public TokenizerResult tokenize(String text, String srcLanguage, String tgtLanguage, @Nullable Translator.HyLanguageInfo srcLangInfo, @Nullable Translator.HyLanguageInfo tgtLangInfo, boolean excludePrompt) {
5858
if(mode != HY_MT) {
5959
//for madlad we add <2tgtLanguage> at the beginning of the text (srcLanguage is not specified)
6060
if (mode == MADLAD || mode == MADLAD_FIXED) {
@@ -158,15 +158,15 @@ public TokenizerResult tokenize(String text, String srcLanguage, String tgtLangu
158158
} else {
159159
prompt = "<|hy_begin▁of▁sentence|><|hy_User|>Translate the following segment into "+tgtLangEnName+", without additional explanation.\n\n"+text+"<|hy_Assistant|>";
160160
}
161-
Encoding result = hfTokenizer.encode(new String[]{prompt}, false, false);
161+
Encoding result = hfTokenizer.encode(new String[]{excludePrompt ? text : prompt}, false, false);
162162

163163
//we convert inputIds and attention mask from long[] to int[]
164164
long[] inputIdsLong = result.getIds();
165165
int[] inputIds = new int[inputIdsLong.length];
166166
for (int i = 0; i < inputIds.length; i++) {
167167
inputIds[i] = (int) inputIdsLong[i];
168168
}
169-
long[] attentionMaskLong = result.getIds();
169+
long[] attentionMaskLong = result.getAttentionMask();
170170
int[] attentionMask = new int[attentionMaskLong.length];
171171
for (int i = 0; i < attentionMask.length; i++) {
172172
attentionMask[i] = (int) attentionMaskLong[i];

app/src/main/java/nie/translator/rtranslator/voice_translation/neural_networks/translation/Translator.java

Lines changed: 28 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -201,7 +201,7 @@ public void onFailure(int[] reasons, long value) {
201201
}
202202
});
203203
} else {
204-
final OrtSession.SessionOptions.OptLevel optDefaultLevel = OrtSession.SessionOptions.OptLevel.BASIC_OPT;
204+
final OrtSession.SessionOptions.OptLevel optDefaultLevel = OrtSession.SessionOptions.OptLevel.EXTENDED_OPT;
205205
boolean arena = true;
206206

207207
OrtSession.SessionOptions decoderOptions = new OrtSession.SessionOptions();
@@ -668,8 +668,8 @@ private void performTextTranslation(final String textToTranslate, final CustomLo
668668
while (joined) {
669669
joined = false;
670670
for (int i = 1; i < textSplit.size(); i++) {
671-
int numTokens = tokenize(textSplit.get(i - 1), inputLanguage, outputLanguage).getInputIDs().length;
672-
int numTokens2 = tokenize(textSplit.get(i), inputLanguage, outputLanguage).getInputIDs().length;
671+
int numTokens = tokenize(textSplit.get(i - 1), inputLanguage, outputLanguage, true).getInputIDs().length;
672+
int numTokens2 = tokenize(textSplit.get(i), inputLanguage, outputLanguage, true).getInputIDs().length;
673673
if ((numTokens + numTokens2 < maxLength) || (numTokens2 < 5)) {
674674
textSplit.set(i - 1, textSplit.get(i - 1) + textSplit.get(i));
675675
textSplit.remove(i);
@@ -1019,6 +1019,7 @@ public void executeCacheDecoderGreedy(String textToTranslate, TokenizerResult in
10191019
}
10201020
}
10211021

1022+
//todo: now beam search not work for hy-mt, fix that
10221023
public void executeCacheDecoder(String textToTranslate, TokenizerResult input, @Nullable OnnxTensor encoderResult, ArrayList<Integer>[] completeBeamOutput, @Nullable double[] beamsOutputsProbabilities, final CustomLocale outputLanguage, int beamSize, @Nullable final TranslateListener responseListener) {
10231024
int eos;
10241025
if(mode == HY_MT){
@@ -1201,7 +1202,8 @@ public void executeCacheDecoder(String textToTranslate, TokenizerResult input, @
12011202
//based on the logits, we initialize max, completeBeamOutput and beamsOutputsProbabilities
12021203
initBeamSearchData(logits, beamSize, max, completeBeamOutput, beamsOutputsProbabilities);
12031204
}else{
1204-
max[0] = Utils.getIndexOfLargest(logits[0][0]);
1205+
int seqLen = logits[0].length;
1206+
max[0] = Utils.getIndexOfLargest(logits[0][seqLen-1]);
12051207
completeBeamOutput[0].add(max[0]);
12061208
}
12071209

@@ -1223,7 +1225,8 @@ public void executeCacheDecoder(String textToTranslate, TokenizerResult input, @
12231225
int[] maxProbabilities = new int[beamSize];
12241226
cacheContainer = updateBeamSearchData(logits, beamSize, eos, result, j, nLayers, nHeads, hiddenSizeAttention, cacheContainer, maxProbabilities, beamMax, max, completeBeamOutput, beamsOutputsProbabilities);
12251227
}else{
1226-
max[0] = Utils.getIndexOfLargest(logits[0][0]);
1228+
int seqLen = logits[0].length;
1229+
max[0] = Utils.getIndexOfLargest(logits[0][seqLen-1]);
12271230
completeBeamOutput[0].add(max[0]);
12281231
}
12291232

@@ -1391,16 +1394,17 @@ private OrtSession.Result batchDecoderKvCache(OrtSession.Result result, OnnxTens
13911394

13921395
private void initBeamSearchData(float [][][] logits, int beamSize, int[] max, ArrayList<Integer>[] completeBeamOutput, double[] beamsOutputsProbabilities){
13931396
//the "beamSize" words with highest probability are inserted into max and added to completeBeamOutput
1397+
int seqLen = logits[0].length;
13941398
ArrayList<Integer> indexesToAvoid = new ArrayList<>();
13951399
for (int i = 0; i < beamSize; i++) {
1396-
max[i] = Utils.getIndexOfLargest(logits[0][0], indexesToAvoid);
1400+
max[i] = Utils.getIndexOfLargest(logits[0][seqLen-1], indexesToAvoid);
13971401
indexesToAvoid.add(max[i]);
13981402
completeBeamOutput[i].add(max[i]);
13991403
}
14001404
//we insert the initial probabilities of the "beamSize" output strings into beamsOutputsProbabilities
14011405
for (int i = 0; i < beamSize; i++) {
1402-
float maxLogit = logits[0][0][max[i]];
1403-
beamsOutputsProbabilities[i] = maxLogit - Utils.logSumExpFast(logits[0][0]);
1406+
float maxLogit = logits[0][seqLen-1][max[i]];
1407+
beamsOutputsProbabilities[i] = maxLogit - Utils.logSumExpFast(logits[0][seqLen-1]);
14041408
}
14051409
}
14061410

@@ -1412,7 +1416,8 @@ private CacheContainerNative updateBeamSearchData(
14121416
for(int k=0; k < beamSize; k++) {
14131417
ArrayList<Integer> indexesToAvoid = new ArrayList<>();
14141418
for (int i = 0; i < beamSize; i++) {
1415-
beamMax[k][i] = Utils.getIndexOfLargest(logits[k][0], indexesToAvoid);
1419+
int seqLen = logits[k].length;
1420+
beamMax[k][i] = Utils.getIndexOfLargest(logits[k][seqLen-1], indexesToAvoid);
14161421
indexesToAvoid.add(beamMax[k][i]);
14171422
}
14181423
}
@@ -1422,9 +1427,10 @@ private CacheContainerNative updateBeamSearchData(
14221427
double[] beamsOutputsProbabilitiesTemp = new double[beamSize*beamSize];
14231428
for(int k=0; k < beamSize; k++) {
14241429
//new version of probability calculation (logSumExp)
1425-
double logSumExp = Utils.logSumExpFast(logits[k][0]);
1430+
int seqLen = logits[k].length;
1431+
double logSumExp = Utils.logSumExpFast(logits[k][seqLen-1]);
14261432
for (int i = 0; i < beamSize; i++) {
1427-
float maxLogit = logits[k][0][beamMax[k][i]];
1433+
float maxLogit = logits[k][seqLen-1][beamMax[k][i]];
14281434
beamsOutputsProbabilitiesTemp[(k*beamSize)+i] = beamsOutputsProbabilities[k] + maxLogit - logSumExp;
14291435
if(beamMax[k][i] == eos){
14301436
beamsOutputsProbabilitiesTemp[(k*beamSize)+i] = beamsOutputsProbabilitiesTemp[(k*beamSize)+i]/EOS_PENALTY;
@@ -1494,7 +1500,17 @@ private TokenizerResult tokenize(String text, final CustomLocale inputLanguage,
14941500
} else if(mode == NLLB_CACHE || mode == NLLB){
14951501
return tokenizer.tokenize(text, getNllbLanguageCode(inputLanguage.getCode()), getNllbLanguageCode(outputLanguage.getCode()));
14961502
}else{ //if mode == HY_MT
1497-
return tokenizer.tokenize(text, inputLanguage.getCode(), outputLanguage.getCode(), getHyLanguageInfo(inputLanguage.getCode()), getHyLanguageInfo(outputLanguage.getCode()));
1503+
return tokenizer.tokenize(text, inputLanguage.getCode(), outputLanguage.getCode(), getHyLanguageInfo(inputLanguage.getCode()), getHyLanguageInfo(outputLanguage.getCode()), false);
1504+
}
1505+
}
1506+
1507+
private TokenizerResult tokenize(String text, final CustomLocale inputLanguage, final CustomLocale outputLanguage, boolean excludePrompt){
1508+
if (mode == MADLAD_CACHE || mode == MADLAD) {
1509+
return tokenizer.tokenize(text, inputLanguage.getCode(), outputLanguage.getCode());
1510+
} else if(mode == NLLB_CACHE || mode == NLLB){
1511+
return tokenizer.tokenize(text, getNllbLanguageCode(inputLanguage.getCode()), getNllbLanguageCode(outputLanguage.getCode()));
1512+
}else{ //if mode == HY_MT
1513+
return tokenizer.tokenize(text, inputLanguage.getCode(), outputLanguage.getCode(), getHyLanguageInfo(inputLanguage.getCode()), getHyLanguageInfo(outputLanguage.getCode()), excludePrompt);
14981514
}
14991515
}
15001516

0 commit comments

Comments
 (0)