You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在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);
})();

关键说明

  1. 长输入问题解决:在分词配置中明确设置truncation: true和max_length: config.max_position_embeddings,和Python代码逻辑保持一致,自动截断超长文本避免报错。
  2. 标签映射一致性:id2label必须和原Python代码中的映射完全匹配,可从模型配置自动读取,也可手动定义。
  3. 异步处理:JavaScript中模型加载、分词、推理都是异步操作,必须用async/await处理流程。
  4. Softmax/Argmax实现:手动实现了基础版本的softmax和argmax逻辑,也可引入mathjs等数值库简化代码。

内容的提问来源于stack exchange,提问作者Reza Hedayati

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.30 04:46:34