RanM/CoreferenceResolution
0
1import os2from fastapi import FastAPI, HTTPException3from fastapi.middleware.cors import CORSMiddleware4from pydantic import BaseModel5import spacy6import re7from typing import List8 9# Set environment variables for writable directories10os.environ['TRANSFORMERS_CACHE'] = '/tmp/transformers_cache'11os.environ['MPLCONFIGDIR'] = '/tmp/.matplotlib'12 13# Initialize FastAPI app14app = FastAPI()15 16# Add CORS middleware17app.add_middleware(18 CORSMiddleware,19 allow_origins=["*"], # Adjust the origins as needed20 allow_credentials=True,21 allow_methods=["*"],22 allow_headers=["*"],23)24 25# Load the spaCy models once26nlp = spacy.load("en_core_web_sm")27nlp_coref = spacy.load("en_coreference_web_trf")28 29REPLACE_PRONOUNS = {"he","his", "she", "her", "they", "He", "His", "She", "Her", "They"}30 31class CorefRequest(BaseModel):32 text: str33 main_characters: List[str]34 35def extract_core_name(mention_text, main_characters):36 words = mention_text.split()37 for character in main_characters:38 if character.lower() in mention_text.lower():39 return character40 return words[-1]41 42def calculate_pronoun_density(text):43 doc = nlp(text)44 pronoun_count = sum(1 for token in doc if token.pos_ == "PRON" and token.text in REPLACE_PRONOUNS)45 named_entity_count = sum(1 for ent in doc.ents if ent.label_ == "PERSON")46 return pronoun_count / max(named_entity_count, 1), named_entity_count47 48def resolve_coreferences_across_text(text, main_characters):49 doc = nlp_coref(text)50 coref_mapping = {}51 for key, cluster in doc.spans.items():52 if re.match(r"coref_clusters_*", key):53 main_mention = cluster[0]54 core_name = extract_core_name(main_mention.text, main_characters)55 if core_name in main_characters:56 for mention in cluster:57 for token in mention:58 if token.text in REPLACE_PRONOUNS:59 core_name_final = core_name if token.text.istitle() else core_name.lower()60 coref_mapping[token.i] = core_name_final61 resolved_tokens = []62 current_sentence_characters = set()63 current_sentence = []64 for i, token in enumerate(doc):65 if token.is_sent_start and current_sentence:66 resolved_tokens.extend(current_sentence)67 current_sentence_characters.clear()68 current_sentence = []69 if i in coref_mapping:70 core_name = coref_mapping[i]71 if core_name not in current_sentence_characters and core_name.lower() not in [t.lower() for t in current_sentence]:72 current_sentence.append(core_name)73 current_sentence_characters.add(core_name)74 else:75 current_sentence.append(token.text)76 else:77 current_sentence.append(token.text)78 resolved_tokens.extend(current_sentence)79 resolved_text = " ".join(resolved_tokens)80 return remove_consecutive_duplicate_phrases(resolved_text)81 82def remove_consecutive_duplicate_phrases(text):83 words = text.split()84 i = 085 while i < len(words) - 1:86 j = i + 187 while j < len(words):88 if words[i:j] == words[j:j + (j - i)]:89 del words[j:j + (j - i)]90 else:91 j += 192 i += 193 return " ".join(words)94 95def process_text(text, main_characters):96 pronoun_density, named_entity_count = calculate_pronoun_density(text)97 min_named_entities = len(main_characters)98 if pronoun_density > 0:99 return resolve_coreferences_across_text(text, main_characters)100 else:101 return text102 103@app.post("/predict")104async def predict(coref_request: CorefRequest):105 resolved_text = process_text(coref_request.text, coref_request.main_characters)106 if resolved_text:107 return {"resolved_text": resolved_text}108 raise HTTPException(status_code=400, detail="Coreference resolution failed")109 110if __name__ == "__main__":111 import uvicorn112 uvicorn.run(app, host="0.0.0.0", port=int(os.getenv("PORT", 7860)))113 