th3james/assembly-decoder
0
1import sys2from threading import Thread3 4import gradio as gr5import spaces6from transformers import (7 AutoModelForCausalLM,8 AutoTokenizer,9 TextIteratorStreamer,10 BitsAndBytesConfig,11)12 13 14import torch15 16 17MODEL = "microsoft/Phi-3.5-mini-instruct"18 19if torch.cuda.is_available():20 device = "cuda"21elif sys.platform == "darwin" and torch.backends.mps.is_available():22 device = "mps"23else:24 device = "cpu"25 26 27# TODO understand this28quantization_config = BitsAndBytesConfig(29 load_in_4bit=True,30 bnb_4bit_compute_dtype=torch.bfloat16,31 bnb_4bit_use_double_quant=True,32 bnb_4bit_quant_type="nf4",33)34 35tokenizer = AutoTokenizer.from_pretrained(MODEL)36model = AutoModelForCausalLM.from_pretrained(37 MODEL,38 torch_dtype=torch.bfloat16,39 device_map="auto",40 quantization_config=quantization_config,41)42 43 44@spaces.GPU()45def stream_chat(46 message: str,47 history: list,48 system_prompt: str,49 temperature: float = 0.8,50 max_new_tokens: int = 1024,51 top_p: float = 1.0,52 top_k: int = 20,53 penalty: float = 1.2,54):55 print(f"message: {message}")56 print(f"history: {history}")57 58 conversation = [{"role": "system", "content": system_prompt}]59 for prompt, answer in history:60 conversation.extend(61 [62 {"role": "user", "content": prompt},63 {"role": "assistant", "content": answer},64 ]65 )66 67 conversation.append({"role": "user", "content": message})68 69 input_ids = tokenizer.apply_chat_template(70 conversation, add_generation_prompt=True, return_tensors="pt"71 ).to(model.device)72 73 streamer = TextIteratorStreamer(74 tokenizer, timeout=60.0, skip_prompt=True, skip_special_tokens=True75 )76 77 generate_kwargs = dict(78 input_ids=input_ids,79 max_new_tokens=max_new_tokens,80 do_sample=False if temperature == 0 else True,81 top_p=top_p,82 top_k=top_k,83 temperature=temperature,84 eos_token_id=[128001, 128008, 128009],85 streamer=streamer,86 )87 88 with torch.no_grad():89 thread = Thread(target=model.generate, kwargs=generate_kwargs)90 thread.start()91 92 buffer = ""93 for new_text in streamer:94 buffer += new_text95 yield buffer96 97 98"""99For information on how to customize the ChatInterface, peruse the gradio docs: https://www.gradio.app/docs/chatinterface100"""101demo = gr.ChatInterface(102 stream_chat,103 additional_inputs=[104 gr.Textbox(value="You are an ARM Assembly language decoder. You receive a line of Arm assembly and respond with a description of what the instruction does.", label="System message"),105 gr.Slider(minimum=0.1, maximum=4.0, value=0.7, step=0.1, label="Temperature"),106 gr.Slider(minimum=1, maximum=2048, value=512, step=1, label="Max new tokens"),107 gr.Slider(108 minimum=0.1,109 maximum=1.0,110 value=0.95,111 step=0.05,112 label="Top-p (nucleus sampling)",113 ),114 ],115)116 117 118if __name__ == "__main__":119 demo.launch()120 