import uuid import spaces import torch import gradio as gr from transformers import AutoTokenizer, AutoModelForCausalLM from storage import save_record, maybe_periodic_push, now_iso MODEL_NAME = "Lyte/Nanochat-Moroccan-Instruct-0.7B" tokenizer = AutoTokenizer.from_pretrained( MODEL_NAME, trust_remote_code=True, ) dtype = torch.float16 if torch.cuda.is_available() else torch.float32 model = AutoModelForCausalLM.from_pretrained( MODEL_NAME, trust_remote_code=True, torch_dtype=dtype, ) device = "cuda" if torch.cuda.is_available() else "cpu" model = model.to(device) model.eval() def make_turn(user_text, assistant_text): return { "turn_id": str(uuid.uuid4()), "user_message_id": str(uuid.uuid4()), "assistant_message_id": str(uuid.uuid4()), "user": user_text, "assistant": assistant_text, "timestamp": now_iso(), } def build_prompt_from_turns(turns, user_message): parts = ["<|bos|>"] for turn in turns: user_text = turn.get("user", "") assistant_text = turn.get("assistant", "") if user_text: parts.append(f"<|user_start|>{user_text}<|user_end|>") if assistant_text: parts.append(f"<|assistant_start|>{assistant_text}<|assistant_end|>") parts.append(f"<|user_start|>{user_message}<|user_end|>") parts.append("<|assistant_start|>") return "".join(parts) def extract_new_assistant_text(full_text): text = full_text if "<|assistant_start|>" in text: text = text.split("<|assistant_start|>")[-1] for stop_token in ["<|assistant_end|>", "<|user_start|>", "<|user_end|>"]: if stop_token in text: text = text.split(stop_token, 1)[0] return text.strip() def turns_to_chatbot_pairs(turns): messages = [] for turn in turns: messages.append({"role": "user", "content": turn["user"]}) messages.append({"role": "assistant", "content": turn["assistant"]}) return messages def turns_to_feedback_choices(turns): choices = [] for idx, turn in enumerate(turns, start=1): preview = (turn.get("user", "") or "").replace("\n", " ").strip() if len(preview) > 60: preview = preview[:60] + "..." label = f"Turn {idx}: {preview}" if preview else f"Turn {idx}" choices.append((label, turn["turn_id"])) return choices def serialize_conversation(turns): conversation = [] for turn in turns: conversation.append({ "role": "user", "content": turn["user"], "turn_id": turn["turn_id"], "user_message_id": turn["user_message_id"], }) conversation.append({ "role": "assistant", "content": turn["assistant"], "turn_id": turn["turn_id"], "assistant_message_id": turn["assistant_message_id"], }) return conversation def selected_turn_from_id(turns, selected_turn_id): for turn in turns or []: if turn["turn_id"] == selected_turn_id: return turn return None def save_and_maybe_push(record): save_record(record) return maybe_periodic_push(MODEL_NAME) @spaces.GPU def chat_fn(message, turns, session_id, conversation_id, max_new_tokens, temperature, top_k, top_p, repetition_penalty): turns = turns or [] if not session_id: session_id = str(uuid.uuid4()) if not conversation_id: conversation_id = str(uuid.uuid4()) if not message.strip(): chatbot_view = turns_to_chatbot_pairs(turns) choices = turns_to_feedback_choices(turns) selected = choices[-1][1] if choices else None return ( chatbot_view, turns, session_id, conversation_id, "", gr.update(choices=choices, value=selected), "Please enter a message." ) prompt_text = build_prompt_from_turns(turns, message) inputs = tokenizer(prompt_text, return_tensors="pt") inputs = {k: v.to(model.device) for k, v in inputs.items()} input_len = inputs["input_ids"].shape[-1] with torch.no_grad(): output = model.generate( **inputs, max_new_tokens=int(max_new_tokens), do_sample=True, temperature=float(temperature), top_k=int(top_k), top_p=float(top_p), repetition_penalty=float(repetition_penalty), min_p=0.01, use_cache=False, ) generated_tokens = output[0][input_len:] assistant_text = tokenizer.decode(generated_tokens, skip_special_tokens=False) assistant_text = extract_new_assistant_text(assistant_text) new_turn = make_turn(message, assistant_text) updated_turns = turns + [new_turn] record = { "event": "generation", "timestamp": now_iso(), "session_id": session_id, "conversation_id": conversation_id, "turn_id": new_turn["turn_id"], "user_message_id": new_turn["user_message_id"], "assistant_message_id": new_turn["assistant_message_id"], "message": new_turn["user"], "response": new_turn["assistant"], "conversation": serialize_conversation(updated_turns), "generation_config": { "max_new_tokens": int(max_new_tokens), "temperature": float(temperature), "top_k": int(top_k), "top_p": float(top_p), "repetition_penalty": float(repetition_penalty), }, } push_result = save_and_maybe_push(record) status = "Response generated." if push_result is not None: if push_result["ok"]: status += f" Auto-pushed dataset at {push_result['record_count']} records." else: status += f" Auto-push failed: {push_result['message']}" chatbot_view = turns_to_chatbot_pairs(updated_turns) choices = turns_to_feedback_choices(updated_turns) selected = new_turn["turn_id"] return ( chatbot_view, updated_turns, session_id, conversation_id, "", gr.update(choices=choices, value=selected), status ) @spaces.GPU def regenerate_last(turns, session_id, conversation_id, max_new_tokens, temperature, top_k, top_p, repetition_penalty): turns = turns or [] if not turns: chatbot_view = turns_to_chatbot_pairs(turns) choices = turns_to_feedback_choices(turns) return ( chatbot_view, turns, session_id or str(uuid.uuid4()), conversation_id or str(uuid.uuid4()), gr.update(choices=choices, value=None), "No turn to regenerate." ) base_turns = turns[:-1] last_user_message = turns[-1]["user"] result = chat_fn( last_user_message, base_turns, session_id, conversation_id, max_new_tokens, temperature, top_k, top_p, repetition_penalty, ) chatbot_view, updated_turns, session_id, conversation_id, _, dropdown_update, _ = result return ( chatbot_view, updated_turns, session_id, conversation_id, dropdown_update, "Last response regenerated." ) def undo_last(turns, session_id, conversation_id): turns = turns or [] if not turns: return ( [], [], session_id or str(uuid.uuid4()), conversation_id or str(uuid.uuid4()), gr.update(choices=[], value=None), "Nothing to undo." ) updated_turns = turns[:-1] record = { "event": "undo_last_turn", "timestamp": now_iso(), "session_id": session_id, "conversation_id": conversation_id, "remaining_conversation": serialize_conversation(updated_turns), } save_and_maybe_push(record) chatbot_view = turns_to_chatbot_pairs(updated_turns) choices = turns_to_feedback_choices(updated_turns) selected = choices[-1][1] if choices else None return ( chatbot_view, updated_turns, session_id, conversation_id, gr.update(choices=choices, value=selected), "Last turn removed." ) def submit_feedback(turns, selected_turn_id, vote_type, session_id, conversation_id): if not turns: return "No turns available to rate." if not selected_turn_id: return "Please select a turn to rate." turn = selected_turn_from_id(turns, selected_turn_id) if turn is None: return "Selected turn not found." record = { "event": "feedback", "timestamp": now_iso(), "session_id": session_id, "conversation_id": conversation_id, "turn_id": turn["turn_id"], "user_message_id": turn["user_message_id"], "assistant_message_id": turn["assistant_message_id"], "vote": vote_type, "rated_user_message": turn["user"], "rated_assistant_message": turn["assistant"], "conversation": serialize_conversation(turns), } push_result = save_and_maybe_push(record) status = f"Saved: {vote_type}." if push_result is not None: if push_result["ok"]: status += f" Auto-pushed dataset at {push_result['record_count']} records." else: status += f" Auto-push failed: {push_result['message']}" return status def flag_response(turns, selected_turn_id, flag_type, session_id, conversation_id): if not turns: return "No turns available to flag." if not selected_turn_id: return "Please select a turn to flag." turn = selected_turn_from_id(turns, selected_turn_id) if turn is None: return "Selected turn not found." record = { "event": "flag", "timestamp": now_iso(), "session_id": session_id, "conversation_id": conversation_id, "turn_id": turn["turn_id"], "user_message_id": turn["user_message_id"], "assistant_message_id": turn["assistant_message_id"], "flag": flag_type, "flagged_user_message": turn["user"], "flagged_assistant_message": turn["assistant"], "conversation": serialize_conversation(turns), } push_result = save_and_maybe_push(record) status = f"Saved flag: {flag_type}." if push_result is not None: if push_result["ok"]: status += f" Auto-pushed dataset at {push_result['record_count']} records." else: status += f" Auto-push failed: {push_result['message']}" return status def clear_chat(): new_session_id = str(uuid.uuid4()) new_conversation_id = str(uuid.uuid4()) return ( [], [], new_session_id, new_conversation_id, gr.update(choices=[], value=None), "Chat cleared." ) examples = [ ["شنو أشهر زيت فالمغرب؟"], ["كيفاش نطيب الطاجين؟"], ["شحال كتساوي واحد زائد واحد؟"], ["شحال عدد سكان المغرب؟"], ["امتا خدا المغرب الاستقلال ديالو و شكون ملك لي كان هاداك الوقت؟"], ["بغيت نفهم كيفاش كتخدم السيارة الكهربائية و شنو الفرق بينها و بين السيارة العادية"], ["واش علميا ممكن نعاودو تجربة الصعود للقمر وشنو خاص يتحقق باش نعادو هاد التجربة"], ] custom_css = """ #chatbot { min-height: 620px; } .gradio-container { max-width: 1050px !important; margin: auto; } /* Keep interface English; only message text flows RTL */ #chatbot .message, #chatbot .message-wrap, #chatbot .bubble, #chatbot .prose, #chatbot [data-testid="user"], #chatbot [data-testid="bot"] { direction: rtl; text-align: right; } """ with gr.Blocks() as demo: gr.Markdown("# Nanochat Moroccan Instruct") gr.Markdown("Multi-turn chat with ratings, moderation flags, undo, regenerate, and automatic dataset sync.") gr.HTML(f"") turns_state = gr.State([]) session_state = gr.State(str(uuid.uuid4())) conversation_state = gr.State(str(uuid.uuid4())) chatbot = gr.Chatbot(label="Chat", height=620, elem_id="chatbot") status = gr.Markdown() with gr.Row(): msg = gr.Textbox( label="Message", placeholder="Type your message here...", lines=2, scale=8, ) send_btn = gr.Button("Send", variant="primary", scale=1) with gr.Accordion("Generation Settings", open=False): max_new_tokens = gr.Slider(32, 512, value=256, step=1, label="max_new_tokens") temperature = gr.Slider(0.1, 1.5, value=0.6, step=0.1, label="temperature") top_k = gr.Slider(1, 300, value=200, step=1, label="top_k") top_p = gr.Slider(0.1, 1.0, value=0.85, step=0.01, label="top_p") repetition_penalty = gr.Slider(1.0, 2.0, value=1.1, step=0.05, label="repetition_penalty") with gr.Row(): regenerate_btn = gr.Button("🔄 Regenerate last") undo_btn = gr.Button("↩️ Undo last") clear_btn = gr.Button("🗑️ Clear chat") with gr.Group(): gr.Markdown("### Feedback and moderation") feedback_turn = gr.Dropdown( label="Select a response", choices=[], value=None, interactive=True, ) with gr.Row(): like_btn = gr.Button("👍 Like") dislike_btn = gr.Button("👎 Dislike") unsafe_btn = gr.Button("🚩 Flag unsafe") bad_btn = gr.Button("⚠️ Flag bad output") gr.Examples(examples=examples, inputs=[msg]) send_btn.click( fn=chat_fn, inputs=[ msg, turns_state, session_state, conversation_state, max_new_tokens, temperature, top_k, top_p, repetition_penalty, ], outputs=[ chatbot, turns_state, session_state, conversation_state, msg, feedback_turn, status, ], ) msg.submit( fn=chat_fn, inputs=[ msg, turns_state, session_state, conversation_state, max_new_tokens, temperature, top_k, top_p, repetition_penalty, ], outputs=[ chatbot, turns_state, session_state, conversation_state, msg, feedback_turn, status, ], ) regenerate_btn.click( fn=regenerate_last, inputs=[ turns_state, session_state, conversation_state, max_new_tokens, temperature, top_k, top_p, repetition_penalty, ], outputs=[ chatbot, turns_state, session_state, conversation_state, feedback_turn, status, ], ) undo_btn.click( fn=undo_last, inputs=[turns_state, session_state, conversation_state], outputs=[ chatbot, turns_state, session_state, conversation_state, feedback_turn, status, ], ) clear_btn.click( fn=clear_chat, outputs=[ chatbot, turns_state, session_state, conversation_state, feedback_turn, status, ], ) like_btn.click( fn=lambda turns, selected_turn_id, session_id, conversation_id: submit_feedback( turns, selected_turn_id, "like", session_id, conversation_id ), inputs=[turns_state, feedback_turn, session_state, conversation_state], outputs=[status], ) dislike_btn.click( fn=lambda turns, selected_turn_id, session_id, conversation_id: submit_feedback( turns, selected_turn_id, "dislike", session_id, conversation_id ), inputs=[turns_state, feedback_turn, session_state, conversation_state], outputs=[status], ) unsafe_btn.click( fn=lambda turns, selected_turn_id, session_id, conversation_id: flag_response( turns, selected_turn_id, "unsafe", session_id, conversation_id ), inputs=[turns_state, feedback_turn, session_state, conversation_state], outputs=[status], ) bad_btn.click( fn=lambda turns, selected_turn_id, session_id, conversation_id: flag_response( turns, selected_turn_id, "bad_output", session_id, conversation_id ), inputs=[turns_state, feedback_turn, session_state, conversation_state], outputs=[status], ) if __name__ == "__main__": demo.launch(share=True)