Shrikrishna's picture
Update app.py
2258940 verified
Raw
History Blame Contribute Delete
5.34 kB
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"
)