File size: 8,343 Bytes
23f8b10
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
4cc585e
23f8b10
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b5935b0
23f8b10
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
---
license: apache-2.0
language:
- en
tags:
- prompt-routing
- category-classifier
- deberta-v3
- llm-router
- cost-optimization
datasets:
- RowRed/ComplexityRouter
- OpenAssistant/oasst2
base_model:
- microsoft/deberta-v3-base
---

# CategoryRouter: A Category‑based LLM Router

**CategoryRouter** is a lightweight prompt category classifier finetuned from **microsoft/deberta-v3-base**. It assigns prompts to one of **6 categories**, making it useful for routing queries to the appropriate specialist LLM or service.

## Model Details

### Model Description

- **Model type:** Text Classification (multi‑class, 6 classes)
- **Language:** English
- **Backbone:** microsoft/deberta-v3-base
- **License:** Apache‑2.0
- **Finetuned from model:** microsoft/deberta-v3-base
- **Training data:** OASST2 + synthetic augmentations + manually created prompts  
  Labels generated by **Qwen3.5‑4B** (non‑thinking mode).

### Model Sources

- **Dataset repository:** https://huggingface.co/datasets/RowRed/ComplexityRouter

## Categories

| Label | Category | Example Prompt |
|-------|----------|----------------|
| 0 | Coding | “Write a Python function to merge two sorted lists.” |
| 1 | Information | “What is the normal force? Explain with an example.” |
| 2 | Guidance | “How do I change a car tire?” |
| 3 | Media Generation | “Generate an image prompt for a sunset.” |
| 4 | Writing | “Write a short story about a robot learning to paint.” |
| 5 | Other | “Invent a new holiday.” |

> **Note:** The original ComplexityRouter dataset had 7 categories; “Math” prompts were merged into “Coding” to retain only 6 distinct categories.

## Uses

### Direct Use

Route prompts to the appropriate LLM backend:

| Category | Suggested Backend |
|----------|-------------------|
| Coding | Code‑specialized model |
| Information | General‑knowledge model |
| Guidance | Instruction‑tuned model |
| Media Generation | Multimodal / image‑generation model |
| Writing | Creative writing model |
| Other | Catch‑all (generic model) |

### Out‑of‑Scope Use

- Multi‑turn conversations (single prompts only).
- Non‑English prompts.
- Prompts requiring image or multimodal understanding.

## Bias, Risks, and Limitations

- Training data is synthetic and may not represent all real‑world prompt distributions.
- Categories with low support (Media Generation, Writing) have lower per‑class F1 scores – boundary cases are inherently ambiguous.
- The model may struggle with very domain‑specific technical jargon.
- The “Other” category is a catch‑all; many “Information” prompts leak into it and vice‑versa.
- Performance may degrade on prompts very different from the training distribution.

## Notice
This is my first attempt making a widespread finetune ( just the part 2 version of it :) ). There are probably lots of issues, but thought the idea was sound. I might make a second (hopefully better) version eventually, but am not sure where to get lots of high-quality open source data.

## Training Details

### Training Data

| Split       | Samples | Source File          | Notes |
|-------------|---------|----------------------|-------|
| Training    | 2,800   | TRAINING.jsonl       | Used for model training |
| Validation  | 600     | TRAINING.jsonl       | Used for early stopping / hyperparameter tuning |
| Test (internal) | 600 | TRAINING.jsonl      | Used for in‑distribution evaluation |
| Test (held‑out) | 400 | TEST.jsonl   | Fully independent test set (reported results) |

**Total unique prompts:** 4,400

Class distribution (training):
Coding: 508 (18.1%) • Information: 937 (33.5%) • Guidance: 609 (21.8%) • Media Generation: 111 (4.0%) • Writing: 88 (3.1%) • Other: 547 (19.6%)

### Training Procedure

- Hardware: NVIDIA T4 (16 GB VRAM, Google Colab)
- Framework: PyTorch 2.11 + Hugging Face Transformers
- Optimizer: AdamW (lr=2e-5, weight_decay=0.01)
- Scheduler: Linear warmup (10% of steps) → linear decay
- Loss: Weighted Cross‑Entropy (sqrt‑scaled class weights) + label smoothing (0.15)
- Batch size: 16 (effective 32 with gradient accumulation)
- Epochs: 7 (early stopping patience = 3, best epoch = 7)
- Training time: ~16 minutes
- Class balancing: sqrt‑scaled class weights + weighted random sampler

## Evaluation Results

### Internal Test (600 held‑out samples from training split)

| Metric | Value |
|--------|-------|
| Exact Match Accuracy | 67.0% |
| Macro F1 | 0.617 |
| Weighted F1 | 0.665 |

**Per‑Class Performance (internal test, 600 samples)**

| Category | Precision | Recall | F1 | Support |
|----------|-----------|--------|----|---------|
| Coding | 0.714 | 0.688 | 0.701 | 109 |
| Information | 0.740 | 0.766 | 0.753 | 201 |
| Guidance | 0.707 | 0.800 | 0.751 | 130 |
| Media Generation | 0.619 | 0.542 | 0.578 | 24 |
| Writing | 0.571 | 0.421 | 0.485 | 19 |
| Other | 0.457 | 0.410 | 0.432 | 117 |

**Confusion Matrix (internal test)**

```
           Pred  Pred  Pred  Pred  Pred  Pred
           Cod   Inf   Gui   Med   Wri   Oth
True Cod    75     8    12     1     1    12
True Inf     5   154    13     0     1    28
True Gui     9     6   104     0     1    10
True Med     3     2     0    13     2     4
True Wri     0     0     3     5     8     3
True Oth    13    38    15     2     1    48
```

### Held‑Out Test (400 independent samples)

| Metric | Value |
|--------|-------|
| Exact Match Accuracy | 76.5% |
| Macro F1 | – (per‑category below) |

**Per‑Category Breakdown**

| Category | Exact Count | Exact Accuracy |
|----------|-------------|----------------|
| Coding | 35/44 | 79.5% |
| Information | 184/229 | 80.3% |
| Guidance | 55/63 | 87.3% |
| Media Generation | 8/14 | 57.1% |
| Writing | 9/20 | 45.0% |
| Other | 15/30 | 50.0% |

## How to Get Started with the Model

```python
from transformers import AutoTokenizer, AutoModel
import torch
import torch.nn as nn

class CategoryRouter(nn.Module):
    def __init__(self, backbone="microsoft/deberta-v3-base", num_labels=6):
        super().__init__()
        self.backbone = AutoModel.from_pretrained(backbone)
        hidden_size = self.backbone.config.hidden_size
        self.classifier = nn.Sequential(
            nn.Dropout(0.1),
            nn.Linear(hidden_size, 256),
            nn.GELU(),
            nn.Dropout(0.1),
            nn.Linear(256, num_labels),
        )

    def forward(self, input_ids, attention_mask):
        outputs = self.backbone(input_ids=input_ids, attention_mask=attention_mask)
        cls_output = outputs.last_hidden_state[:, 0, :]
        return self.classifier(cls_output)

# Load
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
tokenizer = AutoTokenizer.from_pretrained("RowRed/CategoryRouter")
model = CategoryRouter()
model.load_state_dict(
    torch.load("pytorch_model.bin", map_location=device),
    strict=False
)
model.to(device)
model.eval()

# Predict
prompts = ["Write a Python function", "What is the normal force?"]
encoded = tokenizer(prompts, padding=True, truncation=True, return_tensors="pt").to(device)
with torch.no_grad():
    logits = model(encoded["input_ids"], encoded["attention_mask"])
    probs = torch.softmax(logits, dim=-1)
    predictions = torch.argmax(probs, dim=-1)

category_map = ["Coding", "Information", "Guidance", "Media Generation", "Writing", "Other"]
for prompt, idx in zip(prompts, predictions):
    print(f"Category: {category_map[idx.item()]} – {prompt}")
```

## Citation
If you use this model, please cite:

```bibtex
@software{CategoryRouter,
  author = {RowRed},
  title = {CategoryRouter},
  year = {2026},
  url = {https://huggingface.co/RowRed/CategoryRouter}
}
```

Additionally, acknowledge the base dataset and labeling model:

```bibtex
@dataset{oasst2,
  author = {OpenAssistant Contributors},
  title = {Open Assistant Conversations Dataset Release 2},
  year = {2023},
  url = {https://huggingface.co/datasets/OpenAssistant/oasst2}
}

@software{qwen3.5-4b,
  author = {Qwen Team},
  title = {Qwen3.5-4B},
  year = {2026},
  url = {https://huggingface.co/Qwen/Qwen3.5-4B}
}
```

## License
This model is released under Apache‑2.0.
The backbone (microsoft/deberta-v3-base) is MIT‑licensed.
The training dataset is derived from OASST2 (Apache‑2.0) and Qwen3.5‑4B outputs (Apache‑2.0).