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

如何使用Keras的predict_generator预测无标签文本数据?

解决NLP模型预测时的生成器与直接预测问题

问题核心

你的自定义generator函数未适配无标签的测试场景,传入Y_test=None时触发位置错误;直接调用model.predict(X_test)失败,本质是测试数据未完成与训练数据一致的预处理流程(如分词序列化、序列填充、独热编码等)。


1. 修改自定义生成器,兼容无标签模式

调整生成器逻辑,增加对Y=None的判断分支,确保测试阶段仅输出预处理后的特征数据:

def generator(X, Y=None, tokenizer=None, onehot=None, label_encoder=None, n_classes=None, batch_size=32):
    while True:
        for i in range(0, len(X), batch_size):
            # 批量提取文本数据
            batch_X = X[i:i+batch_size]
            # 执行训练时的预处理步骤
            batch_X_seq = tokenizer.texts_to_sequences(batch_X)
            # 注意:maxlen要与训练时设置的序列长度完全一致
            batch_X_padded = pad_sequences(batch_X_seq, maxlen=你的最大序列长度)
            if onehot:
                batch_X_onehot = onehot(batch_X_padded)  # 按你的独热编码逻辑处理
            
            # 训练/验证模式:返回特征与标签
            if Y is not None:
                batch_Y = Y[i:i+batch_size]
                batch_Y_encoded = label_encoder.transform(batch_Y)
                batch_Y_onehot = to_categorical(batch_Y_encoded, num_classes=n_classes)
                yield (batch_X_onehot, batch_Y_onehot)
            # 测试预测模式:仅返回特征
            else:
                yield batch_X_onehot

2. 用修改后的生成器执行预测

初始化测试生成器时传入Y_test=None,并指定预测步数避免重复计算:

test_generator = generator(X_test, Y_test=None, tokenizer=tokenizer, onehot=onehot, label_encoder=label_encoder, n_classes=n_classes, batch_size=batch_size)
# 计算总预测步数
steps = len(X_test) // batch_size
if len(X_test) % batch_size != 0:
    steps += 1
# 生成预测结果
predictions = model.predict_generator(test_generator, steps=steps)
# 将预测概率转换为标签(分类任务)
predicted_labels = label_encoder.inverse_transform(np.argmax(predictions, axis=1))

3. 修复直接预测的问题

若无需生成器,需先将测试文本转换成模型可接收的格式:

# 完全复刻训练时的预处理流程
X_test_seq = tokenizer.texts_to_sequences(X_test)
X_test_padded = pad_sequences(X_test_seq, maxlen=你的最大序列长度)
if onehot:
    X_test_onehot = onehot(X_test_padded)
# 直接执行预测
predictions = model.predict(X_test_onehot)
predicted_labels = label_encoder.inverse_transform(np.argmax(predictions, axis=1))

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 23:57:44