| import gradio as gr |
| import torch |
| from transformers import AutoModelForCausalLM, AutoTokenizer, pipeline |
| import os |
|
|
| |
| if torch.cuda.is_available(): |
| print(f"Using GPU: {torch.cuda.get_device_name(0)}") |
| device = "cuda" |
| else: |
| print("GPU not available, using CPU") |
| device = "cpu" |
|
|
| |
| MODEL_ID = "Ansah-AI/E1-4BIT-GGUF" |
| MODEL_REVISION = "main" |
| |
| |
| def load_model(): |
| print(f"Loading model: {MODEL_ID}") |
| |
| |
| model = AutoModelForCausalLM.from_pretrained( |
| MODEL_ID, |
| revision=MODEL_REVISION, |
| device_map="auto", |
| load_in_4bit=True, |
| trust_remote_code=True |
| ) |
| |
| |
| tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, trust_remote_code=True) |
| |
| print("Model and tokenizer loaded successfully!") |
| |
| return model, tokenizer |
|
|
| |
| def generate_text(prompt, max_length=256, temperature=0.7, top_p=0.9, top_k=40): |
| |
| global model, tokenizer |
| |
| |
| print(f"Generating with parameters: max_length={max_length}, temp={temperature}, top_p={top_p}, top_k={top_k}") |
| |
| |
| text_generator = pipeline( |
| "text-generation", |
| model=model, |
| tokenizer=tokenizer, |
| device_map="auto" |
| ) |
| |
| |
| generation_config = { |
| "max_length": max_length, |
| "temperature": temperature, |
| "top_p": top_p, |
| "top_k": top_k, |
| "num_return_sequences": 1, |
| "do_sample": temperature > 0.1, |
| "pad_token_id": tokenizer.eos_token_id |
| } |
| |
| try: |
| result = text_generator( |
| prompt, |
| **generation_config |
| ) |
| |
| |
| return result[0]["generated_text"] |
| |
| except Exception as e: |
| return f"Error generating text: {str(e)}" |
|
|
| |
| def main(): |
| |
| global model, tokenizer |
| model, tokenizer = load_model() |
| |
| |
| with gr.Blocks(title="E1-4BIT-GGUF Model Interface") as demo: |
| gr.Markdown("# E1-4BIT-GGUF Model Interface") |
| gr.Markdown("Enter your prompt below to generate text using the Ansah-AI/E1-4BIT-GGUF model.") |
| |
| with gr.Row(): |
| with gr.Column(scale=4): |
| prompt_input = gr.Textbox( |
| label="Prompt", |
| placeholder="Enter your prompt here...", |
| lines=5 |
| ) |
| |
| with gr.Column(scale=1): |
| max_length = gr.Slider( |
| minimum=64, |
| maximum=2048, |
| value=256, |
| step=32, |
| label="Max Length" |
| ) |
| |
| temperature = gr.Slider( |
| minimum=0.1, |
| maximum=1.5, |
| value=0.7, |
| step=0.1, |
| label="Temperature" |
| ) |
| |
| top_p = gr.Slider( |
| minimum=0.1, |
| maximum=1.0, |
| value=0.9, |
| step=0.05, |
| label="Top P" |
| ) |
| |
| top_k = gr.Slider( |
| minimum=1, |
| maximum=100, |
| value=40, |
| step=1, |
| label="Top K" |
| ) |
| |
| generate_button = gr.Button("Generate") |
| output_text = gr.Textbox(label="Generated Text", lines=10) |
| |
| |
| generate_button.click( |
| fn=generate_text, |
| inputs=[prompt_input, max_length, temperature, top_p, top_k], |
| outputs=output_text |
| ) |
| |
| |
| gr.Examples( |
| examples=[ |
| ["Write a short story about a space explorer discovering a new planet."], |
| ["Explain quantum computing to a high school student."], |
| ["Create a recipe for a chocolate cake."] |
| ], |
| inputs=prompt_input |
| ) |
| |
| |
| demo.launch(share=True) |
|
|
| if __name__ == "__main__": |
| main() |