(data: list[dict], model, processor, label_tree: LabelTree)
| 107 | return None |
| 108 | |
| 109 | def 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) |
no test coverage detected