hotdog / app.py
suiovo's picture
add
d412efa
Raw
History Blame Contribute Delete
2.88 kB
import gradio as gr
import torch
from transformers import AutoTokenizer, AutoModelForSequenceClassification
# 加载本地模型和分词器
# 请将下面的路径替换为你的模型实际保存路径
MODEL_PATH = "./fake_news_model" # 模型保存的文件夹路径
# 加载分词器和模型
tokenizer = AutoTokenizer.from_pretrained(MODEL_PATH)
model = AutoModelForSequenceClassification.from_pretrained(MODEL_PATH)
# 设置设备(GPU如果可用,否则用CPU)
device = "cuda" if torch.cuda.is_available() else "cpu"
model.to(device)
model.eval()
def predict_news(text):
"""
预测新闻是否为虚假新闻
返回: 0表示真实新闻, 1表示虚假新闻, 以及对应的概率
"""
if not text.strip():
return "请输入新闻文本", 0.0
# 对文本进行分词处理
inputs = tokenizer(
text,
truncation=True,
padding=True,
max_length=512,
return_tensors="pt"
).to(device)
# 模型预测
with torch.no_grad():
outputs = model(**inputs)
logits = outputs.logits
probabilities = torch.softmax(logits, dim=1)
# 获取预测结果和概率
predicted_class = torch.argmax(probabilities, dim=1).item()
confidence = probabilities[0][predicted_class].item() * 100
# 转换为可读性更好的结果
result = "虚假新闻" if predicted_class == 1 else "真实新闻"
return f"{result} (可信度: {confidence:.2f}%)", predicted_class
# 创建Gradio界面
with gr.Blocks(title="虚假新闻检测") as demo:
gr.Markdown("# 📰 虚假新闻检测工具")
gr.Markdown("输入新闻文本,系统将判断其为真实新闻(0)或虚假新闻(1)")
with gr.Row():
with gr.Column(scale=3):
input_text = gr.Textbox(
label="请输入新闻文本",
lines=10,
placeholder="在这里粘贴新闻内容..."
)
submit_btn = gr.Button("检测", variant="primary")
with gr.Column(scale=1):
output_result = gr.Textbox(label="检测结果", interactive=False)
output_label = gr.Number(label="分类标签 (0=真实, 1=虚假)", interactive=False)
# 设置按钮点击事件
submit_btn.click(
fn=predict_news,
inputs=input_text,
outputs=[output_result, output_label]
)
# 添加清除按钮
clear_btn = gr.Button("清除")
clear_btn.click(
fn=lambda: ("", "", 0),
inputs=[],
outputs=[input_text, output_result, output_label]
)
# 允许用户标记结果是否正确
demo.allow_flagging = "manual"
demo.flagging_options = ["结果正确", "结果错误"]
if __name__ == "__main__":
demo.launch(share=True) # 如果需要分享到互联网,将share改为True