Views
No views yet
DistilBERT 占卜问题检测模型,可用于判断输入文本是否为符合塔罗占卜的问题。1 import torch
2 from transformers import DistilBertTokenizer, DistilBertForSequenceClassification
3
4 # 1. 加载模型
5 model_path = "./distilbert-question-detector"
6 tokenizer = DistilBertTokenizer.from_pretrained(model_path)
7 model = DistilBertForSequenceClassification.from_pretrained(model_path)
8 model.eval()
9
10 # 2. 进行推理
11 text = "Is this a question?"
12 inputs = tokenizer(text, return_tensors="pt", padding=True, truncation=True, max_length=128)
13
14 with torch.no_grad():
15 outputs = model(**inputs)
16 logits = outputs.logits
17 probabilities = torch.nn.functional.softmax(logits, dim=-1)
18
19 predicted_class = torch.argmax(probabilities, dim=-1).item()
20
21 print(f"Probabilities: {probabilities}")
22 print(f"Predicted class: {predicted_class}") # 1 代表是疑问句,0 代表不是1 from fastapi import FastAPI
2 import torch
3 from transformers import DistilBertTokenizer, DistilBertForSequenceClassification
4
5 app = FastAPI()
6
7 # 加载模型
8 model_path = "./distilbert-question-detector/checkpoint-5150"
9 tokenizer = DistilBertTokenizer.from_pretrained(model_path)
10 model = DistilBertForSequenceClassification.from_pretrained(model_path)
11 model.eval()
12
13 @app.post("/predict/")
14 async def predict(text: str):
15 inputs = tokenizer(text, return_tensors="pt", padding=True, truncation=True, max_length=128)
16
17 with torch.no_grad():
18 outputs = model(**inputs)
19 logits = outputs.logits
20 probabilities = torch.nn.functional.softmax(logits, dim=-1)
21
22 predicted_class = torch.argmax(probabilities, dim=-1).item()
23 return {"text": text, "probabilities": probabilities.tolist(), "predicted_class": predicted_class}1curl -X 'POST' \
2 'http://127.0.0.1:8000/predict/' \
3 -H 'Content-Type: application/json' \
4 -d '{"text": "Is this a valid question?"}'1{
2 "text": "Is this a valid question?",
3 "probabilities": [[0.9266, 0.0734]],
4 "predicted_class": 0
5}