Function Calling 推理示例函数(仅生成文本,不生成语音) Args: model_path: 模型路径 audio_path: 输入音频路径
(model_path, audio_path)
| 72 | |
| 73 | |
| 74 | def infer_function_calling_example(model_path, audio_path): |
| 75 | """ |
| 76 | Function Calling 推理示例函数(仅生成文本,不生成语音) |
| 77 | |
| 78 | Args: |
| 79 | model_path: 模型路径 |
| 80 | audio_path: 输入音频路径 |
| 81 | """ |
| 82 | config = AutoConfig.from_pretrained(model_path) |
| 83 | processor = AutoProcessor.from_pretrained(model_path) |
| 84 | model = AutoModelForSeq2SeqLM.from_pretrained(model_path, config=config, torch_dtype=torch.bfloat16, |
| 85 | device_map=device) |
| 86 | |
| 87 | # 生成参数 |
| 88 | model.sp_gen_kwargs.update({ |
| 89 | 'text_greedy': True, |
| 90 | 'disable_speech': True, |
| 91 | }) |
| 92 | |
| 93 | # 构建audio样例 |
| 94 | audio = [librosa.load(audio_path, sr=16000)[0]] |
| 95 | |
| 96 | example_tools = [ |
| 97 | {"type": "function", |
| 98 | "function": {"name": "get_weather", "description": "查询天气", |
| 99 | "parameters": {"type": "object", "properties": { |
| 100 | "location": {"type": "string", "description": "地点", "default": "当前位置"}, |
| 101 | "time": {"type": "string", "description": "时间", "default": "当前时间"}},"required": []}}}, |
| 102 | {"type": "function", |
| 103 | "function": {"name": "check_battery", "description": "电量查询,例如:现在还剩多少电", |
| 104 | "parameters": {"type": "object", "properties": {}, "required": []}}} |
| 105 | ] |
| 106 | |
| 107 | example_tools_definition = "\n".join([json.dumps(tool_item, ensure_ascii=False) for tool_item in example_tools]) |
| 108 | system_prompt = FUNCTION_CALLING_PROMPT.replace("{tools_definition}", example_tools_definition) |
| 109 | conversation = [ |
| 110 | {"role": "system", "content": system_prompt}, |
| 111 | {"role": "user", "content": AUDIO_TEMPLATE}, |
| 112 | ] |
| 113 | |
| 114 | text = processor.apply_chat_template(conversation, add_generation_prompt=True, tokenize=False) |
| 115 | inputs = processor(text=text, audio=audio, return_tensors="pt", return_token_type_ids=False).to(model.device) |
| 116 | generate_ids, _ = model.generate(**inputs) |
| 117 | generate_ids = generate_ids[:, inputs.input_ids.size(1):] |
| 118 | generate_text = processor.decode(generate_ids[0], skip_special_tokens=True) |
| 119 | |
| 120 | print("generate_text: ", generate_text) |
| 121 | |
| 122 | |
| 123 | if __name__ == "__main__": |