Spaces:
Sleeping
Sleeping
File size: 9,733 Bytes
d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d47f0c8 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc 09dbec4 d98e3cc | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 | from __future__ import annotations
import argparse
import json
import sys
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from typing import Any
CURRENT_DIR = Path(__file__).resolve().parent
PROJECT_ROOT = CURRENT_DIR.parents[1]
if str(CURRENT_DIR) not in sys.path:
sys.path.append(str(CURRENT_DIR))
if str(PROJECT_ROOT) not in sys.path:
sys.path.append(str(PROJECT_ROOT))
from emotion_classifier import EmotionClassifier
from intent_classifier import IntentClassifier
from language_classifier import LanguageDetector
from response_generator import ResponseGenerator
from safety_router import crisis_reply, detect_crisis
from src.retrieval.retrieval_engine import RetrievalEngine
class ChatbotPipeline:
def __init__(self, retrieval_source: str = "both", top_k: int = 5, retrieval_collection: str | None = None) -> None:
self.retrieval_source = retrieval_source
self.top_k = top_k
self.retrieval_collection = retrieval_collection
self.language_detector = LanguageDetector()
self.language_detector.load_model()
self.emotion_classifier = EmotionClassifier()
self.intent_classifier = IntentClassifier()
self.retrieval_engine = None
self.response_generator = ResponseGenerator()
def run(self, user_message: str, history: list[dict[str, str]] | None = None) -> dict[str, Any]:
clean_message = user_message.strip()
if not clean_message:
return {"response": "Please enter a message.", "state": {}}
conversation_history = self._prepare_history(clean_message, history or [])
state = self._analyze(clean_message, conversation_history)
state["conversation_history"] = conversation_history
language_code = state["language"].get("language_code", "en")
if state["safety"]["is_crisis"]:
state["route"] = "crisis"
return {"response": crisis_reply(language_code), "state": state}
intent = state["intent"]["intent"]
use_retrieval = intent == "asking_mental_health_question"
state["route"] = "rag" if use_retrieval else "direct_response"
state["retrieval"] = {
"enabled": use_retrieval,
"source": self.retrieval_source,
"top_k": self.top_k,
"results": [],
}
if use_retrieval:
retrieval_query = state["intent"].get("retrieval_query") or clean_message
state["retrieval"]["query"] = retrieval_query
try:
state["retrieval"]["results"] = self._retrieve(retrieval_query)
except Exception as error:
state["retrieval"]["error"] = f"{type(error).__name__}"
try:
generated = self.response_generator.generate(state)
state["llm_review"] = {
"language": generated.get("language_review", {}),
"emotion": generated.get("emotion_review", {}),
"intent": generated.get("intent_review", {}),
}
corrected_intent = state["llm_review"]["intent"].get("corrected_intent")
state["final_intent"] = corrected_intent or state["intent"].get("intent")
state["final_route"] = "rag" if state["final_intent"] == "asking_mental_health_question" else "direct_response"
if state["final_intent"] == "asking_mental_health_question":
state["suggested_questions"] = generated.get("suggested_questions", [])
else:
state["suggested_questions"] = []
response = generated.get("answer") or "I am here with you, but I could not generate a complete response."
except RuntimeError as error:
state["llm_review"] = {"language": {}, "emotion": {}, "intent": {}}
state["suggested_questions"] = []
state["generation_error"] = str(error)
response = (
"I'm here with you \u2764\ufe0f. I can listen, help you slow things down, and support you with mental-health questions. "
"The advanced response generator is not available right now, so please try again shortly. "
"If this feels urgent or unsafe, contact local emergency support or someone you trust right away."
)
except Exception as error:
state["llm_review"] = {"language": {}, "emotion": {}, "intent": {}}
state["suggested_questions"] = []
state["generation_error"] = f"{type(error).__name__}: {error}"
response = (
"I am here with you, but I could not complete a full answer at the moment. "
"Try again shortly, or contact a trusted person or professional support if you need help now."
)
return {"response": response, "suggested_questions": state.get("suggested_questions", []), "state": state}
def _analyze(self, message: str, history: list[dict[str, str]]) -> dict[str, Any]:
safety = detect_crisis(message)
with ThreadPoolExecutor(max_workers=3) as executor:
language_future = executor.submit(self._safe_language, message)
emotion_future = executor.submit(self._safe_emotion, message)
if safety["is_crisis"]:
intent_future = None
else:
intent_future = executor.submit(self._safe_intent, message, history)
language = language_future.result()
emotion = emotion_future.result()
intent = self._crisis_intent(message) if intent_future is None else intent_future.result()
return {
"user_message": message,
"language": language,
"emotion": emotion,
"intent": intent,
"safety": safety,
"retrieval": {"enabled": False, "results": []},
}
def _crisis_intent(self, message: str) -> dict[str, Any]:
return {
"intent": "asking_mental_health_question",
"confidence": 0.0,
"confidence_margin": 0.0,
"intent_scores": {},
"reason": "Crisis guardrail matched before live intent classification.",
"retrieval_query": message,
"contextual_follow_up": False,
"interaction_type": "standalone",
"classification_skipped": True,
}
def _safe_intent(self, message: str, history: list[dict[str, str]]) -> dict[str, Any]:
try:
return self.intent_classifier.classify(message, history=history)
except Exception as error:
return self._fallback_intent(message, error)
def set_retrieval_collection(self, collection_name: str | None) -> None:
if collection_name != self.retrieval_collection:
self.retrieval_collection = collection_name
self.retrieval_engine = None
def _safe_language(self, message: str) -> dict[str, Any]:
try:
return self.language_detector.predict_with_confidence(message)
except Exception as error:
return {
"language_code": "en",
"language_name": "English",
"confidence": 0.0,
"is_confident": False,
"message": f"Language detection unavailable: {type(error).__name__}",
}
def _safe_emotion(self, message: str) -> dict[str, Any]:
try:
return self.emotion_classifier.predict_with_confidence(message)
except Exception as error:
return {
"emotion": "unknown",
"confidence": 0.0,
"is_confident": False,
"scores": {},
"message": f"Emotion detection unavailable: {type(error).__name__}",
}
def _fallback_intent(self, message: str, error: Exception) -> dict[str, Any]:
return {
"intent": "out_of_scope",
"confidence": 0.0,
"confidence_margin": 0.0,
"intent_scores": {},
"reason": f"Intent classification unavailable: {type(error).__name__}.",
"retrieval_query": message,
"contextual_follow_up": False,
"interaction_type": "standalone",
}
@staticmethod
def _prepare_history(message: str, history: list[dict[str, str]]) -> list[dict[str, str]]:
clean_history = [
{"role": item.get("role", ""), "content": str(item.get("content", "")).strip()}
for item in history
if item.get("role") in {"user", "assistant"} and str(item.get("content", "")).strip()
]
if clean_history and clean_history[-1]["role"] == "user" and clean_history[-1]["content"] == message:
clean_history.pop()
return clean_history[-8:]
def _retrieve(self, message: str) -> list[dict[str, Any]]:
if self.retrieval_engine is None:
self.retrieval_engine = RetrievalEngine(collection_name=self.retrieval_collection)
return self.retrieval_engine.search(message, source=self.retrieval_source, top_k=self.top_k)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Run the integrated mental-health chatbot pipeline.")
parser.add_argument("message", nargs="?", default="I feel anxious and cannot sleep.")
parser.add_argument("--source", choices=["both", "cci", "amod"], default="both")
parser.add_argument("--top-k", type=int, default=5)
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
pipeline = ChatbotPipeline(retrieval_source=args.source, top_k=args.top_k)
output = pipeline.run(args.message)
print(json.dumps(output, indent=2, ensure_ascii=False))
|