CoolFace
Apppublic

junhyunpark01/Onpremise_LLM_Normal_Detection

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py87 linesDownload Raw Back to root
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