| 71 | self.num_split = getattr(config, 'num_split', self.num_split) |
| 72 | |
| 73 | async def run(self, messages: List[Message]): |
| 74 | if self.memory_called: |
| 75 | return messages |
| 76 | query = None |
| 77 | system = None |
| 78 | for message in messages: |
| 79 | if message.role == 'system': |
| 80 | system = message.content |
| 81 | if message.role == 'user': |
| 82 | query = message.content |
| 83 | if system is None: |
| 84 | system = query |
| 85 | break |
| 86 | |
| 87 | assert query is not None |
| 88 | arguments = [] |
| 89 | for n in range(self.num_split): |
| 90 | inputs = { |
| 91 | 'system': self.div_system1, |
| 92 | 'query': query, |
| 93 | } |
| 94 | arguments.append(inputs) |
| 95 | |
| 96 | arguments = { |
| 97 | 'tasks': arguments, |
| 98 | 'execution_mode': 'sequential', |
| 99 | } |
| 100 | |
| 101 | results = await self.split_task.call_tool( |
| 102 | '', tool_name='', tool_args=arguments) |
| 103 | pattern = r'<result>(.*?)</result>' |
| 104 | all_keywords = [] |
| 105 | for keywords in re.findall(pattern, results, re.DOTALL): |
| 106 | all_keywords.extend([ |
| 107 | keyword.strip() for keyword in keywords.split(',') |
| 108 | if keyword.strip() |
| 109 | ]) |
| 110 | |
| 111 | arguments = [] |
| 112 | _query = ','.join(set(all_keywords)) |
| 113 | logger.info(f'Diversity first round keywords: {_query}') |
| 114 | for n in range(self.num_split): |
| 115 | inputs = { |
| 116 | 'system': self.div_system2, |
| 117 | 'query': _query, |
| 118 | } |
| 119 | arguments.append(inputs) |
| 120 | |
| 121 | arguments = { |
| 122 | 'tasks': arguments, |
| 123 | 'execution_mode': 'sequential', |
| 124 | } |
| 125 | |
| 126 | results = await self.split_task.call_tool( |
| 127 | '', tool_name='', tool_args=arguments) |
| 128 | pattern = r'<result>(.*?)</result>' |
| 129 | all_keywords = [] |
| 130 | for keywords in re.findall(pattern, results, re.DOTALL): |