lucaspetti/function-gemma
0
1import os, json, gradio as gr, torch2from transformers import AutoProcessor, AutoModelForCausalLM3 4hf_token = os.getenv("HF_TOKEN")5model_id = "google/functiongemma-270m-it"6 7# 1. Load for CPU specifically8processor = AutoProcessor.from_pretrained(model_id, token=hf_token)9model = AutoModelForCausalLM.from_pretrained(10 model_id, 11 torch_dtype=torch.float32, # CPU prefers float3212 device_map={"": "cpu"}, # Forces everything onto CPU, avoiding "meta device"13 token=hf_token14)15 16def process_request(user_prompt, developer_prompt, tools_json):17 try:18 tools = json.loads(tools_json) if tools_json.strip() else []19 20 # FunctionGemma format21 messages = [22 {"role": "developer", "content": developer_prompt},23 {"role": "user", "content": user_prompt}24 ]25 26 # 2. Ensure inputs are on CPU27 inputs = processor.apply_chat_template(28 messages, tools=tools, add_generation_prompt=True, 29 return_dict=True, return_tensors="pt"30 ).to("cpu")31 32 with torch.no_grad():33 outputs = model.generate(**inputs, max_new_tokens=128, do_sample=False)34 35 input_len = inputs.input_ids.shape[1]36 decoded = processor.decode(outputs[0][input_len:], skip_special_tokens=True)37 return decoded if decoded.strip() else "Model returned an empty string."38 39 except Exception as e:40 return f"Error: {str(e)}"41 42demo = gr.Interface(43 fn=process_request,44 inputs=[gr.Textbox(label="User Prompt"), gr.Textbox(label="Developer Prompt"), gr.Textbox(label="Tools (JSON Array)")],45 outputs=gr.Code(label="Model Output"),46 title="FunctionGemma CPU Fixed"47)48 49demo.launch(share=True)