MCPcopy Create free account
hub / github.com/FunAudioLLM/Fun-Audio-Chat / infer_function_calling_example

Function infer_function_calling_example

examples/infer_s2t.py:74–120  ·  view source on GitHub ↗

Function Calling 推理示例函数(仅生成文本,不生成语音) Args: model_path: 模型路径 audio_path: 输入音频路径

(model_path, audio_path)

Source from the content-addressed store, hash-verified

72
73
74def 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
123if __name__ == "__main__":

Callers

nothing calls this directly

Calls 1

decodeMethod · 0.80

Tested by

no test coverage detected