Spaces:
Running
Running
| import streamlit as st | |
| import numpy as np | |
| import pickle | |
| import tensorflow as tf | |
| import keras | |
| from tensorflow.keras.models import load_model | |
| from tensorflow.keras.preprocessing.sequence import pad_sequences | |
| from tensorflow.keras.layers import GRU, Bidirectional, Embedding, Dense, TimeDistributed | |
| import pandas as pd | |
| # ========================================== | |
| # 🔧 HOTFIX: Keras 3 Compatibility Patch | |
| # ========================================== | |
| # We define a custom GRU that ignores 'time_major' and other old arguments | |
| class CompatibleGRU(GRU): | |
| def __init__(self, *args, **kwargs): | |
| # Remove arguments that Keras 3 no longer supports | |
| if 'time_major' in kwargs: | |
| kwargs.pop('time_major') | |
| if 'implementation' in kwargs: | |
| kwargs.pop('implementation') | |
| if 'reset_after' in kwargs: | |
| # reset_after is supported but sometimes causes conflicts in restoration | |
| # We keep it unless it causes issues, but usually time_major is the culprit. | |
| pass | |
| super().__init__(*args, **kwargs) | |
| # We map the model's "GRU" layer to our safe "CompatibleGRU" | |
| custom_objects = { | |
| 'GRU': CompatibleGRU, | |
| 'Bidirectional': Bidirectional, | |
| 'Embedding': Embedding, | |
| 'Dense': Dense, | |
| 'TimeDistributed': TimeDistributed | |
| } | |
| # ========================================== | |
| # 📥 Data Loading | |
| # ========================================== | |
| def load_assets(): | |
| # Attempt to load with our patched custom objects | |
| try: | |
| model = load_model("bigru_pos_tagger.h5", custom_objects=custom_objects, compile=False) | |
| except Exception as e: | |
| st.error(f"Standard load failed: {e}. Trying legacy method...") | |
| # Fallback to legacy loading if the direct load fails | |
| model = keras.src.legacy.saving.legacy_h5_format.load_model_from_hdf5( | |
| "bigru_pos_tagger.h5", custom_objects=custom_objects | |
| ) | |
| with open("word2idx_bigru.pkl", "rb") as f: | |
| word2idx = pickle.load(f) | |
| with open("idx2tag_bigru.pkl", "rb") as f: | |
| idx2tag = pickle.load(f) | |
| return model, word2idx, idx2tag | |
| # Initialize global variables | |
| try: | |
| model, word2idx, idx2tag = load_assets() | |
| except Exception as e: | |
| st.error(f"Critical Error Loading Model: {e}") | |
| st.stop() | |
| MAX_LEN = 134 | |
| # ========================================== | |
| # 🧠 Processing Functions | |
| # ========================================== | |
| def pos_tag_sentence(sentence): | |
| tokens = sentence.split() | |
| if not tokens: | |
| return [] | |
| # word2idx is accessible here | |
| token_ids = [word2idx.get(token, 1) for token in tokens] | |
| token_ids_padded = pad_sequences([token_ids], maxlen=MAX_LEN, padding="post") | |
| # verbose=0 prevents console spam | |
| predictions = model.predict(token_ids_padded, verbose=0)[0] | |
| predicted_tags = [idx2tag[np.argmax(tag)] for tag in predictions][:len(tokens)] | |
| return list(zip(tokens, predicted_tags)) | |
| def format_tagged_line(tagged_list): | |
| return " ".join([f"{word}\\{tag}" for word, tag in tagged_list]) | |
| # ========================================== | |
| # 🖥️ Streamlit UI | |
| # ========================================== | |
| st.set_page_config(page_title="Konkani POS Tagger", layout="wide") | |
| st.title("Automatic Konkani POS Tagger") | |
| tab1, tab2 = st.tabs(["Single Sentence", "Batch File Upload"]) | |
| # --- TAB 1: Single Sentence --- | |
| with tab1: | |
| user_input = st.text_area("Enter a Konkani sentence:", height=100, key="single_input") | |
| if st.button("Tag Sentence"): | |
| if user_input.strip(): | |
| tagged = pos_tag_sentence(user_input.strip()) | |
| st.subheader("Tagged Result:") | |
| st.text(format_tagged_line(tagged)) | |
| st.table(pd.DataFrame(tagged, columns=["Word", "Tag"])) | |
| else: | |
| st.warning("Please enter text first.") | |
| # --- TAB 2: Batch File --- | |
| with tab2: | |
| st.subheader("Upload Text File") | |
| uploaded_file = st.file_uploader("Upload a .txt file (One sentence per line)", type=["txt"]) | |
| if uploaded_file is not None: | |
| if st.button("Process and Download"): | |
| # Read input file | |
| input_text = uploaded_file.getvalue().decode("utf-8") | |
| lines = input_text.splitlines() | |
| output_lines = [] | |
| progress_bar = st.progress(0) | |
| status_text = st.empty() | |
| for i, line in enumerate(lines): | |
| if line.strip(): | |
| tagged = pos_tag_sentence(line.strip()) | |
| output_lines.append(format_tagged_line(tagged)) | |
| else: | |
| output_lines.append("") | |
| # Update progress | |
| progress = (i + 1) / len(lines) | |
| progress_bar.progress(progress) | |
| status_text.text(f"Processing line {i+1} of {len(lines)}...") | |
| result_text = "\n".join(output_lines) | |
| status_text.text("Processing complete!") | |
| st.success(f"Successfully processed {len(lines)} lines!") | |
| st.download_button( | |
| label="📥 Download Tagged File", | |
| data=result_text, | |
| file_name="tagged_konkani_output.txt", | |
| mime="text/plain" | |
| ) |