(question: str, contexts: str=None, sample_num: int=3)
| 233 | return result, log |
| 234 | |
| 235 | async def plugin(question: str, contexts: str=None, sample_num: int=3): |
| 236 | # Create tasks for parallel execution |
| 237 | async def process_sample(): |
| 238 | # Get decompose result |
| 239 | decompose_args = {"contexts": contexts} if module == "multi-hop" else {} |
| 240 | decompose_result = await decompose(question, **decompose_args) |
| 241 | |
| 242 | # Separate independent and dependent sub-questions |
| 243 | independent_subqs = [sub_q for sub_q in decompose_result["sub-questions"] if len(sub_q["depend"]) == 0] |
| 244 | dependent_subqs = [sub_q for sub_q in decompose_result["sub-questions"] if sub_q not in independent_subqs] |
| 245 | |
| 246 | # Get contraction result |
| 247 | merging_args = { |
| 248 | "question": question, |
| 249 | "decompose_result": decompose_result, |
| 250 | "independent_subqs": independent_subqs, |
| 251 | "dependent_subqs": dependent_subqs |
| 252 | } |
| 253 | if module == "multi-hop": |
| 254 | merging_args["contexts"] = contexts |
| 255 | |
| 256 | contractd_thought, contractd_question, contraction_result = await merging(**merging_args) |
| 257 | |
| 258 | return { |
| 259 | "decompose_result": decompose_result, |
| 260 | "contractd_thought": contractd_thought, |
| 261 | "contractd_question": contractd_question, |
| 262 | "contraction_result": contraction_result |
| 263 | } |
| 264 | |
| 265 | # Execute all samples in parallel |
| 266 | tasks = [process_sample() for _ in range(sample_num)] |
| 267 | all_results = await asyncio.gather(*tasks) |
| 268 | |
| 269 | # Get direct result for original question |
| 270 | direct_args = (question, contexts) if module == "multi-hop" else (question,) |
| 271 | direct_result = await direct(*direct_args) |
| 272 | |
| 273 | # Get ensemble result from all contracted results plus direct result |
| 274 | all_responses = [direct_result["response"]] + [r["contraction_result"]["response"] for r in all_results] |
| 275 | ensemble_args = [question, all_responses] |
| 276 | if module == "multi-hop": |
| 277 | ensemble_args.append(contexts) |
| 278 | |
| 279 | ensemble_result = await ensemble(*ensemble_args) |
| 280 | ensemble_answer = ensemble_result.get("answer", "") |
| 281 | |
| 282 | # Calculate scores for each contracted result |
| 283 | scores = [] |
| 284 | token_counts = [] |
| 285 | |
| 286 | for result in all_results: |
| 287 | contraction_result = result["contraction_result"] |
| 288 | # Calculate score compared to ensemble answer |
| 289 | scores.append(score(contraction_result["answer"], ensemble_answer)) |
| 290 | |
| 291 | # Estimate token count for the response |
| 292 | token_counts.append(len(contraction_result.get("response", "").split())) |
no test coverage detected