-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathclassifier_inference.py
More file actions
68 lines (56 loc) · 2.12 KB
/
Copy pathclassifier_inference.py
File metadata and controls
68 lines (56 loc) · 2.12 KB
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
from transformers import pipeline
import os
# Model path
MODEL_PATH = 'outputs/classifier'
# Global pipeline instance
_clf_pipe = None
def get_classifier():
global _clf_pipe
if _clf_pipe is None:
if not os.path.exists(MODEL_PATH):
raise FileNotFoundError(f"Classifier model not found at {MODEL_PATH}. Run training first.")
print(f"Loading Temporal Classifier from {MODEL_PATH}...")
_clf_pipe = pipeline(
'text-classification',
model=MODEL_PATH,
tokenizer=MODEL_PATH,
top_k=None # Returns scores for all labels
)
return _clf_pipe
def classify(claims: list[str]) -> list[dict]:
"""
Classifies a list of claims into stable, volatile, or uncertain.
Returns a list of dictionaries with labels, probabilities, and volatility scores.
"""
if not claims:
return []
clf = get_classifier()
results = clf(claims, truncation=True, max_length=128)
output = []
for claim_text, scores in zip(claims, results):
# Convert list of {label, score} to a map
score_map = {s['label']: s['score'] for s in scores}
# Get top label
top_label = max(score_map, key=score_map.get)
# Compute volatility score: volatile weight = 1.0, uncertain weight = 0.5
volatility = (score_map.get('volatile', 0) +
0.5 * score_map.get('uncertain', 0))
output.append({
'claim': claim_text,
'temporal_type': top_label,
'type_probs': score_map,
'volatility_score': round(volatility, 4),
})
return output
if __name__ == "__main__":
test_claims = [
"Sundar Pichai is the CEO of Google",
"The capital of France is Paris",
"What is the current price of Bitcoin?"
]
print("Running classification test...")
results = classify(test_claims)
for res in results:
print(f"\nClaim: {res['claim']}")
print(f"Type: {res['temporal_type']} (Volatility: {res['volatility_score']})")
print(f"Probs: {res['type_probs']}")