import gradio as gr
import torch
import spaces
import sys
from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer
from threading import Thread
import transformers.modeling_rope_utils as rope_utils
from huggingface_hub import hf_hub_download
# 1. Monkey-patch dynamic_rope_update if missing
if not hasattr(rope_utils, 'dynamic_rope_update'):
def dynamic_rope_update(rope_forward):
return rope_forward
rope_utils.dynamic_rope_update = dynamic_rope_update
model_id = "deepgrove/maple-preview"
print("🔄 Loading tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True)
# Explicitly load chat template
try:
template_path = hf_hub_download(repo_id=model_id, filename="chat_template.jinja")
with open(template_path, "r", encoding="utf-8") as f:
tokenizer.chat_template = f.read()
print("✅ Chat template loaded successfully.")
except Exception as e:
print(f"⚠️ Warning: Could not load chat template: {e}")
print("🔄 Loading Maple-Preview...")
model = AutoModelForCausalLM.from_pretrained(
model_id,
device_map="auto",
torch_dtype=torch.bfloat16,
trust_remote_code=True,
)
# 2. Aggressively replace the buggy fa3 wrapper in EVERY loaded module namespace
def custom_flash_attention_forward(
query_states, key_states, value_states, attention_mask=None, query_length=None,
is_causal=True, dropout=0.0, position_ids=None, softmax_scale=None, **kwargs
):
from flash_attn import flash_attn_func, flash_attn_varlen_func
if query_states.dim() == 3:
cu_seq_lens_q = torch.tensor([0, query_states.shape[0]], dtype=torch.int32, device=query_states.device)
cu_seq_lens_k = torch.tensor([0, key_states.shape[0]], dtype=torch.int32, device=key_states.device)
return flash_attn_varlen_func(
query_states, key_states, value_states,
cu_seqlens_q=cu_seq_lens_q, cu_seqlens_k=cu_seq_lens_k,
max_seqlen_q=query_states.shape[0], max_seqlen_k=key_states.shape[0],
dropout_p=dropout, softmax_scale=softmax_scale, causal=is_causal
)
return flash_attn_func(
query_states, key_states, value_states,
dropout_p=dropout, softmax_scale=softmax_scale, causal=is_causal
)
patched_count = 0
for name, module in list(sys.modules.items()):
if 'maple' in name and hasattr(module, '_flash_attention_forward'):
setattr(module, '_flash_attention_forward', custom_flash_attention_forward)
patched_count += 1
print(f"✅ Replaced _flash_attention_forward in {patched_count} modules.")
@spaces.GPU
def predict(message, history):
messages = history + [{"role": "user", "content": message}]
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
generation_kwargs = dict(
**inputs,
streamer=streamer,
max_new_tokens=2048,
do_sample=True,
temperature=0.7,
top_p=0.9,
repetition_penalty=1.15,
no_repeat_ngram_size=3,
pad_token_id=tokenizer.pad_token_id,
eos_token_id=tokenizer.eos_token_id,
)
thread = Thread(target=model.generate, kwargs=generation_kwargs)
thread.start()
generated_text = ""
for new_text in streamer:
generated_text += new_text
# Format thinking tags into a collapsible HTML dropdown element
formatted_text = generated_text
if "" in formatted_text:
if "" in formatted_text:
formatted_text = formatted_text.replace("", "🧠 Thought Process
\n\n").replace(" ", "\n\n\n\n---")
else:
formatted_text = formatted_text.replace("", "🧠 Thought Process (Thinking...)
\n\n") + "\n\n "
yield formatted_text
# 3. Custom Chatbot component with Copy & Copy All features enabled
chatbot_ui = gr.Chatbot(
type="messages",
show_copy_button=True, # Adds a copy button to individual messages
show_copy_all_button=True, # Adds a button to copy the entire conversation history
allow_file_downloads=True,
)
demo = gr.ChatInterface(
fn=predict,
type="messages",
chatbot=chatbot_ui,
title="🍁 Maple-Preview",
description="A 20B-A1B ternary-weight reasoning LLM by DeepGrove. Powered by ZeroGPU.",
theme="soft",
)
if __name__ == "__main__":
demo.queue().launch()