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

如何将Rasa NLU与GPT-J集成?现有尝试遇阻求助

如何将GPT-J与Rasa NLU集成?

一、修复现有Fallback Action代码问题

你的action.py存在几个关键错误,导致GPT-J调用失败:

  1. 未导入UserUtteranceReverted事件,无法正确触发fallback逻辑
  2. 未设置openai.api_key,OpenAI客户端无法正常鉴权
  3. __init__方法名拼写错误(写成了init)

修复后的action.py代码如下:

import openai
from typing import Any, Text, Dict, List

from rasa_sdk import Action, Tracker
from rasa_sdk.executor import CollectingDispatcher
from rasa_sdk.events import UserUtteranceReverted  # 补充导入事件

def gptj(text):
    api_key = "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"
    openai.api_key = api_key  # 必须设置API密钥
    try:
        response = openai.Completion.create(
            model="gpt-j",
            prompt="\n\n" + text,
            temperature=0,
            logprobs=10,
            max_tokens=150,
            top_p=0,
            frequency_penalty=0,
            presence_penalty=0,
            stop=[" \n\n"]
        )
        return response['choices'][0]['text'].strip()
    except Exception as e:
        # 添加异常捕获,便于排查问题
        print(f"GPT-J调用失败: {str(e)}")
        return "抱歉,暂时无法为你提供回复。"

class ActionDefaultFallback(Action):
    def __init__(self):  # 修正方法名
        super().__init__()

    def name(self) -> Text:
        return "action_default_fallback"

    async def run(self, dispatcher, tracker, domain):
        query = tracker.latest_message['text']
        response_text = gptj(query)
        dispatcher.utter_message(text=response_text)
        return [UserUtteranceReverted()]

同时需要确保Rasa项目的domain.yml中配置fallback动作及触发规则:

actions:
  - action_default_fallback

policies:
  - name: RulePolicy
    core_fallback_threshold: 0.3
    core_fallback_action_name: "action_default_fallback"
    enable_fallback_prediction: True

二、验证Fallback集成

  1. 重启action服务:rasa run actions
  2. 启动Rasa shell:rasa shell
  3. 输入未被训练意图覆盖的问题,触发fallback动作,此时应返回GPT-J生成的内容

三、自定义Featurizer实现GPT-J用于意图识别

如果需要将GPT-J作为意图识别的特征提取器(替代不支持GPT-J的LanguageModelFeaturizer),可自定义Featurizer组件:

# 在Rasa项目根目录创建custom_featurizers.py
from rasa.nlu.featurizers.dense_featurizer.language_model_featurizer import LanguageModelFeaturizer
from rasa.shared.nlu.training_data.message import Message
from typing import Any, List, Optional
import torch
from transformers import AutoTokenizer, AutoModel

class GPTJFeaturizer(LanguageModelFeaturizer):
    def __init__(self, component_config: Optional[Dict[Text, Any]] = None) -> None:
        super().__init__(component_config)
        self.model_name = "EleutherAI/gpt-j-6B"
        # 硬件不足可使用量化版本:"EleutherAI/gpt-j-6B-int4"
        self.tokenizer = AutoTokenizer.from_pretrained(self.model_name)
        self.model = AutoModel.from_pretrained(self.model_name)
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        self.model.to(self.device)

    def _encode_text(self, texts: List[Text]) -> torch.Tensor:
        inputs = self.tokenizer(
            texts,
            return_tensors="pt",
            padding=True,
            truncation=True,
            max_length=self.component_config.get("max_sequence_length", 512)
        ).to(self.device)
        with torch.no_grad():
            outputs = self.model(**inputs)
        # 均值池化获取句子嵌入
        embeddings = self._mean_pooling(outputs, inputs['attention_mask'])
        return embeddings.cpu().numpy()

    @staticmethod
    def _mean_pooling(model_output, attention_mask):
        token_embeddings = model_output[0]
        input_mask = attention_mask.unsqueeze(-1).expand(token_embeddings.size()).float()
        return torch.sum(token_embeddings * input_mask, 1) / torch.clamp(input_mask.sum(1), min=1e-9)

在config.yml中配置自定义Featurizer:

language: zh
pipeline:
  - name: "custom_featurizers.GPTJFeaturizer"
    max_sequence_length: 512
  - name: "DIETClassifier"
    epochs: 100

四、其他可行方案

单独编写API的思路也成立:将GPT-J封装为独立的HTTP API服务,在Rasa的Action中通过requests库调用该API。这种方式便于单独维护GPT-J服务,也能灵活调整模型部署方式。

内容的提问来源于stack exchange,提问作者Pei Shan Lim

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 17:45:44