From f7a59cc3aa2ed01e533b8d75a06e4d2130ea31d2 Mon Sep 17 00:00:00 2001 From: dengwenjie <775747758@qq.com> Date: Tue, 17 Jun 2025 11:03:34 +0800 Subject: [PATCH] =?UTF-8?q?=E6=96=B0=E5=A2=9E=E7=BF=BB=E8=AF=91=E6=A8=A1?= =?UTF-8?q?=E5=9D=97=EF=BC=9A=E6=94=AF=E6=8C=81200=E5=A4=9A=E8=AF=AD?= =?UTF-8?q?=E8=A8=80=E7=BF=BB=E8=AF=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../examples/nlp/TextTranslationCPU.java | 57 ---- .../nlp/generate/BatchTensorList.java | 174 ------------ .../examples/nlp/generate/CausalLMOutput.java | 32 --- .../nlp/generate/GreedyBatchTensorList.java | 84 ------ .../examples/nlp/generate/SearchConfig.java | 104 ------- .../nlp/model/Decoder2Translator.java | 45 ---- .../examples/nlp/model/DecoderTranslator.java | 44 --- .../smartai/examples/nlp/model/NllbModel.java | 183 ------------- .../java/cn/smartjavaai/common/entity/R.java | 1 + .../common/pool/CommonPredictorFactory.java | 38 +++ .../common/pool/PredictorFactory.java | 1 + .../common/pool/ZooModelFactory.java | 35 --- smartjavaai-translate/pom.xml | 8 +- ...earchConfig.java => NllbSearchConfig.java} | 6 +- ...onfig.java => TranslationModelConfig.java} | 15 +- .../translation/entity/DirectionInfo.java | 41 --- .../translation/entity/IdCardInfo.java | 12 - .../translation/entity/TranslateParam.java | 50 ++++ .../translation/enums/LanguageCode.java | 236 ++++++++++++++++ ...ModeEnum.java => TranslationModeEnum.java} | 8 +- .../factory/TranslationModelFactory.java | 37 ++- .../translation/model/NllbModel.java | 255 ++++++++++++++++++ .../translation/model/TranslationModel.java | 35 +++ .../model/common/TracedTranslationModel.java | 254 ----------------- .../model/common/TranslationCommonModel.java | 32 --- .../common/translator/EncoderTranslator.java | 49 ---- .../NllbDecoder2Translator.java} | 8 +- .../NllbDecoderTranslator.java} | 8 +- .../translator/NllbEncoderTranslator.java | 8 +- .../translation/utils/TokenUtils.java | 7 +- 30 files changed, 661 insertions(+), 1206 deletions(-) delete mode 100644 examples/src/main/java/smartai/examples/nlp/TextTranslationCPU.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/generate/BatchTensorList.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/generate/CausalLMOutput.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/generate/GreedyBatchTensorList.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/generate/SearchConfig.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/model/Decoder2Translator.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/model/DecoderTranslator.java delete mode 100644 examples/src/main/java/smartai/examples/nlp/model/NllbModel.java create mode 100644 smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java delete mode 100644 smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/ZooModelFactory.java rename smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/{SearchConfig.java => NllbSearchConfig.java} (97%) rename smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/{MachineTranslationModelConfig.java => TranslationModelConfig.java} (53%) delete mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/DirectionInfo.java delete mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/IdCardInfo.java create mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java create mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java rename smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/{MachineTranslationModeEnum.java => TranslationModeEnum.java} (70%) create mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java create mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java delete mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TracedTranslationModel.java delete mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TranslationCommonModel.java delete mode 100644 smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/EncoderTranslator.java rename smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/{common/translator/Decoder2Translator.java => translator/NllbDecoder2Translator.java} (85%) rename smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/{common/translator/DecoderTranslator.java => translator/NllbDecoderTranslator.java} (85%) rename examples/src/main/java/smartai/examples/nlp/model/EncoderTranslator.java => smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbEncoderTranslator.java (87%) diff --git a/examples/src/main/java/smartai/examples/nlp/TextTranslationCPU.java b/examples/src/main/java/smartai/examples/nlp/TextTranslationCPU.java deleted file mode 100644 index 67d5710..0000000 --- a/examples/src/main/java/smartai/examples/nlp/TextTranslationCPU.java +++ /dev/null @@ -1,57 +0,0 @@ -package smartai.examples.nlb; - -import ai.djl.Device; -import ai.djl.ModelException; -import ai.djl.translate.TranslateException; -import smartai.examples.nlb.generate.SearchConfig; -import smartai.examples.nlb.model.NllbModel; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - -import java.io.IOException; - -/** - * 文本翻译,支持202种语言互译 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public final class TextTranslationCPU { - - private static final Logger logger = LoggerFactory.getLogger(TextTranslationCPU.class); - - private TextTranslationCPU() { - } - - public static void main(String[] args) throws ModelException, IOException, - TranslateException { - - SearchConfig config = new SearchConfig(); - // 设置输出文字的最大长度 - config.setMaxSeqLength(128); - // 设置源语言:中文 "zho_Hans": 256200 - config.setSrcLangId(256200); - // 设置目标语言:英文 "eng_Latn": 256047 - config.setForcedBosTokenId(256047); - config.setForcedBosTokenId(256201); - - // 输入文字 - String input = "智利北部的丘基卡马塔矿是世界上最大的露天矿之一,长约4公里,宽3公里,深1公里。"; - - String modelPath = "E:\\ai\\models\\nlp\\"; - String cpuModelName = "traced_translation_cpu.pt"; - String gpuModelName = "traced_translation_gpu.pt"; - try (NllbModel nllbModel = new NllbModel(config, modelPath, cpuModelName, Device.cpu())) { - - System.setProperty("ai.djl.pytorch.graph_optimizer", "false"); - - // 运行模型,获取翻译结果 - String result = nllbModel.translate(input); - - logger.info("result========={}", result); - } finally { - System.clearProperty("ai.djl.pytorch.graph_optimizer"); - } - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/generate/BatchTensorList.java b/examples/src/main/java/smartai/examples/nlp/generate/BatchTensorList.java deleted file mode 100644 index 03ed957..0000000 --- a/examples/src/main/java/smartai/examples/nlp/generate/BatchTensorList.java +++ /dev/null @@ -1,174 +0,0 @@ -/* - * Copyright 2023 Amazon.com, Inc. or its affiliates. All Rights Reserved. - * - * Licensed under the Apache License, Version 2.0 (the "License"). You may not use this file except in compliance - * with the License. A copy of the License is located at - * - * http://aws.amazon.com/apache2.0/ - * - * or in the "license" file accompanying this file. This file is distributed on an "AS IS" BASIS, WITHOUT WARRANTIES - * OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific language governing permissions - * and limitations under the License. - */ -package smartai.examples.nlb.generate; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; - -/** - * BatchTensorList represents a search state, and the NDArrays inside are updated in each iteration - * of the autoregressive loop. - * - *

It is a struct consisting of NDArrays, whose first dimension is batch, and also contains - * sequence dimension (whose position in tensor's shape is specified by seqDimOrder). The SeqBatcher - * batch operations will operate on these two dimensions. - */ -public abstract class BatchTensorList { - // [batch, seq_past]. seq-dim-size == |past_seq| + |inputIds|. Will grow. - private NDArray pastOutputIds; - - // [batch, seq_past] - // The cache of past attentionMask. seq-dim-size == |past_seq| + |inputIds|. Will grow. - private NDArray pastAttentionMask; - - // (k, v) * numLayer, - // kv: [batch, heads, seq_past, kvfeature] - // The cache of past sequence. seq-dim-size == |past_seq| + |inputIds|. Will grow. - private NDList pastKeyValues; - - // Sequence dimension order among all dimensions for each element in the batch list. - private long[] seqDimOrder; - - BatchTensorList() {} - - /** - * Constructs a new {@code BatchTensorList} instance. - * - * @param list the NDList that contains the serialized version of the batch tensors - * @param seqDimOrder the sequence dimension order that specifies where the sequence dimension - * is in a tensor's shape - */ - BatchTensorList(NDList list, long[] seqDimOrder) { - this.seqDimOrder = seqDimOrder; - pastOutputIds = list.get(0); - pastAttentionMask = list.get(1); - pastKeyValues = list.subNDList(2); - } - - /** - * Constructs a new {@code BatchTensorList} instance. - * - * @param pastOutputIds past output token ids - * @param pastAttentionMask past attention mask - * @param pastKeyValues past kv cache - * @param seqDimOrder the sequence dimension order that specifies where the sequence dimension - * is in a tensor's shape - */ - BatchTensorList( - NDArray pastOutputIds, - NDArray pastAttentionMask, - NDList pastKeyValues, - long[] seqDimOrder) { - this.pastKeyValues = pastKeyValues; - this.pastOutputIds = pastOutputIds; - this.pastAttentionMask = pastAttentionMask; - this.seqDimOrder = seqDimOrder; - } - - /** - * Constructs a new {@code BatchTensorList} instance from the serialized version of the batch - * tensors. - * - *

The pastOutputIds has to be the first in the output list. - * - * @param inputList the serialized version of the batch tensors - * @param seqDimOrder the sequence dimension order that specifies where the sequence dimension - * is in a tensor's shape - * @return BatchTensorList - */ - public abstract BatchTensorList fromList(NDList inputList, long[] seqDimOrder); - - /** - * Returns the serialized version of the BatchTensorList. The pastOutputIds has to be the first - * in the output list. - * - * @return the NDList that contains the serialized BatchTensorList - */ - public abstract NDList getList(); - - /** - * Returns the sequence dimension order which specifies where the sequence dimension is in a - * tensor's shape. - * - * @return the sequence dimension order which specifies where the sequence dimension is in a - * tensor's shape - */ - public long[] getSeqDimOrder() { - return seqDimOrder; - } - - /** - * Returns the value of the pastOutputIds. - * - * @return the value of pastOutputIds - */ - public NDArray getPastOutputIds() { - return pastOutputIds; - } - - /** - * Sets the past output token ids. - * - * @param pastOutputIds the past output token ids - */ - public void setPastOutputIds(NDArray pastOutputIds) { - this.pastOutputIds = pastOutputIds; - } - - /** - * Returns the value of the pastAttentionMask. - * - * @return the value of pastAttentionMask - */ - public NDArray getPastAttentionMask() { - return pastAttentionMask; - } - - /** - * Sets the attention mask. - * - * @param pastAttentionMask the attention mask - */ - public void setPastAttentionMask(NDArray pastAttentionMask) { - this.pastAttentionMask = pastAttentionMask; - } - - /** - * Returns the value of the pastKeyValues. - * - * @return the value of pastKeyValues - */ - public NDList getPastKeyValues() { - return pastKeyValues; - } - - /** - * Sets the kv cache. - * - * @param pastKeyValues the kv cache - */ - public void setPastKeyValues(NDList pastKeyValues) { - this.pastKeyValues = pastKeyValues; - } - - /** - * Sets the sequence dimension order which specifies where the sequence dimension is in a - * tensor's shape. - * - * @param seqDimOrder the sequence dimension order which specifies where the sequence dimension - * is in a tensor's shape - */ - public void setSeqDimOrder(long[] seqDimOrder) { - this.seqDimOrder = seqDimOrder; - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/generate/CausalLMOutput.java b/examples/src/main/java/smartai/examples/nlp/generate/CausalLMOutput.java deleted file mode 100644 index 80b9aa4..0000000 --- a/examples/src/main/java/smartai/examples/nlp/generate/CausalLMOutput.java +++ /dev/null @@ -1,32 +0,0 @@ -package smartai.examples.nlb.generate; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -/** - * 解码输出对象 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class CausalLMOutput { - private NDArray logits; - private NDList pastKeyValuesList; - - public CausalLMOutput(NDArray logits, NDList pastKeyValues) { - this.logits = logits; - this.pastKeyValuesList = pastKeyValues; - } - - public NDArray getLogits() { - return logits; - } - - public void setLogits(NDArray logits) { - this.logits = logits; - } - - public NDList getPastKeyValuesList() { - return pastKeyValuesList; - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/generate/GreedyBatchTensorList.java b/examples/src/main/java/smartai/examples/nlp/generate/GreedyBatchTensorList.java deleted file mode 100644 index 388de14..0000000 --- a/examples/src/main/java/smartai/examples/nlp/generate/GreedyBatchTensorList.java +++ /dev/null @@ -1,84 +0,0 @@ -package smartai.examples.nlb.generate; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; - -/** - * 贪婪搜索张量对象列表 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class GreedyBatchTensorList extends BatchTensorList { - // [batch, 1] - private NDArray nextInputIds; - - private NDArray pastOutputIds; - - private NDArray encoderHiddenStates; - private NDArray attentionMask; - private NDList pastKeyValues; - - public GreedyBatchTensorList( - NDArray nextInputIds, - NDArray pastOutputIds, - NDList pastKeyValues, - NDArray encoderHiddenStates, - NDArray attentionMask) { - this.nextInputIds = nextInputIds; - this.pastKeyValues = pastKeyValues; - this.pastOutputIds = pastOutputIds; - this.attentionMask = attentionMask; - this.encoderHiddenStates = encoderHiddenStates; - } - - public GreedyBatchTensorList() {} - - public BatchTensorList fromList(NDList inputList, long[] seqDimOrder) { - return new GreedyBatchTensorList(); - } - - public NDList getList() { - return new NDList(); - } - - public NDArray getNextInputIds() { - return nextInputIds; - } - - public void setNextInputIds(NDArray nextInputIds) { - this.nextInputIds = nextInputIds; - } - public NDArray getPastOutputIds() { - return pastOutputIds; - } - - public void setPastOutputIds(NDArray pastOutputIds) { - this.pastOutputIds = pastOutputIds; - } - - public NDList getPastKeyValues() { - return pastKeyValues; - } - - public void setPastKeyValues(NDList pastKeyValues) { - this.pastKeyValues = pastKeyValues; - } - - public NDArray getEncoderHiddenStates() { - return encoderHiddenStates; - } - - public void setEncoderHiddenStates(NDArray encoderHiddenStates) { - this.encoderHiddenStates = encoderHiddenStates; - } - - public NDArray getAttentionMask() { - return attentionMask; - } - - public void setAttentionMask(NDArray attentionMask) { - this.attentionMask = attentionMask; - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/generate/SearchConfig.java b/examples/src/main/java/smartai/examples/nlp/generate/SearchConfig.java deleted file mode 100644 index a040b50..0000000 --- a/examples/src/main/java/smartai/examples/nlp/generate/SearchConfig.java +++ /dev/null @@ -1,104 +0,0 @@ -package smartai.examples.nlb.generate; -/** - * 配置信息 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class SearchConfig { - - private int maxSeqLength; - private long padTokenId; - private long eosTokenId; - private long bosTokenId; - private long decoderStartTokenId; - private float encoderRepetitionPenalty; - private long forcedBosTokenId; - private long srcLangId; - private float lengthPenalty; - public SearchConfig() { - this.maxSeqLength = 512; - this.eosTokenId = 2; - this.bosTokenId = 0; - this.padTokenId = 1; - this.decoderStartTokenId = 2; - this.encoderRepetitionPenalty = 1.0f; - this.srcLangId = 0; - this.forcedBosTokenId = 0; - this.lengthPenalty = 1.0f; - - } - - public long getSrcLangId() { - return srcLangId; - } - - public void setSrcLangId(long srcLangId) { - this.srcLangId = srcLangId; - } - - public void setEosTokenId(long eosTokenId) { - this.eosTokenId = eosTokenId; - } - - public int getMaxSeqLength() { - return maxSeqLength; - } - - public void setMaxSeqLength(int maxSeqLength) { - this.maxSeqLength = maxSeqLength; - } - - public long getPadTokenId() { - return padTokenId; - } - - public void setPadTokenId(long padTokenId) { - this.padTokenId = padTokenId; - } - - public long getEosTokenId() { - return eosTokenId; - } - - public long getDecoderStartTokenId() { - return decoderStartTokenId; - } - - public void setDecoderStartTokenId(long decoderStartTokenId) { - this.decoderStartTokenId = decoderStartTokenId; - } - - public float getEncoderRepetitionPenalty() { - return encoderRepetitionPenalty; - } - - public void setEncoderRepetitionPenalty(float encoderRepetitionPenalty) { - this.encoderRepetitionPenalty = encoderRepetitionPenalty; - } - - public long getForcedBosTokenId() { - return forcedBosTokenId; - } - - public void setForcedBosTokenId(long forcedBosTokenId) { - this.forcedBosTokenId = forcedBosTokenId; - } - - public float getLengthPenalty() { - return lengthPenalty; - } - - public void setLengthPenalty(float lengthPenalty) { - this.lengthPenalty = lengthPenalty; - } - - public long getBosTokenId() { - return bosTokenId; - } - - public void setBosTokenId(long bosTokenId) { - this.bosTokenId = bosTokenId; - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/model/Decoder2Translator.java b/examples/src/main/java/smartai/examples/nlp/model/Decoder2Translator.java deleted file mode 100644 index 2625547..0000000 --- a/examples/src/main/java/smartai/examples/nlp/model/Decoder2Translator.java +++ /dev/null @@ -1,45 +0,0 @@ -package smartai.examples.nlb.model; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.translate.NoBatchifyTranslator; -import ai.djl.translate.TranslatorContext; -import smartai.examples.nlb.generate.CausalLMOutput; - -/** - * 解碼器,參數支持 pastKeyValues - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class Decoder2Translator implements NoBatchifyTranslator { - private String tupleName; - - public Decoder2Translator() { - tupleName = "past_key_values(" + 12 + ',' + 4 + ')'; - } - - @Override - public NDList processInput(TranslatorContext ctx, NDList input) { - - NDArray placeholder = ctx.getNDManager().create(0); - placeholder.setName("module_method:decoder2"); - - input.add(placeholder); - - return input; - } - - @Override - public CausalLMOutput processOutput(TranslatorContext ctx, NDList output) { - NDArray logitsOutput = output.get(0); - NDList pastKeyValuesOutput = output.subNDList(1, 12 * 4 + 1); - - for (NDArray array : pastKeyValuesOutput) { - array.setName(tupleName); - } - - return new CausalLMOutput(logitsOutput, pastKeyValuesOutput); - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/model/DecoderTranslator.java b/examples/src/main/java/smartai/examples/nlp/model/DecoderTranslator.java deleted file mode 100644 index 9211a82..0000000 --- a/examples/src/main/java/smartai/examples/nlp/model/DecoderTranslator.java +++ /dev/null @@ -1,44 +0,0 @@ -package smartai.examples.nlb.model; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.translate.NoBatchifyTranslator; -import ai.djl.translate.TranslatorContext; -import smartai.examples.nlb.generate.CausalLMOutput; -/** - * 解碼器,參數沒有 pastKeyValues - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class DecoderTranslator implements NoBatchifyTranslator { - private String tupleName; - - public DecoderTranslator() { - tupleName = "past_key_values(" + 12 + ',' + 4 + ')'; - } - - @Override - public NDList processInput(TranslatorContext ctx, NDList input) { - - NDArray placeholder = ctx.getNDManager().create(0); - placeholder.setName("module_method:decoder"); - - input.add(placeholder); - - return input; - } - - @Override - public CausalLMOutput processOutput(TranslatorContext ctx, NDList output) { - NDArray logitsOutput = output.get(0); - NDList pastKeyValuesOutput = output.subNDList(1, 12 * 4 + 1); - - for (NDArray array : pastKeyValuesOutput) { - array.setName(tupleName); - } - - return new CausalLMOutput(logitsOutput, pastKeyValuesOutput); - } -} \ No newline at end of file diff --git a/examples/src/main/java/smartai/examples/nlp/model/NllbModel.java b/examples/src/main/java/smartai/examples/nlp/model/NllbModel.java deleted file mode 100644 index a987ab2..0000000 --- a/examples/src/main/java/smartai/examples/nlp/model/NllbModel.java +++ /dev/null @@ -1,183 +0,0 @@ -package smartai.examples.nlb.model; - -import ai.djl.Device; -import ai.djl.ModelException; -import ai.djl.huggingface.tokenizers.Encoding; -import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; -import ai.djl.inference.Predictor; -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.ndarray.NDManager; -import ai.djl.ndarray.index.NDIndex; -import ai.djl.repository.zoo.Criteria; -import ai.djl.repository.zoo.ZooModel; -import ai.djl.translate.NoopTranslator; -import ai.djl.translate.TranslateException; -import smartai.examples.nlb.generate.CausalLMOutput; -import smartai.examples.nlb.generate.GreedyBatchTensorList; -import smartai.examples.nlb.generate.SearchConfig; -import smartai.examples.nlb.tokenizer.TokenUtils; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; - -import java.io.IOException; -import java.nio.file.Paths; -import java.util.Arrays; -/** - * 模型载入及推理 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class NllbModel implements AutoCloseable { - private static final Logger logger = LoggerFactory.getLogger(NllbModel.class); - private SearchConfig config; - private ZooModel nllbModel; - private HuggingFaceTokenizer tokenizer; - private Predictor encoderPredictor; - private Predictor decoderPredictor; - private Predictor decoder2Predictor; - private NDManager manager; - - public NllbModel(SearchConfig config, String modelPath, String modelName, Device device) throws ModelException, IOException { - this.config = config; - Criteria criteria = - Criteria.builder() - .setTypes(NDList.class, NDList.class) - .optModelPath(Paths.get(modelPath + modelName)) - .optEngine("PyTorch") - .optDevice(device) - .optTranslator(new NoopTranslator()) - .build(); - - manager = NDManager.newBaseManager(device); - nllbModel = criteria.loadModel(); - tokenizer = HuggingFaceTokenizer.newInstance(Paths.get(modelPath + "tokenizer.json")); - encoderPredictor = nllbModel.newPredictor(new EncoderTranslator()); - decoderPredictor = nllbModel.newPredictor(new DecoderTranslator()); - decoder2Predictor = nllbModel.newPredictor(new Decoder2Translator()); - } - - public NDArray encoder(long[] ids) throws TranslateException { - return encoderPredictor.predict(ids); - } - - public CausalLMOutput decoder(NDList input) throws TranslateException { - return decoderPredictor.predict(input); - } - - public CausalLMOutput decoder2(NDList input) throws TranslateException { - return decoder2Predictor.predict(input); - } - - @Override - public void close() { - encoderPredictor.close(); - decoderPredictor.close(); - decoder2Predictor.close(); - nllbModel.close(); - manager.close(); - tokenizer.close(); - } - - public String translate(String input) throws TranslateException { - - Encoding encoding = tokenizer.encode(input); - long[] ids = encoding.getIds(); - // 1. Encoder - long[] inputIds = new long[ids.length]; - // 设置源语言编码 - inputIds[0] = config.getSrcLangId(); - for (int i = 0; i < ids.length - 1; i++) { - inputIds[i + 1] = ids[i]; - } - logger.info("inputIds: " + Arrays.toString(inputIds)); - long[] attentionMask = encoding.getAttentionMask(); - NDArray attentionMaskArray = manager.create(attentionMask).expandDims(0); - - NDArray encoderHiddenStates = encoder(inputIds); - - NDArray decoder_input_ids = manager.create(new long[]{config.getDecoderStartTokenId()}).reshape(1, 1); - NDList decoderInput = new NDList(decoder_input_ids, encoderHiddenStates, attentionMaskArray); - - // 2. Initial Decoder - CausalLMOutput modelOutput = decoder(decoderInput); - modelOutput.getLogits().attach(manager); - modelOutput.getPastKeyValuesList().attach(manager); - - GreedyBatchTensorList searchState = - new GreedyBatchTensorList(null, decoder_input_ids, modelOutput.getPastKeyValuesList(), encoderHiddenStates, attentionMaskArray); - - while (true) { -// try (NDScope ignore = new NDScope()) { - NDArray pastOutputIds = searchState.getPastOutputIds(); - - if (searchState.getNextInputIds() != null) { - decoderInput = new NDList(searchState.getNextInputIds(), searchState.getEncoderHiddenStates(), searchState.getAttentionMask()); - decoderInput.addAll(searchState.getPastKeyValues()); - // 3. Decoder loop - modelOutput = decoder2(decoderInput); - } - - NDArray outputIds = greedyStepGen(config, pastOutputIds, modelOutput.getLogits()); - - searchState.setNextInputIds(outputIds); - pastOutputIds = pastOutputIds.concat(outputIds, 1); - searchState.setPastOutputIds(pastOutputIds); - - searchState.setPastKeyValues(modelOutput.getPastKeyValuesList()); - - // memory management -// NDScope.unregister(outputIds, pastOutputIds); -// } - - // Termination Criteria - long id = searchState.getNextInputIds().toLongArray()[0]; - if (config.getEosTokenId() == id) { - searchState.setNextInputIds(null); - break; - } - if (searchState.getPastOutputIds() != null && searchState.getPastOutputIds().getShape().get(1) + 1 >= config.getMaxSeqLength()) { - break; - } - } - - if (searchState.getNextInputIds() == null) { - NDArray resultIds = searchState.getPastOutputIds(); - String result = TokenUtils.decode(config, tokenizer, resultIds); - return result; - } else { - NDArray resultIds = searchState.getPastOutputIds(); // .concat(searchState.getNextInputIds(), 1) - String result = TokenUtils.decode(config, tokenizer, resultIds); - return result; - } - - } - - public NDArray greedyStepGen(SearchConfig config, NDArray pastOutputIds, NDArray next_token_scores) { - next_token_scores = next_token_scores.get(":, -1, :"); - - NDArray new_next_token_scores = manager.create(next_token_scores.getShape(), next_token_scores.getDataType()); - next_token_scores.copyTo(new_next_token_scores); - - // LogitsProcessor 1. ForcedBOSTokenLogitsProcessor - // 设置目标语言 - long cur_len = pastOutputIds.getShape().getLastDimension(); - if (cur_len == 1) { - long num_tokens = new_next_token_scores.getShape().getLastDimension(); - for (long i = 0; i < num_tokens; i++) { - if (i != config.getForcedBosTokenId()) { - new_next_token_scores.set(new NDIndex(":," + i), Float.NEGATIVE_INFINITY); - } - } - new_next_token_scores.set(new NDIndex(":," + config.getForcedBosTokenId()), 0); - } - - NDArray probs = new_next_token_scores.softmax(-1); - NDArray next_tokens = probs.argMax(-1); - - return next_tokens.expandDims(0); - } - -} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/R.java b/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/R.java index 02ac9cb..9a316af 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/R.java +++ b/smartjavaai-common/src/main/java/cn/smartjavaai/common/entity/R.java @@ -49,6 +49,7 @@ public class R { INVALID_IMAGE(1, "图像无效"), FILE_NOT_FOUND(2, "图像文件不存在"), NO_FACE_DETECTED(3, "未检测到人脸"), + PARAM_ERROR(4, "参数错误"), Unknown(-1, "未知错误"); private final int code; diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java b/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java new file mode 100644 index 0000000..6f3550f --- /dev/null +++ b/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/CommonPredictorFactory.java @@ -0,0 +1,38 @@ +package cn.smartjavaai.common.pool; + +import ai.djl.inference.Predictor; +import ai.djl.repository.zoo.ZooModel; +import ai.djl.translate.NoBatchifyTranslator; +import org.apache.commons.pool2.BasePooledObjectFactory; +import org.apache.commons.pool2.PooledObject; +import org.apache.commons.pool2.impl.DefaultPooledObject; + +/** + * @author dwj + * @date 2025/6/14 + */ +public class CommonPredictorFactory extends BasePooledObjectFactory> { + + private final ZooModel model; + private final NoBatchifyTranslator translator; + + public CommonPredictorFactory(ZooModel model, NoBatchifyTranslator translator) { + this.model = model; + this.translator = translator; + } + + @Override + public Predictor create() { + return model.newPredictor(translator); + } + + @Override + public PooledObject> wrap(Predictor predictor) { + return new DefaultPooledObject<>(predictor); + } + + @Override + public void destroyObject(PooledObject> p) { + p.getObject().close(); + } +} diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java b/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java index 582ddca..b71e513 100644 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java +++ b/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/PredictorFactory.java @@ -2,6 +2,7 @@ package cn.smartjavaai.common.pool; import ai.djl.inference.Predictor; import ai.djl.repository.zoo.ZooModel; +import ai.djl.translate.Translator; import org.apache.commons.pool2.BasePooledObjectFactory; import org.apache.commons.pool2.PooledObject; import org.apache.commons.pool2.impl.DefaultPooledObject; diff --git a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/ZooModelFactory.java b/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/ZooModelFactory.java deleted file mode 100644 index 4c8bdaf..0000000 --- a/smartjavaai-common/src/main/java/cn/smartjavaai/common/pool/ZooModelFactory.java +++ /dev/null @@ -1,35 +0,0 @@ -package cn.smartjavaai.common.pool; - -import ai.djl.inference.Predictor; -import ai.djl.repository.zoo.ZooModel; -import org.apache.commons.pool2.BasePooledObjectFactory; -import org.apache.commons.pool2.PooledObject; -import org.apache.commons.pool2.impl.DefaultPooledObject; - -/** - * ZooModel 工厂类 - * @author lwx - * @date 2025/6/06 - */ -public class ZooModelFactory extends BasePooledObjectFactory> { - private final ZooModel model; - - public ZooModelFactory(ZooModel model) { - this.model = model; - } - - @Override - public ZooModel create() { - return model; - } - - @Override - public PooledObject> wrap(ZooModel predictor) { - return new DefaultPooledObject<>(predictor); - } - - @Override - public void destroyObject(PooledObject> p) { - p.getObject().close(); - } -} diff --git a/smartjavaai-translate/pom.xml b/smartjavaai-translate/pom.xml index a957198..9916afc 100644 --- a/smartjavaai-translate/pom.xml +++ b/smartjavaai-translate/pom.xml @@ -18,15 +18,9 @@ smartjavaai-common ${project.version} - - junit - junit - 4.13.1 - compile - - 1.0.15 + 1.0.16 smartjavaai-ocr SmartJavaAI https://github.com/geekwenjie/SmartJavaAI diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/SearchConfig.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java similarity index 97% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/SearchConfig.java rename to smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java index 06799af..4fd3522 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/SearchConfig.java +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/NllbSearchConfig.java @@ -4,7 +4,7 @@ package cn.smartjavaai.translation.config; * @author lwx * @date 2025/6/05 */ -public class SearchConfig { +public class NllbSearchConfig { private int maxSeqLength; private long padTokenId; @@ -15,7 +15,7 @@ public class SearchConfig { private long forcedBosTokenId; private long srcLangId; private float lengthPenalty; - public SearchConfig() { + public NllbSearchConfig() { this.maxSeqLength = 512; this.eosTokenId = 2; this.bosTokenId = 0; @@ -99,4 +99,4 @@ public class SearchConfig { public void setBosTokenId(long bosTokenId) { this.bosTokenId = bosTokenId; } -} \ No newline at end of file +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/MachineTranslationModelConfig.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java similarity index 53% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/MachineTranslationModelConfig.java rename to smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java index 1b9ee40..55e6b7f 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/MachineTranslationModelConfig.java +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/config/TranslationModelConfig.java @@ -2,7 +2,7 @@ package cn.smartjavaai.translation.config; import cn.smartjavaai.common.enums.DeviceEnum; -import cn.smartjavaai.translation.enums.MachineTranslationModeEnum; +import cn.smartjavaai.translation.enums.TranslationModeEnum; import lombok.Data; /** @@ -11,11 +11,11 @@ import lombok.Data; * @date 2025/6/05 */ @Data -public class MachineTranslationModelConfig { +public class TranslationModelConfig { /** * 翻译模型 */ - private MachineTranslationModeEnum modelEnum; + private TranslationModeEnum modelEnum; /** * 设备类型 @@ -26,14 +26,7 @@ public class MachineTranslationModelConfig { * 翻译模型路径 */ private String modelPath; - /** - * 翻译模型路径 - */ - private String modelName; - /** - * 翻译模型配置 - */ - private SearchConfig searchConfig; + diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/DirectionInfo.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/DirectionInfo.java deleted file mode 100644 index c0d7be2..0000000 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/DirectionInfo.java +++ /dev/null @@ -1,41 +0,0 @@ -package cn.smartjavaai.translation.entity; - -/** - * 方向检测结果 - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class DirectionInfo { - - /** - * 方向 0 90 180 270 - */ - private String name; - - /** - * 置信度 - */ - private Double prob; - - public DirectionInfo(String name, Double prob) { - this.name = name; - this.prob = prob; - } - - public String getName() { - return name; - } - - public void setName(String name) { - this.name = name; - } - - public Double getProb() { - return prob; - } - - public void setProb(Double prob) { - this.prob = prob; - } -} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/IdCardInfo.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/IdCardInfo.java deleted file mode 100644 index 8eae0b9..0000000 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/IdCardInfo.java +++ /dev/null @@ -1,12 +0,0 @@ -package cn.smartjavaai.translation.entity; - -/** - * 身份证信息 - * @author dwj - * @date 2025/5/22 - */ -public class IdCardInfo { - - - -} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java new file mode 100644 index 0000000..6d426c9 --- /dev/null +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/entity/TranslateParam.java @@ -0,0 +1,50 @@ +package cn.smartjavaai.translation.entity; + +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.translation.enums.LanguageCode; +import lombok.Data; +import org.apache.commons.lang3.StringUtils; + +/** + * 翻译参数 + * @author dwj + * @date 2025/6/16 + */ +@Data +public class TranslateParam { + + /** + * 输入文本 + */ + private String input; + + /** + * 源语言 + */ + private LanguageCode sourceLanguage; + + /** + * 目标语言 + */ + private LanguageCode targetLanguage; + + /** + * 参数校验方法 + * @return 如果参数有误返回 R.fail,否则返回 R.ok(null) + */ + public R validate() { + if (StringUtils.isBlank(input)) { + return R.fail(R.Status.PARAM_ERROR.getCode(), "输入文本不能为空"); + } + if (sourceLanguage == null) { + return R.fail(R.Status.PARAM_ERROR.getCode(), "源语言不能为空"); + } + if (targetLanguage == null) { + return R.fail(R.Status.PARAM_ERROR.getCode(), "目标语言不能为空"); + } + return R.ok(null); + } + + + +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java new file mode 100644 index 0000000..7a00059 --- /dev/null +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/LanguageCode.java @@ -0,0 +1,236 @@ +package cn.smartjavaai.translation.enums; + +/** + * 语言枚举 + * @author dwj + * @date 2025/6/16 + */ +public enum LanguageCode { + ACE_ARAB("ace_Arab", 256001), // 亚齐语(阿拉伯文) + ACE_LATN("ace_Latn", 256002), // 亚齐语(拉丁文) + ACM_ARAB("acm_Arab", 256003), // 美索不达米亚阿拉伯语(阿拉伯文) + ACQ_ARAB("acq_Arab", 256004), // 南也门阿拉伯语(阿拉伯文) + AEB_ARAB("aeb_Arab", 256005), // 突尼斯阿拉伯语(阿拉伯文) + AFR_LATN("afr_Latn", 256006), // 南非荷兰语(拉丁文) + AJP_ARAB("ajp_Arab", 256007), // 南黎凡特阿拉伯语(阿拉伯文) + AKA_LATN("aka_Latn", 256008), // 阿坎语(拉丁文) + AMH_ETHI("amh_Ethi", 256009), // 阿姆哈拉语(吉兹字母) + APC_ARAB("apc_Arab", 256010), // 北黎凡特阿拉伯语(阿拉伯文) + ARB_ARAB("arb_Arab", 256011), // 标准阿拉伯语(阿拉伯文) + ARS_ARAB("ars_Arab", 256012), // 纳吉迪阿拉伯语(阿拉伯文) + ARY_ARAB("ary_Arab", 256013), // 摩洛哥阿拉伯语(阿拉伯文) + ARZ_ARAB("arz_Arab", 256014), // 埃及阿拉伯语(阿拉伯文) + ASM_BENG("asm_Beng", 256015), // 阿萨姆语(孟加拉文) + AST_LATN("ast_Latn", 256016), // 阿斯图里亚斯语(拉丁文) + AWA_DEVA("awa_Deva", 256017), // 阿瓦德语(天城文) + AYR_LATN("ayr_Latn", 256018), // 南艾马拉语(拉丁文) + AZB_ARAB("azb_Arab", 256019), // 南阿塞拜疆语(阿拉伯文) + AZJ_LATN("azj_Latn", 256020), // 北阿塞拜疆语(拉丁文) + BAK_CYRL("bak_Cyrl", 256021), // 巴什基尔语(西里尔文) + BAM_LATN("bam_Latn", 256022), // 班巴拉语(拉丁文) + BAN_LATN("ban_Latn", 256023), // 巴厘语(拉丁文) + BEL_CYRL("bel_Cyrl", 256024), // 白俄罗斯语(西里尔文) + BEM_LATN("bem_Latn", 256025), // 本巴语(拉丁文) + BEN_BENG("ben_Beng", 256026), // 孟加拉语(孟加拉文) + BHO_DEVA("bho_Deva", 256027), // 博杰普尔语(天城文) + BJN_ARAB("bjn_Arab", 256028), // 班贾尔语(阿拉伯文) + BJN_LATN("bjn_Latn", 256029), // 班贾尔语(拉丁文) + BOD_TIBT("bod_Tibt", 256030), // 藏语(藏文) + BOS_LATN("bos_Latn", 256031), // 波斯尼亚语(拉丁文) + BUG_LATN("bug_Latn", 256032), // 布吉语(拉丁文) + BUL_CYRL("bul_Cyrl", 256033), // 保加利亚语(西里尔文) + CAT_LATN("cat_Latn", 256034), // 加泰罗尼亚语(拉丁文) + CEB_LATN("ceb_Latn", 256035), // 宿务语(拉丁文) + CES_LATN("ces_Latn", 256036), // 捷克语(拉丁文) + CJK_LATN("cjk_Latn", 256037), // 琼卡语(拉丁文) + CKB_ARAB("ckb_Arab", 256038), // 中库尔德语(阿拉伯文) + CRH_LATN("crh_Latn", 256039), // 克里米亚鞑靼语(拉丁文) + CYM_LATN("cym_Latn", 256040), // 威尔士语(拉丁文) + DAN_LATN("dan_Latn", 256041), // 丹麦语(拉丁文) + DEU_LATN("deu_Latn", 256042), // 德语(拉丁文) + DIK_LATN("dik_Latn", 256043), // 南丁卡语(拉丁文) + DYU_LATN("dyu_Latn", 256044), // 迪尤拉语(拉丁文) + DZO_TIBT("dzo_Tibt", 256045), // 宗喀语(藏文) + ELL_GREK("ell_Grek", 256046), // 希腊语(希腊文) + ENG_LATN("eng_Latn", 256047), // 英语(拉丁文) + EPO_LATN("epo_Latn", 256048), // 世界语(拉丁文) + EST_LATN("est_Latn", 256049), // 爱沙尼亚语(拉丁文) + EUS_LATN("eus_Latn", 256050), // 巴斯克语(拉丁文) + EWE_LATN("ewe_Latn", 256051), // 埃维语(拉丁文) + FAO_LATN("fao_Latn", 256052), // 法罗语(拉丁文) + PES_ARAB("pes_Arab", 256053), // 波斯语(阿拉伯文) + FIJ_LATN("fij_Latn", 256054), // 斐济语(拉丁文) + FIN_LATN("fin_Latn", 256055), // 芬兰语(拉丁文) + FON_LATN("fon_Latn", 256056), // 丰语(拉丁文) + FRA_LATN("fra_Latn", 256057), // 法语(拉丁文) + FUR_LATN("fur_Latn", 256058), // 弗留利语(拉丁文) + FUV_LATN("fuv_Latn", 256059), // 富拉语(拉丁文) + GLA_LATN("gla_Latn", 256060), // 苏格兰盖尔语(拉丁文) + GLE_LATN("gle_Latn", 256061), // 爱尔兰语(拉丁文) + GLG_LATN("glg_Latn", 256062), // 加利西亚语(拉丁文) + GRN_LATN("grn_Latn", 256063), // 瓜拉尼语(拉丁文) + GUJ_GUJR("guj_Gujr", 256064), // 古吉拉特语(古吉拉特文) + HAT_LATN("hat_Latn", 256065), // 海地克里奥尔语(拉丁文) + HAU_LATN("hau_Latn", 256066), // 豪萨语(拉丁文) + HEB_HEBR("heb_Hebr", 256067), // 希伯来语(希伯来文) + HIN_DEVA("hin_Deva", 256068), // 印地语(天城文) + HNE_DEVA("hne_Deva", 256069), // 丘德语(天城文) + HRV_LATN("hrv_Latn", 256070), // 克罗地亚语(拉丁文) + HUN_LATN("hun_Latn", 256071), // 匈牙利语(拉丁文) + HYE_ARMN("hye_Armn", 256072), // 亚美尼亚语(亚美尼亚文) + IBO_LATN("ibo_Latn", 256073), // 伊博语(拉丁文) + ILO_LATN("ilo_Latn", 256074), // 伊洛卡诺语(拉丁文) + IND_LATN("ind_Latn", 256075), // 印尼语(拉丁文) + ISL_LATN("isl_Latn", 256076), // 冰岛语(拉丁文) + ITA_LATN("ita_Latn", 256077), // 意大利语(拉丁文) + JAV_LATN("jav_Latn", 256078), // 爪哇语(拉丁文) + JPN_JPAN("jpn_Jpan", 256079), // 日语(日文) + KAB_LATN("kab_Latn", 256080), // 卡拜尔语(拉丁文) + KAC_LATN("kac_Latn", 256081), // 克钦语(拉丁文) + KAM_LATN("kam_Latn", 256082), // 卡姆巴语(拉丁文) + KAN_KNDA("kan_Knda", 256083), // 卡纳达语(卡纳达文) + KAS_ARAB("kas_Arab", 256084), // 克什米尔语(阿拉伯文) + KAS_DEVA("kas_Deva", 256085), // 克什米尔语(天城文) + KAT_GEOR("kat_Geor", 256086), // 格鲁吉亚语(格鲁吉亚文) + KNC_ARAB("knc_Arab", 256087), // 中卡努里语(阿拉伯文) + KNC_LATN("knc_Latn", 256088), // 中卡努里语(拉丁文) + KAZ_CYRL("kaz_Cyrl", 256089), // 哈萨克语(西里尔文) + KBP_LATN("kbp_Latn", 256090), // 卡比耶语(拉丁文) + KEA_LATN("kea_Latn", 256091), // 卡布佛得鲁语(拉丁文) + KHM_KHMR("khm_Khmr", 256092), // 高棉语(高棉文) + KIK_LATN("kik_Latn", 256093), // 基库尤语(拉丁文) + KIN_LATN("kin_Latn", 256094), // 卢旺达语(拉丁文) + KIR_CYRL("kir_Cyrl", 256095), // 柯尔克孜语(西里尔文) + KMB_LATN("kmb_Latn", 256096), // 金邦杜语(拉丁文) + KON_LATN("kon_Latn", 256097), // 刚果语(拉丁文) + KOR_HANG("kor_Hang", 256098), // 韩语(韩文) + KMR_LATN("kmr_Latn", 256099), // 北库尔德语(拉丁文) + LAO_LAOO("lao_Laoo", 256100), // 老挝语(老挝文) + LVS_LATN("lvs_Latn", 256101), // 标准拉脱维亚语(拉丁文) + LIJ_LATN("lij_Latn", 256102), // 利古里亚语(拉丁文) + LIM_LATN("lim_Latn", 256103), // 林堡语(拉丁文) + LIN_LATN("lin_Latn", 256104), // 林加拉语(拉丁文) + LIT_LATN("lit_Latn", 256105), // 立陶宛语(拉丁文) + LMO_LATN("lmo_Latn", 256106), // 伦巴第语(拉丁文) + LTG_LATN("ltg_Latn", 256107), // 拉特加莱语(拉丁文) + LTZ_LATN("ltz_Latn", 256108), // 卢森堡语(拉丁文) + LUA_LATN("lua_Latn", 256109), // 卢巴-卢拉语(拉丁文) + LUG_LATN("lug_Latn", 256110), // 卢干达语(拉丁文) + LUO_LATN("luo_Latn", 256111), // 卢奥语(拉丁文) + LUS_LATN("lus_Latn", 256112), // 米佐语(拉丁文) + MAG_DEVA("mag_Deva", 256113), // 摩揭陀语(天城文) + MAI_DEVA("mai_Deva", 256114), // 迈蒂利语(天城文) + MAL_MLYM("mal_Mlym", 256115), // 马拉雅拉姆语(马拉雅拉姆文) + MAR_DEVA("mar_Deva", 256116), // 马拉地语(天城文) + MIN_LATN("min_Latn", 256117), // 米南佳保语(拉丁文) + MKD_CYRL("mkd_Cyrl", 256118), // 马其顿语(西里尔文) + PLT_LATN("plt_Latn", 256119), // 高原马达加斯加语(拉丁文) + MLT_LATN("mlt_Latn", 256120), // 马耳他语(拉丁文) + MNI_BENG("mni_Beng", 256121), // 曼尼普尔语(孟加拉文) + KHK_CYRL("khk_Cyrl", 256122), // 蒙古语(西里尔文) + MOS_LATN("mos_Latn", 256123), // 莫西语(拉丁文) + MRI_LATN("mri_Latn", 256124), // 毛利语(拉丁文) + ZSM_LATN("zsm_Latn", 256125), // 标准马来语(拉丁文) + MYA_MYMR("mya_Mymr", 256126), // 缅甸语(缅甸文) + NLD_LATN("nld_Latn", 256127), // 荷兰语(拉丁文) + NNO_LATN("nno_Latn", 256128), // 新挪威语(拉丁文) + NOB_LATN("nob_Latn", 256129), // 书面挪威语(拉丁文) + NPI_DEVA("npi_Deva", 256130), // 尼泊尔语(天城文) + NSO_LATN("nso_Latn", 256131), // 北索托语(拉丁文) + NUS_LATN("nus_Latn", 256132), // 努埃尔语(拉丁文) + NYA_LATN("nya_Latn", 256133), // 齐切瓦语(拉丁文) + OCI_LATN("oci_Latn", 256134), // 奥克语(拉丁文) + GAZ_LATN("gaz_Latn", 256135), // 西奥莫罗语(拉丁文) + ORY_ORYA("ory_Orya", 256136), // 奥里亚语(奥里亚文) + PAG_LATN("pag_Latn", 256137), // 邦阿西楠语(拉丁文) + PAN_GURU("pan_Guru", 256138), // 旁遮普语(果鲁穆奇文) + PAP_LATN("pap_Latn", 256139), // 帕皮阿门托语(拉丁文) + POL_LATN("pol_Latn", 256140), // 波兰语(拉丁文) + POR_LATN("por_Latn", 256141), // 葡萄牙语(拉丁文) + PRS_ARAB("prs_Arab", 256142), // 达里语(阿拉伯文) + PBT_ARAB("pbt_Arab", 256143), // 南普什图语(阿拉伯文) + QUY_LATN("quy_Latn", 256144), // 丘尤钱卡语(拉丁文) + RON_LATN("ron_Latn", 256145), // 罗马尼亚语(拉丁文) + RUN_LATN("run_Latn", 256146), // 隆迪语(拉丁文) + RUS_CYRL("rus_Cyrl", 256147), // 俄语(西里尔文) + SAG_LATN("sag_Latn", 256148), // 桑戈语(拉丁文) + SAN_DEVA("san_Deva", 256149), // 梵语(天城文) + SAT_BENG("sat_Beng", 256150), // 桑塔利语(孟加拉文) + SCN_LATN("scn_Latn", 256151), // 西西里语(拉丁文) + SHN_MYMR("shn_Mymr", 256152), // 掸语(缅甸文) + SIN_SINH("sin_Sinh", 256153), // 僧伽罗语(僧伽罗文) + SLK_LATN("slk_Latn", 256154), // 斯洛伐克语(拉丁文) + SLV_LATN("slv_Latn", 256155), // 斯洛文尼亚语(拉丁文) + SMO_LATN("smo_Latn", 256156), // 萨摩亚语(拉丁文) + SNA_LATN("sna_Latn", 256157), // 绍纳语(拉丁文) + SND_ARAB("snd_Arab", 256158), // 信德语(阿拉伯文) + SOM_LATN("som_Latn", 256159), // 索马里语(拉丁文) + SOT_LATN("sot_Latn", 256160), // 南索托语(拉丁文) + SPA_LATN("spa_Latn", 256161), // 西班牙语(拉丁文) + ALS_LATN("als_Latn", 256162), // 托萨语(拉丁文) + SRD_LATN("srd_Latn", 256163), // 撒丁语(拉丁文) + SRP_CYRL("srp_Cyrl", 256164), // 塞尔维亚语(西里尔文) + SSW_LATN("ssw_Latn", 256165), // 斯威士语(拉丁文) + SUN_LATN("sun_Latn", 256166), // 巽他语(拉丁文) + SWE_LATN("swe_Latn", 256167), // 瑞典语(拉丁文) + SWH_LATN("swh_Latn", 256168), // 斯瓦希里语(拉丁文) + SZL_LATN("szl_Latn", 256169), // 西里西亚语(拉丁文) + TAM_TAML("tam_Taml", 256170), // 泰米尔语(泰米尔文) + TAT_CYRL("tat_Cyrl", 256171), // 鞑靼语(西里尔文) + TEL_TELU("tel_Telu", 256172), // 泰卢固语(泰卢固文) + TGK_CYRL("tgk_Cyrl", 256173), // 塔吉克语(西里尔文) + TGL_LATN("tgl_Latn", 256174), // 他加禄语(拉丁文) + THA_THAI("tha_Thai", 256175), // 泰语(泰文) + TIR_ETHI("tir_Ethi", 256176), // 提格里尼亚语(吉兹字母) + TAQ_LATN("taq_Latn", 256177), // 塔马舍克语(拉丁文) + TAQ_TFNG("taq_Tfng", 256178), // 塔马舍克语(提非纳文) + TPI_LATN("tpi_Latn", 256179), // 托克皮辛语(拉丁文) + TSN_LATN("tsn_Latn", 256180), // 茨瓦纳语(拉丁文) + TSO_LATN("tso_Latn", 256181), // 聪加语(拉丁文) + TUK_LATN("tuk_Latn", 256182), // 土库曼语(拉丁文) + TUM_LATN("tum_Latn", 256183), // 通布卡语(拉丁文) + TUR_LATN("tur_Latn", 256184), // 土耳其语(拉丁文) + TWI_LATN("twi_Latn", 256185), // 契维语(拉丁文) + TZM_TFNG("tzm_Tfng", 256186), // 塔马齐格特语(提非纳文) + UIG_ARAB("uig_Arab", 256187), // 维吾尔语(阿拉伯文) + UKR_CYRL("ukr_Cyrl", 256188), // 乌克兰语(西里尔文) + UMB_LATN("umb_Latn", 256189), // 翁本杜语(拉丁文) + URD_ARAB("urd_Arab", 256190), // 乌尔都语(阿拉伯文) + UZN_LATN("uzn_Latn", 256191), // 乌兹别克语(拉丁文) + VEC_LATN("vec_Latn", 256192), // 威尼斯语(拉丁文) + VIE_LATN("vie_Latn", 256193), // 越南语(拉丁文) + WAR_LATN("war_Latn", 256194), // 瓦瑞语(拉丁文) + WOL_LATN("wol_Latn", 256195), // 沃洛夫语(拉丁文) + XHO_LATN("xho_Latn", 256196), // 科萨语(拉丁文) + YDD_HEBR("ydd_Hebr", 256197), // 东意第绪语(希伯来文) + YOR_LATN("yor_Latn", 256198), // 约鲁巴语(拉丁文) + YUE_HANT("yue_Hant", 256199), // 粤语(繁体中文) + ZHO_HANS("zho_Hans", 256200), // 中文(简体中文) + ZHO_HANT("zho_Hant", 256201), // 中文(繁体中文) + ZUL_LATN("zul_Latn", 256202); // 祖鲁语(拉丁文) + + private final String code; + private final int id; + + LanguageCode(String code, int id) { + this.code = code; + this.id = id; + } + + public String getCode() { return code; } + public int getId() { return id; } + + public static LanguageCode fromCode(String code) { + for (LanguageCode lang : values()) { + if (lang.code.equals(code)) return lang; + } + throw new IllegalArgumentException("Unknown language code: " + code); + } + + public static LanguageCode fromId(int id) { + for (LanguageCode lang : values()) { + if (lang.id == id) return lang; + } + throw new IllegalArgumentException("Unknown language ID: " + id); + } +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/MachineTranslationModeEnum.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java similarity index 70% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/MachineTranslationModeEnum.java rename to smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java index 1ffa2fd..e558a48 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/MachineTranslationModeEnum.java +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/enums/TranslationModeEnum.java @@ -4,16 +4,16 @@ package cn.smartjavaai.translation.enums; * @author lwx * @date 2025/6/05 */ -public enum MachineTranslationModeEnum { +public enum TranslationModeEnum { - TRACED_TRANSLATION_CPU; + NLLB_MODEL; /** * 根据名称获取枚举 (忽略大小写和下划线变体) */ - public static MachineTranslationModeEnum fromName(String name) { + public static TranslationModeEnum fromName(String name) { String formatted = name.trim().toUpperCase().replaceAll("[-_]", ""); - for (MachineTranslationModeEnum model : values()) { + for (TranslationModeEnum model : values()) { if (model.name().replaceAll("_", "").equals(formatted)) { return model; } diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java index 1cab02b..1d7a832 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/factory/TranslationModelFactory.java @@ -3,10 +3,10 @@ package cn.smartjavaai.translation.factory; import cn.smartjavaai.common.config.Config; -import cn.smartjavaai.translation.config.MachineTranslationModelConfig; +import cn.smartjavaai.translation.config.TranslationModelConfig; import cn.smartjavaai.translation.exception.TranslationException; -import cn.smartjavaai.translation.model.common.TracedTranslationModel; -import cn.smartjavaai.translation.model.common.TranslationCommonModel; +import cn.smartjavaai.translation.model.NllbModel; +import cn.smartjavaai.translation.model.TranslationModel; import lombok.extern.slf4j.Slf4j; import java.util.Map; @@ -23,14 +23,14 @@ public class TranslationModelFactory { // 使用 volatile 和双重检查锁定来确保线程安全的单例模式 private static volatile TranslationModelFactory instance; - private static final ConcurrentHashMap commonDetModelMap = new ConcurrentHashMap<>(); + private static final ConcurrentHashMap modelMap = new ConcurrentHashMap<>(); /** * 检测模型注册表 */ - private static final Map> commonDetRegistry = + private static final Map> modelRegistry = new ConcurrentHashMap<>(); @@ -48,42 +48,42 @@ public class TranslationModelFactory { /** - * 注册通用检测模型 + * 注册翻译模型 * @param name * @param clazz */ - private static void registerCommonDetModel(String name, Class clazz) { - commonDetRegistry.put(name.toLowerCase(), clazz); + private static void registerCommonDetModel(String name, Class clazz) { + modelRegistry.put(name.toLowerCase(), clazz); } /** - * 获取检测模型(通过配置) + * 获取翻译模型(通过配置) * @param config * @return */ - public TranslationCommonModel getDetModel(MachineTranslationModelConfig config) { + public TranslationModel getModel(TranslationModelConfig config) { if(Objects.isNull(config) || Objects.isNull(config.getModelEnum())){ throw new TranslationException("未配置OCR模型"); } - return commonDetModelMap.computeIfAbsent(config.getModelEnum().name(), k -> { - return createCommonDetModel(config); + return modelMap.computeIfAbsent(config.getModelEnum().name(), k -> { + return createModel(config); }); } /** - * 创建OCR通用检测模型 + * 创建翻译模型 * @param config * @return */ - private TranslationCommonModel createCommonDetModel(MachineTranslationModelConfig config) { - Class clazz = commonDetRegistry.get(config.getModelEnum().name().toLowerCase()); + private TranslationModel createModel(TranslationModelConfig config) { + Class clazz = modelRegistry.get(config.getModelEnum().name().toLowerCase()); if(clazz == null){ throw new TranslationException("Unsupported model"); } - TranslationCommonModel model = null; + TranslationModel model = null; try { - model = (TranslationCommonModel) clazz.newInstance(); + model = (TranslationModel) clazz.newInstance(); } catch (InstantiationException | IllegalAccessException e) { throw new TranslationException(e); } @@ -94,8 +94,7 @@ public class TranslationModelFactory { // 初始化默认算法 static { - registerCommonDetModel("TRACED_TRANSLATION_CPU", TracedTranslationModel.class); - + registerCommonDetModel("NLLB_MODEL", NllbModel.class); log.info("缓存目录:{}", Config.getCachePath()); } diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java new file mode 100644 index 0000000..fa35bfb --- /dev/null +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/NllbModel.java @@ -0,0 +1,255 @@ +package cn.smartjavaai.translation.model; + +import ai.djl.Device; +import ai.djl.MalformedModelException; +import ai.djl.engine.Engine; +import ai.djl.huggingface.tokenizers.Encoding; +import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; +import ai.djl.inference.Predictor; +import ai.djl.ndarray.NDArray; +import ai.djl.ndarray.NDList; +import ai.djl.ndarray.NDManager; +import ai.djl.ndarray.index.NDIndex; +import ai.djl.repository.zoo.Criteria; +import ai.djl.repository.zoo.ModelNotFoundException; +import ai.djl.repository.zoo.ModelZoo; +import ai.djl.repository.zoo.ZooModel; +import ai.djl.translate.NoopTranslator; +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.common.enums.DeviceEnum; +import cn.smartjavaai.common.pool.CommonPredictorFactory; +import cn.smartjavaai.translation.config.TranslationModelConfig; +import cn.smartjavaai.translation.config.NllbSearchConfig; +import cn.smartjavaai.translation.entity.CausalLMOutput; +import cn.smartjavaai.translation.entity.GreedyBatchTensorList; +import cn.smartjavaai.translation.entity.TranslateParam; +import cn.smartjavaai.translation.exception.TranslationException; +import cn.smartjavaai.translation.model.translator.NllbDecoder2Translator; +import cn.smartjavaai.translation.model.translator.NllbDecoderTranslator; +import cn.smartjavaai.translation.model.translator.NllbEncoderTranslator; +import cn.smartjavaai.translation.utils.TokenUtils; +import lombok.extern.slf4j.Slf4j; +import org.apache.commons.lang3.StringUtils; +import org.apache.commons.pool2.ObjectPool; +import org.apache.commons.pool2.impl.GenericObjectPool; + +import java.io.File; +import java.io.IOException; +import java.nio.file.Path; +import java.nio.file.Paths; +import java.util.Objects; + +/** + * 机器翻译通用检测模型 + * + * @author lwx + * @date 2025/6/05 + */ +@Slf4j +public class NllbModel implements TranslationModel{ + + private ObjectPool> encodePredictorPool; + + private ObjectPool> decodePredictorPool; + + private ObjectPool> decode2PredictorPool; + + private ZooModel nllbModel; + private HuggingFaceTokenizer tokenizer; + + private NllbSearchConfig searchConfig; + private TranslationModelConfig config; + + + + @Override + public void loadModel(TranslationModelConfig config) { + if (StringUtils.isBlank(config.getModelPath())) { + throw new TranslationException("modelPath is null"); + } + Device device = null; + if (!Objects.isNull(config.getDevice())) { + device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(); + } + this.config = config; + Path modelPath = Paths.get(config.getModelPath()); + //this.searchConfig = config.getSearchConfig(); + Criteria criteria = + Criteria.builder() + .setTypes(NDList.class, NDList.class) + .optModelPath(modelPath) + .optEngine("PyTorch") + .optDevice(device) + .optTranslator(new NoopTranslator()) + .build(); + try { + nllbModel = ModelZoo.loadModel(criteria); + encodePredictorPool = new GenericObjectPool<>(new CommonPredictorFactory(nllbModel,new NllbEncoderTranslator())); + decodePredictorPool = new GenericObjectPool<>(new CommonPredictorFactory(nllbModel,new NllbDecoderTranslator())); + decode2PredictorPool = new GenericObjectPool<>(new CommonPredictorFactory(nllbModel,new NllbDecoder2Translator())); + Path tokenizerPath = modelPath.getParent().resolve("tokenizer.json"); + tokenizer = HuggingFaceTokenizer.newInstance(tokenizerPath); + //初始化searchConfig + this.searchConfig = new NllbSearchConfig(); + log.info("当前设备: " + nllbModel.getNDManager().getDevice()); + log.info("当前引擎: " + Engine.getInstance().getEngineName()); + } catch (IOException | ModelNotFoundException | MalformedModelException e) { + throw new TranslationException("模型加载失败", e); + } + } + + @Override + public R translate(TranslateParam translateParam) { + if(translateParam == null){ + return R.fail(R.Status.PARAM_ERROR); + } + //验证 + R validateResult = translateParam.validate(); + if(!validateResult.isSuccess()){ + return validateResult; + } + //补充参数 + this.searchConfig.setSrcLangId(translateParam.getSourceLanguage().getId()); + this.searchConfig.setForcedBosTokenId(translateParam.getTargetLanguage().getId()); + return R.ok(translateLanguage(translateParam)); + } + + private String translateLanguage(TranslateParam translateParam) { + Predictor encoderPredictor = null; + Predictor decoderPredictor = null; + Predictor decoder2Predictor = null; + try (NDManager manager = NDManager.newBaseManager()) { + encoderPredictor = (Predictor)encodePredictorPool.borrowObject(); + decoderPredictor = (Predictor)decodePredictorPool.borrowObject(); + decoder2Predictor = (Predictor)decode2PredictorPool.borrowObject(); + + Encoding encoding = tokenizer.encode(translateParam.getInput()); + long[] ids = encoding.getIds(); + // 1. Encoder + long[] inputIds = new long[ids.length]; + // 设置源语言编码 + inputIds[0] = searchConfig.getSrcLangId(); + for (int i = 0; i < ids.length - 1; i++) { + inputIds[i + 1] = ids[i]; + } + + long[] attentionMask = encoding.getAttentionMask(); + NDArray attentionMaskArray = manager.create(attentionMask).expandDims(0); + + NDArray encoderHiddenStates = encoderPredictor.predict(inputIds); + + NDArray decoder_input_ids = manager.create(new long[]{searchConfig.getDecoderStartTokenId()}).reshape(1, 1); + NDList decoderInput = new NDList(decoder_input_ids, encoderHiddenStates, attentionMaskArray); + + // 2. Initial Decoder + CausalLMOutput modelOutput = decoderPredictor.predict(decoderInput); + modelOutput.getLogits().attach(manager); + modelOutput.getPastKeyValuesList().attach(manager); + + GreedyBatchTensorList searchState = + new GreedyBatchTensorList(null, decoder_input_ids, modelOutput.getPastKeyValuesList(), encoderHiddenStates, attentionMaskArray); + + while (true) { +// try (NDScope ignore = new NDScope()) { + NDArray pastOutputIds = searchState.getPastOutputIds(); + if (searchState.getNextInputIds() != null) { + decoderInput = new NDList(searchState.getNextInputIds(), searchState.getEncoderHiddenStates(), searchState.getAttentionMask()); + decoderInput.addAll(searchState.getPastKeyValues()); + // 3. Decoder loop + modelOutput = decoder2Predictor.predict(decoderInput); + } + NDArray outputIds = greedyStepGen(searchConfig, pastOutputIds, modelOutput.getLogits(), manager); + + searchState.setNextInputIds(outputIds); + pastOutputIds = pastOutputIds.concat(outputIds, 1); + searchState.setPastOutputIds(pastOutputIds); + + searchState.setPastKeyValues(modelOutput.getPastKeyValuesList()); + + long id = searchState.getNextInputIds().toLongArray()[0]; + if (searchConfig.getEosTokenId() == id) { + searchState.setNextInputIds(null); + break; + } + if (searchState.getPastOutputIds() != null && searchState.getPastOutputIds().getShape().get(1) + 1 >= searchConfig.getMaxSeqLength()) { + break; + } + } + + if (searchState.getNextInputIds() == null) { + NDArray resultIds = searchState.getPastOutputIds(); + String result = TokenUtils.decode(searchConfig, tokenizer, resultIds); + return result; + } else { + NDArray resultIds = searchState.getPastOutputIds(); // .concat(searchState.getNextInputIds(), 1) + String result = TokenUtils.decode(searchConfig, tokenizer, resultIds); + return result; + } + } catch (Exception e) { + throw new TranslationException("翻译错误", e); + } finally { + if (encoderPredictor != null) { + try { + encodePredictorPool.returnObject(encoderPredictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + encoderPredictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + if (decoderPredictor != null) { + try { + decodePredictorPool.returnObject(decoderPredictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + decoderPredictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + if (decoder2Predictor != null) { + try { + decode2PredictorPool.returnObject(decoder2Predictor); //归还 + } catch (Exception e) { + log.warn("归还Predictor失败", e); + try { + decoder2Predictor.close(); // 归还失败才销毁 + } catch (Exception ex) { + log.error("关闭Predictor失败", ex); + } + } + } + } + } + + public NDArray greedyStepGen(NllbSearchConfig config, NDArray pastOutputIds, NDArray next_token_scores, NDManager manager) { + next_token_scores = next_token_scores.get(":, -1, :"); + + NDArray new_next_token_scores = manager.create(next_token_scores.getShape(), next_token_scores.getDataType()); + next_token_scores.copyTo(new_next_token_scores); + + // LogitsProcessor 1. ForcedBOSTokenLogitsProcessor + // 设置目标语言 + long cur_len = pastOutputIds.getShape().getLastDimension(); + if (cur_len == 1) { + long num_tokens = new_next_token_scores.getShape().getLastDimension(); + for (long i = 0; i < num_tokens; i++) { + if (i != config.getForcedBosTokenId()) { + new_next_token_scores.set(new NDIndex(":," + i), Float.NEGATIVE_INFINITY); + } + } + new_next_token_scores.set(new NDIndex(":," + config.getForcedBosTokenId()), 0); + } + + NDArray probs = new_next_token_scores.softmax(-1); + NDArray next_tokens = probs.argMax(-1); + + return next_tokens.expandDims(0); + } + +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java new file mode 100644 index 0000000..a663f02 --- /dev/null +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/TranslationModel.java @@ -0,0 +1,35 @@ +package cn.smartjavaai.translation.model; + +import cn.smartjavaai.common.entity.R; +import cn.smartjavaai.translation.config.TranslationModelConfig; +import cn.smartjavaai.translation.entity.TranslateParam; + +/** + * 机器翻译通用检测模型 + * @author lwx + * @date 2025/6/05 + */ +public interface TranslationModel { + + /** + * 加载模型 + * @param config + */ + void loadModel(TranslationModelConfig config); // 加载模型 + + + /** + * 机器翻译 + * @param translateParam 翻译参数 + * @return + */ + default R translate(TranslateParam translateParam) { + throw new UnsupportedOperationException("默认不支持该功能"); + } + + + + + + +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TracedTranslationModel.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TracedTranslationModel.java deleted file mode 100644 index af01b8f..0000000 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TracedTranslationModel.java +++ /dev/null @@ -1,254 +0,0 @@ -package cn.smartjavaai.translation.model.common; - -import ai.djl.Device; -import ai.djl.MalformedModelException; -import ai.djl.engine.Engine; -import ai.djl.huggingface.tokenizers.Encoding; -import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; -import ai.djl.inference.Predictor; -import ai.djl.modality.cv.Image; -import ai.djl.modality.cv.output.DetectedObjects; -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.ndarray.NDManager; -import ai.djl.ndarray.index.NDIndex; -import ai.djl.repository.zoo.Criteria; -import ai.djl.repository.zoo.ModelNotFoundException; -import ai.djl.repository.zoo.ModelZoo; -import ai.djl.repository.zoo.ZooModel; -import ai.djl.translate.NoopTranslator; -import ai.djl.translate.TranslateException; -import cn.smartjavaai.common.config.Config; -import cn.smartjavaai.common.enums.DeviceEnum; -import cn.smartjavaai.common.pool.PredictorFactory; -import cn.smartjavaai.common.pool.ZooModelFactory; -import cn.smartjavaai.translation.config.MachineTranslationModelConfig; -import cn.smartjavaai.translation.config.SearchConfig; -import cn.smartjavaai.translation.entity.CausalLMOutput; -import cn.smartjavaai.translation.entity.GreedyBatchTensorList; -import cn.smartjavaai.translation.enums.MachineTranslationModeEnum; -import cn.smartjavaai.translation.exception.TranslationException; -import cn.smartjavaai.translation.factory.TranslationModelFactory; -import cn.smartjavaai.translation.model.common.translator.Decoder2Translator; -import cn.smartjavaai.translation.model.common.translator.DecoderTranslator; -import cn.smartjavaai.translation.model.common.translator.EncoderTranslator; -import cn.smartjavaai.translation.utils.TokenUtils; -import lombok.extern.slf4j.Slf4j; -import org.apache.commons.lang3.StringUtils; -import org.apache.commons.pool2.ObjectPool; -import org.apache.commons.pool2.impl.GenericObjectPool; -import org.junit.Test; - -import java.io.IOException; -import java.nio.file.Paths; -import java.util.Objects; - -/** - * 机器翻译通用检测模型 - * - * @author lwx - * @date 2025/6/05 - */ -@Slf4j -public class TracedTranslationModel implements TranslationCommonModel { - - private ObjectPool> detPredictorPool; - private ZooModel nllbModel; - private HuggingFaceTokenizer tokenizer; - private Predictor encoderPredictor; - private Predictor decoderPredictor; - private Predictor decoder2Predictor; - private SearchConfig searchConfig; - private MachineTranslationModelConfig config; - private NDManager manager; - - @Test - public void detect(){ - Config.setCachePath("E:\\ai\\models\\libs"); - MachineTranslationModelConfig config = new MachineTranslationModelConfig(); - SearchConfig searchConfig = new SearchConfig(); - // 设置输出文字的最大长度 - searchConfig.setMaxSeqLength(128); - // 设置源语言:中文 "zho_Hans": 256200 - searchConfig.setSrcLangId(256200); - // 设置目标语言:英文 "eng_Latn": 256047 - searchConfig.setForcedBosTokenId(256047); - config.setSearchConfig(searchConfig); - config.setDevice(DeviceEnum.CPU); - config.setModelEnum(MachineTranslationModeEnum.TRACED_TRANSLATION_CPU); - config.setModelPath("E:\\ai\\models\\nlp\\"); - config.setModelName("traced_translation_cpu.pt"); - // 输入文字 - String input2 = "智利北部的丘基卡马塔矿是世界上最大的露天矿之一,长约4公里,宽3公里,深1公里。"; - String input = "你好,欢迎使用SmartJavaAI!"; - TranslationCommonModel detModel = TranslationModelFactory.getInstance().getDetModel(config); - // detModel.loadModel(config); - String translate = detModel.translate(input); - System.out.println("识别结果 translate"+translate); - } - @Override - public void loadModel(MachineTranslationModelConfig config) { - if (StringUtils.isBlank(config.getModelPath())) { - throw new TranslationException("modelPath is null"); - } - Device device = null; - if (!Objects.isNull(config.getDevice())) { - device = config.getDevice() == DeviceEnum.CPU ? Device.cpu() : Device.gpu(); - } - this.config = config; - this.searchConfig = config.getSearchConfig(); - Criteria criteria = - Criteria.builder() - .setTypes(NDList.class, NDList.class) - .optModelPath(Paths.get(config.getModelPath()+config.getModelName())) - .optEngine("PyTorch") - .optDevice(device) - .optTranslator(new NoopTranslator()) - .build(); - try { - nllbModel = ModelZoo.loadModel(criteria); - // 创建池子:每个线程独享 Predictor - this.detPredictorPool = new GenericObjectPool<>(new ZooModelFactory<>(nllbModel)); - log.info("当前设备: " + nllbModel.getNDManager().getDevice()); - log.info("当前引擎: " + Engine.getInstance().getEngineName()); - } catch (IOException | ModelNotFoundException | MalformedModelException e) { - throw new TranslationException("模型加载失败", e); - } - } - - @Override - public String translate(String input) throws TranslationException { - ZooModel zooModel = null; - try (NDManager manager = NDManager.newBaseManager()) { - zooModel = detPredictorPool.borrowObject(); - tokenizer = HuggingFaceTokenizer.newInstance(Paths.get(config.getModelPath() + "tokenizer.json")); - encoderPredictor = zooModel.newPredictor(new EncoderTranslator()); - decoderPredictor = zooModel.newPredictor(new DecoderTranslator()); - decoder2Predictor = zooModel.newPredictor(new Decoder2Translator()); - this.manager=manager; - return translateLanguage(input); - } catch (Exception e) { - throw new TranslationException("翻译错误", e); - } finally { - if (zooModel != null) { - try { - detPredictorPool.returnObject(zooModel); //归还 - } catch (Exception e) { - log.warn("归还Predictor失败", e); - try { - zooModel.close(); // 归还失败才销毁 - } catch (Exception ex) { - log.error("关闭Predictor失败", ex); - } - } - } - - } - } - - private String translateLanguage(String input) throws TranslateException { - - Encoding encoding = tokenizer.encode(input); - long[] ids = encoding.getIds(); - // 1. Encoder - long[] inputIds = new long[ids.length]; - // 设置源语言编码 - inputIds[0] = searchConfig.getSrcLangId(); - for (int i = 0; i < ids.length - 1; i++) { - inputIds[i + 1] = ids[i]; - } - - long[] attentionMask = encoding.getAttentionMask(); - NDArray attentionMaskArray = manager.create(attentionMask).expandDims(0); - - NDArray encoderHiddenStates = encoder(inputIds); - - NDArray decoder_input_ids = manager.create(new long[]{searchConfig.getDecoderStartTokenId()}).reshape(1, 1); - NDList decoderInput = new NDList(decoder_input_ids, encoderHiddenStates, attentionMaskArray); - - // 2. Initial Decoder - CausalLMOutput modelOutput = decoder(decoderInput); - modelOutput.getLogits().attach(manager); - modelOutput.getPastKeyValuesList().attach(manager); - - GreedyBatchTensorList searchState = - new GreedyBatchTensorList(null, decoder_input_ids, modelOutput.getPastKeyValuesList(), encoderHiddenStates, attentionMaskArray); - - while (true) { -// try (NDScope ignore = new NDScope()) { - NDArray pastOutputIds = searchState.getPastOutputIds(); - - if (searchState.getNextInputIds() != null) { - decoderInput = new NDList(searchState.getNextInputIds(), searchState.getEncoderHiddenStates(), searchState.getAttentionMask()); - decoderInput.addAll(searchState.getPastKeyValues()); - // 3. Decoder loop - modelOutput = decoder2(decoderInput); - } - - NDArray outputIds = greedyStepGen(searchConfig, pastOutputIds, modelOutput.getLogits()); - - searchState.setNextInputIds(outputIds); - pastOutputIds = pastOutputIds.concat(outputIds, 1); - searchState.setPastOutputIds(pastOutputIds); - - searchState.setPastKeyValues(modelOutput.getPastKeyValuesList()); - - long id = searchState.getNextInputIds().toLongArray()[0]; - if (searchConfig.getEosTokenId() == id) { - searchState.setNextInputIds(null); - break; - } - if (searchState.getPastOutputIds() != null && searchState.getPastOutputIds().getShape().get(1) + 1 >= searchConfig.getMaxSeqLength()) { - break; - } - } - - if (searchState.getNextInputIds() == null) { - NDArray resultIds = searchState.getPastOutputIds(); - String result = TokenUtils.decode(searchConfig, tokenizer, resultIds); - return result; - } else { - NDArray resultIds = searchState.getPastOutputIds(); // .concat(searchState.getNextInputIds(), 1) - String result = TokenUtils.decode(searchConfig, tokenizer, resultIds); - return result; - } - - } - - public NDArray greedyStepGen(SearchConfig config, NDArray pastOutputIds, NDArray next_token_scores) { - next_token_scores = next_token_scores.get(":, -1, :"); - - NDArray new_next_token_scores = manager.create(next_token_scores.getShape(), next_token_scores.getDataType()); - next_token_scores.copyTo(new_next_token_scores); - - // LogitsProcessor 1. ForcedBOSTokenLogitsProcessor - // 设置目标语言 - long cur_len = pastOutputIds.getShape().getLastDimension(); - if (cur_len == 1) { - long num_tokens = new_next_token_scores.getShape().getLastDimension(); - for (long i = 0; i < num_tokens; i++) { - if (i != config.getForcedBosTokenId()) { - new_next_token_scores.set(new NDIndex(":," + i), Float.NEGATIVE_INFINITY); - } - } - new_next_token_scores.set(new NDIndex(":," + config.getForcedBosTokenId()), 0); - } - - NDArray probs = new_next_token_scores.softmax(-1); - NDArray next_tokens = probs.argMax(-1); - - return next_tokens.expandDims(0); - } - - public NDArray encoder(long[] ids) throws TranslateException { - return encoderPredictor.predict(ids); - } - - public CausalLMOutput decoder(NDList input) throws TranslateException { - return decoderPredictor.predict(input); - } - - public CausalLMOutput decoder2(NDList input) throws TranslateException { - return decoder2Predictor.predict(input); - } -} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TranslationCommonModel.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TranslationCommonModel.java deleted file mode 100644 index 513331d..0000000 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/TranslationCommonModel.java +++ /dev/null @@ -1,32 +0,0 @@ -package cn.smartjavaai.translation.model.common; - -import cn.smartjavaai.translation.config.MachineTranslationModelConfig; - -/** - * 机器翻译通用检测模型 - * @author lwx - * @date 2025/6/05 - */ -public interface TranslationCommonModel { - - /** - * 加载模型 - * @param config - */ - void loadModel(MachineTranslationModelConfig config); // 加载模型 - - /** - * 机器翻译 - * @param input 翻译内容 - * @return - */ - default String translate(String input) { - throw new UnsupportedOperationException("默认不支持该功能"); - } - - - - - - -} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/EncoderTranslator.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/EncoderTranslator.java deleted file mode 100644 index 8c99619..0000000 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/EncoderTranslator.java +++ /dev/null @@ -1,49 +0,0 @@ -package cn.smartjavaai.translation.model.common.translator; - -import ai.djl.ndarray.NDArray; -import ai.djl.ndarray.NDList; -import ai.djl.ndarray.NDManager; -import ai.djl.translate.NoBatchifyTranslator; -import ai.djl.translate.TranslatorContext; - -import java.util.Arrays; - -/** - * 编码器前后处理 - * - * @author Calvin - * @mail 179209347@qq.com - * @website www.aias.top - */ -public class EncoderTranslator implements NoBatchifyTranslator { - - - public EncoderTranslator() { - } - - @Override - public NDList processInput(TranslatorContext ctx, long[] input) throws Exception { - NDManager manager = ctx.getNDManager(); - - NDArray inputIdArray = manager.create(input).expandDims(0); - inputIdArray.setName("input_ids"); - - long[] attentionMask = new long[input.length]; - Arrays.fill(attentionMask, 1); - NDArray attentionMaskArray = manager.create(attentionMask).expandDims(0); - attentionMaskArray.setName("attention_mask"); - - NDArray placeholder = ctx.getNDManager().create(0); - placeholder.setName("module_method:encoder"); - - return new NDList(inputIdArray, attentionMaskArray, placeholder); - } - - @Override - public NDArray processOutput(TranslatorContext ctx, NDList list) { - NDArray encoderHiddenStates = list.get(0); - encoderHiddenStates.detach(); - return encoderHiddenStates; - } - -} \ No newline at end of file diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/Decoder2Translator.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java similarity index 85% rename from smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/Decoder2Translator.java rename to smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java index 2c32518..6e18541 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/common/translator/Decoder2Translator.java +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/model/translator/NllbDecoder2Translator.java @@ -1,4 +1,4 @@ -package cn.smartjavaai.translation.model.common.translator; +package cn.smartjavaai.translation.model.translator; import ai.djl.ndarray.NDArray; import ai.djl.ndarray.NDList; @@ -14,10 +14,10 @@ import cn.smartjavaai.translation.entity.CausalLMOutput; * @mail 179209347@qq.com * @website www.aias.top */ -public class Decoder2Translator implements NoBatchifyTranslator { +public class NllbDecoder2Translator implements NoBatchifyTranslator { private String tupleName; - public Decoder2Translator() { + public NllbDecoder2Translator() { tupleName = "past_key_values(" + 12 + ',' + 4 + ')'; } @@ -43,4 +43,4 @@ public class Decoder2Translator implements NoBatchifyTranslator { +public class NllbDecoderTranslator implements NoBatchifyTranslator { private String tupleName; - public DecoderTranslator() { + public NllbDecoderTranslator() { tupleName = "past_key_values(" + 12 + ',' + 4 + ')'; } @@ -42,4 +42,4 @@ public class DecoderTranslator implements NoBatchifyTranslator { +public class NllbEncoderTranslator implements NoBatchifyTranslator { - public EncoderTranslator() { + public NllbEncoderTranslator() { } @Override @@ -46,4 +46,4 @@ public class EncoderTranslator implements NoBatchifyTranslator return encoderHiddenStates; } -} \ No newline at end of file +} diff --git a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java index 582804e..d613432 100644 --- a/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java +++ b/smartjavaai-translate/src/main/java/cn/smartjavaai/translation/utils/TokenUtils.java @@ -2,7 +2,7 @@ package cn.smartjavaai.translation.utils; import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; import ai.djl.ndarray.NDArray; -import cn.smartjavaai.translation.config.SearchConfig; +import cn.smartjavaai.translation.config.NllbSearchConfig; import java.util.ArrayList; @@ -25,7 +25,7 @@ public final class TokenUtils { * @param output * @return */ - public static String decode(SearchConfig config, HuggingFaceTokenizer tokenizer, NDArray output) { + public static String decode(NllbSearchConfig config, HuggingFaceTokenizer tokenizer, NDArray output) { long[] outputIds = output.toLongArray(); ArrayList outputIdsList = new ArrayList<>(); @@ -42,7 +42,6 @@ public final class TokenUtils { ids[i] = objArr[i]; } String text = tokenizer.decode(ids); - return text; } -} \ No newline at end of file +}