deprem-ml/ner-active-learning
5
1import os2import gradio as gr3from gradio import FlaggingCallback4from gradio.components import IOComponent5 6from transformers import pipeline7 8from typing import List, Optional, Any9 10import argilla as rg11 12import os13 14 15 16nlp = pipeline("ner", model="deprem-ml/deprem-ner")17 18examples = [19 ["Lütfen yardım Akevler mahallesi Rüzgar sokak Tuncay apartmanı zemin kat Antakya akrabalarım göçük altında #hatay #Afad"]20]21 22def create_record(input_text, feedback):23 # define the record status based on feedback24 # default means it needs to be reviewed --> "Incorrect" or "Ambiguous"25 # validated means it's correct and has been checked --> "Correct"26 status = "Validated" if feedback == "Doğru" else "Default"27 28 # Making the prediction29 predictions = nlp(input_text, aggregation_strategy="first")30 31 # Creating the predicted entities as a list of tuples (entity, start_char, end_char, score)32 prediction = [(pred["entity_group"], pred["start"], pred["end"], pred["score"]) for pred in predictions]33 34 # Create word tokens35 batch_encoding = nlp.tokenizer(input_text)36 word_ids = sorted(set(batch_encoding.word_ids()) - {None})37 words = []38 for word_id in word_ids:39 char_span = batch_encoding.word_to_chars(word_id)40 words.append(input_text[char_span.start:char_span.end])41 42 # Building a TokenClassificationRecord43 record = rg.TokenClassificationRecord(44 text=input_text,45 tokens=words,46 prediction=prediction,47 prediction_agent="deprem-ml/deprem-ner",48 status=status,49 metadata={"feedback": feedback}50 )51 print(record)52 return record53 54class ArgillaLogger(FlaggingCallback):55 def __init__(self, api_url, api_key, dataset_name):56 rg.init(api_url=api_url, api_key=api_key)57 self.dataset_name = dataset_name58 def setup(self, components: List[IOComponent], flagging_dir: str):59 pass60 def flag(61 self,62 flag_data: List[Any],63 flag_option: Optional[str] = None,64 flag_index: Optional[int] = None,65 username: Optional[str] = None,66 ) -> int:67 text = flag_data[0]68 inference = flag_data[1]69 rg.log(name=self.dataset_name, records=create_record(text, flag_option))70 71 72 73gr.Interface.load(74 "models/deprem-ml/deprem-ner",75 examples=examples,76 title = "NER Adres Aktif Öğrenme Arayüzü",77 description = "Aşağıda veri girişi yapıp modelin çıktısına göre Doğru/Yanlış/Belirsiz olarak işaretleyerek modelimizi değerlendirmemize yardımcı olabilirsiniz. Not: flag'lere bir kez tıklamanız yeterlidir. Şu an arayüzü flag alındığında size feedback verecek şekilde düzeltiyoruz. ",78 allow_flagging="manual",79 flagging_callback=ArgillaLogger(80 api_url="https://sandbox.argilla.io", 81 api_key=os.getenv("TEAM_API_KEY"), 82 dataset_name="ner-flags"83 ),84 flagging_options=["Doğru", "Yanlış", "Belirsiz"]85).launch()