@@ -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