Team Ai
Apppublic

nxphi47/test-zero-gpu

sourceHugging Faceapache-2.0updated 3y agoView on Hugging Face
0likes
app.py146 linesDownload Raw Back to root
1import os2from threading import Thread3from typing import Iterator4 5import gradio as gr6import spaces7import torch8from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer9 10MAX_MAX_NEW_TOKENS = 204811DEFAULT_MAX_NEW_TOKENS = 102412MAX_INPUT_TOKEN_LENGTH = int(os.getenv("MAX_INPUT_TOKEN_LENGTH", "4096"))13 14DESCRIPTION = """\15# Llama-2 13B Chat16This Space demonstrates model [Llama-2-13b-chat](https://huggingface.co/meta-llama/Llama-2-13b-chat) by Meta, a Llama 2 model with 13B parameters fine-tuned for chat instructions. Feel free to play with it, or duplicate to run generations without a queue! If you want to run your own service, you can also [deploy the model on Inference Endpoints](https://huggingface.co/inference-endpoints).17๐Ÿ”Ž For more details about the Llama 2 family of models and how to use them with `transformers`, take a look [at our blog post](https://huggingface.co/blog/llama2).18๐Ÿ”จ Looking for an even more powerful model? Check out the large [**70B** model demo](https://huggingface.co/spaces/ysharma/Explore_llamav2_with_TGI).19๐Ÿ‡ For a smaller model that you can run on many GPUs, check our [7B model demo](https://huggingface.co/spaces/huggingface-projects/llama-2-7b-chat).20"""21 22LICENSE = """23<p/>24---25As a derivate work of [Llama-2-13b-chat](https://huggingface.co/meta-llama/Llama-2-7b-chat) by Meta,26this demo is governed by the original [license](https://huggingface.co/spaces/huggingface-projects/llama-2-13b-chat/blob/main/LICENSE.txt) and [acceptable use policy](https://huggingface.co/spaces/huggingface-projects/llama-2-13b-chat/blob/main/USE_POLICY.md).27"""28 29if not torch.cuda.is_available():30    DESCRIPTION += "\n<p>Running on CPU ๐Ÿฅถ This demo does not work on CPU.</p>"31 32 33if torch.cuda.is_available():34    # model_id = "meta-llama/Llama-2-7b-chat-hf"35    model_id = "mistralai/Mistral-7B-Instruct-v0.2"36    model_id = "SeaLLMs/SeaLLM-7B-v2"37    model = AutoModelForCausalLM.from_pretrained(model_id, device_map="cuda", torch_dtype=torch.bfloat16)38    tokenizer = AutoTokenizer.from_pretrained(model_id)39    tokenizer.use_default_system_prompt = False40    print(f"{model}")41    print(f"{model.device}")42 43 44@spaces.GPU45def generate(46    message: str,47    chat_history: list[tuple[str, str]],48    system_prompt: str,49    max_new_tokens: int = 1024,50    temperature: float = 0.6,51    top_p: float = 0.9,52    top_k: int = 50,53    repetition_penalty: float = 1.2,54) -> Iterator[str]:55    conversation = []56    # if system_prompt:57    #     conversation.append({"role": "system", "content": system_prompt})58    for user, assistant in chat_history:59        conversation.extend([{"role": "user", "content": user}, {"role": "assistant", "content": assistant}])60    conversation.append({"role": "user", "content": message})61 62    input_ids = tokenizer.apply_chat_template(conversation, return_tensors="pt")63    if input_ids.shape[1] > MAX_INPUT_TOKEN_LENGTH:64        input_ids = input_ids[:, -MAX_INPUT_TOKEN_LENGTH:]65        gr.Warning(f"Trimmed input from conversation as it was longer than {MAX_INPUT_TOKEN_LENGTH} tokens.")66    input_ids = input_ids.to(model.device)67 68    streamer = TextIteratorStreamer(tokenizer, timeout=10.0, skip_prompt=True, skip_special_tokens=True)69    generate_kwargs = dict(70        {"input_ids": input_ids},71        streamer=streamer,72        max_new_tokens=max_new_tokens,73        do_sample=True,74        top_p=top_p,75        top_k=top_k,76        temperature=temperature,77        num_beams=1,78        repetition_penalty=repetition_penalty,79    )80    t = Thread(target=model.generate, kwargs=generate_kwargs)81    t.start()82 83    outputs = []84    for text in streamer:85        outputs.append(text)86        yield "".join(outputs)87 88 89chat_interface = gr.ChatInterface(90    fn=generate,91    additional_inputs=[92        gr.Textbox(label="System prompt", lines=6),93        gr.Slider(94            label="Max new tokens",95            minimum=1,96            maximum=MAX_MAX_NEW_TOKENS,97            step=1,98            value=DEFAULT_MAX_NEW_TOKENS,99        ),100        gr.Slider(101            label="Temperature",102            minimum=0.1,103            maximum=4.0,104            step=0.1,105            value=0.6,106        ),107        gr.Slider(108            label="Top-p (nucleus sampling)",109            minimum=0.05,110            maximum=1.0,111            step=0.05,112            value=0.9,113        ),114        gr.Slider(115            label="Top-k",116            minimum=1,117            maximum=1000,118            step=1,119            value=50,120        ),121        gr.Slider(122            label="Repetition penalty",123            minimum=1.0,124            maximum=2.0,125            step=0.05,126            value=1.2,127        ),128    ],129    stop_btn=None,130    examples=[131        ["Hello there! How are you doing?"],132        ["Can you explain briefly to me what is the Python programming language?"],133        ["Explain the plot of Cinderella in a sentence."],134        ["How many hours does it take a man to eat a Helicopter?"],135        ["Write a 100-word article on 'Benefits of Open-Source in AI research'"],136    ],137)138 139with gr.Blocks(css="style.css") as demo:140    gr.Markdown(DESCRIPTION)141    gr.DuplicateButton(value="Duplicate Space for private use", elem_id="duplicate-button")142    chat_interface.render()143    gr.Markdown(LICENSE)144 145if __name__ == "__main__":146    demo.queue(max_size=20).launch()