junhyunpark01/Onpremise_LLM_Normal_Detection
0
1import gradio as gr2import torch3from transformers import AutoTokenizer, BertForSequenceClassification, AutoModel4from torch import nn5import re6 7 8def paragraph_leveling(text):9 model_name = "contrastive_encoder_sentence"10 model = AutoModel.from_pretrained(model_name)11 tokenizer = AutoTokenizer.from_pretrained('zzxslp/RadBERT-RoBERTa-4m')12 13 class MLP(nn.Module):14 def __init__(self, target_size=3, input_size=768):15 super(MLP, self).__init__()16 self.num_classes = target_size17 self.input_size = input_size18 self.fc1 = nn.Linear(input_size, target_size)19 20 def forward(self, x):21 out = self.fc1(x)22 return out23 24 classifier = MLP(target_size=3, input_size=768)25 classifier.load_state_dict(torch.load('fine_tunning_classifier', map_location=torch.device('cpu')))26 classifier.eval()27 28 output_list = []29 text_list = text.split(".")30 result = []31 32 output_list.append(("\n", None))33 34 for idx_sentence in text_list:35 train_encoding = tokenizer(36 idx_sentence,37 return_tensors='pt',38 padding='max_length',39 truncation=True,40 max_length=120)41 output = model(**train_encoding)42 output = classifier(output[1])43 output = output[0]44 45 if output.argmax(-1) == 0:46 output_list.append((idx_sentence, 'abnormal'))47 result.append(0)48 elif output.argmax(-1) == 1:49 output_list.append((idx_sentence, 'normal'))50 result.append(1)51 else:52 output_list.append((idx_sentence, 'uncertain'))53 result.append(2)54 55 output_list.append(('\n', None))56 if 0 in result:57 output_list.append(('FINAL LABEL: ', None))58 output_list.append(('ABNORMAL', 'abnormal'))59 60 else:61 output_list.append(('FINAL LABEL: ', None))62 output_list.append(('NORMAL', 'normal'))63 64 return output_list65 66 67demo = gr.Interface(68 paragraph_leveling,69 [70 gr.Textbox(71 label="Medical Report",72 info="You may put radiology medical report. Each sentence should be seperate with period mark.",73 lines=20,74 value=" ",75 ),76 ],77 gr.HighlightedText(78 label="labeling",79 show_legend = True,80 show_label = True,81 color_map={"abnormal": "violet", "normal": "lightgreen", "uncertain": "lightgray"}),82 theme=gr.themes.Base()83)84if __name__ == "__main__":85 demo.launch()86 87 