#!pip install -U transformers gradio pillow import torch import gradio as gr from transformers import AutoProcessor, AutoModelForImageTextToText from PIL import Image # ── 載入模型 ────────────────────────────────────────────────────────────────── MODEL_ID = "google/gemma-4-12B" print(f"載入模型:{MODEL_ID} ...") processor = AutoProcessor.from_pretrained(MODEL_ID) model = AutoModelForImageTextToText.from_pretrained( MODEL_ID, torch_dtype=torch.float16, device_map="auto", ) model.eval() print("模型載入完成!") # ── 推論函式 ────────────────────────────────────────────────────────────────── def build_messages(history: list[dict], user_text: str, image=None) -> list[dict]: """將 Gradio 歷史記錄轉換成 HuggingFace messages 格式。""" messages = [] # 加入歷史對話 for turn in history: messages.append({"role": "user", "content": [{"type": "text", "text": turn["content"] if turn["role"] == "user" else ""}]}) messages.append({"role": "assistant", "content": [{"type": "text", "text": turn["content"] if turn["role"] == "assistant" else ""}]}) # 加入本輪使用者訊息 user_content = [] if image is not None: user_content.append({"type": "image", "image": image}) user_content.append({"type": "text", "text": user_text}) messages.append({"role": "user", "content": user_content}) return messages def chat( user_message: str, image, history: list[dict], max_new_tokens: int, temperature: float, top_p: float, ): """主對話函式,回傳更新後的歷史與清空的輸入。""" if not user_message.strip() and image is None: return history, history, gr.update(value=""), gr.update(value=None) # 建構 messages messages = build_messages(history, user_message, image) # Processor 前處理 pil_images = [image] if image is not None else None inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_tensors="pt", return_dict=True, images=pil_images, ).to(model.device, dtype=torch.float16) # 生成回覆 with torch.inference_mode(): output_ids = model.generate( **inputs, max_new_tokens=max_new_tokens, do_sample=temperature > 0, temperature=temperature if temperature > 0 else 1.0, top_p=top_p, ) # 只取新生成的部分 input_len = inputs["input_ids"].shape[-1] generated_ids = output_ids[:, input_len:] response = processor.decode(generated_ids[0], skip_special_tokens=True).strip() # 更新歷史(使用 Gradio messages 格式) history = history + [ {"role": "user", "content": user_message}, {"role": "assistant", "content": response}, ] return history, history, gr.update(value=""), gr.update(value=None) def clear_history(): return [], [], gr.update(value=""), gr.update(value=None) # ── Gradio UI ───────────────────────────────────────────────────────────────── with gr.Blocks( title="Gemma-4 Chat", theme=gr.themes.Soft(primary_hue="emerald"), css=""" #chatbot { height: 550px; } .send-btn { background: #10b981 !important; color: white !important; } footer { display: none !important; } """, ) as demo: # ── 標題 ── gr.Markdown( """ # 🤖 Gemma-4 多模態對話助理 支援純文字對話,亦可上傳圖片進行圖文問答。 """ ) # ── 對話狀態 ── state = gr.State([]) # 儲存 messages 歷史 with gr.Row(): # ── 左欄:聊天視窗 ── with gr.Column(scale=3): chatbot = gr.Chatbot( elem_id="chatbot", label="對話視窗", type="messages", # 使用 dict 格式 avatar_images=(None, "https://huggingface.co/datasets/huggingface/brand-assets/resolve/main/hf-logo.svg"), ) with gr.Row(): user_input = gr.Textbox( placeholder="輸入訊息,按 Enter 或點擊送出…", show_label=False, lines=2, scale=5, ) send_btn = gr.Button("送出 ▶", elem_classes="send-btn", scale=1) # ── 右欄:設定與圖片上傳 ── with gr.Column(scale=1): image_input = gr.Image( label="上傳圖片(選填)", type="pil", height=220, ) gr.Markdown("### ⚙️ 生成參數") max_new_tokens = gr.Slider(64, 2048, value=512, step=64, label="最大生成長度") temperature = gr.Slider(0.0, 2.0, value=0.7, step=0.05, label="Temperature(0 = 確定性)") top_p = gr.Slider(0.1, 1.0, value=0.9, step=0.05, label="Top-p") clear_btn = gr.Button("🗑️ 清除對話", variant="secondary") # ── 事件綁定 ── send_inputs = [user_input, image_input, state, max_new_tokens, temperature, top_p] send_outputs = [chatbot, state, user_input, image_input] send_btn.click(chat, inputs=send_inputs, outputs=send_outputs) user_input.submit(chat, inputs=send_inputs, outputs=send_outputs) clear_btn.click(clear_history, outputs=[chatbot, state, user_input, image_input]) # ── 啟動 ───────────────────────────────────────────────────────────────────── if __name__ == "__main__": demo.launch( server_name="0.0.0.0", # 允許外部連線(Colab / 遠端伺服器用) server_port=7860, share=False, # 設為 True 可產生公開連結(Colab 需要) inbrowser=True, )