# 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( """
Bidirectional Formal <--> Casual Transformation · Research Sandbox
quality_filter=0.75 and max_candidates=12 – but your mileage may vary!