Mirror of Khriis/RECCON
Browse files- .gitignore +29 -0
- README.md +123 -0
- config.json +26 -0
- handler.py +288 -0
- model.safetensors +3 -0
- special_tokens_map.json +7 -0
- tokenizer_config.json +58 -0
- vocab.txt +0 -0
.gitignore
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Training-only files (not needed for inference)
|
| 2 |
+
optimizer.pt
|
| 3 |
+
scheduler.pt
|
| 4 |
+
training_args.bin
|
| 5 |
+
eval_results.log
|
| 6 |
+
|
| 7 |
+
# Python artifacts
|
| 8 |
+
__pycache__/
|
| 9 |
+
*.pyc
|
| 10 |
+
*.pyo
|
| 11 |
+
*.pyd
|
| 12 |
+
.Python
|
| 13 |
+
|
| 14 |
+
# Testing artifacts
|
| 15 |
+
emotional_trigger_debug.log
|
| 16 |
+
utterance_by_utterance_debug.log
|
| 17 |
+
|
| 18 |
+
# Local development
|
| 19 |
+
.env
|
| 20 |
+
.venv/
|
| 21 |
+
venv/
|
| 22 |
+
*.local
|
| 23 |
+
|
| 24 |
+
# IDE
|
| 25 |
+
.vscode/
|
| 26 |
+
.idea/
|
| 27 |
+
*.swp
|
| 28 |
+
|
| 29 |
+
file_structure.txt
|
README.md
ADDED
|
@@ -0,0 +1,123 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
language:
|
| 3 |
+
- en
|
| 4 |
+
tags:
|
| 5 |
+
- psychology
|
| 6 |
+
- emotion-recognition
|
| 7 |
+
- nlp
|
| 8 |
+
- question-answering
|
| 9 |
+
- trigger-extraction
|
| 10 |
+
datasets:
|
| 11 |
+
- daily_dialog
|
| 12 |
+
---
|
| 13 |
+
|
| 14 |
+
# RECCON: Emotional Trigger Extraction Model
|
| 15 |
+
|
| 16 |
+
**RECCON** (Recognizing Emotion Cause in CONversations) is a model designed to identify and extract the specific text spans (triggers) within a conversation that correspond to a labeled emotion.
|
| 17 |
+
|
| 18 |
+
This repository contains the weights and custom inference handler to deploy RECCON as a **Hugging Face Inference Endpoint**.
|
| 19 |
+
|
| 20 |
+
## 🧠 Model Details
|
| 21 |
+
|
| 22 |
+
- **Task**: Extractive Question Answering (Span Extraction)
|
| 23 |
+
- **Base Model**: `SpanBERT` (without context)
|
| 24 |
+
- **Training Dataset**: [RECCON Dataset](https://github.com/declare-lab/RECCON) (derived from DailyDialog)
|
| 25 |
+
- **Paper**: [Recognizing Emotion Cause in Conversations (Poria et al., 2021)](https://arxiv.org/abs/2012.11820)
|
| 26 |
+
|
| 27 |
+
## 🚀 Deployment (Inference Endpoints)
|
| 28 |
+
|
| 29 |
+
This repository is structured to be deployed directly to [Hugging Face Inference Endpoints](https://ui.endpoints.huggingface.co/).
|
| 30 |
+
|
| 31 |
+
### Prerequisites
|
| 32 |
+
Ensure the following files are present in the root of this repository:
|
| 33 |
+
1. `handler.py`: The custom inference logic (included).
|
| 34 |
+
2. `requirements.txt`: Dependencies (included).
|
| 35 |
+
3. `model.safetensors` (or `pytorch_model.bin`): The model weights.
|
| 36 |
+
4. `config.json`: The BERT model configuration.
|
| 37 |
+
5. `tokenizer.json` / `vocab.json`: Tokenizer files.
|
| 38 |
+
|
| 39 |
+
### Configuration
|
| 40 |
+
When creating the endpoint:
|
| 41 |
+
- **Task**: Select **Custom** or **Question Answering**.
|
| 42 |
+
- **Container Type**: The custom `handler.py` will automatically be detected and used.
|
| 43 |
+
|
| 44 |
+
## 💻 API Usage
|
| 45 |
+
|
| 46 |
+
The endpoint accepts a JSON payload containing an `utterance` and its associated `emotion`. It returns the specific phrase(s) that triggered that emotion.
|
| 47 |
+
|
| 48 |
+
### Request Format
|
| 49 |
+
|
| 50 |
+
**Single Input:**
|
| 51 |
+
```json
|
| 52 |
+
{
|
| 53 |
+
"inputs": {
|
| 54 |
+
"utterance": "I'm so excited about the promotion!",
|
| 55 |
+
"emotion": "happiness"
|
| 56 |
+
}
|
| 57 |
+
}
|
| 58 |
+
```
|
| 59 |
+
|
| 60 |
+
**Batch Input (Recommended):**
|
| 61 |
+
```json
|
| 62 |
+
{
|
| 63 |
+
"inputs": [
|
| 64 |
+
{
|
| 65 |
+
"utterance": "I'm so excited about the promotion!",
|
| 66 |
+
"emotion": "happiness"
|
| 67 |
+
},
|
| 68 |
+
{
|
| 69 |
+
"utterance": "I really miss my family back home.",
|
| 70 |
+
"emotion": "sadness"
|
| 71 |
+
}
|
| 72 |
+
]
|
| 73 |
+
}
|
| 74 |
+
```
|
| 75 |
+
|
| 76 |
+
### Response Format
|
| 77 |
+
|
| 78 |
+
The model returns a list of objects containing the extracted triggers.
|
| 79 |
+
|
| 80 |
+
```json
|
| 81 |
+
[
|
| 82 |
+
{
|
| 83 |
+
"utterance": "I'm so excited about the promotion!",
|
| 84 |
+
"emotion": "happiness",
|
| 85 |
+
"triggers": [
|
| 86 |
+
"excited about the promotion"
|
| 87 |
+
]
|
| 88 |
+
},
|
| 89 |
+
{
|
| 90 |
+
"utterance": "I really miss my family back home.",
|
| 91 |
+
"emotion": "sadness",
|
| 92 |
+
"triggers": [
|
| 93 |
+
"miss my family"
|
| 94 |
+
]
|
| 95 |
+
}
|
| 96 |
+
]
|
| 97 |
+
```
|
| 98 |
+
|
| 99 |
+
## 🛠️ logic (handler.py)
|
| 100 |
+
|
| 101 |
+
The custom handler performs the following steps:
|
| 102 |
+
1. **Preprocessing**: Formats the input into a Question-Answering format: *"Extract the exact short phrase (<= 8 words) from the target utterance that most strongly signals the emotion {emotion}..."*
|
| 103 |
+
2. **Inference**: Runs the RoBERTa model to predict start and end logits.
|
| 104 |
+
3. **Post-processing**:
|
| 105 |
+
* Extracts the best text span.
|
| 106 |
+
* Filters out stopwords.
|
| 107 |
+
* Ensures the trigger is a valid substring of the original text.
|
| 108 |
+
* Deduplicates overlapping triggers.
|
| 109 |
+
|
| 110 |
+
## 📚 Citation
|
| 111 |
+
|
| 112 |
+
If you use this model, please cite the original paper:
|
| 113 |
+
|
| 114 |
+
```bibtex
|
| 115 |
+
@article{poria2021recognizing,
|
| 116 |
+
title={Recognizing Emotion Cause in Conversations},
|
| 117 |
+
author={Poria, Soujanya and Majumder, Navonil and Hazarika, Devamanyu and Ghosal, Deepanway and Bhardwaj, Rishabh and Jian, Samson Yu Bai and Hong, Pengfei and Ghosh, Romila and Roy, Abhinaba and Chhaya, Niyati and others},
|
| 118 |
+
journal={Cognitive Computation},
|
| 119 |
+
pages={1--16},
|
| 120 |
+
year={2021},
|
| 121 |
+
publisher={Springer}
|
| 122 |
+
}
|
| 123 |
+
```
|
config.json
ADDED
|
@@ -0,0 +1,26 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"BertForQuestionAnswering"
|
| 4 |
+
],
|
| 5 |
+
"attention_probs_dropout_prob": 0.1,
|
| 6 |
+
"classifier_dropout": null,
|
| 7 |
+
"directionality": "bidi",
|
| 8 |
+
"dtype": "float32",
|
| 9 |
+
"hidden_act": "gelu",
|
| 10 |
+
"hidden_dropout_prob": 0.1,
|
| 11 |
+
"hidden_size": 768,
|
| 12 |
+
"initializer_range": 0.02,
|
| 13 |
+
"intermediate_size": 3072,
|
| 14 |
+
"layer_norm_eps": 1e-12,
|
| 15 |
+
"max_position_embeddings": 512,
|
| 16 |
+
"model_type": "bert",
|
| 17 |
+
"num_attention_heads": 12,
|
| 18 |
+
"num_hidden_layers": 12,
|
| 19 |
+
"output_past": true,
|
| 20 |
+
"pad_token_id": 0,
|
| 21 |
+
"position_embedding_type": "absolute",
|
| 22 |
+
"transformers_version": "4.57.6",
|
| 23 |
+
"type_vocab_size": 2,
|
| 24 |
+
"use_cache": true,
|
| 25 |
+
"vocab_size": 28996
|
| 26 |
+
}
|
handler.py
ADDED
|
@@ -0,0 +1,288 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import torch
|
| 2 |
+
import logging
|
| 3 |
+
import re
|
| 4 |
+
from typing import Dict, List, Any
|
| 5 |
+
from transformers import pipeline, AutoModelForQuestionAnswering, AutoTokenizer
|
| 6 |
+
|
| 7 |
+
# Configure logging
|
| 8 |
+
logging.basicConfig(level=logging.INFO)
|
| 9 |
+
logger = logging.getLogger(__name__)
|
| 10 |
+
|
| 11 |
+
class EndpointHandler:
|
| 12 |
+
def __init__(self, path=""):
|
| 13 |
+
"""
|
| 14 |
+
Initialize the RECCON emotional trigger extraction model using native transformers.
|
| 15 |
+
Args:
|
| 16 |
+
path: Path to model directory (provided by HuggingFace Inference Endpoints)
|
| 17 |
+
"""
|
| 18 |
+
logger.info("Initializing RECCON Trigger Extraction endpoint...")
|
| 19 |
+
|
| 20 |
+
# Detect device (CUDA/CPU)
|
| 21 |
+
cuda_available = torch.cuda.is_available()
|
| 22 |
+
if not cuda_available:
|
| 23 |
+
logger.warning("GPU not detected. Running on CPU. Inference will be slower.")
|
| 24 |
+
|
| 25 |
+
# In 'pipeline', device is an integer (-1 for CPU, 0+ for GPU)
|
| 26 |
+
self.device_id = 0 if cuda_available else -1
|
| 27 |
+
|
| 28 |
+
# Determine model path
|
| 29 |
+
model_path = path if path and path != "." else "."
|
| 30 |
+
logger.info(f"Loading model from {model_path}...")
|
| 31 |
+
|
| 32 |
+
try:
|
| 33 |
+
# Load tokenizer and model explicitly to ensure correct loading
|
| 34 |
+
tokenizer = AutoTokenizer.from_pretrained(model_path)
|
| 35 |
+
model, loading_info = AutoModelForQuestionAnswering.from_pretrained(
|
| 36 |
+
model_path,
|
| 37 |
+
output_loading_info=True
|
| 38 |
+
)
|
| 39 |
+
|
| 40 |
+
logger.warning("RECCON load info - missing_keys: %s", loading_info.get("missing_keys"))
|
| 41 |
+
logger.warning("RECCON load info - unexpected_keys: %s", loading_info.get("unexpected_keys"))
|
| 42 |
+
logger.warning("RECCON load info - error_msgs: %s", loading_info.get("error_msgs"))
|
| 43 |
+
logger.warning("Loaded model class: %s", model.__class__.__name__)
|
| 44 |
+
logger.warning("Loaded model name_or_path: %s", getattr(model.config, "_name_or_path", None))
|
| 45 |
+
|
| 46 |
+
# Initialize the pipeline
|
| 47 |
+
# top_k=20 matches your previous 'n_best_size=20' logic
|
| 48 |
+
self.pipe = pipeline(
|
| 49 |
+
"question-answering",
|
| 50 |
+
model=model,
|
| 51 |
+
tokenizer=tokenizer,
|
| 52 |
+
device=self.device_id,
|
| 53 |
+
top_k=20,
|
| 54 |
+
handle_impossible_answer=False
|
| 55 |
+
)
|
| 56 |
+
logger.info("Model loaded successfully.")
|
| 57 |
+
except Exception as e:
|
| 58 |
+
logger.error(f"Failed to load model: {e}")
|
| 59 |
+
raise
|
| 60 |
+
|
| 61 |
+
# Question template (must match training)
|
| 62 |
+
self.question_template = (
|
| 63 |
+
"Extract the exact short phrase (<= 8 words) from the target "
|
| 64 |
+
"utterance that most strongly signals the emotion {emotion}. "
|
| 65 |
+
"Return only a substring of the target utterance."
|
| 66 |
+
)
|
| 67 |
+
|
| 68 |
+
def __call__(self, data: Dict[str, Any]) -> List[Dict[str, Any]]:
|
| 69 |
+
"""
|
| 70 |
+
Process inference request.
|
| 71 |
+
"""
|
| 72 |
+
# Extract inputs
|
| 73 |
+
inputs = data.pop("inputs", data)
|
| 74 |
+
|
| 75 |
+
# Normalize to list format
|
| 76 |
+
if isinstance(inputs, dict):
|
| 77 |
+
inputs = [inputs]
|
| 78 |
+
|
| 79 |
+
if not inputs:
|
| 80 |
+
return [{"error": "No inputs provided", "triggers": []}]
|
| 81 |
+
|
| 82 |
+
# Validate and format inputs for the pipeline
|
| 83 |
+
pipeline_inputs = []
|
| 84 |
+
valid_indices = []
|
| 85 |
+
|
| 86 |
+
for i, item in enumerate(inputs):
|
| 87 |
+
utterance = item.get("utterance", "").strip()
|
| 88 |
+
emotion = item.get("emotion", "")
|
| 89 |
+
|
| 90 |
+
if not utterance:
|
| 91 |
+
logger.warning(f"Empty utterance at index {i}")
|
| 92 |
+
continue
|
| 93 |
+
|
| 94 |
+
# Format as QA task
|
| 95 |
+
question = self.question_template.format(emotion=emotion)
|
| 96 |
+
|
| 97 |
+
# The pipeline expects a list of dicts with 'question' and 'context'
|
| 98 |
+
pipeline_inputs.append({
|
| 99 |
+
'question': question,
|
| 100 |
+
'context': utterance
|
| 101 |
+
})
|
| 102 |
+
valid_indices.append(i)
|
| 103 |
+
|
| 104 |
+
# Run prediction
|
| 105 |
+
results = []
|
| 106 |
+
|
| 107 |
+
if not pipeline_inputs:
|
| 108 |
+
# All inputs were invalid
|
| 109 |
+
for item in inputs:
|
| 110 |
+
results.append({
|
| 111 |
+
"utterance": item.get("utterance", ""),
|
| 112 |
+
"emotion": item.get("emotion", ""),
|
| 113 |
+
"error": "Missing or empty utterance",
|
| 114 |
+
"triggers": []
|
| 115 |
+
})
|
| 116 |
+
return results
|
| 117 |
+
|
| 118 |
+
try:
|
| 119 |
+
# Run inference (batch_size helps with multiple inputs)
|
| 120 |
+
predictions = self.pipe(pipeline_inputs, batch_size=8)
|
| 121 |
+
|
| 122 |
+
# If batch_size=1 or single input, pipeline might return a single list/dict
|
| 123 |
+
# We ensure it's a list of lists (since top_k > 1)
|
| 124 |
+
if isinstance(predictions, dict): # Single input result
|
| 125 |
+
predictions = [predictions] # Wrap in list
|
| 126 |
+
elif isinstance(predictions, list) and len(predictions) > 0 and isinstance(predictions[0], dict):
|
| 127 |
+
# This happens if we have multiple inputs but top_k=1 (which is not the case here),
|
| 128 |
+
# OR if we have a single input and top_k > 1.
|
| 129 |
+
# If we have multiple inputs and top_k > 1, it returns a list of lists.
|
| 130 |
+
if len(pipeline_inputs) == 1:
|
| 131 |
+
predictions = [predictions]
|
| 132 |
+
# If multiple inputs and list of dicts, that implies top_k=1.
|
| 133 |
+
# But we set top_k=20. So it should be list of lists.
|
| 134 |
+
|
| 135 |
+
logger.debug(f"Raw predictions: {predictions}")
|
| 136 |
+
|
| 137 |
+
# Post-process results
|
| 138 |
+
pred_idx = 0
|
| 139 |
+
for i, item in enumerate(inputs):
|
| 140 |
+
utterance = item.get("utterance", "").strip()
|
| 141 |
+
emotion = item.get("emotion", "")
|
| 142 |
+
|
| 143 |
+
if i not in valid_indices:
|
| 144 |
+
results.append({
|
| 145 |
+
"utterance": utterance,
|
| 146 |
+
"emotion": emotion,
|
| 147 |
+
"error": "Missing or empty utterance",
|
| 148 |
+
"triggers": []
|
| 149 |
+
})
|
| 150 |
+
else:
|
| 151 |
+
# Get prediction for this item
|
| 152 |
+
# Because top_k=20, 'current_preds' is a list of dicts: [{'answer': '...', 'score': ...}, ...]
|
| 153 |
+
current_preds = predictions[pred_idx]
|
| 154 |
+
|
| 155 |
+
|
| 156 |
+
# Ensure it is a list
|
| 157 |
+
if isinstance(current_preds, dict):
|
| 158 |
+
current_preds = [current_preds]
|
| 159 |
+
|
| 160 |
+
logger.info(
|
| 161 |
+
"RECCON raw spans (answer, score): %s",
|
| 162 |
+
[(p.get("answer"), p.get("score", 0.0), 3) for p in current_preds[:5]]
|
| 163 |
+
)
|
| 164 |
+
|
| 165 |
+
def is_good_span(ans: str) -> bool:
|
| 166 |
+
if not ans:
|
| 167 |
+
return False
|
| 168 |
+
a = ans.strip()
|
| 169 |
+
if len(a) < 3:
|
| 170 |
+
return False
|
| 171 |
+
# reject pure punctuation
|
| 172 |
+
if all(ch in ".,!?;:-—'\"()[]{}" for ch in a):
|
| 173 |
+
return False
|
| 174 |
+
# require at least one letter
|
| 175 |
+
if not any(ch.isalpha() for ch in a):
|
| 176 |
+
return False
|
| 177 |
+
return True
|
| 178 |
+
|
| 179 |
+
raw_answers = [p.get("answer", "") for p in current_preds]
|
| 180 |
+
raw_answers = [a for a in raw_answers if is_good_span(a)]
|
| 181 |
+
triggers = self._clean_spans(raw_answers, utterance)
|
| 182 |
+
|
| 183 |
+
results.append({
|
| 184 |
+
"utterance": utterance,
|
| 185 |
+
"emotion": emotion,
|
| 186 |
+
"triggers": triggers
|
| 187 |
+
})
|
| 188 |
+
pred_idx += 1
|
| 189 |
+
|
| 190 |
+
logger.debug(f"Cleaned results: {results}")
|
| 191 |
+
return results
|
| 192 |
+
|
| 193 |
+
except Exception as e:
|
| 194 |
+
logger.error(f"Model prediction failed: {e}")
|
| 195 |
+
return [{
|
| 196 |
+
"utterance": item.get("utterance", ""),
|
| 197 |
+
"emotion": item.get("emotion", ""),
|
| 198 |
+
"error": str(e),
|
| 199 |
+
"triggers": []
|
| 200 |
+
} for item in inputs]
|
| 201 |
+
|
| 202 |
+
def _clean_spans(self, spans: List[str], target_text: str) -> List[str]:
|
| 203 |
+
"""
|
| 204 |
+
Clean and filter extracted trigger spans.
|
| 205 |
+
(Logic preserved exactly as provided)
|
| 206 |
+
"""
|
| 207 |
+
target_text = target_text or ""
|
| 208 |
+
target_lower = target_text.lower()
|
| 209 |
+
|
| 210 |
+
def _norm(s: str) -> str:
|
| 211 |
+
s = (s or "").strip().lower()
|
| 212 |
+
s = re.sub(r"\s+", " ", s)
|
| 213 |
+
s = re.sub(r"^[^\w]+|[^\w]+$", "", s)
|
| 214 |
+
return s
|
| 215 |
+
|
| 216 |
+
def _extract_from_target(target: str, phrase_lower: str) -> str:
|
| 217 |
+
idx = target.lower().find(phrase_lower)
|
| 218 |
+
if idx >= 0:
|
| 219 |
+
return target[idx:idx+len(phrase_lower)]
|
| 220 |
+
return phrase_lower
|
| 221 |
+
|
| 222 |
+
STOP = {
|
| 223 |
+
"a", "an", "the", "and", "or", "but", "so", "to", "of", "in", "on", "at",
|
| 224 |
+
"with", "for", "from", "is", "am", "are", "was", "were", "be", "been",
|
| 225 |
+
"being", "i", "you", "he", "she", "it", "we", "they", "my", "your", "his",
|
| 226 |
+
"her", "their", "our", "me", "him", "her", "them", "this", "that", "these",
|
| 227 |
+
"those"
|
| 228 |
+
}
|
| 229 |
+
|
| 230 |
+
candidates = []
|
| 231 |
+
for s in spans:
|
| 232 |
+
s = (s or "").strip()
|
| 233 |
+
if not s:
|
| 234 |
+
continue
|
| 235 |
+
s_norm = _norm(s)
|
| 236 |
+
if not s_norm:
|
| 237 |
+
continue
|
| 238 |
+
if target_text and s_norm not in target_lower:
|
| 239 |
+
continue
|
| 240 |
+
tokens = s_norm.split()
|
| 241 |
+
if len(tokens) > 8 or len(s_norm) > 80:
|
| 242 |
+
continue
|
| 243 |
+
if len(tokens) == 1 and (tokens[0] in STOP or len(tokens[0]) <= 2):
|
| 244 |
+
continue
|
| 245 |
+
candidates.append({
|
| 246 |
+
"norm": s_norm,
|
| 247 |
+
"tokens": tokens,
|
| 248 |
+
"tok_len": len(tokens),
|
| 249 |
+
"char_len": len(s_norm)
|
| 250 |
+
})
|
| 251 |
+
|
| 252 |
+
# Prioritize short, focused emotional keywords (1-3 words)
|
| 253 |
+
short_candidates = [c for c in candidates if 1 <= c["tok_len"] <= 3]
|
| 254 |
+
if short_candidates:
|
| 255 |
+
candidates = short_candidates
|
| 256 |
+
|
| 257 |
+
# Sort by SHORTEST spans first (most focused keywords)
|
| 258 |
+
candidates.sort(key=lambda x: (x["tok_len"], x["char_len"]), reverse=False)
|
| 259 |
+
kept_norms = []
|
| 260 |
+
for c in list(candidates):
|
| 261 |
+
n = c["norm"]
|
| 262 |
+
if any(n in kn or kn in n for kn in kept_norms):
|
| 263 |
+
continue
|
| 264 |
+
kept_norms.append(n)
|
| 265 |
+
|
| 266 |
+
cleaned = [_extract_from_target(target_text, n) for n in kept_norms]
|
| 267 |
+
|
| 268 |
+
if not cleaned and spans:
|
| 269 |
+
tt_tokens = target_lower.split()
|
| 270 |
+
best = None
|
| 271 |
+
for s in spans:
|
| 272 |
+
words = [w for w in (s or '').lower().strip().split() if w]
|
| 273 |
+
for L in range(min(8, len(words)), 0, -1):
|
| 274 |
+
for i in range(len(words) - L + 1):
|
| 275 |
+
phrase = words[i:i+L]
|
| 276 |
+
for j in range(len(tt_tokens) - L + 1):
|
| 277 |
+
if tt_tokens[j:j+L] == phrase:
|
| 278 |
+
cand = " ".join(phrase)
|
| 279 |
+
best = cand
|
| 280 |
+
break
|
| 281 |
+
if best:
|
| 282 |
+
break
|
| 283 |
+
if best:
|
| 284 |
+
break
|
| 285 |
+
if best:
|
| 286 |
+
return [_extract_from_target(target_text, best)]
|
| 287 |
+
|
| 288 |
+
return cleaned[:3]
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:a6672b27522322c199b40ef0d8d8ea2300a745a02942b74e0f48f14a5fa61cbc
|
| 3 |
+
size 430908208
|
special_tokens_map.json
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"cls_token": "[CLS]",
|
| 3 |
+
"mask_token": "[MASK]",
|
| 4 |
+
"pad_token": "[PAD]",
|
| 5 |
+
"sep_token": "[SEP]",
|
| 6 |
+
"unk_token": "[UNK]"
|
| 7 |
+
}
|
tokenizer_config.json
ADDED
|
@@ -0,0 +1,58 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"added_tokens_decoder": {
|
| 3 |
+
"0": {
|
| 4 |
+
"content": "[PAD]",
|
| 5 |
+
"lstrip": false,
|
| 6 |
+
"normalized": false,
|
| 7 |
+
"rstrip": false,
|
| 8 |
+
"single_word": false,
|
| 9 |
+
"special": true
|
| 10 |
+
},
|
| 11 |
+
"100": {
|
| 12 |
+
"content": "[UNK]",
|
| 13 |
+
"lstrip": false,
|
| 14 |
+
"normalized": false,
|
| 15 |
+
"rstrip": false,
|
| 16 |
+
"single_word": false,
|
| 17 |
+
"special": true
|
| 18 |
+
},
|
| 19 |
+
"101": {
|
| 20 |
+
"content": "[CLS]",
|
| 21 |
+
"lstrip": false,
|
| 22 |
+
"normalized": false,
|
| 23 |
+
"rstrip": false,
|
| 24 |
+
"single_word": false,
|
| 25 |
+
"special": true
|
| 26 |
+
},
|
| 27 |
+
"102": {
|
| 28 |
+
"content": "[SEP]",
|
| 29 |
+
"lstrip": false,
|
| 30 |
+
"normalized": false,
|
| 31 |
+
"rstrip": false,
|
| 32 |
+
"single_word": false,
|
| 33 |
+
"special": true
|
| 34 |
+
},
|
| 35 |
+
"103": {
|
| 36 |
+
"content": "[MASK]",
|
| 37 |
+
"lstrip": false,
|
| 38 |
+
"normalized": false,
|
| 39 |
+
"rstrip": false,
|
| 40 |
+
"single_word": false,
|
| 41 |
+
"special": true
|
| 42 |
+
}
|
| 43 |
+
},
|
| 44 |
+
"clean_up_tokenization_spaces": true,
|
| 45 |
+
"cls_token": "[CLS]",
|
| 46 |
+
"do_basic_tokenize": true,
|
| 47 |
+
"do_lower_case": false,
|
| 48 |
+
"extra_special_tokens": {},
|
| 49 |
+
"mask_token": "[MASK]",
|
| 50 |
+
"model_max_length": 1000000000000000019884624838656,
|
| 51 |
+
"never_split": null,
|
| 52 |
+
"pad_token": "[PAD]",
|
| 53 |
+
"sep_token": "[SEP]",
|
| 54 |
+
"strip_accents": null,
|
| 55 |
+
"tokenize_chinese_chars": true,
|
| 56 |
+
"tokenizer_class": "BertTokenizer",
|
| 57 |
+
"unk_token": "[UNK]"
|
| 58 |
+
}
|
vocab.txt
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|