hrnrxb's picture
Update app.py
81e5ce9 verified
Raw History Blame Contribute Delete
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())
@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(
"""
<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)