(dataset: str, model: str, start: int = 0, end: int = -1)
| 189 | |
| 190 | |
| 191 | async def optimize_dataset(dataset: str, model: str, start: int = 0, end: int = -1): |
| 192 | # Optimize dataset questions and save to new file |
| 193 | print(f"Optimizing {dataset} dataset questions from index {start} to {end}") |
| 194 | timestamp = time.time() |
| 195 | |
| 196 | # Set model and module |
| 197 | set_model(model) |
| 198 | config = DATASET_CONFIGS[dataset] |
| 199 | set_module(config.module_type) |
| 200 | |
| 201 | # Load test set |
| 202 | testset = load_data(dataset, "test")[start:None if end == -1 else end] |
| 203 | question_key = config.question_key |
| 204 | if isinstance(question_key, list): |
| 205 | question_key = question_key[0] |
| 206 | |
| 207 | # Create tasks |
| 208 | async def process_item(item): |
| 209 | try: |
| 210 | if config.requires_context(): |
| 211 | from experiment.prompter.multihop import contexts |
| 212 | optimized_question = await plugin(item[question_key], contexts(item, dataset)) |
| 213 | else: |
| 214 | optimized_question = await plugin(item[question_key]) |
| 215 | |
| 216 | # Create new entry |
| 217 | new_item = item.copy() |
| 218 | new_item["original_question"] = item[question_key] |
| 219 | new_item[question_key] = optimized_question |
| 220 | return new_item |
| 221 | except Exception as e: |
| 222 | print(f"Error processing item: {e}") |
| 223 | return item # Return original item on error |
| 224 | |
| 225 | # Process all items in parallel |
| 226 | tasks = [process_item(item) for item in testset] |
| 227 | optimized_data = await tqdm.gather(*tasks, desc=f"Optimizing {dataset} questions") |
| 228 | |
| 229 | # Ensure output directory exists |
| 230 | os.makedirs(f"experiment/data/{dataset}", exist_ok=True) |
| 231 | |
| 232 | # Save optimized dataset |
| 233 | output_path = f"experiment/data/{dataset}/contracted.json" |
| 234 | save_json(output_path, optimized_data) |
| 235 | |
| 236 | elapsed_time = time.time() - timestamp |
| 237 | print(f"Optimized dataset saved to {output_path}") |
| 238 | print(f"Time taken: {duration_formatter(elapsed_time)}") |
| 239 | |
| 240 | return optimized_data |
| 241 | |
| 242 | async def main(): |
| 243 | # Main function |
no test coverage detected