Spaces:
Running
Running
| import torch | |
| from transformers import AutoTokenizer, AutoModelForCausalLM | |
| import gradio as gr | |
| MODEL_ID = "finnianx/Gros-Michel-Instruct" | |
| tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) | |
| dtype = torch.float32 | |
| model = AutoModelForCausalLM.from_pretrained(MODEL_ID, torch_dtype=dtype).eval() | |
| css = """ | |
| .gradio-container footer { | |
| display: none !important; | |
| } | |
| h1 { | |
| color: white !important; | |
| font-size: 28px; | |
| font-weight: 600; | |
| } | |
| h1::after { | |
| content: "Gros-Michel Instruct 🍌"; | |
| color: yellow; | |
| } | |
| :root{ | |
| --bg:#000;--fg:#fff;--yellow:#f5d742;--muted:#555; | |
| --border:#222;--surface:#0a0a0a; | |
| } | |
| body{ | |
| font-family:'SF Mono','Fira Code','Cascadia Code',monospace; | |
| background:var(--bg);color:var(--fg);min-height:100vh; | |
| } | |
| button { | |
| background-color: yellow !important; | |
| color: black !important; | |
| border-radius: 5px !important; | |
| border: 1px solid transparent !important; | |
| transition: all 0.2s ease; | |
| } | |
| button:hover { | |
| background-color: black !important; | |
| color: yellow !important; | |
| border: 1px solid white !important; | |
| } | |
| #input{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| #input:focus{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| #input:focus-within{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| .gradio-container textarea{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| .gradio-container textarea:focus-within{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| .gradio-container textarea:focus{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| textarea:focus{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| textarea{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| #output{ | |
| background-color: var(--bg); | |
| color:var(--fg); | |
| } | |
| .gradio-container .examples .example, | |
| .gradio-container .dataset .example { | |
| font-size: 17px !important; | |
| padding: 14px 18px !important; | |
| min-height: 50px !important; | |
| line-height: 1.4 !important; | |
| } | |
| .gradio-container .examples, | |
| .gradio-container .dataset { | |
| gap: 10px !important; | |
| } | |
| #input, | |
| #input > div { | |
| background-color: var(--bg) !important; | |
| border: 1px solid var(--border) !important; | |
| } | |
| #input textarea { | |
| background-color: var(--bg) !important; | |
| color: var(--fg) !important; | |
| } | |
| #input:focus-within, | |
| #input > div:focus-within { | |
| background-color: var(--bg) !important; | |
| box-shadow: none !important; | |
| border: 1px solid var(--yellow) !important; /* optional highlight */ | |
| } | |
| #input textarea:focus { | |
| outline: none !important; | |
| box-shadow: none !important; | |
| background-color: var(--bg) !important; | |
| } | |
| #output, | |
| #output > div { | |
| background-color: var(--bg) !important; | |
| border: 1px solid var(--border) !important; | |
| } | |
| #output textarea { | |
| background-color: var(--bg) !important; | |
| color: var(--fg) !important; | |
| } | |
| /* stop grey focus */ | |
| #output:focus-within, | |
| #output > div:focus-within { | |
| background-color: var(--bg) !important; | |
| box-shadow: none !important; | |
| border: 1px solid var(--border) !important; | |
| } | |
| #output textarea:focus { | |
| outline: none !important; | |
| box-shadow: none !important; | |
| background-color: var(--bg) !important; | |
| } | |
| /* kill Gradio loading grey overlay */ | |
| #output[data-loading="true"], | |
| #output[data-loading="true"] > div { | |
| background-color: var(--bg) !important; | |
| } | |
| /* sometimes applied as a class */ | |
| #output .loading, | |
| #output .generating { | |
| background-color: var(--bg) !important; | |
| } | |
| /* remove shimmer effect */ | |
| #output [class*="loading"] { | |
| background: none !important; | |
| animation: none !important; | |
| } | |
| """ | |
| if torch.cuda.is_available(): | |
| model.to("cuda") | |
| def generate(prompt_text): | |
| messages = [{"role": "user", "content": prompt_text}] | |
| text = tokenizer.apply_chat_template( | |
| messages, | |
| add_generation_prompt=True, | |
| tokenize=False, | |
| ) | |
| input_ids = tokenizer(text, return_tensors="pt").input_ids.to(model.device) | |
| with torch.no_grad(): | |
| output_ids = model.generate( | |
| input_ids, | |
| max_new_tokens=512, | |
| do_sample=True, | |
| temperature=0.5, | |
| top_p=0.9, | |
| repetition_penalty=1.2, | |
| pad_token_id=tokenizer.pad_token_id, | |
| eos_token_id=tokenizer.eos_token_id, | |
| ) | |
| prompt_len = input_ids.shape[1] | |
| generated_text = tokenizer.decode(output_ids[0][prompt_len:], skip_special_tokens=True).strip() | |
| return generated_text | |
| demo = gr.Interface( | |
| fn=generate, | |
| css=css, | |
| inputs = gr.Textbox(lines=5, label="Input", elem_id="input"), | |
| outputs = gr.Textbox(lines=10, label="Output", elem_id="output"), | |
| title="Chat with ", | |
| examples=[ | |
| ["Write a haiku about bananas"], | |
| ["Write a short story about a robot"], | |
| ["What is the capital of spain?"] | |
| ] | |
| ) | |
| demo.launch() |