SayItBetter / app.py
jang17's picture
Update app.py
668e068 verified
Raw
History Blame Contribute Delete
5.34 kB
import re
import gradio as gr
from transformers import T5ForConditionalGeneration, T5Tokenizer
from difflib import SequenceMatcher
# ----------------- Load Model -----------------
model_id = "jang17/AGLI_model" # Your model repo
tokenizer = T5Tokenizer.from_pretrained(model_id)
model = T5ForConditionalGeneration.from_pretrained(model_id)
# ----------------- Rule-Based Pre-Processor -----------------
def apply_rules(text: str) -> str:
if not text:
return text
# 1. Strip & collapse internal whitespace
text = text.strip()
text = re.sub(r' {2,}', ' ', text)
# 2. Remove space before punctuation
text = re.sub(r'\s+([.,!?;:])', r'\1', text)
# 3. Ensure one space after punctuation (but not at end-of-string)
text = re.sub(r'([.,!?;:])(?=[^\s])', r'\1 ', text)
# 4. Capitalize first character
text = text[0].upper() + text[1:] if len(text) > 1 else text.upper()
# 5. Capitalize letter after terminal punctuation
text = re.sub(
r'([.!?])\s+([a-z])',
lambda m: m.group(1) + ' ' + m.group(2).upper(),
text
)
# 6. Capitalize standalone 'i'
text = re.sub(r'\bi\b', 'I', text)
# 7. Add terminal period if missing
if text and text[-1] not in '.!?':
text += '.'
return text
# ----------------- Grammar Correction Function -----------------
def correct_grammar(text):
original_text = text.strip()
if not original_text:
return "", "", []
# Apply rules before T5
rule_corrected = apply_rules(original_text)
inputs = tokenizer(
"grammar: " + rule_corrected,
return_tensors="pt",
max_length=256,
truncation=True
)
outputs = model.generate(
**inputs,
max_length=256,
num_beams=6,
early_stopping=False
)
corrected = tokenizer.decode(outputs[0], skip_special_tokens=True)
# Apply rules after T5
corrected = apply_rules(corrected)
original_words = original_text.split()
corrected_words = corrected.split()
matcher = SequenceMatcher(None, original_words, corrected_words)
opcodes = list(matcher.get_opcodes())
# Build initial groups (merge gaps ≤1 token)
groups = []
current = None
for idx, (tag, i1, i2, j1, j2) in enumerate(opcodes):
if tag != "equal":
if current is None:
current = [i1, i2, j1, j2]
else:
current[1] = i2
current[3] = j2
else:
gap = i2 - i1
next_is_change = idx + 1 < len(opcodes) and opcodes[idx + 1][0] != "equal"
if current is not None and gap <= 1 and next_is_change:
current[1] = i2
current[3] = j2
else:
if current is not None:
groups.append(current)
current = None
if current is not None:
groups.append(current)
# Merge groups that share words (word-order / move detection)
merged = True
while merged:
merged = False
skip = set()
for i in range(len(groups)):
if i in skip:
continue
g1 = groups[i]
g1_orig = set(original_words[g1[0]:g1[1]])
g1_corr = set(corrected_words[g1[2]:g1[3]])
for j in range(i + 1, len(groups)):
if j in skip:
continue
g2 = groups[j]
g2_orig = set(original_words[g2[0]:g2[1]])
g2_corr = set(corrected_words[g2[2]:g2[3]])
if g1_orig & g2_corr or g1_corr & g2_orig:
groups[i] = [
min(g1[0], g2[0]), max(g1[1], g2[1]),
min(g1[2], g2[2]), max(g1[3], g2[3])
]
g1 = groups[i]
g1_orig = set(original_words[g1[0]:g1[1]])
g1_corr = set(corrected_words[g1[2]:g1[3]])
skip.add(j)
merged = True
groups = [g for i, g in enumerate(groups) if i not in skip]
errors = []
for i1, i2, j1, j2 in groups:
orig_part = " ".join(original_words[i1:i2])
corr_part = " ".join(corrected_words[j1:j2])
if orig_part or corr_part:
errors.append({
"original": orig_part,
"correction": corr_part
})
return original_text, corrected, errors
# ----------------- Gradio Interface -----------------
demo = gr.Interface(
fn=correct_grammar,
inputs=gr.Textbox(
label="Enter text to correct",
lines=5,
placeholder="Type or paste your sentence/paragraph here..."
),
outputs=[
gr.Textbox(label="Original Text"),
gr.Textbox(label="Corrected Text"),
gr.JSON(label="Errors")
],
title="AGLI Grammar Correction",
description="Fine-tuned T5 model for grammar correction. Paste text and get corrections + detected errors.",
examples=[
["He go to school yesterday but forget his book."],
["The childrens plays happy in the parks."],
["She don't likes apple but she love oranges."],
["I have cat."]
],
flagging_mode="never"
)
# ----------------- Launch -----------------
if __name__ == "__main__":
demo.launch()