medmekk/BitNet.cpp
18
1import gradio as gr2import subprocess3import os4import time5from transformers import AutoTokenizer, AutoModelForCausalLM6import logging7from starlette.middleware.sessions import SessionMiddleware8 9 10# Configure logging11logging.basicConfig(level=logging.INFO)12 13# Path to the cloned repository14BITNET_REPO_PATH = "/home/user/app/BitNet"15SETUP_SCRIPT = os.path.join(BITNET_REPO_PATH, "setup_env.py")16INFERENCE_SCRIPT = os.path.join(BITNET_REPO_PATH, "run_inference.py")17 18# Function to set up the environment by running setup.py19def setup_bitnet(model_name):20 try:21 result = subprocess.run(22 f"python {SETUP_SCRIPT} --hf-repo {model_name} -q i2_s",23 shell=True,24 cwd=BITNET_REPO_PATH,25 capture_output=True,26 text=True27 )28 if result.returncode == 0:29 return "Setup completed successfully!"30 else:31 return f"Error in setup: {result.stderr}"32 except Exception as e:33 return str(e)34 35# Function to run inference using the `run_inference.py` file36def run_inference(model_name, input_text, num_tokens=6):37 try:38 # Call the `run_inference.py` script with the model and input39 40 model_name = model_name.split("/")[1]41 start_time = time.time()42 if input_text is None or input_text == "": 43 return "Please provide an input text for the model"44 result = subprocess.run(45 f"python run_inference.py -m models/{model_name}/ggml-model-i2_s.gguf -p \"{input_text}\" -n {num_tokens} -temp 0",46 shell=True,47 cwd=BITNET_REPO_PATH,48 capture_output=True,49 text=True50 )51 end_time = time.time()52 53 if result.returncode == 0:54 inference_time = round(end_time - start_time, 2)55 return result.stdout, f"Inference took {inference_time} seconds."56 else:57 return f"Error during inference: {result.stderr}", None58 except Exception as e:59 return str(e), None60 61def run_transformers(model_name, input_text, num_tokens):62 63 # if oauth_token is None : 64 # return "Error : To Compare please login to your HF account and make sure you have access to the used Llama models"65 # Load the model and tokenizer dynamically if needed (commented out for performance)66 # if model_name=="TinyLlama/TinyLlama-1.1B-Chat-v1.0" : 67 print(input_text)68 if input_text is None or input_text == "": 69 return "Please provide an input text for the model", None70 tokenizer = AutoTokenizer.from_pretrained(model_name)71 model = AutoModelForCausalLM.from_pretrained(model_name)72 73 # Encode the input text74 input_ids = tokenizer.encode(input_text, return_tensors="pt")75 76 # Start time for inference77 start_time = time.time()78 79 # Generate output with the specified number of tokens80 output = model.generate(input_ids, max_length=len(input_ids[0]) + num_tokens, num_return_sequences=1)81 82 # Calculate inference time83 inference_time = time.time() - start_time84 85 # Decode the generated output86 generated_text = tokenizer.decode(output[0], skip_special_tokens=True)87 88 return generated_text, f"{inference_time:.2f} seconds"89 90# Gradio Interface91def interface():92 with gr.Blocks(theme=gr.themes.Ocean()) as demo:93 94 # Header95 gr.Markdown(96 """97 <h1 style="text-align: center; color: #7AB8E5;">BitNet.cpp Speed Demonstration 💻</h1>98 <p style="text-align: center; color: #6A1B9A;">Compare the speed and performance of BitNet with popular Transformer models.</p>99 """,100 elem_id="header"101 )102 103 # Instructions104 gr.Markdown(105 """106 ### Instructions for Using the BitNet.cpp Speed Demonstration107 1. **Set Up Your Project**: Begin by selecting the model you wish to use. Please note that this process may take a few minutes to complete.108 2. **Select Token Count**: Choose the number of tokens you want to generate for your inference.109 3. **Input Your Text**: Enter the text you wish to analyze, then compare the performance of BitNet with popular Transformer models.110 """,111 elem_id="instructions"112 )113 114 # Model Selection and Setup115 with gr.Column(elem_id="container"):116 gr.Markdown("<h2 style='color: #5CA2D3; text-align: center;'>Model Selection and Setup</h2>")117 with gr.Row():118 model_dropdown = gr.Dropdown(119 label="Select Model",120 choices=[121 "HF1BitLLM/Llama3-8B-1.58-100B-tokens", 122 "1bitLLM/bitnet_b1_58-3B", 123 "1bitLLM/bitnet_b1_58-large"124 ],125 value="HF1BitLLM/Llama3-8B-1.58-100B-tokens",126 interactive=True127 )128 setup_button = gr.Button("Run Setup")129 setup_status = gr.Textbox(label="Setup Status", interactive=False, placeholder="Setup status will appear here...")130 131 # Inference Section132 with gr.Column(elem_id="container"):133 gr.Markdown("<h2 style='color: #5CA2D3; text-align: center;'>BitNet Inference</h2>")134 with gr.Row():135 num_tokens = gr.Slider(136 minimum=1, maximum=100, 137 label="Number of Tokens to Generate", 138 value=50, step=1139 )140 input_text = gr.Textbox(141 label="Input Text", 142 placeholder="Enter your input text here...",143 value="Who is Zeus?"144 )145 with gr.Row():146 infer_button = gr.Button("Run Inference")147 result_output = gr.Textbox(label="Output", interactive=False, placeholder="Inference output will appear here...")148 time_output = gr.Textbox(label="Inference Time", interactive=False, placeholder="Inference time will appear here...")149 150 # Comparison with Transformers Section151 with gr.Column(elem_id="container"):152 gr.Markdown("<h2 style='color: #5CA2D3; text-align: center;'>Compare with Transformers</h2>")153 with gr.Row():154 transformer_model_dropdown = gr.Dropdown(155 label="Select Transformers Model",156 choices=["TinyLlama/TinyLlama_v1.1"],157 value="TinyLlama/TinyLlama_v1.1",158 interactive=True159 )160 input_text_tr = gr.Textbox(label="Input Text", placeholder="Enter your input text here...", value="Who is Zeus?")161 with gr.Row():162 compare_button = gr.Button("Run Transformers Inference")163 transformer_result_output = gr.Textbox(label="Transformers Output", interactive=False, placeholder="Transformers output will appear here...")164 transformer_time_output = gr.Textbox(label="Transformers Inference Time", interactive=False, placeholder="Transformers inference time will appear here...")165 166 # Actions167 setup_button.click(setup_bitnet, inputs=model_dropdown, outputs=setup_status)168 infer_button.click(run_inference, inputs=[model_dropdown, input_text, num_tokens], outputs=[result_output, time_output])169 compare_button.click(run_transformers, inputs=[transformer_model_dropdown, input_text_tr, num_tokens], outputs=[transformer_result_output, transformer_time_output])170 171 return demo172 173demo = interface()174 175# # Access FastAPI app instance from Gradio176# fastapi_app = demo.app 177 178# # Add SessionMiddleware to enable session management179# fastapi_app.add_middleware(SessionMiddleware, secret_key="secret_key") # Use a secure, random secret key180 181# # Launch the app182demo.launch()183 184# from fastapi import FastAPI185 186# app = FastAPI()187 188# # Add SessionMiddleware for sessions handling189# app.add_middleware(SessionMiddleware, secret_key="secure_secret_key")190 191# # Mount Gradio app to FastAPI at the root192# app.mount("/", demo)