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 @keras.saving.register_keras_serializable(package="Complex") 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 # ========================================== @st.cache_resource 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" )