MCPcopy Create free account
hub / github.com/AlayaLab/Hive / batch_query

Function batch_query

pipeline/code/05_leaf_label_qwen.py:109–199  ·  view source on GitHub ↗
(data: list[dict], model, processor, label_tree: LabelTree)

Source from the content-addressed store, hash-verified

107 return None
108
109def batch_query(data: list[dict], model, processor, label_tree: LabelTree):
110 results = [None] * len(data)
111 conversations = []
112
113 for index, audio_data in enumerate(data):
114 text_label = audio_data['text_label'][0]
115 audio = audio_data['audio_path']
116 leaf_nodes = label_tree.get_leaf_nodes_by_name(text_label)
117 if leaf_nodes == []:
118 root_leaf_node = label_tree.get_root_leaf_node_by_name(text_label)
119 audio_data_copy = audio_data.copy()
120 audio_data_copy['text_label'] = [root_leaf_node]
121 results[index] = audio_data_copy
122 continue
123 if len(leaf_nodes) == 1:
124 results[index] = audio_data
125 continue
126 leaf_labels = leaf_nodes
127 prompt = f"Please analyze this audio and determine which type of sound it contains from {leaf_labels}. After analysis, please return only the index of the corresponding label (the minimum index value is 0). Please note that you should only return a single number, without any other symbols. At the same time, if you believe the content of the audio does not match any of the sounds described above, please return -1. Here is the index value you return:"
128 conversation = [
129 {
130 "role": "user",
131 "content": [
132 {"type": "audio", "audio": audio},
133 {"type": "text", "text": {prompt}}
134 ],
135 }
136 ]
137 conversations.append(conversation)
138
139 def process_text(texts: list[str]):
140 results = []
141 for text in texts:
142 parts = text.split('assistant\n')
143 if len(parts) != 2:
144 return None
145 try:
146 index = int(parts[1])
147 results.append(index)
148 except ValueError:
149 logger.info(f"Current model output: {texts}")
150 return None
151 return results
152
153 if not conversations:
154 return results,[]
155
156 USE_AUDIO_IN_VIDEO = True
157 text = processor.apply_chat_template(conversations, add_generation_prompt=True, tokenize=False)
158 audios, images, videos = process_mm_info(conversations, use_audio_in_video=USE_AUDIO_IN_VIDEO)
159 inputs = processor(text=text, audio=audios, images=images, videos=videos, return_tensors="pt", padding=True, use_audio_in_video=USE_AUDIO_IN_VIDEO)
160 inputs = inputs.to(model.device).to(model.dtype)
161
162 text_ids = model.generate(**inputs, use_audio_in_video=USE_AUDIO_IN_VIDEO, return_audio=False)
163
164 text = processor.batch_decode(text_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)
165
166 middle_results = process_text(text)

Callers 1

Calls 5

process_textFunction · 0.70
appendMethod · 0.45
toMethod · 0.45

Tested by

no test coverage detected