CoolFace
Apppublic

RanM/CoreferenceResolution

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