在Transformers.js中实现文本分类推理并配置截断参数
用Hugging Face Transformers.js实现Python文本分类推理逻辑(解决长输入截断问题)
原Python代码功能概述
这段Python代码实现了文本分类的推理流程:
- 检查模型、分词器、标签映射表是否就绪
- 对输入文本批量分词,启用padding、截断(truncation),限制最大长度
- 模型推理得到logits,通过softmax和argmax获取预测标签ID
- 将标签ID映射为可读标签,返回(输入文本,对应标签)的列表
Transformers.js 实现代码
首先安装依赖:
npm install @xenova/transformers
以下是对应功能的JavaScript实现,已处理长输入截断问题:
import { AutoTokenizer, AutoModelForSequenceClassification } from '@xenova/transformers'; // 初始化模型、分词器和标签映射(对应Python中的类属性) let model; let tokenizer; let id2label; let config; // 加载模型和分词器的异步函数 async function loadModel(modelName) { // 加载分词器 tokenizer = await AutoTokenizer.from_pretrained(modelName); // 加载分类模型 model = await AutoModelForSequenceClassification.from_pretrained(modelName); // 获取模型配置(用于读取max_position_embeddings) config = await model.config; // 标签映射:优先从模型配置读取,也可手动定义(需和Python端一致) id2label = config.id2label || { // 示例:替换为你的实际标签映射 0: '负面', 1: '中性', 2: '正面' }; } // 文本分类推理函数,对应Python的text_classification_inference async function textClassificationInference(inputText) { // 检查组件是否就绪 if (!model || !tokenizer || !id2label) { console.error('模型、分词器或标签映射未初始化!'); return; } // 分词配置:明确设置截断、padding和最大长度,解决长输入报错 const encoding = await tokenizer(inputText, { padding: true, truncation: true, max_length: config.max_position_embeddings, return_tensor: 'pt' // 返回PyTorch格式张量适配模型输入 }); // 模型推理获取logits const outputs = await model(encoding); const logits = outputs.logits; // 实现softmax和argmax逻辑,获取预测标签ID const softmax = (arr) => { const expValues = arr.map(x => Math.exp(x)); const sumExp = expValues.reduce((a, b) => a + b, 0); return expValues.map(x => x / sumExp); }; const predictionIds = []; for (const logit of logits) { const probabilities = softmax(logit.tolist()); const predId = probabilities.indexOf(Math.max(...probabilities)); predictionIds.push(predId); } // 映射标签ID到文本,组装结果 const outputPredictions = inputText.map((text, index) => { return [text, id2label[predictionIds[index]]]; }); return outputPredictions; } // 使用示例 (async () => { // 替换为你实际使用的模型名称(如原Python代码中的波斯语模型) await loadModel('你的模型名称'); const testTexts = [ '这是一条普通测试文本', '这是一段超长的测试文本,用来验证截断功能是否正常,确保超过模型最大输入长度的文本不会触发报错,同时保证推理结果的准确性...' ]; const results = await textClassificationInference(testTexts); console.log(results); })();
关键说明
- 长输入问题解决:在分词配置中明确设置
truncation: true和max_length: config.max_position_embeddings,和Python代码逻辑保持一致,自动截断超长文本避免报错。 - 标签映射一致性:
id2label必须和原Python代码中的映射完全匹配,可从模型配置自动读取,也可手动定义。 - 异步处理:JavaScript中模型加载、分词、推理都是异步操作,必须用
async/await处理流程。 - Softmax/Argmax实现:手动实现了基础版本的softmax和argmax逻辑,也可引入
mathjs等数值库简化代码。
内容的提问来源于stack exchange,提问作者Reza Hedayati
相关产品推荐
相关产品推荐

