mothy-08 commited on
Commit
d533730
·
1 Parent(s): 8096acb

changed model

Browse files
api/config.py CHANGED
@@ -4,4 +4,6 @@ from transformers import pipeline
4
 
5
  @lru_cache(maxsize=1)
6
  def get_classifier():
7
- return pipeline("text-classification", "mothy-08/drbftcm")
 
 
 
4
 
5
  @lru_cache(maxsize=1)
6
  def get_classifier():
7
+ return pipeline(
8
+ "text-classification", "mothy-08/deberta-v3-xsmall-finetuned-content-moderator"
9
+ )
tests/__pycache__/test_moderator.cpython-313-pytest-9.0.1.pyc CHANGED
Binary files a/tests/__pycache__/test_moderator.cpython-313-pytest-9.0.1.pyc and b/tests/__pycache__/test_moderator.cpython-313-pytest-9.0.1.pyc differ
 
tests/test_moderator.py CHANGED
@@ -5,12 +5,12 @@ from api.content_moderator import ContentModerator
5
 
6
  def mock_classifier(text):
7
  if text == "Hello":
8
- return [{"label": "safe", "score": 0.99}]
9
- return [{"label": "unsafe", "score": 0.99}]
10
 
11
 
12
  def get_mock_moderator():
13
- return ContentModerator(classifier=mock_classifier)
14
 
15
 
16
  app.dependency_overrides[get_moderator] = get_mock_moderator
@@ -27,14 +27,14 @@ def test_predict_valid_text():
27
  response = client.post("/predict", json={"text": "Hello"})
28
  assert response.status_code == 200
29
  data = response.json()
30
- assert data["label"] == "safe"
31
  assert data["score"] == 0.99
32
 
33
 
34
  def test_predict_normalization():
35
  response = client.post("/predict", json={"text": "Hello"})
36
  assert response.status_code == 200
37
- assert response.json()["label"] == "safe"
38
 
39
 
40
  def test_predict_empty_text():
 
5
 
6
  def mock_classifier(text):
7
  if text == "Hello":
8
+ return [{"label": "SAFE", "score": 0.99}]
9
+ return [{"label": "UNSAFE", "score": 0.99}]
10
 
11
 
12
  def get_mock_moderator():
13
+ return ContentModerator(mock_classifier)
14
 
15
 
16
  app.dependency_overrides[get_moderator] = get_mock_moderator
 
27
  response = client.post("/predict", json={"text": "Hello"})
28
  assert response.status_code == 200
29
  data = response.json()
30
+ assert data["label"] == "SAFE"
31
  assert data["score"] == 0.99
32
 
33
 
34
  def test_predict_normalization():
35
  response = client.post("/predict", json={"text": "Hello"})
36
  assert response.status_code == 200
37
+ assert response.json()["label"] == "SAFE"
38
 
39
 
40
  def test_predict_empty_text():