Spaces:
Sleeping
Sleeping
Download app.py from hrnrxb/Text_Style_Transfer: direct link, hf CLI and curl.
- Browser
- Download file 21.2 kB
-
https://huggingface.co/spaces/hrnrxb/Text_Style_Transfer/resolve/main/app.py
- Command line
-
hf download hf://spaces/hrnrxb/Text_Style_Transfer/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hrnrxb/Text_Style_Transfer/resolve/main/app.py
21.2 kB
| # 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()) | |
| 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( | |
| """ | |
| <div class="title-header"> | |
| <h1>📝 Text Style Transfer Playground</h1> | |
| <p class="subtitle">Bidirectional Formal <--> Casual Transformation · Research Sandbox</p> | |
| <p style="text-align: center; margin-top: 12px;"> | |
| <a href="https://github.com/hrnrxb/TST__Bridging-Casual-and-Formal-Domains" target="_blank" style="color: #2563eb; text-decoration: none; font-weight: 600; border: 1px solid #bfdbfe; padding: 6px 16px; border-radius: 20px; background-color: #eff6ff; display: inline-block; transition: all 0.2s;"> | |
| ⭐ View Source on GitHub | |
| </a> | |
| </p> | |
| </div> | |
| """ | |
| ) | |
| 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( | |
| """ | |
| <div class="guide-box"> | |
| <strong>🎯 Quality Filter Threshold</strong><br> | |
| Controls how strict the model is when evaluating its own candidates. | |
| <ul> | |
| <li><strong>Lower (0.1–0.5)</strong> → More attempts, risk of gibberish/hallucinations.</li> | |
| <li><strong>Mid (0.6–0.75)</strong> → Sweet spot: good balance of quality and transfer rate.</li> | |
| <li><strong>Higher (0.8–1.0)</strong> → Only near‑perfect outputs, but many no‑ops.</li> | |
| </ul> | |
| <strong>🎲 Max Candidates</strong><br> | |
| Number of paraphrases generated internally before picking the best. | |
| <ul> | |
| <li><strong>1–3</strong> → Fast, but less chance of a great paraphrase.</li> | |
| <li><strong>5–8</strong> → Good balance of speed and quality.</li> | |
| <li><strong>10+</strong> → Slowest, but best chance to find a semantically strong output.</li> | |
| </ul> | |
| <strong>⚠️ Model Limitation</strong><br> | |
| 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 <code>quality_filter=0.75</code> and <code>max_candidates=12</code> – but your mileage may vary! | |
| </div> | |
| """ | |
| ) | |
| 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( | |
| """ | |
| --- | |
| <div style="text-align: center; color: #94a3b8; font-size: 0.8rem;"> | |
| Research Artifact · Interactive Evaluation Sandbox · TST for Formal/Casual Registers | |
| </div> | |
| """ | |
| ) | |
| # ---- 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) |