diff --git a/tests/test_runs/test_run_functions.py b/tests/test_runs/test_run_functions.py index 38128de30..68146b16c 100644 --- a/tests/test_runs/test_run_functions.py +++ b/tests/test_runs/test_run_functions.py @@ -224,7 +224,9 @@ def _assert_predictions_equal(self, predictions, predictions_prime): def _rerun_model_and_compare_predictions(self, run_id, model_prime, seed, create_task_obj): run = openml.runs.get_run(run_id) - # TODO: assert holdout task + # assert for holdout task + task = openml.tasks.get_task(run.task_id) + assert task.estimation_procedure["type"] == "holdout" # downloads the predictions of the old task file_id = run.output_files["predictions"] @@ -341,12 +343,18 @@ def _remove_random_state(flow): assert isinstance(run.dataset_id, int) # This is only a smoke check right now - # TODO add a few asserts here run._to_xml() + assert isinstance(run._to_xml(), str) + if run.trace is not None: # This is only a smoke check right now - # TODO add a few asserts here + assert run.trace.run_id == run.run_id + assert run.trace.trace_iterations is not None + for trace_iteration in run.trace.trace_iterations.values(): + assert (trace_iteration.setup_string is not None) != (trace_iteration.parameters is not None) + run.trace.trace_to_arff() + assert isinstance(run.trace, dict) # check arff output assert len(run.data_content) == num_instances @@ -394,13 +402,18 @@ def _remove_random_state(flow): downloaded = openml.runs.get_run(run_.run_id) assert "openml-python" in downloaded.tags - # TODO make sure that these attributes are instantiated when # downloading a run? Or make sure that the trace object is created when # running a flow on a task (and not only the arff object is created, # so that the two objects can actually be compared): # downloaded_run_trace = downloaded._generate_trace_arff_dict() # self.assertEqual(run_trace, downloaded_run_trace) + + downloaded_run_trace = downloaded.trace + assert downloaded_run_trace is not None + assert run.trace is not None + self.assertEqual(run.trace, downloaded_run_trace) + return run def _check_sample_evaluations(