""" 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, )