File size: 5,848 Bytes
21dca9d
 
0b55df4
 
 
 
 
 
 
 
 
 
21dca9d
0b55df4
 
 
62a5c8a
0b55df4
 
 
 
 
 
62a5c8a
0b55df4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
62a5c8a
0b55df4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
62a5c8a
0b55df4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
---
license: apache-2.0
base_model: answerdotai/ModernBERT-large
tags:
- onnx
- sequence-classification
- text-classification
- query-classifier
- modernbert
- llm-routing
pipeline_tag: text-classification
library_name: onnxruntime
---

# Query Classifier ONNX (ModernBERT-large Fine-Tuned)

A CPU-optimized ONNX sequence classification model fine-tuned from **ModernBERT-large** to classify user queries into **`low`**, **`medium`**, and **`hard`** complexity levels. It is designed for LLM routing, enabling applications to dispatch simple requests to smaller, lower-cost models while reserving larger reasoning models for complex, multi-step tasks.

---

## πŸ“Œ Model Details

- **Base Architecture**: [`answerdotai/ModernBERT-large`](https://huggingface.co/answerdotai/ModernBERT-large) (`ModernBertForSequenceClassification`)
- **Training Data**: Fine-tuned on a manually curated dataset of **1,800 user queries** spanning factual QA, programming, reasoning, mathematics, and technical problem solving.
- **Model Format**: Int8 quantized ONNX (`model.onnx`, ~397 MB).
- **Inference Runtime**: Pure **ONNX Runtime + Hugging Face Rust `tokenizers`** β€” zero PyTorch or `transformers` runtime dependencies are required for inference.

---

## 🏷️ Class Labels & Definitions

| Class ID | Label | Description & Examples |
| :--- | :--- | :--- |
| **0** | **`low`** | Simple factual questions, basic definitions, direct short lookup, unit conversions.<br>*Example:* "What is the capital of France?", "How many days in a leap year?" |
| **1** | **`medium`** | Explanations, code snippets, summarizing text, multi-step instructions.<br>*Example:* "Explain how key-value storage works in Redis", "Write a python script to parse CSV files." |
| **2** | **`hard`** | Complex algorithms, multi-file code synthesis, lock-free concurrency, advanced mathematics.<br>*Example:* "Implement a lock-free SPMC queue in C++ using atomics", "Calculate the integral of x^2 * sin(x) dx." |

---

## πŸ“Š Benchmark & Evaluation Results

Evaluated on an independent **held-out test set of 301 unseen queries** (balanced across `low`, `medium`, and `hard` buckets):

### Accuracy & F1 Breakdown

- **Overall Accuracy**: **88.04%**
- **Macro F1-Score**: **0.8797**

| Class | Precision | Recall | F1-Score | Support |
| :--- | :--- | :--- | :--- | :--- |
| **`low`** | **100.00%** | 81.00% | **0.8950** | 100 |
| **`medium`** | **82.18%** | 83.00% | **0.8259** | 100 |
| **`hard`** | **84.87%** | **100.00%** | **0.9182** | 101 |
| **Overall** | **88.04% Acc** | β€” | **0.8797 Macro-F1** | **301 total** |

### 🎯 Confusion Matrix (Actual vs. Predicted)

| Actual \ Predicted | `low` | `medium` | `hard` |
| :--- | :--- | :--- | :--- |
| **`low`** | **81** | 18 | 1 |
| **`medium`** | **0** | **83** | 17 |
| **`hard`** | **0** | **0** | **101** |

### Latency & Throughput Profile (CPU Execution)

| Metric | Measurement |
| :--- | :--- |
| **Median Latency (p50)** | **67.68 ms** |
| **Mean Latency** | **81.44 ms** |
| **p90 Latency** | **134.69 ms** |
| **p99 Latency** | **161.74 ms** |
| **Throughput** | **12.3 queries/sec** |

### πŸ§ͺ Reproducing Evaluation & Running Local Tests

The held-out test dataset and evaluation script are provided directly in the repository [`tests/`](tests) folder:
- [`tests/test_dataset.csv`](tests/test_dataset.csv): 301 held-out test queries (100 `low`, 100 `medium`, 101 `hard`).
- [`tests/evaluate.py`](tests/evaluate.py): Standalone evaluation script for computing accuracy, F1, confusion matrix, and latency profile.

To run the benchmark suite locally:
```bash
python tests/evaluate.py
```

---

## πŸš€ Quickstart: Pure ONNX Runtime Usage

The model can be used directly with ONNX Runtime without requiring PyTorch or Transformers for inference.

### 1. Installation

```bash
pip install onnxruntime tokenizers numpy huggingface-hub
```

### 2. Python Code Example

```python
import numpy as np
import onnxruntime as ort
from tokenizers import Tokenizer
from huggingface_hub import hf_hub_download

# Repository details
REPO_ID = "prvn-ramesh/query-classifier-onnx"

# 1. Download model artifacts from Hugging Face Hub
model_path = hf_hub_download(repo_id=REPO_ID, filename="model.onnx")
tok_path = hf_hub_download(repo_id=REPO_ID, filename="tokenizer.json")

# 2. Load tokenizer and initialize ONNX Runtime session
tokenizer = Tokenizer.from_file(tok_path)
session = ort.InferenceSession(model_path, providers=["CPUExecutionProvider"])
label_map = {0: "low", 1: "medium", 2: "hard"}

def classify_query(text: str):
    # Tokenize input string
    encoded = tokenizer.encode(text)
    input_ids = np.array([encoded.ids], dtype=np.int64)
    attention_mask = np.array([encoded.attention_mask], dtype=np.int64)

    # Run inference session
    outputs = session.run(None, {
        "input_ids": input_ids,
        "attention_mask": attention_mask
    })
    
    # Softmax over logits
    logits = outputs[0][0]
    exp_logits = np.exp(logits - np.max(logits))
    probs = exp_logits / np.sum(exp_logits)
    
    pred_idx = int(np.argmax(probs))
    
    return {
        "label": label_map[pred_idx],
        "confidence": float(probs[pred_idx]),
        "scores": {label_map[i]: float(probs[i]) for i in range(len(probs))}
    }

# Example usage
query = "Write a lock-free multi-threaded SPMC queue in C++ using atomics"
result = classify_query(query)

print(f"Query: {query}")
print(f"Predicted Class: {result['label'].upper()} (Confidence: {result['confidence']:.2%})")
print(f"All Scores: {result['scores']}")
```

---

## πŸ™Œ Citation & Acknowledgements

- Base model: [ModernBERT](https://huggingface.co/answerdotai/ModernBERT-large) developed by **Answer.AI** and **LightOn**.
- Frameworks: [ONNX Runtime](https://onnxruntime.ai/) and [Hugging Face Tokenizers](https://github.com/huggingface/tokenizers).