lone-wolf-beta's picture
not the last one
a03810a
Raw
History Blame Contribute Delete
4.4 kB
#importing libraries
import torch
import streamlit as st
from transformers import BertForQuestionAnswering
from transformers import BertTokenizer
import textwrap
#loading model
model = BertForQuestionAnswering.from_pretrained('bert-large-uncased-whole-word-masking-finetuned-squad')
#loading tokenizer
tokenizer = BertTokenizer.from_pretrained('bert-large-uncased-whole-word-masking-finetuned-squad')
#function for question answering
def answer_question(question, answer_text):
'''
Takes a `question` string and an `answer_text` string (which contains the
answer), and identifies the words within the `answer_text` that are the
answer. Prints them out.
'''
# ======== Tokenize ========
# Apply the tokenizer to the input text, treating them as a text-pair.
input_ids = tokenizer.encode(question, answer_text)
# Report how long the input sequence is.
print('Query has {:,} tokens.\n'.format(len(input_ids)))
# ======== Set Segment IDs ========
# Search the input_ids for the first instance of the `[SEP]` token.
sep_index = input_ids.index(tokenizer.sep_token_id)
# The number of segment A tokens includes the [SEP] token istelf.
num_seg_a = sep_index + 1
# The remainder are segment B.
num_seg_b = len(input_ids) - num_seg_a
# Construct the list of 0s and 1s.
segment_ids = [0]*num_seg_a + [1]*num_seg_b
# There should be a segment_id for every input token.
assert len(segment_ids) == len(input_ids)
# ======== Evaluate ========
# Run our example through the model.
outputs = model(torch.tensor([input_ids]), # The tokens representing our input text.
token_type_ids=torch.tensor([segment_ids]), # The segment IDs to differentiate question from answer_text
return_dict=True)
start_scores = outputs.start_logits
end_scores = outputs.end_logits
# ======== Reconstruct Answer ========
# Find the tokens with the highest `start` and `end` scores.
answer_start = torch.argmax(start_scores)
answer_end = torch.argmax(end_scores)
# Get the string versions of the input tokens.
tokens = tokenizer.convert_ids_to_tokens(input_ids)
# Start with the first token.
answer = tokens[answer_start]
# Select the remaining answer tokens and join them with whitespace.
for i in range(answer_start + 1, answer_end + 1):
# If it's a subword token, then recombine it with the previous token.
if tokens[i][0:2] == '##':
answer += tokens[i][2:]
# Otherwise, add a space then the token.
else:
answer += ' ' + tokens[i]
#print('Answer: "' + answer + '"')
return answer
# Wrap text to 80 characters.
wrapper = textwrap.TextWrapper(width=80)
# context=st.text_area('Enter the context of your Topic')
# question = st.text_area('Enter your question...')
# answer_key = 'answer_0'
# answer=answer_question(question, context)
# st.text(answer)
# #the loop will answer the questions until user ends the program by typing
# #'q' in the text field
# while question.lower()!='q':
# question = st.text_area('Enter your question...', key=answer_key)
# #answer=answer_question(question, bert_abstract)
# answer=answer_question(question, context)
# st.text(answer)
# if question.lower() =='q':
# break
# context = st.text_area('Enter the context of your Topic')
# question_key = 'question_0'
# question = st.text_area('Enter your question...', key=question_key)
# while question.lower() != 'q':
# answer = answer_question(question, context)
# st.text(answer)
# # Display the second question area only after the answer is generated
# question_key = f'question_{int(question_key.split("_")[1]) + 1}'
# question = st.text_area('Enter your next question...', key=question_key)
# if question.lower() == 'q':
# break
context = st.text_area('Enter the context of your Topic')
question_key = 'question_0'
while True:
question = st.text_area('Enter your question...', key=question_key)
if question.lower() == 'q':
break
answer = answer_question(question, context)
st.text(answer)
# Display the text area for the next question only after the current question is answered
question_key = f'question_{int(question_key.split("_")[1]) + 1}'