Team Ai
Apppublic

LULDev/CodeGemma

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
app.py112 linesDownload Raw Back to root
1import gradio as gr2import os3import spaces4from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer5from threading import Thread6 7 8# Set an environment variable9HF_TOKEN = os.environ.get("HF_TOKEN", None)10 11DESCRIPTION = '''12<div>13<h1 style="text-align: center;">CodeGemma</h1>14 15<p>This Space demonstrates model <a href="https://huggingface.co/google/codegemma-7b-it">CodeGemma-7b-it</a> by Google. CodeGemma is a collection of lightweight open code models built on top of Gemma. Feel free to play with it, or duplicate to run privately!</p>16 17<p>🔎 For more details about the CodeGemma release and how to use the models with <code>transformers</code>, take a look <a href="https://huggingface.co/blog/codegemma">at our blog post</a>.</p>18</div>19'''20 21PLACEHOLDER = """22<div style="opacity: 0.65;">23    <img src="https://ysharma-dummy-chat-app.hf.space/file=/tmp/gradio/7dd7659cff2eab51f0f5336f378edfca01dd16fa/gemma_lockup_vertical_full-color_rgb.png" style="width:30%;">24    <br><b>CodeGemma-7B-IT Chatbot</b>25</div>26"""27 28    29# Load the tokenizer and model30tokenizer = AutoTokenizer.from_pretrained("google/codegemma-7b-it")31model = AutoModelForCausalLM.from_pretrained("google/codegemma-7b-it", device_map="auto")32 33 34@spaces.GPU(duration=120)35def codegemma(message: str, 36              history: list, 37              temperature: float, 38              max_new_tokens: int39             ) -> str:40    """41    Generate a streaming response using the CodeGemma model.42    Args:43        message (str): The input message.44        history (list): The conversation history used by ChatInterface.45        temperature (float): The temperature for generating the response.46        max_new_tokens (int): The maximum number of new tokens to generate.47    Returns:48        str: The generated response.49    """50    conversation = []51    for user, assistant in history:52        conversation.extend([{"role": "user", "content": user}, {"role": "assistant", "content": assistant}])53    conversation.append({"role": "user", "content": message})54 55    input_ids = tokenizer.apply_chat_template(conversation, return_tensors="pt").to(model.device)56    57    streamer = TextIteratorStreamer(tokenizer, timeout=10.0, skip_prompt=True, skip_special_tokens=True)58 59    generate_kwargs = dict(60        input_ids= input_ids,61        streamer=streamer,62        max_new_tokens=max_new_tokens,63        do_sample=True,64        temperature=temperature,65    )66    # This will enforce greedy generation (do_sample=False) when the temperature is passed 0, avoiding the crash.             67    if temperature == 0:68        generate_kwargs['do_sample'] = False69        70    t = Thread(target=model.generate, kwargs=generate_kwargs)71    t.start()72 73    outputs = []74    for text in streamer:75        outputs.append(text)76        yield "".join(outputs)77        78 79# Gradio block80chatbot=gr.Chatbot(placeholder=PLACEHOLDER,height=500)81 82with gr.Blocks(fill_height=True) as demo:83    84    gr.HTML(DESCRIPTION)85    86    gr.ChatInterface(87        fn=codegemma,88        chatbot=chatbot,89        fill_height=True,90        additional_inputs_accordion=gr.Accordion(label="⚙️ Parameters", open=False, render=False),91        additional_inputs=[92            gr.Slider(minimum=0,93                      maximum=1, 94                      step=0.1,95                      value=0.95, 96                      label="Temperature", 97                      render=False),98            gr.Slider(minimum=128, 99                      maximum=4096,100                      step=1,101                      value=512, 102                      label="Max new tokens", 103                      render=False ),104            ],105        examples=[106            ["Write a Python function to calculate the nth fibonacci number."]107            ],108        cache_examples=False,109                     )110    111if __name__ == "__main__":112    demo.launch()