初始化项目,由ModelHub XC社区提供模型
Model: feanet/eou-detector-russian Source: Original Platform
This commit is contained in:
36
.gitattributes
vendored
Normal file
36
.gitattributes
vendored
Normal file
@@ -0,0 +1,36 @@
|
|||||||
|
*.7z filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.arrow filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.bin filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.ftz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.gz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.h5 filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.joblib filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.model filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.npy filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.npz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.onnx filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.ot filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.parquet filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pb filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pickle filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pkl filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pt filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.pth filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.rar filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
||||||
|
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tar filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tflite filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.tgz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.wasm filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.xz filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.zip filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*.zst filter=lfs diff=lfs merge=lfs -text
|
||||||
|
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
||||||
|
tokenizer.json filter=lfs diff=lfs merge=lfs -text
|
||||||
204
README.md
Normal file
204
README.md
Normal file
@@ -0,0 +1,204 @@
|
|||||||
|
---
|
||||||
|
language:
|
||||||
|
- ru
|
||||||
|
license: apache-2.0
|
||||||
|
base_model: Qwen/Qwen2.5-0.5B-Instruct
|
||||||
|
tags:
|
||||||
|
- end-of-utterance
|
||||||
|
- dialog
|
||||||
|
- call-center
|
||||||
|
- conversational-ai
|
||||||
|
- russian
|
||||||
|
- speech
|
||||||
|
- voice-activity
|
||||||
|
pipeline_tag: text-generation
|
||||||
|
---
|
||||||
|
|
||||||
|
# EOU Detector — Russian Call-Center Dialog
|
||||||
|
|
||||||
|
End-of-Utterance (EOU) detector for Russian conversational speech, fine-tuned from
|
||||||
|
[Qwen2.5-0.5B-Instruct](https://huggingface.co/Qwen/Qwen2.5-0.5B-Instruct) on 200k real
|
||||||
|
call-center dialogs.
|
||||||
|
|
||||||
|
The model predicts **P(`<|im_end|>`)** at the last token position — the probability that
|
||||||
|
the current speaker has finished their utterance. No classification head; the LM vocabulary
|
||||||
|
does the detection.
|
||||||
|
|
||||||
|
Developed at [Simplexphone](https://simplexphone.com) — real-time voice AI for call centers.
|
||||||
|
|
||||||
|
## Performance
|
||||||
|
|
||||||
|
Evaluated on 200 stratified samples (100 positive EOU + 100 negative) from held-out call-center data:
|
||||||
|
|
||||||
|
| Metric | Value |
|
||||||
|
|---|---|
|
||||||
|
| F1 | **0.851** |
|
||||||
|
| False Alarm (1 − Precision) | 22.3% |
|
||||||
|
| False Rejection (1 − Recall) | 6.0% |
|
||||||
|
| Optimal threshold | 0.077 |
|
||||||
|
| GPU latency (H100, batch=1) | ~10 ms |
|
||||||
|
| CPU latency (Xeon 28-core, batch=1) | ~55 ms |
|
||||||
|
|
||||||
|
## Usage
|
||||||
|
|
||||||
|
```python
|
||||||
|
from transformers import AutoTokenizer, AutoModelForCausalLM
|
||||||
|
import torch
|
||||||
|
|
||||||
|
tok = AutoTokenizer.from_pretrained("feanet/eou-detector-russian")
|
||||||
|
model = AutoModelForCausalLM.from_pretrained(
|
||||||
|
"feanet/eou-detector-russian", torch_dtype=torch.float32
|
||||||
|
)
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
EOU_ID = tok.convert_tokens_to_ids("<|im_end|>")
|
||||||
|
THRESHOLD = 0.077
|
||||||
|
|
||||||
|
def eou_probability(history: list[dict], current_text: str) -> float:
|
||||||
|
"""
|
||||||
|
history: list of {"role": "user"|"assistant", "content": "..."}
|
||||||
|
current_text: the utterance to score (last client turn)
|
||||||
|
Returns P(end-of-utterance) in [0, 1].
|
||||||
|
"""
|
||||||
|
msgs = history + [{"role": "user", "content": current_text}]
|
||||||
|
prompt = tok.apply_chat_template(msgs, add_generation_prompt=False, tokenize=False)
|
||||||
|
prompt = prompt[: prompt.rfind("<|im_end|>")] # strip trailing EOU token
|
||||||
|
enc = tok(prompt, return_tensors="pt", truncation=True, max_length=512)
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = model(**enc).logits
|
||||||
|
return torch.softmax(logits[0, -1, :], dim=-1)[EOU_ID].item()
|
||||||
|
|
||||||
|
|
||||||
|
# Example
|
||||||
|
history = [{"role": "assistant", "content": "добрый день чем могу помочь"}]
|
||||||
|
|
||||||
|
print(eou_probability(history, "спасибо до свидания")) # → ~0.8 (farewell, EOU)
|
||||||
|
print(eou_probability(history, "хотел уточнить по")) # → ~0.02 (incomplete, not EOU)
|
||||||
|
```
|
||||||
|
|
||||||
|
### ONNX / production deployment
|
||||||
|
|
||||||
|
For lower-latency production use, export to ONNX:
|
||||||
|
|
||||||
|
```python
|
||||||
|
import torch, torch.nn as nn
|
||||||
|
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||||
|
|
||||||
|
class EOUModel(nn.Module):
|
||||||
|
def __init__(self, model, eou_id):
|
||||||
|
super().__init__()
|
||||||
|
self.lm = model
|
||||||
|
self.eou_id = eou_id
|
||||||
|
|
||||||
|
def forward(self, input_ids):
|
||||||
|
logits = self.lm(input_ids).logits
|
||||||
|
return torch.softmax(logits[:, -1, :], dim=-1)[:, self.eou_id]
|
||||||
|
|
||||||
|
tok = AutoTokenizer.from_pretrained("feanet/eou-detector-russian")
|
||||||
|
model = AutoModelForCausalLM.from_pretrained("feanet/eou-detector-russian",
|
||||||
|
torch_dtype=torch.float32).eval()
|
||||||
|
eou_model = EOUModel(model, tok.convert_tokens_to_ids("<|im_end|>")).eval()
|
||||||
|
|
||||||
|
dummy = tok(["хорошо спасибо"], return_tensors="pt")["input_ids"]
|
||||||
|
torch.onnx.export(eou_model, (dummy,), "model.onnx",
|
||||||
|
input_names=["input_ids"], output_names=["eou_prob"],
|
||||||
|
dynamic_axes={"input_ids": {0: "batch", 1: "seq_len"},
|
||||||
|
"eou_prob": {0: "batch"}},
|
||||||
|
opset_version=18)
|
||||||
|
```
|
||||||
|
|
||||||
|
ONNX batch=1 GPU latency: **~6 ms** (H100).
|
||||||
|
|
||||||
|
## Training
|
||||||
|
|
||||||
|
**Data:** 200,667 Russian call-center dialog files in `[HH:MM] A/B: text` format.
|
||||||
|
Speaker A = customer (`user`), Speaker B = operator (`assistant`).
|
||||||
|
|
||||||
|
**Method:** Causal language modelling on the full Qwen2.5 chat template.
|
||||||
|
The model learns to predict `<|im_end|>` at natural turn boundaries as part of
|
||||||
|
standard next-token prediction — no artificial labels.
|
||||||
|
|
||||||
|
Full loss (not masked to EOU positions only) is essential: masking causes catastrophic
|
||||||
|
overfitting where the model memorises positions rather than learning turn-end signals.
|
||||||
|
|
||||||
|
**Key training details:**
|
||||||
|
- Sequences slid into 512-token windows (stride 256) → 341k training chunks
|
||||||
|
- Optimizer: AdamW, lr=1e-5, cosine schedule, 5% warmup
|
||||||
|
- Precision: bf16 on 1× H100 80 GB
|
||||||
|
- Early stopping on eval loss, patience=3
|
||||||
|
- Best checkpoint: step 29,326 (~2 epochs)
|
||||||
|
- Weight untying applied before training (safetensors requirement for Qwen)
|
||||||
|
|
||||||
|
## C++ / ONNX Runtime
|
||||||
|
|
||||||
|
Dependencies: [onnxruntime](https://github.com/microsoft/onnxruntime),
|
||||||
|
[tokenizers-cpp](https://github.com/mlc-ai/tokenizers-cpp) (reads `tokenizer.json` directly),
|
||||||
|
[ICU](https://icu.unicode.org/) for NFKC normalisation.
|
||||||
|
|
||||||
|
```cpp
|
||||||
|
#include <onnxruntime_cxx_api.h>
|
||||||
|
#include <tokenizers_cpp.h>
|
||||||
|
|
||||||
|
// Build the Qwen chat-template prompt manually and strip the trailing <|im_end|>
|
||||||
|
// token — the model scores P(<|im_end|>) as the *next* token at that position.
|
||||||
|
//
|
||||||
|
// Template token IDs (Qwen2.5 vocab):
|
||||||
|
// <|im_start|>=151644 <|im_end|>=151645 \n=198
|
||||||
|
// system=8948 user=872 assistant=77091
|
||||||
|
//
|
||||||
|
// ONNX interface:
|
||||||
|
// input "input_ids" INT64 [1, seq_len]
|
||||||
|
// output "eou_prob" FLOAT [1]
|
||||||
|
|
||||||
|
static constexpr float THRESHOLD = 0.0766f;
|
||||||
|
static constexpr int MAX_TOKENS = 512;
|
||||||
|
|
||||||
|
Ort::Env env(ORT_LOGGING_LEVEL_WARNING, "eou");
|
||||||
|
Ort::Session session(env, "model.onnx", Ort::SessionOptions{});
|
||||||
|
|
||||||
|
auto tokenizer = tokenizers::Tokenizer::FromBlobJSON(
|
||||||
|
ReadFile("tokenizer.json")); // your file-read helper
|
||||||
|
|
||||||
|
// Build input_ids: [system block] + turns + [user open, current text]
|
||||||
|
// then truncate to MAX_TOKENS from the right.
|
||||||
|
std::vector<int64_t> ids = BuildPromptIds(tokenizer, history, current_text);
|
||||||
|
|
||||||
|
Ort::MemoryInfo mem("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
|
||||||
|
std::array<int64_t, 2> shape{1, (int64_t)ids.size()};
|
||||||
|
auto input_tensor = Ort::Value::CreateTensor<int64_t>(
|
||||||
|
mem, ids.data(), ids.size(), shape.data(), shape.size());
|
||||||
|
|
||||||
|
const char* input_names[] = {"input_ids"};
|
||||||
|
const char* output_names[] = {"eou_prob"};
|
||||||
|
auto output = session.Run(Ort::RunOptions{}, input_names, &input_tensor, 1,
|
||||||
|
output_names, 1);
|
||||||
|
|
||||||
|
float prob = output[0].GetTensorData<float>()[0];
|
||||||
|
bool eou = prob >= THRESHOLD; // FA=22.3% FR=6.0% F1=0.851
|
||||||
|
```
|
||||||
|
|
||||||
|
Full header-only class with preprocessing, tokenisation, and GPU support:
|
||||||
|
[`eou_detector.h`](https://huggingface.co/feanet/eou-detector-russian/blob/main/eou_detector.h)
|
||||||
|
|
||||||
|
### Latency
|
||||||
|
|
||||||
|
| Runtime | Hardware | Batch | Latency |
|
||||||
|
|---|---|---|---|
|
||||||
|
| PyTorch FP32 | H100 80 GB | 1 | ~23 ms |
|
||||||
|
| ONNX Runtime FP32 | H100 80 GB | 1 | **6 ms** |
|
||||||
|
| ONNX Runtime FP32 | Xeon 28-core | 1 | ~55 ms |
|
||||||
|
| ONNX Runtime FP32 | H100 80 GB | 128 | 14 ms (9 k items/s) |
|
||||||
|
|
||||||
|
## Intended use
|
||||||
|
|
||||||
|
- Voice assistant / IVR systems: detect when the caller has finished speaking
|
||||||
|
before routing to ASR or NLU
|
||||||
|
- Call-center analytics: segment transcripts by speaker turn
|
||||||
|
- Real-time dialog systems needing a language-aware alternative to silence-based VAD
|
||||||
|
|
||||||
|
## Limitations
|
||||||
|
|
||||||
|
- Trained on Russian call-center speech transcripts; performance on other domains is good
|
||||||
|
on other languages is not good
|
||||||
|
- Scores ASR transcript text, not audio — a separate VAD/ASR stage is needed upstream
|
||||||
|
- Short utterances (< 3 tokens) may score unreliably
|
||||||
24
added_tokens.json
Normal file
24
added_tokens.json
Normal file
@@ -0,0 +1,24 @@
|
|||||||
|
{
|
||||||
|
"</tool_call>": 151658,
|
||||||
|
"<tool_call>": 151657,
|
||||||
|
"<|box_end|>": 151649,
|
||||||
|
"<|box_start|>": 151648,
|
||||||
|
"<|endoftext|>": 151643,
|
||||||
|
"<|file_sep|>": 151664,
|
||||||
|
"<|fim_middle|>": 151660,
|
||||||
|
"<|fim_pad|>": 151662,
|
||||||
|
"<|fim_prefix|>": 151659,
|
||||||
|
"<|fim_suffix|>": 151661,
|
||||||
|
"<|im_end|>": 151645,
|
||||||
|
"<|im_start|>": 151644,
|
||||||
|
"<|image_pad|>": 151655,
|
||||||
|
"<|object_ref_end|>": 151647,
|
||||||
|
"<|object_ref_start|>": 151646,
|
||||||
|
"<|quad_end|>": 151651,
|
||||||
|
"<|quad_start|>": 151650,
|
||||||
|
"<|repo_name|>": 151663,
|
||||||
|
"<|video_pad|>": 151656,
|
||||||
|
"<|vision_end|>": 151653,
|
||||||
|
"<|vision_pad|>": 151654,
|
||||||
|
"<|vision_start|>": 151652
|
||||||
|
}
|
||||||
54
chat_template.jinja
Normal file
54
chat_template.jinja
Normal file
@@ -0,0 +1,54 @@
|
|||||||
|
{%- if tools %}
|
||||||
|
{{- '<|im_start|>system\n' }}
|
||||||
|
{%- if messages[0]['role'] == 'system' %}
|
||||||
|
{{- messages[0]['content'] }}
|
||||||
|
{%- else %}
|
||||||
|
{{- 'You are Qwen, created by Alibaba Cloud. You are a helpful assistant.' }}
|
||||||
|
{%- endif %}
|
||||||
|
{{- "\n\n# Tools\n\nYou may call one or more functions to assist with the user query.\n\nYou are provided with function signatures within <tools></tools> XML tags:\n<tools>" }}
|
||||||
|
{%- for tool in tools %}
|
||||||
|
{{- "\n" }}
|
||||||
|
{{- tool | tojson }}
|
||||||
|
{%- endfor %}
|
||||||
|
{{- "\n</tools>\n\nFor each function call, return a json object with function name and arguments within <tool_call></tool_call> XML tags:\n<tool_call>\n{\"name\": <function-name>, \"arguments\": <args-json-object>}\n</tool_call><|im_end|>\n" }}
|
||||||
|
{%- else %}
|
||||||
|
{%- if messages[0]['role'] == 'system' %}
|
||||||
|
{{- '<|im_start|>system\n' + messages[0]['content'] + '<|im_end|>\n' }}
|
||||||
|
{%- else %}
|
||||||
|
{{- '<|im_start|>system\nYou are Qwen, created by Alibaba Cloud. You are a helpful assistant.<|im_end|>\n' }}
|
||||||
|
{%- endif %}
|
||||||
|
{%- endif %}
|
||||||
|
{%- for message in messages %}
|
||||||
|
{%- if (message.role == "user") or (message.role == "system" and not loop.first) or (message.role == "assistant" and not message.tool_calls) %}
|
||||||
|
{{- '<|im_start|>' + message.role + '\n' + message.content + '<|im_end|>' + '\n' }}
|
||||||
|
{%- elif message.role == "assistant" %}
|
||||||
|
{{- '<|im_start|>' + message.role }}
|
||||||
|
{%- if message.content %}
|
||||||
|
{{- '\n' + message.content }}
|
||||||
|
{%- endif %}
|
||||||
|
{%- for tool_call in message.tool_calls %}
|
||||||
|
{%- if tool_call.function is defined %}
|
||||||
|
{%- set tool_call = tool_call.function %}
|
||||||
|
{%- endif %}
|
||||||
|
{{- '\n<tool_call>\n{"name": "' }}
|
||||||
|
{{- tool_call.name }}
|
||||||
|
{{- '", "arguments": ' }}
|
||||||
|
{{- tool_call.arguments | tojson }}
|
||||||
|
{{- '}\n</tool_call>' }}
|
||||||
|
{%- endfor %}
|
||||||
|
{{- '<|im_end|>\n' }}
|
||||||
|
{%- elif message.role == "tool" %}
|
||||||
|
{%- if (loop.index0 == 0) or (messages[loop.index0 - 1].role != "tool") %}
|
||||||
|
{{- '<|im_start|>user' }}
|
||||||
|
{%- endif %}
|
||||||
|
{{- '\n<tool_response>\n' }}
|
||||||
|
{{- message.content }}
|
||||||
|
{{- '\n</tool_response>' }}
|
||||||
|
{%- if loop.last or (messages[loop.index0 + 1].role != "tool") %}
|
||||||
|
{{- '<|im_end|>\n' }}
|
||||||
|
{%- endif %}
|
||||||
|
{%- endif %}
|
||||||
|
{%- endfor %}
|
||||||
|
{%- if add_generation_prompt %}
|
||||||
|
{{- '<|im_start|>assistant\n' }}
|
||||||
|
{%- endif %}
|
||||||
54
config.json
Normal file
54
config.json
Normal file
@@ -0,0 +1,54 @@
|
|||||||
|
{
|
||||||
|
"architectures": [
|
||||||
|
"Qwen2ForCausalLM"
|
||||||
|
],
|
||||||
|
"attention_dropout": 0.0,
|
||||||
|
"bos_token_id": 151643,
|
||||||
|
"dtype": "bfloat16",
|
||||||
|
"eos_token_id": 151645,
|
||||||
|
"hidden_act": "silu",
|
||||||
|
"hidden_size": 896,
|
||||||
|
"initializer_range": 0.02,
|
||||||
|
"intermediate_size": 4864,
|
||||||
|
"layer_types": [
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention",
|
||||||
|
"full_attention"
|
||||||
|
],
|
||||||
|
"max_position_embeddings": 32768,
|
||||||
|
"max_window_layers": 21,
|
||||||
|
"model_type": "qwen2",
|
||||||
|
"num_attention_heads": 14,
|
||||||
|
"num_hidden_layers": 24,
|
||||||
|
"num_key_value_heads": 2,
|
||||||
|
"rms_norm_eps": 1e-06,
|
||||||
|
"rope_scaling": null,
|
||||||
|
"rope_theta": 1000000.0,
|
||||||
|
"sliding_window": null,
|
||||||
|
"tie_word_embeddings": true,
|
||||||
|
"transformers_version": "4.57.6",
|
||||||
|
"use_cache": true,
|
||||||
|
"use_sliding_window": false,
|
||||||
|
"vocab_size": 151936
|
||||||
|
}
|
||||||
247
eou_detector.h
Normal file
247
eou_detector.h
Normal file
@@ -0,0 +1,247 @@
|
|||||||
|
#pragma once
|
||||||
|
#include <onnxruntime_cxx_api.h>
|
||||||
|
#include <tokenizers_cpp.h> // mlc-ai/tokenizers-cpp — reads tokenizer.json
|
||||||
|
|
||||||
|
#include <string>
|
||||||
|
#include <vector>
|
||||||
|
#include <memory>
|
||||||
|
#include <fstream>
|
||||||
|
#include <sstream>
|
||||||
|
#include <algorithm>
|
||||||
|
#include <cctype>
|
||||||
|
#include <unicode/unistr.h> // ICU — for NFKC + case fold
|
||||||
|
#include <unicode/normalizer2.h>
|
||||||
|
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
// EOUDetector — End-of-Utterance probability via Qwen2.5-0.5B ONNX model
|
||||||
|
//
|
||||||
|
// ONNX interface:
|
||||||
|
// input "input_ids" INT64 [1, seq_len] (seq_len dynamic)
|
||||||
|
// output "eou_prob" FLOAT [1]
|
||||||
|
//
|
||||||
|
// Usage:
|
||||||
|
// EOUDetector eou("path/to/model.onnx", "path/to/tokenizer.json");
|
||||||
|
// float p = eou.score(history, current_text);
|
||||||
|
// if (p >= EOUDetector::THRESHOLD) { /* fire endpoint */ }
|
||||||
|
// ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
class EOUDetector {
|
||||||
|
public:
|
||||||
|
// Tuned on 200 labelled samples: FA=22.3%, FR=6.0%, F1=0.851
|
||||||
|
static constexpr float THRESHOLD = 0.0766f;
|
||||||
|
static constexpr int MAX_TOKENS = 512;
|
||||||
|
|
||||||
|
struct Turn {
|
||||||
|
enum class Role { User, Assistant } role;
|
||||||
|
std::string text; // raw text — preprocessing applied internally
|
||||||
|
};
|
||||||
|
|
||||||
|
// modelPath — directory containing model.onnx AND model.onnx.data
|
||||||
|
// (ORT finds the .data file automatically; both must coexist)
|
||||||
|
// tokenizerJson — path to tokenizer.json from the model directory
|
||||||
|
EOUDetector(const std::string& modelPath,
|
||||||
|
const std::string& tokenizerJson,
|
||||||
|
bool useGpu = false)
|
||||||
|
: env_(ORT_LOGGING_LEVEL_WARNING, "eou")
|
||||||
|
, memoryInfo_("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault)
|
||||||
|
{
|
||||||
|
// ── session ──────────────────────────────────────────────────────────
|
||||||
|
Ort::SessionOptions opts;
|
||||||
|
opts.SetIntraOpNumThreads(1); // per-request; tune for your CPU core budget
|
||||||
|
opts.SetGraphOptimizationLevel(GraphOptimizationLevel::ORT_ENABLE_ALL);
|
||||||
|
if (useGpu) {
|
||||||
|
OrtCUDAProviderOptions cuda{};
|
||||||
|
cuda.device_id = 0;
|
||||||
|
opts.AppendExecutionProvider_CUDA(cuda);
|
||||||
|
}
|
||||||
|
|
||||||
|
std::string onnxPath = modelPath + "/model.onnx";
|
||||||
|
session_ = std::make_unique<Ort::Session>(env_, onnxPath.c_str(), opts);
|
||||||
|
|
||||||
|
// ── cache I/O names (avoids per-call allocation) ─────────────────────
|
||||||
|
inputName_ = session_->GetInputNameAllocated(0, allocator_).get();
|
||||||
|
outputName_ = session_->GetOutputNameAllocated(0, allocator_).get();
|
||||||
|
|
||||||
|
// ── tokenizer ────────────────────────────────────────────────────────
|
||||||
|
std::ifstream f(tokenizerJson);
|
||||||
|
std::string json((std::istreambuf_iterator<char>(f)),
|
||||||
|
std::istreambuf_iterator<char>());
|
||||||
|
tokenizer_ = tokenizers::Tokenizer::FromBlobJSON(json);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Returns P(end-of-utterance) in [0, 1].
|
||||||
|
// history — prior turns in chronological order, alternating user/assistant
|
||||||
|
// (may be empty for first turn)
|
||||||
|
// currentText — the latest partial or complete utterance being evaluated
|
||||||
|
float score(const std::vector<Turn>& history,
|
||||||
|
const std::string& currentText)
|
||||||
|
{
|
||||||
|
std::vector<int64_t> ids = buildInputIds(history, currentText);
|
||||||
|
|
||||||
|
// ── dims ─────────────────────────────────────────────────────────────
|
||||||
|
int64_t seqLen = static_cast<int64_t>(ids.size());
|
||||||
|
std::vector<int64_t> inputDims {1, seqLen};
|
||||||
|
std::vector<int64_t> outputDims {1};
|
||||||
|
|
||||||
|
// ── I/O name pointers ─────────────────────────────────────────────────
|
||||||
|
std::vector<const char*> inputNames {inputName_.c_str()};
|
||||||
|
std::vector<const char*> outputNames {outputName_.c_str()};
|
||||||
|
|
||||||
|
// ── input tensor (INT64) ──────────────────────────────────────────────
|
||||||
|
std::vector<Ort::Value> inputTensors;
|
||||||
|
inputTensors.push_back(
|
||||||
|
Ort::Value::CreateTensor<int64_t>(
|
||||||
|
memoryInfo_,
|
||||||
|
ids.data(), ids.size(),
|
||||||
|
inputDims.data(), inputDims.size()));
|
||||||
|
|
||||||
|
// ── output tensor (FLOAT32) ───────────────────────────────────────────
|
||||||
|
std::vector<float> outputValues(1);
|
||||||
|
std::vector<Ort::Value> outputTensors;
|
||||||
|
outputTensors.push_back(
|
||||||
|
Ort::Value::CreateTensor<float>(
|
||||||
|
memoryInfo_,
|
||||||
|
outputValues.data(), outputValues.size(),
|
||||||
|
outputDims.data(), outputDims.size()));
|
||||||
|
|
||||||
|
// ── run ───────────────────────────────────────────────────────────────
|
||||||
|
session_->Run(
|
||||||
|
Ort::RunOptions{nullptr},
|
||||||
|
inputNames.data(), inputTensors.data(), inputTensors.size(),
|
||||||
|
outputNames.data(), outputTensors.data(), outputTensors.size());
|
||||||
|
|
||||||
|
return outputValues[0];
|
||||||
|
}
|
||||||
|
|
||||||
|
bool isEndOfUtterance(const std::vector<Turn>& history,
|
||||||
|
const std::string& currentText)
|
||||||
|
{
|
||||||
|
return score(history, currentText) >= THRESHOLD;
|
||||||
|
}
|
||||||
|
|
||||||
|
private:
|
||||||
|
// ── Qwen2.5 chat-template token IDs (hardcoded — never change) ───────────
|
||||||
|
//
|
||||||
|
// <|im_start|>system\n
|
||||||
|
// You are Qwen, created by Alibaba Cloud. You are a helpful assistant.
|
||||||
|
// <|im_end|>\n
|
||||||
|
static const std::vector<int64_t> SYSTEM_BLOCK;
|
||||||
|
static const std::vector<int64_t> USER_OPEN; // <|im_start|>user\n
|
||||||
|
static const std::vector<int64_t> USER_CLOSE; // <|im_end|>\n
|
||||||
|
static const std::vector<int64_t> ASST_OPEN; // <|im_start|>assistant\n
|
||||||
|
static const std::vector<int64_t> ASST_CLOSE; // <|im_end|>\n
|
||||||
|
|
||||||
|
// ── members ───────────────────────────────────────────────────────────────
|
||||||
|
Ort::Env env_;
|
||||||
|
Ort::AllocatorWithDefaultOptions allocator_;
|
||||||
|
Ort::MemoryInfo memoryInfo_;
|
||||||
|
std::unique_ptr<Ort::Session> session_;
|
||||||
|
std::unique_ptr<tokenizers::Tokenizer> tokenizer_;
|
||||||
|
std::string inputName_;
|
||||||
|
std::string outputName_;
|
||||||
|
|
||||||
|
// ── text preprocessing ────────────────────────────────────────────────────
|
||||||
|
// Matches training: NFKC → lowercase → strip punctuation (keep ' -)
|
||||||
|
// → collapse whitespace
|
||||||
|
std::string preprocess(const std::string& raw)
|
||||||
|
{
|
||||||
|
// NFKC normalise + fold case via ICU
|
||||||
|
icu::UnicodeString u = icu::UnicodeString::fromUTF8(raw);
|
||||||
|
UErrorCode err = U_ZERO_ERROR;
|
||||||
|
const icu::Normalizer2* nfkc =
|
||||||
|
icu::Normalizer2::getNFKCInstance(err);
|
||||||
|
u = nfkc->normalize(u, err);
|
||||||
|
u.foldCase(U_FOLD_CASE_DEFAULT);
|
||||||
|
|
||||||
|
std::string s;
|
||||||
|
u.toUTF8String(s);
|
||||||
|
|
||||||
|
// strip punctuation (keep apostrophe and hyphen), collapse spaces
|
||||||
|
std::string out;
|
||||||
|
out.reserve(s.size());
|
||||||
|
bool lastWasSpace = true;
|
||||||
|
for (unsigned char c : s) {
|
||||||
|
if (c == '\'' || c == '-') {
|
||||||
|
out += c; lastWasSpace = false;
|
||||||
|
} else if (std::ispunct(c)) {
|
||||||
|
// skip
|
||||||
|
} else if (std::isspace(c)) {
|
||||||
|
if (!lastWasSpace) { out += ' '; lastWasSpace = true; }
|
||||||
|
} else {
|
||||||
|
out += c; lastWasSpace = false;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
while (!out.empty() && out.back() == ' ') out.pop_back();
|
||||||
|
return out;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── BPE-encode text, return int64 ids ────────────────────────────────────
|
||||||
|
std::vector<int64_t> encode(const std::string& text)
|
||||||
|
{
|
||||||
|
auto ids32 = tokenizer_->Encode(text, /*add_special_tokens=*/false);
|
||||||
|
std::vector<int64_t> ids64(ids32.begin(), ids32.end());
|
||||||
|
return ids64;
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── append helper ─────────────────────────────────────────────────────────
|
||||||
|
static void append(std::vector<int64_t>& dst, const std::vector<int64_t>& src)
|
||||||
|
{
|
||||||
|
dst.insert(dst.end(), src.begin(), src.end());
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── build full input_ids for ONNX ────────────────────────────────────────
|
||||||
|
std::vector<int64_t> buildInputIds(const std::vector<Turn>& history,
|
||||||
|
const std::string& currentText)
|
||||||
|
{
|
||||||
|
std::vector<int64_t> ids;
|
||||||
|
ids.reserve(MAX_TOKENS);
|
||||||
|
|
||||||
|
// system block (always present)
|
||||||
|
append(ids, SYSTEM_BLOCK);
|
||||||
|
|
||||||
|
// prior turns — merge consecutive same-role turns
|
||||||
|
std::vector<Turn> merged;
|
||||||
|
for (const auto& t : history) {
|
||||||
|
std::string clean = preprocess(t.text);
|
||||||
|
if (!merged.empty() && merged.back().role == t.role)
|
||||||
|
merged.back().text += " " + clean;
|
||||||
|
else
|
||||||
|
merged.push_back({t.role, clean});
|
||||||
|
}
|
||||||
|
|
||||||
|
for (const auto& t : merged) {
|
||||||
|
if (t.role == Turn::Role::User) {
|
||||||
|
append(ids, USER_OPEN);
|
||||||
|
append(ids, encode(t.text));
|
||||||
|
append(ids, USER_CLOSE);
|
||||||
|
} else {
|
||||||
|
append(ids, ASST_OPEN);
|
||||||
|
append(ids, encode(t.text));
|
||||||
|
append(ids, ASST_CLOSE);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// current (last) user turn — NO closing <|im_end|>
|
||||||
|
append(ids, USER_OPEN);
|
||||||
|
append(ids, encode(preprocess(currentText)));
|
||||||
|
|
||||||
|
// truncate to MAX_TOKENS from the right (keep most recent context)
|
||||||
|
if ((int)ids.size() > MAX_TOKENS)
|
||||||
|
ids = std::vector<int64_t>(ids.end() - MAX_TOKENS, ids.end());
|
||||||
|
|
||||||
|
return ids;
|
||||||
|
}
|
||||||
|
};
|
||||||
|
|
||||||
|
// ── static member definitions ─────────────────────────────────────────────────
|
||||||
|
|
||||||
|
const std::vector<int64_t> EOUDetector::SYSTEM_BLOCK = {
|
||||||
|
151644, 8948, 198, // <|im_start|>system\n
|
||||||
|
2610, 525, 1207, 16948, 11, 3465, 553, 54364, 14817, // You are Qwen, created by Alibaba Cloud.
|
||||||
|
13, 1446, 525, 264, 10950, 17847, 13, // You are a helpful assistant.
|
||||||
|
151645, 198 // <|im_end|>\n
|
||||||
|
};
|
||||||
|
const std::vector<int64_t> EOUDetector::USER_OPEN = {151644, 872, 198};
|
||||||
|
const std::vector<int64_t> EOUDetector::USER_CLOSE = {151645, 198};
|
||||||
|
const std::vector<int64_t> EOUDetector::ASST_OPEN = {151644, 77091, 198};
|
||||||
|
const std::vector<int64_t> EOUDetector::ASST_CLOSE = {151645, 198};
|
||||||
14
generation_config.json
Normal file
14
generation_config.json
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
{
|
||||||
|
"bos_token_id": 151643,
|
||||||
|
"do_sample": true,
|
||||||
|
"eos_token_id": [
|
||||||
|
151645,
|
||||||
|
151643
|
||||||
|
],
|
||||||
|
"pad_token_id": 151643,
|
||||||
|
"repetition_penalty": 1.1,
|
||||||
|
"temperature": 0.7,
|
||||||
|
"top_k": 20,
|
||||||
|
"top_p": 0.8,
|
||||||
|
"transformers_version": "4.57.6"
|
||||||
|
}
|
||||||
151388
merges.txt
Normal file
151388
merges.txt
Normal file
File diff suppressed because it is too large
Load Diff
3
model.safetensors
Normal file
3
model.safetensors
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
version https://git-lfs.github.com/spec/v1
|
||||||
|
oid sha256:798d5e7bda7c7e33d63264c6fab9bededfd1e6763c9b3347e7bb59c37eabdd3c
|
||||||
|
size 1260367448
|
||||||
31
special_tokens_map.json
Normal file
31
special_tokens_map.json
Normal file
@@ -0,0 +1,31 @@
|
|||||||
|
{
|
||||||
|
"additional_special_tokens": [
|
||||||
|
"<|im_start|>",
|
||||||
|
"<|im_end|>",
|
||||||
|
"<|object_ref_start|>",
|
||||||
|
"<|object_ref_end|>",
|
||||||
|
"<|box_start|>",
|
||||||
|
"<|box_end|>",
|
||||||
|
"<|quad_start|>",
|
||||||
|
"<|quad_end|>",
|
||||||
|
"<|vision_start|>",
|
||||||
|
"<|vision_end|>",
|
||||||
|
"<|vision_pad|>",
|
||||||
|
"<|image_pad|>",
|
||||||
|
"<|video_pad|>"
|
||||||
|
],
|
||||||
|
"eos_token": {
|
||||||
|
"content": "<|im_end|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false
|
||||||
|
},
|
||||||
|
"pad_token": {
|
||||||
|
"content": "<|endoftext|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false
|
||||||
|
}
|
||||||
|
}
|
||||||
3
tokenizer.json
Normal file
3
tokenizer.json
Normal file
@@ -0,0 +1,3 @@
|
|||||||
|
version https://git-lfs.github.com/spec/v1
|
||||||
|
oid sha256:9c5ae00e602b8860cbd784ba82a8aa14e8feecec692e7076590d014d7b7fdafa
|
||||||
|
size 11421896
|
||||||
207
tokenizer_config.json
Normal file
207
tokenizer_config.json
Normal file
@@ -0,0 +1,207 @@
|
|||||||
|
{
|
||||||
|
"add_bos_token": false,
|
||||||
|
"add_prefix_space": false,
|
||||||
|
"added_tokens_decoder": {
|
||||||
|
"151643": {
|
||||||
|
"content": "<|endoftext|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151644": {
|
||||||
|
"content": "<|im_start|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151645": {
|
||||||
|
"content": "<|im_end|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151646": {
|
||||||
|
"content": "<|object_ref_start|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151647": {
|
||||||
|
"content": "<|object_ref_end|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151648": {
|
||||||
|
"content": "<|box_start|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151649": {
|
||||||
|
"content": "<|box_end|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151650": {
|
||||||
|
"content": "<|quad_start|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151651": {
|
||||||
|
"content": "<|quad_end|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151652": {
|
||||||
|
"content": "<|vision_start|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151653": {
|
||||||
|
"content": "<|vision_end|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151654": {
|
||||||
|
"content": "<|vision_pad|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151655": {
|
||||||
|
"content": "<|image_pad|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151656": {
|
||||||
|
"content": "<|video_pad|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": true
|
||||||
|
},
|
||||||
|
"151657": {
|
||||||
|
"content": "<tool_call>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151658": {
|
||||||
|
"content": "</tool_call>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151659": {
|
||||||
|
"content": "<|fim_prefix|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151660": {
|
||||||
|
"content": "<|fim_middle|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151661": {
|
||||||
|
"content": "<|fim_suffix|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151662": {
|
||||||
|
"content": "<|fim_pad|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151663": {
|
||||||
|
"content": "<|repo_name|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
},
|
||||||
|
"151664": {
|
||||||
|
"content": "<|file_sep|>",
|
||||||
|
"lstrip": false,
|
||||||
|
"normalized": false,
|
||||||
|
"rstrip": false,
|
||||||
|
"single_word": false,
|
||||||
|
"special": false
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"additional_special_tokens": [
|
||||||
|
"<|im_start|>",
|
||||||
|
"<|im_end|>",
|
||||||
|
"<|object_ref_start|>",
|
||||||
|
"<|object_ref_end|>",
|
||||||
|
"<|box_start|>",
|
||||||
|
"<|box_end|>",
|
||||||
|
"<|quad_start|>",
|
||||||
|
"<|quad_end|>",
|
||||||
|
"<|vision_start|>",
|
||||||
|
"<|vision_end|>",
|
||||||
|
"<|vision_pad|>",
|
||||||
|
"<|image_pad|>",
|
||||||
|
"<|video_pad|>"
|
||||||
|
],
|
||||||
|
"bos_token": null,
|
||||||
|
"clean_up_tokenization_spaces": false,
|
||||||
|
"eos_token": "<|im_end|>",
|
||||||
|
"errors": "replace",
|
||||||
|
"extra_special_tokens": {},
|
||||||
|
"model_max_length": 131072,
|
||||||
|
"pad_token": "<|endoftext|>",
|
||||||
|
"split_special_tokens": false,
|
||||||
|
"tokenizer_class": "Qwen2Tokenizer",
|
||||||
|
"unk_token": null
|
||||||
|
}
|
||||||
1
vocab.json
Normal file
1
vocab.json
Normal file
File diff suppressed because one or more lines are too long
Reference in New Issue
Block a user