MCPcopy Create free account
hub / github.com/InternScience/InternAgent / optimize_dataset

Function optimize_dataset

tasks/AutoTTS/code/main.py:191–240  ·  view source on GitHub ↗
(dataset: str, model: str, start: int = 0, end: int = -1)

Source from the content-addressed store, hash-verified

189
190
191async 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
242async def main():
243 # Main function

Callers 1

mainFunction · 0.70

Calls 6

set_modelFunction · 0.90
set_moduleFunction · 0.90
load_dataFunction · 0.90
save_jsonFunction · 0.90
duration_formatterFunction · 0.90
process_itemFunction · 0.70

Tested by

no test coverage detected