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

如何向微调后的BERT问题分类模型传入字符串列表?

批量处理BERT模型输入字符串列表的解决方案

问题背景

现有微调后的BERT问题分类模型仅支持单个字符串输入,尝试直接传入字符串列表时,tokenizer返回的结果异常(所有输入对应相同的无效token),且不想通过循环逐个处理(耗时较长)。当前代码如下:

questionclassification_model = tf.keras.models.load_model('/content/drive/MyDrive/questionclassification_model')
tokenizer = BertTokenizer.from_pretrained('bert-base-cased')

def prepare_data(input_text):
    token = tokenizer.encode_plus(
        input_text,
        max_length=256, 
        truncation=True, 
        padding='max_length', 
        add_special_tokens=True,
        return_tensors='tf'
    )
    return {
        'input_ids': tf.cast(token['input_ids'], tf.float64),
        'attention_mask': tf.cast(token['attention_mask'], tf.float64)
    }

def make_prediction(model, processed_data, classes=['Easy', 'Medium', 'Hard']):
    probs = model.predict(processed_data)[0]
    return classes[np.argmax(probs)],probs;

传入列表时的错误示例:

input_text = ["What is gandhi commonly considered to be?,Father of the nation in india","What is the long-term warming of the planets overall temperature called?, Global Warming"]
processed_data = prepare_data(input_text)

返回的异常结果(关键部分):

{'input_ids': <tf.Tensor: shape=(1, 256), dtype=float64, numpy=array([[101., 100., 100., 102., ...]]), 'attention_mask': <tf.Tensor: shape=(1, 256), dtype=float64, numpy=array([[1., 1., 1., 1., ...]])>}

问题原因

tokenizer.encode_plus仅针对单个文本设计,传入列表时会被当作一个整体字符串处理,导致生成错误的无效token。批量处理需使用tokenizer原生支持列表输入的__call__方法(直接调用tokenizer)。

解决方案

修改prepare_data函数用tokenizer直接批量处理文本,同时调整make_prediction函数适配批量结果返回:

import numpy as np
import tensorflow as tf
from transformers import BertTokenizer

questionclassification_model = tf.keras.models.load_model('/content/drive/MyDrive/questionclassification_model')
tokenizer = BertTokenizer.from_pretrained('bert-base-cased')

def prepare_data(input_texts):
    # 直接调用tokenizer处理批量文本
    tokens = tokenizer(
        input_texts,
        max_length=256, 
        truncation=True, 
        padding='max_length', 
        add_special_tokens=True,
        return_tensors='tf'
    )
    return {
        'input_ids': tf.cast(tokens['input_ids'], tf.float64),
        'attention_mask': tf.cast(tokens['attention_mask'], tf.float64)
    }

def make_prediction(model, processed_data, classes=['Easy', 'Medium', 'Hard']):
    probs = model.predict(processed_data)
    # 批量生成预测结果
    predictions = [classes[np.argmax(p)] for p in probs]
    return predictions, probs

使用示例

input_texts = [
    "What is gandhi commonly considered to be?,Father of the nation in india",
    "What is the long-term warming of the planets overall temperature called?, Global Warming"
]
processed_data = prepare_data(input_texts)
predictions, probs = make_prediction(questionclassification_model, processed_data)

# 输出每个文本的预测结果
for text, pred, prob in zip(input_texts, predictions, probs):
    print(f"文本: {text}")
    print(f"预测类别: {pred}")
    print(f"类别概率: {prob}\n")

关键说明

  • tokenizer(input_texts)自动处理批量输入,生成(batch_size, max_length)形状的input_ids和attention_mask,完全匹配模型批量输入要求。
  • 修改后的make_prediction直接返回所有输入对应的预测类别列表和概率矩阵,无需循环逐个处理,效率大幅提升。

内容的提问来源于stack exchange,提问作者Abin Jilson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 21:57:12