import spaces # must come before torch / transformers import threading import gradio as gr import torch from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer MODEL_ID = "JetBrains/Mellum2.1-12B-A2.5B-Thinking" tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) model = AutoModelForCausalLM.from_pretrained( MODEL_ID, dtype=torch.bfloat16, attn_implementation="sdpa" ).to("cuda").eval() def _text(content) -> str: """Flatten Gradio message content (str or list of parts) to plain text.""" if isinstance(content, str): return content if isinstance(content, dict): return content.get("text") or "" if isinstance(content, (list, tuple)): return "".join(_text(c) for c in content) return "" def _build_messages(message, history, system_prompt): msgs = [] if system_prompt and system_prompt.strip(): msgs.append({"role": "system", "content": system_prompt.strip()}) for m in history or []: # skip the collapsible "Thinking" bubbles we emit; only keep real turns if (m.get("metadata") or {}).get("title"): continue text = _text(m.get("content")) if text: msgs.append({"role": m["role"], "content": text}) msgs.append({"role": "user", "content": _text(message)}) return msgs def _duration(message, history, system_prompt="", enable_thinking=True, max_new_tokens=4096, *args, **kwargs): return min(240, 30 + int(max_new_tokens) // 25) @spaces.GPU(duration=_duration) def chat( message: str, history: list, system_prompt: str = "", enable_thinking: bool = True, max_new_tokens: int = 4096, temperature: float = 0.6, top_p: float = 0.95, top_k: int = 20, ): """Chat with Mellum2.1 Thinking, streaming its reasoning and final answer. Args: message: The user's message. history: Prior conversation turns in OpenAI-style message format. system_prompt: Optional system prompt. enable_thinking: Whether the model should reason before answering. max_new_tokens: Maximum number of tokens to generate (thinking + answer). temperature: Sampling temperature. top_p: Nucleus sampling probability. top_k: Top-k sampling cutoff. """ msgs = _build_messages(message, history, system_prompt) input_ids = tokenizer.apply_chat_template( msgs, add_generation_prompt=True, return_tensors="pt", return_dict=True, enable_thinking=enable_thinking, )["input_ids"].to("cuda") streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) gen_kwargs = dict( input_ids=input_ids, attention_mask=torch.ones_like(input_ids), streamer=streamer, max_new_tokens=int(max_new_tokens), do_sample=temperature > 0, temperature=float(temperature) if temperature > 0 else None, top_p=float(top_p), top_k=int(top_k), pad_token_id=tokenizer.pad_token_id, ) threading.Thread(target=model.generate, kwargs=gen_kwargs, daemon=True).start() raw = "" for chunk in streamer: raw += chunk yield _render(raw, enable_thinking) def _render(raw: str, enable_thinking: bool): body = raw.lstrip() if body.startswith(""): body = body[len(""):] if "" in body: thought, answer = body.split("", 1) done = True elif enable_thinking: thought, answer, done = body, "", False else: thought, answer, done = "", body, True out = [] if thought.strip(): out.append(gr.ChatMessage( role="assistant", content=thought.strip(), metadata={"title": "🧠 Thinking" if done else "🧠 Thinking…", "status": "done" if done else "pending"}, )) if answer.strip() or not out: out.append(gr.ChatMessage(role="assistant", content=answer.strip())) return out CSS = """ #col-container { max-width: 1000px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); } """ with gr.Blocks(title="Mellum2.1 Thinking") as demo: with gr.Column(elem_id="col-container"): gr.Markdown( "# 🧠 Mellum2.1 Thinking\n" "Chat with [JetBrains/Mellum2.1-12B-A2.5B-Thinking](https://huggingface.co/JetBrains/Mellum2.1-12B-A2.5B-Thinking), " "a 12B MoE reasoning model (2.5B active) for coding, math and agentic work. " "Expand **Thinking** to see its reasoning." ) gr.ChatInterface( fn=chat, chatbot=gr.Chatbot(height=600), additional_inputs=[ gr.Textbox(label="System prompt", lines=2), gr.Checkbox(label="Enable thinking", value=True), gr.Slider(256, 8192, value=4096, step=256, label="Max new tokens"), gr.Slider(0.0, 1.5, value=0.6, step=0.05, label="Temperature"), gr.Slider(0.05, 1.0, value=0.95, step=0.05, label="Top-p"), gr.Slider(1, 100, value=20, step=1, label="Top-k"), ], additional_inputs_accordion=gr.Accordion("Settings", open=False), examples=[ ["Find the bug in this function and explain the fix: def mean(xs): return sum(xs) / len(xs) - 1"], ["Write a Python function that returns the longest palindromic substring of a string in O(n^2) time."], ["How many positive integers n ≤ 1000 make n² + 1 divisible by 5?"], ["Explain the difference between a mutex and a semaphore with a short Go example."], ], cache_examples=False, ) demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)