import re import nltk import torch import spaces import gradio as gr from threading import Thread from nltk.tag import pos_tag from nltk.chunk import ne_chunk from nltk.tokenize import word_tokenize from peft import PeftModel from transformers import CsmForConditionalGeneration from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer, AutoProcessor try: nltk.data.find('tokenizers/punkt') nltk.data.find('taggers/averaged_perceptron_tagger') nltk.data.find('chunkers/maxent_ne_chunker') nltk.data.find('corpora/words') except LookupError: nltk.download('punkt', quiet=True) nltk.download('averaged_perceptron_tagger', quiet=True) nltk.download('maxent_ne_chunker', quiet=True) nltk.download('words', quiet=True) SYSTEM_PROMPT = """You are an expert creative director specializing in visual descriptions for image generation. Your task: Transform the user's concept into a rich, detailed image description while PRESERVING their core idea. IMPORTANT RULES: 1. Keep ALL key elements (intents, entities) from the original concept 2. Enhance with artistic details, NOT change the fundamental idea 3. Maintain the user's intended subject, action, and setting You should elaborate on: • Visual composition and perspective (bird's eye, close-up, wide angle, etc.) • Artistic style (photorealistic, impressionist, specific artist like Van Gogh, etc.) • Color palette and color temperature • Lighting (golden hour, dramatic shadows, soft diffused, etc.) • Atmosphere and mood • Textures and materials (rough, smooth, metallic, organic, etc.) • Technical details (medium, brushwork, rendering style) • Environmental context (time of day, weather, season, era) • Level of detail and focus points Output format: A single, flowing paragraph that reads naturally as an image prompt.""" prompt_model = AutoModelForCausalLM.from_pretrained("shb777/PromptTuner-v0.1") tokenizer = AutoTokenizer.from_pretrained("shb777/PromptTuner-v0.1") prompt_model.eval() CSM_BASE_MODEL_ID = "sesame/csm-1b" CSM_ADAPTER_ID = "shb777/csm-maya-exp2" SPEAKER_ID = 4 # Was trained on this speaker ID device = "cuda" if torch.cuda.is_available() else "cpu" csm_processor = AutoProcessor.from_pretrained(CSM_BASE_MODEL_ID) csm_model = CsmForConditionalGeneration.from_pretrained(CSM_BASE_MODEL_ID, device_map=device) csm_model = PeftModel.from_pretrained(csm_model, CSM_ADAPTER_ID) csm_model.eval() def extract_key_phrases(text: str) -> list: # We will highlight key phrases in the enhanced prompt phrases = [] try: tokens = word_tokenize(text) tagged = pos_tag(tokens) chunks = ne_chunk(tagged) current_phrase = [] for chunk in chunks: if hasattr(chunk, 'label'): phrase = ' '.join([token for token, _ in chunk.leaves()]) phrases.append(phrase.lower()) elif chunk[1].startswith('NN'): current_phrase.append(chunk[0]) elif chunk[1].startswith('JJ') and current_phrase: current_phrase.append(chunk[0]) else: if current_phrase: phrases.append(' '.join(current_phrase).lower()) current_phrase = [] if current_phrase: phrases.append(' '.join(current_phrase).lower()) for word, tag in tagged: if tag.startswith('JJ') or tag in ('RB', 'RBR', 'RBS'): phrases.append(word.lower()) except Exception: words = re.findall(r'\b[a-zA-Z]{3,}\b', text.lower()) phrases = list(set(words)) multi_word = re.findall(r'\b[a-zA-Z]{3,}(?:\s+[a-zA-Z]{3,}){1,3}\b', text) phrases.extend([mw.lower() for mw in multi_word]) phrases = list(set(phrases)) phrases.sort(key=len, reverse=True) return phrases[:20] def highlight_matches(original_input: str, enhanced_output: str) -> str: if not original_input.strip(): return f'

{enhanced_output}

' key_phrases = extract_key_phrases(original_input) if not key_phrases: return f'

{enhanced_output}

' key_phrases.sort(key=len, reverse=True) output = enhanced_output highlighted_spans = [] for phrase in key_phrases: pattern = re.compile(r'\b' + re.escape(phrase) + r'\b', re.IGNORECASE) def replace_with_highlight(match): matched_text = match.group(0) start = match.start() for h_start, h_end in highlighted_spans: if start >= h_start and start <= h_end: return matched_text highlighted_spans.append((start, match.end())) return f'{matched_text}' output = pattern.sub(replace_with_highlight, output) return f'

{output}

' def enhance_prompt(user_prompt: str): if not user_prompt or not user_prompt.strip(): yield ( 'Please enter a prompt to enhance.', "", gr.update(interactive=True), gr.update(interactive=True) ) return messages = [ {"role": "system", "content": SYSTEM_PROMPT}, {"role": "user", "content": user_prompt} ] prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) inputs = tokenizer(prompt, return_tensors="pt") streamer = TextIteratorStreamer(tokenizer, skip_special_tokens=True, skip_prompt=True) generation_kwargs = { 'max_new_tokens': 512, 'streamer': streamer, 'do_sample': True, 'temperature': 1, 'top_p': 0.95, 'top_k': 64 } placeholder = 'Your enhanced prompt will appear here' yield placeholder, "", gr.update(interactive=False), gr.update(interactive=False) thread = Thread(target=prompt_model.generate, kwargs={**inputs, **generation_kwargs}) thread.start() output = "" for text in streamer: output += text highlighted = highlight_matches(user_prompt, output) yield highlighted, output, gr.update(), gr.update() final_highlighted = highlight_matches(user_prompt, output) yield final_highlighted, output, gr.update(interactive=True), gr.update(interactive=True) @spaces.GPU def generate_tts_gpu(text): conversation = [ {"role": str(SPEAKER_ID), "content": [{"type": "text", "text": text}]}, ] enc = csm_processor.apply_chat_template( conversation, tokenize=True, return_dict=True, return_tensors="pt", ).to(device) gen_kwargs = { "max_new_tokens": 375, # "do_sample": True, # "temperature": 0.7, # "depth_decoder_do_sample": True, # "depth_decoder_temperature": 0.7, # "depth_decoder_top_k": 20, # "depth_decoder_top_p": 0.95, } audio = csm_model.generate( **enc, **gen_kwargs, output_audio=True ) audio_array = audio[0].to(torch.float32).cpu().numpy() return (24000, audio_array) def text_to_speech(text: str): if not text or not text.strip(): raise gr.Error("Please enter text to convert to speech.") if len(text) > 200: raise gr.Error("Text too long. Please limit to 200 characters or split into sentences.") # Preprocess text - remove characters that CSM struggles with text = text.replace('(', '').replace(')', '') text = text.replace('"', '').replace('"', '').replace('"', '') text = text.replace(';', ',') text = text.replace('!', ' ') text = text.replace('[', '').replace(']', '') text = text.replace('/', ' ') try: audio_result = generate_tts_gpu(text) return audio_result except Exception as e: raise gr.Error(f"Failed to generate speech: {str(e)}") with open("style.css", "r") as f: # Load CSS for styling custom_css = f.read() with gr.Blocks(css=custom_css, title="TinkerSpace") as demo: with gr.Tabs(): with gr.Tab("CSM Maya TTS"): with gr.Row(elem_classes=["main-grid"]): with gr.Column(elem_classes=["card"]): gr.HTML('') tts_input = gr.Textbox( placeholder="Enter text to convert to speech...", lines=5, show_label=False, container=False, elem_classes=["input-textarea"] ) with gr.Row(elem_classes=["flex gap-2 mt-6"]): generate_tts_btn = gr.Button( "Generate Speech", variant="primary", scale=2, elem_classes=["btn", "btn-primary"] ) clear_tts_btn = gr.Button( "Clear", scale=1, elem_classes=["btn", "btn-secondary"] ) with gr.Column(elem_classes=["card"]): gr.HTML('') tts_output = gr.Audio( label=None, show_label=False, autoplay=False, interactive=False, elem_classes=["output-container"] ) with gr.Column(elem_classes=["examples-section"]): gr.Examples( examples=[ ['You went to the party, even though I explicitly told you not to?'], ["With a gentle touch and a loving smile, she reassured, 'Dont worry, my love. We'll get through this together, just like we always have. I love you.'"], ["After years of work, Heisenberg finally published a ground-breaking cutting-edge research paper on quantum physics."], ['"Uh, are you sure about this?" Tim asked nervously, looking at the steep slope before them. "Whoa, it\'s higher than I thought," he continued, his voice filled with trepidation.'], ['"Shh, Lucy, shh, we mustn\'t wake your baby brother," Tom whispered, as they tiptoed past the nursery.'] ], inputs=tts_input, label="Examples" ) with gr.Row(): gr.Markdown( "Powered by [CSM Maya](https://huggingface.co/shb777/csm-maya-exp2)" ) with gr.Tab("Prompt Enhancer"): with gr.Row(elem_classes=["main-grid"]): with gr.Column(elem_classes=["card"]): gr.HTML('') input_text = gr.Textbox( placeholder="Describe your image concept... e.g., fox, red tail, blue moon, clouds", lines=5, show_label=False, autofocus=True, container=False, elem_classes=["input-textarea"] ) with gr.Row(elem_classes=["flex gap-2 mt-6"]): enhance_btn = gr.Button( "Enhance Prompt", variant="primary", scale=2, elem_classes=["btn", "btn-primary"] ) clear_btn = gr.Button( "Clear", scale=1, elem_classes=["btn", "btn-secondary"] ) with gr.Column(elem_classes=["card"]): gr.HTML('') output_html = gr.HTML( value='Your enhanced prompt will appear here', elem_classes=["output-container"] ) raw_output = gr.Textbox(visible=False) with gr.Column(elem_classes=["examples-section"]): gr.Examples( examples=[ ["fox, red tail, blue moon, clouds"], ["room with french window, cozy morning vibes, minimal"], ["anime style, sunset, japan"] ], inputs=input_text, label="Examples" ) with gr.Row(): gr.Markdown( "Powered by [PromptTuner](https://huggingface.co/shb777/PromptTuner-v0.1), " "a finetuned gemma3-270M model specifically designed to enhance text prompts " "for text-to-image generation." ) enhance_btn.click( fn=enhance_prompt, inputs=[input_text], outputs=[output_html, raw_output, enhance_btn, clear_btn] ) clear_btn.click( fn=lambda: ( "", 'Your enhanced prompt will appear here', "", gr.update(interactive=True), gr.update(interactive=True) ), inputs=None, outputs=[input_text, output_html, raw_output, enhance_btn, clear_btn] ) generate_tts_btn.click( fn=text_to_speech, inputs=[tts_input], outputs=[tts_output] ) clear_tts_btn.click( fn=lambda: ("", None), inputs=None, outputs=[tts_input, tts_output] ) if __name__ == "__main__": demo.queue(max_size=20).launch(mcp_server=True)