Text_Swipe / app.py
D3V1L1810's picture
Update app.py
68925b3 verified
Raw
History Blame Contribute Delete
2.56 kB
import tensorflow as tf
from transformers import BertTokenizer, TFBertForSequenceClassification
import numpy as np
import json
import requests
import gradio as gr
bert_tokenizer = BertTokenizer.from_pretrained('BinaryTokenizer_ep5') #path + '/BinaryTokenizer_ep5'
bert_model = TFBertForSequenceClassification.from_pretrained('BinaryModel_ep5')
# def send_results_to_api(data, result_url):
# headers = {"Content-Type":"application/json"}
# response = requests.post(result_url, json=data, headers=headers)
# if response.status_code == 200:
# return response.json()
# else:
# return {"error": f"Failed to send results to API: {response.status_code}"}
def predict_text(params):
try:
params = json.loads(params)
except Exception as e:
return {"error": f"Invalid JSON input: {e.msg} at line {e.lineno} column {e.colno}"}
texts = params.get("urls", [])
if not params.get("normalfileID",[]):
file_ids = [None]*len(texts)
else:
file_ids = params.get("normalfileID",[])
# api = params.get("api", "")
# job_id = params.get("job_id", "")
print("JSON : \n",params)
if not texts:
return { "error": f"No Text Given"}
solutions = []
for text, file_id in zip(texts,file_ids):
encoding = bert_tokenizer.encode_plus(
text,
add_special_tokens=True,
max_length=128,
return_token_type_ids=True,
padding='max_length',
truncation=True,
return_attention_mask=True,
return_tensors='tf'
)
input_ids = encoding['input_ids']
token_type_ids = encoding['token_type_ids']
attention_mask = encoding['attention_mask']
pred = bert_model.predict([input_ids, token_type_ids, attention_mask])
logits = pred.logits
pred_label = tf.argmax(logits, axis=1).numpy()[0]
label = {1: 'positive', 0: 'negative'}
result = {'text': text, 'label': [label[pred_label]], "qcUser" : None, "file_id": file_id}
solutions.append(result)
# result_url = f"{api}/{job_id}"
# send_results_to_api(solutions, result_url)
return json.dumps({"solutions": solutions})
inputt = gr.Textbox(label="Parameter in JSON format (e.g., {'texts': ['sample text', 'sample text2']'})")
outputt = gr.JSON()
application = gr.Interface(fn=predict_text, inputs=inputt, outputs=outputt, title="Text Classification with BERT and API Integration")
application.launch()