Download classifier_inference.py from Deepnar/ice-v2-classifier: direct link, hf CLI and curl.
- Browser
- Download file 2.23 kB
-
https://huggingface.co/Deepnar/ice-v2-classifier/resolve/main/classifier_inference.py
- Command line
-
hf download hf://Deepnar/ice-v2-classifier/classifier_inference.py
-
curl -L -o classifier_inference.py https://huggingface.co/Deepnar/ice-v2-classifier/resolve/main/classifier_inference.py
2.23 kB
| """ICE v2 learned-head input and decoding, without database or DI3 rules.""" | |
| import json | |
| from pathlib import Path | |
| import torch | |
| def build_input(prompt: str, context_text: str | None = None) -> str: | |
| if context_text: | |
| return ( | |
| f"Conversation context (summarized):\n{context_text}\n\n" | |
| "Given the above conversation and the user's latest prompt, predict:\n" | |
| "1. TOPIC: what is the subject (Software_&_Tech, Creative_&_Media, etc.)\n" | |
| "2. INTENT: what is the user trying to do (Factual_Retrieval, Troubleshooting, etc.)\n" | |
| "3. CONTEXT RELIANCE: does the user need memory (Zero_Shot, Long_Term_Memory, Real_Time_Search)\n\n" | |
| f"User prompt: {prompt}" | |
| ) | |
| return ( | |
| "Given a user prompt, predict:\n" | |
| "1. TOPIC: what is the subject (Software_&_Tech, Creative_&_Media, etc.)\n" | |
| "2. INTENT: what is the user trying to do (Factual_Retrieval, Troubleshooting, etc.)\n" | |
| "3. CONTEXT RELIANCE: does the user need memory (Zero_Shot, Long_Term_Memory, Real_Time_Search)\n\n" | |
| f"User prompt: {prompt}" | |
| ) | |
| def predict_head(model, embedder, prompt: str, context_text: str | None = None): | |
| config = json.loads(Path(__file__).with_name("config.json").read_text()) | |
| embedding = embedder.encode(build_input(prompt, context_text), convert_to_tensor=True) | |
| if embedding.shape != (384,): | |
| raise ValueError("ICE v2 requires exactly 384 embedding coordinates") | |
| with torch.no_grad(): | |
| outputs = model(embedding.unsqueeze(0).float()) | |
| topic = torch.sigmoid(outputs[0, :11]) | |
| intent = torch.sigmoid(outputs[0, 11:22]) | |
| context = torch.softmax(outputs[0, 22:], dim=0) | |
| def tags(probs, labels): | |
| return [labels[i] for i, p in enumerate(probs) if p > 0.3] or [labels[probs.argmax().item()]] | |
| probabilities = topic.tolist() + intent.tolist() + context.tolist() | |
| return { | |
| "topic_tags": tags(topic, config["topic_labels"]), | |
| "intent_tags": tags(intent, config["intent_labels"]), | |
| "context_reliance": config["context_labels"][context.argmax().item()], | |
| "raw_probs": probabilities, | |
| "max_confidence": max(probabilities), | |
| } | |