# app.py import spaces import os import gradio as gr import time import random import numpy as np import torch import logging from typing import Optional, Literal, Dict, Any from styleformer import Styleformer from sentence_transformers import SentenceTransformer, util from transformers import T5ForConditionalGeneration # ============================================================================== # SECRETS MANAGEMENT (HUGGING FACE SPACES) # ============================================================================== # If sb wants to change models (not public ones) in future.. HF_TOKEN = os.environ.get("HF_TOKEN") # ============================================================================== # 1. LOGGING & T5 PATCH # ============================================================================== # Configure structured logging for production observability logging.basicConfig( level=logging.INFO, format="%(asctime)s | %(levelname)s | %(name)s | %(message)s" ) logger = logging.getLogger("StyleTransferPipeline") # Patch for incompatible transformer lib (ignoring auth kwarg!) original_from_pretrained = T5ForConditionalGeneration.from_pretrained def patched_from_pretrained(*args, **kwargs): # Remove the argument that causes the error kwargs.pop('use_auth_token', None) return original_from_pretrained(*args, **kwargs) # Apply the patch globally T5ForConditionalGeneration.from_pretrained = patched_from_pretrained # ============================================================================== # 2. SYSTEM ARCHITECTURE & ENGINEERING DESIGN (PIPELINE) # ============================================================================== class StyleTransferPipeline: """ A robust, object-oriented inference pipeline for bidirectional text style transfer between formal and casual linguistic registers. Encapsulates dedicated sequence-to-sequence transformer checkpoints for: - Casual-to-Formal (Style ID: 0) - Formal-to-Casual (Style ID: 1) Implements in-memory weight residency, device auto-resolution, and defensive input sanitization. """ SUPPORTED_STYLES = ("casual_to_formal", "formal_to_casual") def __init__(self, quality_filter: float = 0.95, max_candidates: int = 1) -> None: self.quality_filter: float = quality_filter self.max_candidates: int = max_candidates self.device: str = "cuda" if torch.cuda.is_available() else "cpu" self._inference_device_flag: int = torch.cuda.current_device() if self.device == "cuda" else -1 logger.info(f"Initializing StyleTransferPipeline on hardware target: [{self.device.upper()}]") # Model containers self._casual_to_formal_model: Optional[Styleformer] = None self._formal_to_casual_model: Optional[Styleformer] = None self._load_models() def _load_models(self) -> None: try: logger.info("Loading Casual-to-Formal model checkpoint (Style ID: 0)...") self._casual_to_formal_model = Styleformer(style=0) logger.info("Casual-to-Formal model loaded successfully.") except Exception as e: logger.error(f"Failed to load Casual-to-Formal model: {str(e)}") raise RuntimeError(f"Casual-to-Formal initialization error: {e}") try: logger.info("Loading Formal-to-Casual model checkpoint (Style ID: 1)...") self._formal_to_casual_model = Styleformer(style=1) logger.info("Formal-to-Casual model loaded successfully.") except Exception as e: logger.error(f"Failed to load Formal-to-Casual model: {str(e)}") raise RuntimeError(f"Formal-to-Casual initialization error: {e}") def _sanitize_input(self, text: str) -> str: if not isinstance(text, str): raise ValueError(f"Input must be of type 'str', received: {type(text).__name__}") cleaned_text = text.strip() if not cleaned_text: raise ValueError("Input text cannot be empty or purely whitespace.") return cleaned_text def transfer_style( self, text: str, target_style: Literal["casual_to_formal", "formal_to_casual"] = "casual_to_formal" ) -> str: cleaned_text = self._sanitize_input(text) if target_style not in self.SUPPORTED_STYLES: raise ValueError( f"Unsupported target_style '{target_style}'. Must be one of {self.SUPPORTED_STYLES}" ) try: if target_style == "casual_to_formal": if self._casual_to_formal_model is None: raise RuntimeError("Casual-to-Formal model is not initialized.") output = self._casual_to_formal_model.transfer( cleaned_text, inference_on=self._inference_device_flag, quality_filter=self.quality_filter, max_candidates=self.max_candidates ) else: # formal_to_casual if self._formal_to_casual_model is None: raise RuntimeError("Formal-to-Casual model is not initialized.") output = self._formal_to_casual_model.transfer( cleaned_text, inference_on=self._inference_device_flag, quality_filter=self.quality_filter, max_candidates=self.max_candidates ) # Defensive fallback: If model produces None or empty string, retain original text if not output or not output.strip(): logger.warning("Generation yielded empty output or failed quality filter. Preserving source text.") return cleaned_text return output.strip() except Exception as e: logger.error(f"Inference error during [{target_style}] execution on input '{cleaned_text}': {str(e)}") return cleaned_text def get_pipeline_metadata(self) -> Dict[str, Any]: return { "device": self.device, "quality_filter": self.quality_filter, "max_candidates": self.max_candidates, "casual_to_formal_ready": self._casual_to_formal_model is not None, "formal_to_casual_ready": self._formal_to_casual_model is not None, } # ============================================================================== # 3. GLOBAL MODEL INITIALIZATION (Loads into RAM on Space Boot) # ============================================================================== logger.info("Instantiating pipeline (weights will be fetched and loaded into memory once)...") pipeline = StyleTransferPipeline(quality_filter=0.75, max_candidates=12) logger.info("Loading 'all-MiniLM-L6-v2' semantic bi-encoder...") eval_encoder = SentenceTransformer("all-MiniLM-L6-v2") # ============================================================================== # 4. EVALUATION & UI HELPER FUNCTIONS # ============================================================================== def compute_similarity(text1, text2): """Compute cosine similarity between two texts using the loaded encoder.""" if not text1 or not text2: return 0.0 emb1 = eval_encoder.encode(text1, convert_to_tensor=True) emb2 = eval_encoder.encode(text2, convert_to_tensor=True) return float(util.cos_sim(emb1, emb2).item()) @spaces.GPU def process_style_transfer_ui(input_text, transfer_direction, quality_filter, max_candidates): """ Main UI function: applies style transfer, computes similarity, returns all metrics. Temporarily overrides pipeline params, then restores them. """ if not input_text or not input_text.strip(): return "", 0.0, "Error: Empty input", "0.00 ms", "0", "0" direction_map = { "Casual → Formal": "casual_to_formal", "Formal → Casual": "formal_to_casual" } target_style = direction_map.get(transfer_direction, "casual_to_formal") # Save original pipeline settings orig_qf = pipeline.quality_filter orig_mc = pipeline.max_candidates pipeline.quality_filter = quality_filter pipeline.max_candidates = max_candidates start_time = time.perf_counter() try: output_text = pipeline.transfer_style(text=input_text, target_style=target_style) end_time = time.perf_counter() latency_ms = (end_time - start_time) * 1000 # Restore original settings pipeline.quality_filter = orig_qf pipeline.max_candidates = orig_mc sim_score = compute_similarity(input_text, output_text) if output_text.strip() == input_text.strip(): status_msg = "⚠️ No‑op fallback (quality filter rejected all candidates)" else: status_msg = f"✅ Transformed [{transfer_direction}] | Similarity: {sim_score:.3f}" in_chars = str(len(input_text)) out_chars = str(len(output_text)) return output_text, sim_score, status_msg, f"{latency_ms:.2f} ms", in_chars, out_chars except Exception as e: end_time = time.perf_counter() latency_ms = (end_time - start_time) * 1000 pipeline.quality_filter = orig_qf pipeline.max_candidates = orig_mc return input_text, 0.0, f"❌ Error: {str(e)}", f"{latency_ms:.2f} ms", str(len(input_text)), str(len(input_text)) # --- Expanded Example Pool (16 diverse samples) --- EXAMPLE_POOL = [ # Casual → Formal (8 samples) ["Hey, can you send me the docs asap?", "Casual → Formal"], ["Wanna grab a coffee and chat about the project?", "Casual → Formal"], ["The app totally crashed when I clicked the submit button.", "Casual → Formal"], ["Don't spill the beans about the surprise party.", "Casual → Formal"], ["Thanks a ton for having my back during the presentation.", "Casual → Formal"], ["I wrapped up the script, go ahead and check it out.", "Casual → Formal"], ["Sorry for ghosting you, got caught up with random stuff.", "Casual → Formal"], ["Nah, I don't think that idea's gonna fly with the team.", "Casual → Formal"], # Formal → Casual (8 samples) ["I would appreciate it if you could assist me with this task.", "Formal → Casual"], ["Pursuant to our previous correspondence, please find the document attached.", "Formal → Casual"], ["Due to unforeseen circumstances, I must request a rescheduling of our appointment.", "Formal → Casual"], ["I regret to inform you that we are incapable of fulfilling this request at this juncture.", "Formal → Casual"], ["The empirical findings demonstrate substantial alignment with our initial hypothesis.", "Formal → Casual"], ["We extend our deepest gratitude for your invaluable contributions to this endeavor.", "Formal → Casual"], ["It is imperative that all personnel adhere strictly to the established safety protocols.", "Formal → Casual"], ["Kindly notify me of your availability to convene at your earliest convenience.", "Formal → Casual"], ] def load_random_example(): """Pick a random example and return as list.""" return random.choice(EXAMPLE_POOL) # ============================================================================== # 5. GRADIO INTERFACE CONSTRUCTION # ============================================================================== # --- Custom CSS for a polished, modern look --- custom_css = """ .gradio-container { font-family: 'Segoe UI', Roboto, sans-serif; } .title-header { text-align: center; color: #1E293B; margin-bottom: 0.2rem; } .subtitle { text-align: center; color: #475569; font-weight: 300; margin-top: 0; } .main-card { background: #f8fafc; border-radius: 16px; padding: 24px; box-shadow: 0 4px 12px rgba(0,0,0,0.05); } .output-box textarea { background: #ffffff !important; border: 1px solid #e2e8f0 !important; font-weight: 500; } .sim-score { font-size: 1.2rem; font-weight: 600; } .sim-high { color: #22c55e; } .sim-mid { color: #eab308; } .sim-low { color: #ef4444; } .status-text { font-size: 0.9rem; color: #334155; } .copy-btn { margin-left: 8px; } .slider-label { font-weight: 500; color: #0f172a; } .guide-box { background: #f1f5f9; border-radius: 8px; padding: 12px 16px; margin: 8px 0; border-left: 4px solid #3b82f6; } .guide-box strong { color: #0f172a; } """ # --- Build the Interface --- with gr.Blocks(title="Text Style Transfer Playground", css=custom_css, theme=gr.themes.Soft()) as demo: gr.Markdown( """

📝 Text Style Transfer Playground

Bidirectional Formal <--> Casual Transformation · Research Sandbox

⭐ View Source on GitHub

""" ) with gr.Row(): # Left column: Input & Controls with gr.Column(scale=1, elem_classes=["main-card"]): gr.Markdown("### ✏️ Source Text") user_input = gr.Textbox( lines=5, placeholder="Type or paste your text here…", label=None, show_label=False ) with gr.Row(): style_selector = gr.Radio( choices=["Casual → Formal", "Formal → Casual"], value="Casual → Formal", label="Transfer Direction", interactive=True ) # Advanced Settings with detailed guide with gr.Accordion("⚙️ Advanced Settings & How to Tune", open=False): gr.Markdown( """
🎯 Quality Filter Threshold
Controls how strict the model is when evaluating its own candidates. 🎲 Max Candidates
Number of paraphrases generated internally before picking the best. ⚠️ Model Limitation
This is a lightweight T5‑base model (220M parameters). It struggles with idioms, slang, and dense academic jargon. Expect occasional hallucinations, typos, or no‑ops. The best we've found empirically is quality_filter=0.75 and max_candidates=12 – but your mileage may vary!
""" ) quality_slider = gr.Slider( minimum=0.1, maximum=1.0, value=0.75, step=0.05, label="Quality Filter Threshold", info="Higher = stricter, lower = more attempts but risk of hallucinations" ) candidates_slider = gr.Slider( minimum=1, maximum=20, value=12, step=1, label="Max Candidates", info="More = slower but potentially better; we recommend 8–12 for quality" ) reset_settings_btn = gr.Button("↺ Reset to Recommended (0.75 / 12)", variant="secondary", size="sm") with gr.Row(): submit_btn = gr.Button("🚀 Transform", variant="primary", scale=2) clear_btn = gr.Button("🗑️ Clear", variant="secondary", scale=1) random_btn = gr.Button("🎲 Random Example", variant="secondary", scale=1) # Right column: Output & Metrics with gr.Column(scale=1, elem_classes=["main-card"]): gr.Markdown("### ✨ Transformed Output") output_display = gr.Textbox( lines=5, label=None, show_label=False, interactive=False, elem_classes=["output-box"] ) with gr.Row(): sim_score_display = gr.Number( value=0.0, label="Semantic Similarity (cosine)", interactive=False, elem_classes=["sim-score"] ) with gr.Row(): copy_btn = gr.Button("📋 Copy Output", variant="secondary", size="sm", elem_classes=["copy-btn"]) # Diagnostics accordion with more details with gr.Accordion("📊 System Diagnostics", open=False): status_display = gr.Textbox( label="Status", interactive=False, lines=1, elem_classes=["status-text"] ) with gr.Row(): latency_display = gr.Label(label="Inference Latency") in_count = gr.Label(label="Input Characters") out_count = gr.Label(label="Output Characters") with gr.Row(): used_qf = gr.Label(label="Quality Filter Used") used_mc = gr.Label(label="Candidates Used") # Quick Examples section (two rows of 8) gr.Markdown("### 📚 Quick Examples (click to auto‑load and transform)") with gr.Row(): gr.Examples( examples=EXAMPLE_POOL[:8], inputs=[user_input, style_selector], label=None, cache_examples=False ) with gr.Row(): gr.Examples( examples=EXAMPLE_POOL[8:], inputs=[user_input, style_selector], label=None, cache_examples=False ) # Footer gr.Markdown( """ ---
Research Artifact · Interactive Evaluation Sandbox · TST for Formal/Casual Registers
""" ) # ---- Event Handlers ---- def transform_and_log(input_text, direction, qf, mc): """Wrapper to call the main function and return all outputs including used settings.""" output, sim, status, latency, in_chars, out_chars = process_style_transfer_ui( input_text, direction, qf, mc ) return output, sim, status, latency, in_chars, out_chars, f"{qf:.2f}", f"{mc}" # Main submit submit_btn.click( fn=transform_and_log, inputs=[user_input, style_selector, quality_slider, candidates_slider], outputs=[output_display, sim_score_display, status_display, latency_display, in_count, out_count, used_qf, used_mc] ) # Clear button: resets all fields and diagnostics clear_btn.click( fn=lambda: ("", "Casual → Formal", "", 0.0, "Status: Cleared", "0.00 ms", "0", "0", "0.00", "0"), inputs=[], outputs=[user_input, style_selector, output_display, sim_score_display, status_display, latency_display, in_count, out_count, used_qf, used_mc] ) # Random example: load and auto-transform random_btn.click( fn=load_random_example, inputs=[], outputs=[user_input, style_selector] ).then( fn=transform_and_log, inputs=[user_input, style_selector, quality_slider, candidates_slider], outputs=[output_display, sim_score_display, status_display, latency_display, in_count, out_count, used_qf, used_mc] ) # Reset settings to recommended defaults (0.75, 12) reset_settings_btn.click( fn=lambda: (0.75, 12), inputs=[], outputs=[quality_slider, candidates_slider] ) # Copy to clipboard with JS (native clipboard API) copy_btn.click( fn=None, inputs=[], outputs=[], js=""" function() { let outputText = document.querySelector('.output-box textarea'); if (outputText) { navigator.clipboard.writeText(outputText.value); alert('Copied to clipboard!'); } return []; } """ ) # Launch the app on Hugging Face Spaces demo.launch(show_error=True)