"""
Robust Grammar Correction with FLAN-T5.
A Gradio demo that runs a CoLA-based quality gate per sentence and only sends
ungrammatical sentences through a FLAN-T5 grammar corrector. The output is
rendered as an inline diff with two view modes: a default highlight-only view
and a "show changes" view that displays original alongside each edit.
"""
from __future__ import annotations
# `import spaces` MUST happen before any torch/cuda use so the ZeroGPU runtime
# can install its lazy-CUDA shims. Import order in this block is deliberate.
# isort: off
import spaces # noqa: F401, I001 -- side-effectful import; must be first
import torch # noqa: I001
# isort: on
import difflib
import html
import logging
import re
import unicodedata
from dataclasses import dataclass
import gradio as gr
import pysbd
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, pipeline
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
log = logging.getLogger("grammar-demo")
log.info("torch %s | gradio %s", torch.__version__, gr.__version__)
# --------------------------------------------------------------------------------------
# Configuration
# --------------------------------------------------------------------------------------
@dataclass(frozen=True)
class Config:
"""Application-wide constants for the grammar correction demo."""
max_chars: int = 4000
quality_threshold: float = 0.90
batch_size: int = 4
max_new_tokens: int = 256
num_beams: int = 4
gpu_duration: int = 60
CFG = Config()
CHECKER_MODEL = "textattack/roberta-base-CoLA"
CORRECTOR_MODEL = "pszemraj/flan-t5-large-grammar-synthesis"
# --------------------------------------------------------------------------------------
# Model init (CPU at import time -- ZeroGPU forbids CUDA init in the main process)
# --------------------------------------------------------------------------------------
log.info("Loading checker model %s on CPU", CHECKER_MODEL)
checker = pipeline("text-classification", CHECKER_MODEL)
# The `text2text-generation` pipeline was removed in transformers 5.x. Load the
# corrector as a plain seq2seq model + tokenizer so the code works across both 4.x
# and 5.x without version pinning.
log.info("Loading corrector model %s on CPU", CORRECTOR_MODEL)
corrector_tokenizer = AutoTokenizer.from_pretrained(CORRECTOR_MODEL)
corrector_model = AutoModelForSeq2SeqLM.from_pretrained(CORRECTOR_MODEL)
log.info("Models ready")
# Resolve the "acceptable" label dynamically rather than hardcoding "LABEL_1".
# CoLA conventions vary; pick the label whose name contains "accept" if present.
_id2label = {int(k): v for k, v in checker.model.config.id2label.items()}
_ACCEPT_LABEL = next(
(v for v in _id2label.values() if "accept" in v.lower()),
_id2label.get(1, "LABEL_1"),
)
log.info("CoLA acceptable label resolved to %r", _ACCEPT_LABEL)
# Sentence segmenter: tolerates lowercase sentence starts, Mr./Dr., question
# marks, and other things the original regex splitter silently failed on.
SEGMENTER = pysbd.Segmenter(language="en", clean=False)
# --------------------------------------------------------------------------------------
# Text utilities
# --------------------------------------------------------------------------------------
# Word + punctuation + whitespace, with contractions kept as single tokens
# (so "don't" stays one token instead of splitting into "don", "'", "t").
TOKEN_PATTERN = re.compile(r"\w+(?:'\w+)*|[^\w\s]+|\s+")
def normalize_text(text: str) -> str:
"""Apply light, lossless normalization to user input.
Performs NFKC unicode normalization, collapses horizontal whitespace, and
limits consecutive newlines to two. Does not strip URLs, emails, or digits.
Args:
text: Raw user input.
Returns:
Normalized text with leading/trailing whitespace removed.
"""
if not text:
return ""
text = unicodedata.normalize("NFKC", text)
text = re.sub(r"[ \t]+", " ", text)
text = re.sub(r"\n{3,}", "\n\n", text)
return text.strip()
def split_sentences(text: str) -> list[str]:
"""Split text into sentences using a rule-based segmenter.
Args:
text: Input text to segment.
Returns:
List of non-empty, stripped sentence strings.
"""
if not text:
return []
return [s.strip() for s in SEGMENTER.segment(text) if s.strip()]
def tokenize(text: str) -> list[str]:
"""Split text into word, punctuation, and whitespace tokens.
Args:
text: Input string to tokenize.
Returns:
List of token strings preserving all original characters.
"""
return TOKEN_PATTERN.findall(text)
def chunk(items: list, size: int) -> list[list]:
"""Split a list into fixed-size sublists.
Args:
items: The list to partition.
size: Maximum number of elements per sublist.
Returns:
List of sublists, each containing at most ``size`` elements.
"""
return [items[i : i + size] for i in range(0, len(items), size)]
# --------------------------------------------------------------------------------------
# Diff rendering
# --------------------------------------------------------------------------------------
def render_diff(original: str, corrected: str, inline_mode: bool = False) -> str:
"""Render an inline diff between original and corrected text.
Default mode: shows the corrected text with insertions and replacements
highlighted with a green underline. Pure deletions are omitted.
Inline mode: shows the original text struck through followed by an arrow
and the corrected text for every edit. All edits are visible side by side.
Args:
original: The original (pre-correction) text.
corrected: The model-corrected text.
inline_mode: When True, render before/after for each edit.
Returns:
HTML string containing the diff markup.
"""
a = tokenize(original)
b = tokenize(corrected)
sm = difflib.SequenceMatcher(a=a, b=b)
parts: list[str] = []
has_change = False
for tag, i1, i2, j1, j2 in sm.get_opcodes():
old = "".join(a[i1:i2])
new = "".join(b[j1:j2])
if tag == "equal":
parts.append(html.escape(new))
continue
if tag == "insert":
if not new.strip():
# Whitespace-only insertion: emit silently to preserve spacing
parts.append(html.escape(new))
continue
has_change = True
parts.append(f'{html.escape(new)}')
continue
if tag == "delete":
if not old.strip():
continue # pure-whitespace deletion: silent
has_change = True
if inline_mode:
parts.append(f'{html.escape(old)}')
# default mode: omit deletions entirely
continue
if tag == "replace":
if not new.strip() and not old.strip():
parts.append(html.escape(new))
continue
has_change = True
if inline_mode:
parts.append(
''
f'{html.escape(old.strip())}'
'→'
f'{html.escape(new.strip())}'
""
)
else:
parts.append(f'{html.escape(new)}')
body = "".join(parts)
if not has_change:
return (
'
No changes were needed.
'
f'{html.escape(corrected)}
'
)
return f'{body}
'
# --------------------------------------------------------------------------------------
# Quality gate + correction
# --------------------------------------------------------------------------------------
def looks_clean(text: str, threshold: float) -> bool:
"""Run the CoLA quality gate on a single sentence.
Args:
text: A single sentence to classify.
threshold: Minimum confidence score to consider the sentence acceptable.
Returns:
True if CoLA classifies the sentence as acceptable above ``threshold``.
"""
res = checker(text, truncation=True)
label = res[0]["label"]
score = float(res[0]["score"])
return label == _ACCEPT_LABEL and score >= threshold
def _gen_kwargs() -> dict:
"""Build keyword arguments for ``model.generate()``.
Returns:
Dict of generation parameters (beam search, no sampling).
"""
return dict(
max_new_tokens=CFG.max_new_tokens,
num_beams=CFG.num_beams,
do_sample=False,
)
def _to_cuda_if_available() -> None:
"""Move both models to CUDA if a GPU is present.
Called inside the ``@spaces.GPU``-decorated function where the ZeroGPU
runtime has made the device available.
"""
if torch.cuda.is_available():
checker.model.to("cuda")
corrector_model.to("cuda")
def _correct_batch(sentences: list[str]) -> list[str]:
"""Run the seq2seq corrector on a batch of sentences.
Args:
sentences: One or more sentences to correct.
Returns:
List of corrected sentence strings, one per input.
"""
device = corrector_model.device
inputs = corrector_tokenizer(
sentences, return_tensors="pt", padding=True, truncation=True, max_length=512
).to(device)
with torch.inference_mode():
output_ids = corrector_model.generate(**inputs, **_gen_kwargs())
return [
corrector_tokenizer.decode(ids, skip_special_tokens=True).strip()
for ids in output_ids
]
@spaces.GPU(duration=CFG.gpu_duration)
def correct_text(text: str, progress: gr.Progress) -> str:
"""Split text into sentences, gate each with CoLA, and correct the bad ones.
Args:
text: Normalized full input text.
progress: Gradio progress tracker for the UI progress bar.
Returns:
Reassembled text with corrected sentences substituted in.
"""
sents = split_sentences(text)
if not sents:
return ""
_to_cuda_if_available()
progress(0.05, desc="Checking grammar")
needs_fix = [not looks_clean(s, CFG.quality_threshold) for s in sents]
fix_indices = [i for i, n in enumerate(needs_fix) if n]
fix_sents = [sents[i] for i in fix_indices]
corrections: dict[int, str] = {}
if fix_sents:
batches = chunk(fix_sents, CFG.batch_size)
idx_batches = chunk(fix_indices, CFG.batch_size)
for batch, idx_batch in progress.tqdm(
list(zip(batches, idx_batches, strict=True)), desc="Correcting"
):
corrected_batch = _correct_batch(batch)
for idx, corrected in zip(idx_batch, corrected_batch, strict=True):
corrections[idx] = corrected
out_sents = [corrections.get(i, s) for i, s in enumerate(sents)]
text_out = " ".join(out_sents)
text_out = re.sub(r"\s+([,.!?;:])", r"\1", text_out)
text_out = re.sub(r"\s{2,}", " ", text_out)
return text_out.strip()
def process(
text: str,
show_inline: bool,
progress: gr.Progress = gr.Progress(), # noqa: B008
) -> tuple[str, str, str]:
"""Main entry point wired to the Correct grammar button.
Args:
text: Raw user input from the textbox.
show_inline: Current state of the Show changes toggle.
progress: Gradio progress tracker (auto-injected per request).
Returns:
Tuple of (diff_html, cleaned_input, corrected_text). The latter two
populate ``gr.State``s for re-rendering and clipboard copy.
"""
placeholder = 'Enter some text to correct.
'
if not text or not text.strip():
return placeholder, "", ""
text = text[: CFG.max_chars]
cleaned = normalize_text(text)
if not cleaned:
return placeholder, "", ""
corrected = correct_text(cleaned, progress)
diff_html = render_diff(cleaned, corrected, inline_mode=show_inline)
return diff_html, cleaned, corrected
def rerender_diff(cleaned: str, corrected: str, show_inline: bool) -> str:
"""Re-render the diff with a different view mode without re-running the model.
Args:
cleaned: The normalized input text from the last correction run.
corrected: The corrected output text from the last correction run.
show_inline: When True, render inline before/after for each edit.
Returns:
HTML string containing the updated diff markup, or a placeholder if
no prior correction result is available.
"""
if not cleaned or not corrected:
return 'Output will appear here.
'
return render_diff(cleaned, corrected, inline_mode=show_inline)
# --------------------------------------------------------------------------------------
# UI
# --------------------------------------------------------------------------------------
THEME = gr.themes.Default(
primary_hue="emerald",
neutral_hue="slate",
radius_size="md",
spacing_size="sm",
)
CSS = """
:root {
--diff-add: #22c55e;
--diff-del: #ef4444;
--diff-arrow: rgba(148, 163, 184, 0.7);
--col-header-h: 36px;
--box-h: 380px;
}
/* ---- Page-level tighten-up ---- */
.gradio-container {
max-width: 1200px !important;
margin: 0 auto !important;
}
#title-block h1 {
margin: 0 0 8px 0 !important;
font-size: 26px !important;
font-weight: 600 !important;
letter-spacing: -0.02em !important;
}
#subtitle-block {
color: var(--body-text-color-subdued, #888) !important;
font-size: 14px !important;
margin-bottom: 20px !important;
}
#subtitle-block p { margin: 0 !important; }
/* ---- Column header row: matches box-top alignment on both sides ---- */
.col-header {
min-height: var(--col-header-h) !important;
align-items: center !important;
margin: 0 0 8px 0 !important;
gap: 10px !important;
flex-wrap: nowrap !important;
}
/* The "Output" label takes the leftover space, pushing controls to the right */
.col-label-wrap {
flex: 1 1 auto !important;
min-width: 60px !important;
}
.col-label {
font-size: 11px;
font-weight: 600;
text-transform: uppercase;
letter-spacing: 0.08em;
color: var(--body-text-color-subdued, #888);
padding: 0 2px;
line-height: var(--col-header-h);
white-space: nowrap;
}
.col-header .gradio-checkbox,
.col-header .form {
margin: 0 !important;
padding: 0 !important;
}
/* Force the Show changes checkbox wrapper to size to content, not stretch */
#show-changes-wrap {
flex: 0 0 auto !important;
width: auto !important;
max-width: 160px !important;
}
/* ---- Input textbox: drop the in-box label pill, match output height ---- */
#input-box {
height: var(--box-h) !important;
}
#input-box textarea {
min-height: calc(var(--box-h) - 24px) !important;
font-size: 14px !important;
line-height: 1.55 !important;
}
/* ---- Output diff box ---- */
#diff-output {
height: var(--box-h) !important;
overflow-y: auto;
padding: 14px 16px;
border-radius: 8px;
border: 1px solid var(--border-color-primary, rgba(128, 128, 128, 0.25));
background: var(--block-background-fill, transparent);
line-height: 1.65;
font-size: 14px;
}
#diff-output .d-out {
white-space: pre-wrap;
word-wrap: break-word;
}
#diff-output .d-empty {
color: var(--body-text-color-subdued, #888);
font-style: italic;
}
/* ---- Diff token styling ---- */
#diff-output ins.d-ins {
background: transparent;
color: inherit;
text-decoration: underline;
text-decoration-color: var(--diff-add);
text-decoration-thickness: 2px;
text-underline-offset: 3px;
padding: 0;
}
#diff-output del.d-del {
background: transparent;
color: var(--diff-del);
text-decoration: line-through;
text-decoration-color: var(--diff-del);
text-decoration-thickness: 2px;
padding: 0;
margin-right: 2px;
}
#diff-output .d-change { display: inline; }
#diff-output .d-change .d-arrow {
color: var(--diff-arrow);
margin: 0 5px;
font-weight: 500;
user-select: none;
}
/* ---- Footer row: char count on right side of input column ---- */
.col-footer {
min-height: 22px !important;
margin: 6px 0 0 0 !important;
align-items: center !important;
width: 100% !important;
flex-wrap: nowrap !important;
}
.col-footer > * {
flex: 1 1 100% !important;
width: 100% !important;
max-width: 100% !important;
}
.char-count {
text-align: right;
font-size: 12px;
color: var(--body-text-color-subdued, #888);
padding: 0 4px;
line-height: 22px;
width: 100%;
display: block;
}
.char-count.warn { color: #f59e0b; }
/* ---- Show changes: hide the input-style border, make compact ---- */
#show-changes-wrap {
flex: 0 0 auto !important;
width: auto !important;
max-width: 160px !important;
background: transparent !important;
border: none !important;
box-shadow: none !important;
padding: 0 !important;
}
#show-changes-wrap .form,
#show-changes-wrap .block {
background: transparent !important;
border: none !important;
box-shadow: none !important;
padding: 0 !important;
}
#show-changes-wrap label {
background: transparent !important;
border: none !important;
padding: 0 6px !important;
margin: 0 !important;
font-size: 13px !important;
color: var(--body-text-color, #ddd) !important;
cursor: pointer;
white-space: nowrap !important;
}
#show-changes-wrap label > span:not(:first-child) {
margin-left: 6px !important;
}
/* ---- Action buttons row ---- */
#action-row {
margin-top: 8px !important;
gap: 10px !important;
}
#action-row button {
height: 44px !important;
font-size: 14px !important;
font-weight: 500 !important;
}
/* ---- Smaller copy button ---- */
#copy-btn {
min-width: 72px !important;
max-width: 96px !important;
height: var(--col-header-h) !important;
font-size: 13px !important;
font-weight: 500 !important;
padding: 0 12px !important;
}
/* ---- Examples block: subtle ---- */
#examples-block .label {
font-size: 11px !important;
text-transform: uppercase;
letter-spacing: 0.08em;
color: var(--body-text-color-subdued, #888);
}
/* ---- Footer link block ---- */
#footer-models {
margin-top: 24px !important;
padding-top: 16px !important;
border-top: 1px solid var(--border-color-primary, rgba(128, 128, 128, 0.18));
color: var(--body-text-color-subdued, #888);
font-size: 12.5px;
}
#footer-models p { margin: 0 !important; }
"""
EXAMPLES = [
[
"I wen to the store yesturday to bye some food. I needd milk, bread, "
"and a few otter things. The store was really crowed and I had a hard "
"time finding everyting I needed."
],
[
"She don't has time for this kind of mistake, but he are insisting on "
"doing it anyway."
],
[
"The other question is how do I make money off of it? Was thinking of "
"buying some leaps expiring end of 27 as as puts."
],
[
"this paragraph have many problems with grammar and spelling, it also "
"lack proper sentence structure of any kind."
],
]
def update_count(text: str | None) -> str:
"""Return an HTML character counter string for the input textbox.
Args:
text: Current content of the input textbox (may be None on init).
Returns:
HTML div with the current character count and a warning class if the
count exceeds 95 percent of the configured maximum.
"""
n = len(text or "")
cls = " warn" if n >= int(CFG.max_chars * 0.95) else ""
return f'{n:,} / {CFG.max_chars:,}
'
with gr.Blocks(title="Grammar Correction", analytics_enabled=False) as demo:
with gr.Column(elem_id="title-block"):
gr.Markdown("# Robust Grammar Correction with FLAN-T5")
with gr.Column(elem_id="subtitle-block"):
gr.Markdown(
f"Enter text and click **Correct grammar**. Changes are underlined "
f"in the output; toggle **Show changes** to see each edit in "
f"before/after form. Input is truncated to {CFG.max_chars:,} "
f"characters."
)
# State carries the last-processed input and corrected text. Powers the
# show-changes toggle (re-renders the diff without re-running the model)
# and the copy-to-clipboard button (writes corrected text via JS).
cleaned_state = gr.State("")
corrected_state = gr.State("")
with gr.Row(equal_height=True):
# ---------- LEFT COLUMN: input ----------
with gr.Column():
with gr.Row(elem_classes=["col-header"]):
gr.HTML('Input
')
inp = gr.Textbox(
show_label=False,
container=False,
placeholder="Paste or type text to correct...",
lines=15,
max_length=CFG.max_chars,
buttons=["copy"],
value=EXAMPLES[0][0],
elem_id="input-box",
)
with gr.Row(elem_classes=["col-footer"]):
char_md = gr.HTML(value=update_count(EXAMPLES[0][0]))
# ---------- RIGHT COLUMN: output ----------
with gr.Column():
with gr.Row(elem_classes=["col-header"]):
gr.HTML(
'Output
',
elem_classes=["col-label-wrap"],
)
show_changes = gr.Checkbox(
label="Show changes",
value=False,
container=False,
scale=0,
min_width=140,
elem_id="show-changes-wrap",
)
copy_btn = gr.Button(
"Copy",
size="sm",
scale=0,
min_width=80,
elem_id="copy-btn",
)
diff_out = gr.HTML(
elem_id="diff-output",
padding=False,
value='Output will appear here.
',
)
with gr.Row(elem_classes=["col-footer"]):
gr.HTML(" ") # spacer to align with char count
with gr.Row(elem_id="action-row"):
run_btn = gr.Button("Correct grammar", variant="primary", scale=4)
clear_btn = gr.ClearButton(
[inp, diff_out, cleaned_state, corrected_state],
value="Clear",
scale=1,
)
with gr.Column(elem_id="examples-block"):
gr.Examples(
examples=EXAMPLES,
inputs=inp,
label="Examples",
)
with gr.Column(elem_id="footer-models"):
gr.Markdown(
f"**Models:** [`{CHECKER_MODEL}`](https://huggingface.co/{CHECKER_MODEL}) "
f"(per-sentence quality gate) and "
f"[`{CORRECTOR_MODEL}`](https://huggingface.co/{CORRECTOR_MODEL}) "
f"(corrector)."
)
inp.change(
fn=update_count,
inputs=inp,
outputs=char_md,
api_visibility="private",
show_progress="hidden",
)
run_btn.click(
fn=process,
inputs=[inp, show_changes],
outputs=[diff_out, cleaned_state, corrected_state],
api_visibility="undocumented",
)
show_changes.change(
fn=rerender_diff,
inputs=[cleaned_state, corrected_state, show_changes],
outputs=[diff_out],
api_visibility="private",
show_progress="hidden",
)
copy_btn.click(
fn=None,
inputs=[corrected_state],
outputs=[],
js=(
"async (text) => {"
" if (text && text.trim()) {"
" try { await navigator.clipboard.writeText(text); }"
" catch (e) { console.error('Copy failed:', e); }"
" }"
" return [];"
"}"
),
)
if __name__ == "__main__":
demo.queue(max_size=10).launch(
theme=THEME,
css=CSS,
show_error=True,
)