Spaces:
Running on Zero
Running on Zero
| import spaces | |
| import html | |
| import json | |
| import re | |
| import tempfile | |
| from functools import lru_cache | |
| from pathlib import Path | |
| import gradio as gr | |
| import torch | |
| from huggingface_hub import hf_hub_download | |
| from inference import ( | |
| MODEL_REPO_ID, | |
| load_lemmatizer, | |
| load_registry, | |
| ) | |
| REGISTRY = load_registry() | |
| TARGET_COL_IDX = 2 | |
| BATCH_SIZE = 32 | |
| REPLACEMENTS = { | |
| '': 'и', | |
| '': '҃', | |
| '': '҃', | |
| '': 'ч', | |
| '': 'Ѥ', | |
| '': 'н', | |
| '': '҇', | |
| '': '҃', | |
| '': '~', | |
| '': 'ⷦ҇', | |
| '': ' ⷮ', | |
| '': '҆̀', | |
| '': '҆̀', | |
| '': 'ⷹ', | |
| '': 'ч', | |
| '': 'ⷹ', | |
| '': ':', | |
| '': 'Чⷹ', | |
| '': 'о', | |
| '': 'с', | |
| '': 'е', | |
| '': '͠', | |
| '': '·̀', | |
| '': '·̀', | |
| 'ⷭⷭ': '҇', | |
| '': '҇', | |
| '': '҆', | |
| '': 'ⷩ', | |
| '': 'ꙶ', | |
| '': 'оу', | |
| '': 'ꙁ', | |
| '': 'ⷿ', | |
| '': 'ⷿ', | |
| '': 'ѧ', | |
| ' ': 'с', | |
| '': 'ѱ', | |
| '': 'Ѥ', | |
| '': 'р҃', | |
| '꙳': 'у', | |
| '꙯': '҃', | |
| '': '͡', | |
| '': '͠', | |
| 'ъ': 'уъ', | |
| '': 'у', | |
| 'ⷣ': 'ⷣ͡', | |
| 'ⷮ': 'ⷣ͡', | |
| 'ⷯ': '̈͠', | |
| '': 'ꙩ́', | |
| 'ѧ': 'уѧ', | |
| '': '҆', | |
| '': 'ꙺ', | |
| '': 'Ю', | |
| '': 'ꙻ', | |
| 'ⷤ': '', | |
| } | |
| ORDERED_REPLACEMENTS = sorted( | |
| REPLACEMENTS.items(), | |
| key=lambda item: len(item[0]), | |
| reverse=True, | |
| ) | |
| def preprocess_text(text): | |
| text = str(text or "") | |
| for source, replacement in ORDERED_REPLACEMENTS: | |
| text = text.replace(source, replacement) | |
| return text.strip() | |
| def display_name(item): | |
| return f"{item['language']} - {item['treebank']}" | |
| LANGUAGES = sorted( | |
| { | |
| item["language"] | |
| for item in REGISTRY.values() | |
| } | |
| ) | |
| DISPLAY_TO_ID = { | |
| display_name(item): model_id | |
| for model_id, item in REGISTRY.items() | |
| } | |
| def treebank_choices(language): | |
| return sorted( | |
| display_name(item) | |
| for item in REGISTRY.values() | |
| if item["language"] == language | |
| ) | |
| def default_language(): | |
| preferred = "Old Church Slavonic" | |
| if preferred in LANGUAGES: | |
| return preferred | |
| return LANGUAGES[0] if LANGUAGES else None | |
| def default_treebank(language): | |
| choices = treebank_choices(language) | |
| preferred = "Old Church Slavonic - Combined Model" | |
| if preferred in choices: | |
| return preferred | |
| return choices[0] if choices else None | |
| def update_treebanks(language): | |
| choices = treebank_choices(language) | |
| return gr.Dropdown( | |
| choices=choices, | |
| value=default_treebank(language), | |
| ) | |
| def selected_model_id(selected_name): | |
| if ( | |
| not selected_name | |
| or selected_name not in DISPLAY_TO_ID | |
| ): | |
| raise ValueError( | |
| "Please select a valid model." | |
| ) | |
| return DISPLAY_TO_ID[selected_name] | |
| def load_vocab_chars_for_model(model_id): | |
| item = REGISTRY[model_id] | |
| vocab_path = hf_hub_download( | |
| repo_id=MODEL_REPO_ID, | |
| repo_type="model", | |
| filename=item["vocab_file"], | |
| ) | |
| with open(vocab_path, encoding="utf8") as file: | |
| vocab_data = json.load(file) | |
| return set(vocab_data["char2idx"]) - { | |
| "<pad>", | |
| "<sos>", | |
| "<eos>", | |
| "<unk>", | |
| } | |
| def unsupported_input_for_model( | |
| text, | |
| allowed_chars, | |
| max_bad_ratio=0.60, | |
| min_checked_chars=4, | |
| ): | |
| checked = [] | |
| bad = [] | |
| ignored = { | |
| ".", | |
| ",", | |
| ";", | |
| ":", | |
| "!", | |
| "?", | |
| "'", | |
| '"', | |
| "(", | |
| ")", | |
| "[", | |
| "]", | |
| "/", | |
| } | |
| for character in text: | |
| if character.isspace(): | |
| continue | |
| if character.isdigit(): | |
| continue | |
| if character in ignored: | |
| continue | |
| checked.append(character) | |
| if character not in allowed_chars: | |
| bad.append(character) | |
| if len(checked) < min_checked_chars: | |
| return False, [] | |
| bad_ratio = len(bad) / len(checked) | |
| return ( | |
| bad_ratio >= max_bad_ratio, | |
| sorted(set(bad)), | |
| ) | |
| def make_html_table(tokens, lemmas): | |
| if not tokens: | |
| return "" | |
| rows = [] | |
| for token, lemma in zip(tokens, lemmas): | |
| rows.append( | |
| "<tr>" | |
| f"<td>{html.escape(token)}</td>" | |
| f"<td>{html.escape(lemma)}</td>" | |
| "</tr>" | |
| ) | |
| return ( | |
| '<div id="token-card">' | |
| '<table class="lemma-table">' | |
| "<thead>" | |
| "<tr>" | |
| "<th>Word</th>" | |
| "<th>Lemma</th>" | |
| "</tr>" | |
| "</thead>" | |
| "<tbody>" | |
| f"{''.join(rows)}" | |
| "</tbody>" | |
| "</table>" | |
| "</div>" | |
| ) | |
| def lemmatize_sentence(sentence, selected_name): | |
| sentence = str(sentence or "").strip() | |
| if not sentence: | |
| return "", "" | |
| try: | |
| processed_sentence = preprocess_text(sentence) | |
| if not processed_sentence: | |
| return "", "" | |
| model_id = selected_model_id( | |
| selected_name | |
| ) | |
| allowed_chars = load_vocab_chars_for_model( | |
| model_id | |
| ) | |
| is_bad, bad = unsupported_input_for_model( | |
| processed_sentence, | |
| allowed_chars, | |
| ) | |
| if is_bad: | |
| message = ( | |
| "The processed input contains too many " | |
| "characters that are not present in the " | |
| "selected model vocabulary. " | |
| "Unsupported characters: " | |
| + " ".join(bad[:20]) | |
| ) | |
| return ( | |
| "", | |
| f"<p>{html.escape(message)}</p>", | |
| ) | |
| tokens = processed_sentence.split() | |
| lemmatizer = load_lemmatizer( | |
| model_id=model_id, | |
| device="cuda", | |
| ) | |
| lemmas = lemmatizer.lemmatize_sentence( | |
| tokens | |
| ) | |
| result = " ".join(lemmas) | |
| table = make_html_table(tokens, lemmas) | |
| del lemmatizer | |
| torch.cuda.empty_cache() | |
| return result, table | |
| except Exception as error: | |
| torch.cuda.empty_cache() | |
| return ( | |
| "", | |
| f"<p>{html.escape(str(error))}</p>", | |
| ) | |
| def parse_conllu_sentences_from_text(text): | |
| text = text.strip() | |
| sentences = [] | |
| for block in re.split(r"\n\n+", text): | |
| sentence = [] | |
| for line in block.splitlines(): | |
| if not line: | |
| continue | |
| if line.startswith("#"): | |
| continue | |
| columns = line.split("\t") | |
| if len(columns) != 10: | |
| continue | |
| token_id = columns[0] | |
| if "-" in token_id or "." in token_id: | |
| continue | |
| processed_form = preprocess_text( | |
| columns[1] | |
| ) | |
| if processed_form: | |
| sentence.append(processed_form) | |
| if sentence: | |
| sentences.append(sentence) | |
| return sentences | |
| def make_source_for_token( | |
| tokens, | |
| index, | |
| k_context, | |
| sep_char, | |
| ): | |
| form = tokens[index] | |
| left_context = tokens[ | |
| max(0, index - k_context):index | |
| ] | |
| right_context = tokens[ | |
| index + 1:index + 1 + k_context | |
| ] | |
| left = " ".join(left_context).strip() | |
| right = " ".join(right_context).strip() | |
| src_left = left + " " if left else "" | |
| src_right = " " + right if right else "" | |
| return ( | |
| f"{src_left}" | |
| f"{sep_char}" | |
| f"{form}" | |
| f"{sep_char}" | |
| f"{src_right}" | |
| ) | |
| def make_all_sources_from_conllu( | |
| text, | |
| lemmatizer, | |
| ): | |
| sentences = parse_conllu_sentences_from_text( | |
| text | |
| ) | |
| sources = [] | |
| for tokens in sentences: | |
| for index in range(len(tokens)): | |
| source = make_source_for_token( | |
| tokens=tokens, | |
| index=index, | |
| k_context=lemmatizer.k_context, | |
| sep_char=lemmatizer.sep_char, | |
| ) | |
| sources.append(source) | |
| return sources | |
| def predict_sources_batched( | |
| sources, | |
| lemmatizer, | |
| batch_size=BATCH_SIZE, | |
| ): | |
| predictions = [] | |
| if not sources: | |
| return predictions | |
| pad_id = lemmatizer.vocab.char2idx["<pad>"] | |
| sos_id = lemmatizer.vocab.char2idx["<sos>"] | |
| eos_id = lemmatizer.vocab.char2idx["<eos>"] | |
| for start in range( | |
| 0, | |
| len(sources), | |
| batch_size, | |
| ): | |
| batch_sources = sources[ | |
| start:start + batch_size | |
| ] | |
| source_ids_list = [] | |
| source_lengths = [] | |
| for source_string in batch_sources: | |
| source_ids = ( | |
| [sos_id] | |
| + lemmatizer.vocab.encode( | |
| source_string | |
| ) | |
| + [eos_id] | |
| ) | |
| source_ids_list.append(source_ids) | |
| source_lengths.append( | |
| len(source_ids) | |
| ) | |
| maximum_length = max(source_lengths) | |
| padded = [ | |
| ids | |
| + [pad_id] | |
| * (maximum_length - len(ids)) | |
| for ids in source_ids_list | |
| ] | |
| source_tensor = torch.tensor( | |
| padded, | |
| dtype=torch.long, | |
| device=lemmatizer.device, | |
| ) | |
| length_tensor = torch.tensor( | |
| source_lengths, | |
| dtype=torch.long, | |
| device=lemmatizer.device, | |
| ) | |
| batch_predictions = ( | |
| lemmatizer.model.generate( | |
| source_tensor, | |
| length_tensor, | |
| lemmatizer.vocab, | |
| max_len=lemmatizer.max_gen_len, | |
| ) | |
| ) | |
| predictions.extend(batch_predictions) | |
| return predictions | |
| def write_back_conllu(input_text, predictions): | |
| text = input_text.rstrip("\n") | |
| blocks = re.split(r"\n\n+", text) | |
| output_blocks = [] | |
| prediction_index = 0 | |
| for block in blocks: | |
| new_lines = [] | |
| for line in block.split("\n"): | |
| if not line or line.startswith("#"): | |
| new_lines.append(line) | |
| continue | |
| columns = line.split("\t") | |
| if len(columns) != 10: | |
| new_lines.append(line) | |
| continue | |
| token_id = columns[0] | |
| if "-" in token_id or "." in token_id: | |
| new_lines.append(line) | |
| continue | |
| if prediction_index < len(predictions): | |
| prediction = predictions[ | |
| prediction_index | |
| ] | |
| else: | |
| prediction = "_" | |
| columns[TARGET_COL_IDX] = ( | |
| prediction if prediction else "_" | |
| ) | |
| new_lines.append( | |
| "\t".join(columns) | |
| ) | |
| prediction_index += 1 | |
| output_blocks.append( | |
| "\n".join(new_lines) | |
| ) | |
| output_text = ( | |
| "\n\n".join(output_blocks).rstrip() | |
| + "\n\n" | |
| ) | |
| return output_text, prediction_index | |
| def lemmatize_conllu_file( | |
| file_obj, | |
| selected_name, | |
| ): | |
| if file_obj is None: | |
| return ( | |
| gr.update( | |
| value=None, | |
| visible=False, | |
| ), | |
| "Please upload a CoNLL-U file.", | |
| ) | |
| try: | |
| model_id = selected_model_id( | |
| selected_name | |
| ) | |
| input_path = Path(file_obj) | |
| with input_path.open( | |
| encoding="utf8" | |
| ) as file: | |
| text = file.read() | |
| lemmatizer = load_lemmatizer( | |
| model_id=model_id, | |
| device="cuda", | |
| ) | |
| sources = make_all_sources_from_conllu( | |
| text, | |
| lemmatizer, | |
| ) | |
| predictions = predict_sources_batched( | |
| sources, | |
| lemmatizer, | |
| ) | |
| output_text, total = write_back_conllu( | |
| text, | |
| predictions, | |
| ) | |
| safe_name = re.sub( | |
| r"[^A-Za-z0-9_]+", | |
| "", | |
| selected_name.replace( | |
| " ", | |
| "_", | |
| ).replace( | |
| "-", | |
| "_", | |
| ), | |
| ) | |
| output_path = ( | |
| Path(tempfile.gettempdir()) | |
| / ( | |
| f"{input_path.stem}." | |
| f"{safe_name}." | |
| "lemmatized.conllu" | |
| ) | |
| ) | |
| with output_path.open( | |
| "w", | |
| encoding="utf8", | |
| newline="\n", | |
| ) as file: | |
| file.write(output_text) | |
| del lemmatizer | |
| torch.cuda.empty_cache() | |
| message = ( | |
| f"Done. Wrote {total:,} lemma predictions.\n" | |
| "PUA replacements were applied before inference.\n" | |
| "Input column: FORM.\n" | |
| "Updated column: LEMMA.\n" | |
| "Original FORM values and all other columns were preserved." | |
| ) | |
| return ( | |
| gr.update( | |
| value=str(output_path), | |
| visible=True, | |
| ), | |
| message, | |
| ) | |
| except Exception as error: | |
| torch.cuda.empty_cache() | |
| return ( | |
| gr.update( | |
| value=None, | |
| visible=False, | |
| ), | |
| f"Error: {error}", | |
| ) | |
| def reset_download_button(file_obj): | |
| return ( | |
| gr.update( | |
| value=None, | |
| visible=False, | |
| ), | |
| "", | |
| ) | |
| CUSTOM_CSS = """ | |
| body { | |
| background: linear-gradient( | |
| 135deg, | |
| #eaf3ff 0%, | |
| #ffffff 48%, | |
| #dbeafe 100% | |
| ); | |
| } | |
| .gradio-container { | |
| max-width: 980px !important; | |
| margin: auto !important; | |
| font-family: Arial, Helvetica, sans-serif !important; | |
| } | |
| #main-card { | |
| background: #ffffff; | |
| border: 1px solid #bfdbfe; | |
| border-radius: 26px; | |
| padding: 30px; | |
| box-shadow: 0 20px 50px rgba(15, 23, 42, 0.16); | |
| } | |
| #title { | |
| text-align: center; | |
| color: #020617; | |
| font-size: 2.5rem; | |
| font-weight: 900; | |
| } | |
| #subtitle { | |
| text-align: center; | |
| color: #1e40af; | |
| font-size: 1.05rem; | |
| line-height: 1.55; | |
| } | |
| #badge-row { | |
| text-align: center; | |
| margin-bottom: 1.2rem; | |
| } | |
| #badge-row span { | |
| display: inline-block; | |
| background: #eff6ff; | |
| color: #1e3a8a; | |
| border: 1px solid #bfdbfe; | |
| border-radius: 999px; | |
| padding: 7px 13px; | |
| margin: 4px; | |
| font-weight: 700; | |
| } | |
| button { | |
| background: linear-gradient( | |
| 90deg, | |
| #020617, | |
| #1d4ed8 | |
| ) !important; | |
| color: #ffffff !important; | |
| border-radius: 16px !important; | |
| font-weight: 900 !important; | |
| } | |
| #token-card { | |
| background: #f8fbff; | |
| border: 1.5px solid #1d4ed8; | |
| border-radius: 18px; | |
| padding: 18px; | |
| } | |
| .lemma-table { | |
| width: 100%; | |
| border-collapse: collapse; | |
| } | |
| .lemma-table th { | |
| background: #1d4ed8; | |
| color: white; | |
| padding: 12px; | |
| text-align: left; | |
| } | |
| .lemma-table td { | |
| padding: 12px; | |
| border-bottom: 1px solid #bfdbfe; | |
| } | |
| footer { | |
| display: none !important; | |
| } | |
| """ | |
| DEFAULT_LANGUAGE = default_language() | |
| DEFAULT_TREEBANK = default_treebank( | |
| DEFAULT_LANGUAGE | |
| ) | |
| APP_THEME = gr.themes.Soft( | |
| primary_hue="blue", | |
| secondary_hue="sky", | |
| neutral_hue="slate", | |
| ) | |
| with gr.Blocks( | |
| title="OCS Combined Lemmatizer", | |
| ) as demo: | |
| with gr.Column(elem_id="main-card"): | |
| gr.Markdown( | |
| "# Old Church Slavonic Lemmatizer Demo", | |
| elem_id="title", | |
| ) | |
| gr.Markdown( | |
| "Paste a tokenized sentence or upload a " | |
| "CoNLL-U file to generate lemma predictions. " | |
| "PUA characters are normalized before tokenization.", | |
| elem_id="subtitle", | |
| ) | |
| gr.HTML( | |
| """ | |
| <div id="badge-row"> | |
| <span>Old Church Slavonic</span> | |
| <span>Combined Model</span> | |
| <span>PUA Normalization</span> | |
| <span>Context-aware</span> | |
| <span>ZeroGPU</span> | |
| </div> | |
| """ | |
| ) | |
| with gr.Row(): | |
| language_input = gr.Dropdown( | |
| label="Language", | |
| choices=LANGUAGES, | |
| value=DEFAULT_LANGUAGE, | |
| ) | |
| treebank_input = gr.Dropdown( | |
| label="Model", | |
| choices=treebank_choices( | |
| DEFAULT_LANGUAGE | |
| ), | |
| value=DEFAULT_TREEBANK, | |
| ) | |
| with gr.Tab("Sentence"): | |
| sentence_input = gr.Textbox( | |
| label="Input sentence", | |
| lines=5, | |
| placeholder=( | |
| "Enter a whitespace-tokenized sentence" | |
| ), | |
| ) | |
| sentence_button = gr.Button( | |
| "Lemmatize sentence" | |
| ) | |
| sentence_output = gr.Textbox( | |
| label="Lemmatized sentence", | |
| lines=5, | |
| ) | |
| token_output = gr.HTML() | |
| with gr.Tab("CoNLL-U file"): | |
| conllu_input = gr.File( | |
| label="Upload CoNLL-U file", | |
| file_types=[".conllu", ".txt"], | |
| type="filepath", | |
| ) | |
| conllu_button = gr.Button( | |
| "Lemmatize CoNLL-U file" | |
| ) | |
| conllu_output = gr.DownloadButton( | |
| label="Download result", | |
| value=None, | |
| visible=False, | |
| ) | |
| conllu_message = gr.Textbox( | |
| label="Status", | |
| lines=5, | |
| ) | |
| language_input.change( | |
| fn=update_treebanks, | |
| inputs=language_input, | |
| outputs=treebank_input, | |
| ) | |
| sentence_button.click( | |
| fn=lemmatize_sentence, | |
| inputs=[ | |
| sentence_input, | |
| treebank_input, | |
| ], | |
| outputs=[ | |
| sentence_output, | |
| token_output, | |
| ], | |
| ) | |
| conllu_input.change( | |
| fn=reset_download_button, | |
| inputs=conllu_input, | |
| outputs=[ | |
| conllu_output, | |
| conllu_message, | |
| ], | |
| ) | |
| conllu_button.click( | |
| fn=lemmatize_conllu_file, | |
| inputs=[ | |
| conllu_input, | |
| treebank_input, | |
| ], | |
| outputs=[ | |
| conllu_output, | |
| conllu_message, | |
| ], | |
| ) | |
| demo.queue() | |
| demo.launch( | |
| server_name="0.0.0.0", | |
| server_port=7860, | |
| ssr_mode=False, | |
| theme=APP_THEME, | |
| css=CUSTOM_CSS, | |
| footer_links=[], | |
| ) |