File size: 4,398 Bytes
98683d0
0cfc2e9
 
 
98683d0
 
0cfc2e9
98683d0
0cfc2e9
 
98683d0
0cfc2e9
 
98683d0
0cfc2e9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
98683d0
0cfc2e9
 
 
13ab63d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
a03810a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
13ab63d
 
 
a03810a
 
 
 
 
 
13ab63d
e93043a
13ab63d
a03810a
13ab63d
e93043a
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
#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}'