Skip to content

Commit 116a0a6

Browse files
vertex-sdk-botcopybara-github
authored andcommitted
feat: Populate task_unique_name from initial pipeline run in Pipeline Task Rerun Configs for pipeline job rerun
PiperOrigin-RevId: 781659183
1 parent 6e5c421 commit 116a0a6

2 files changed

Lines changed: 116 additions & 1 deletion

File tree

google/cloud/aiplatform/preview/pipelinejob/pipeline_jobs.py

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -578,12 +578,34 @@ def rerun(
578578
pipeline_job.original_pipeline_job_id = int(
579579
original_pipeline_job.labels["vertex-ai-pipelines-run-billing-id"]
580580
)
581+
original_pipeline_task_details = (
582+
original_pipeline_job.job_detail.task_details
583+
)
581584
except Exception as e:
582585
raise ValueError(
583586
f"Failed to get original pipeline job: {original_pipelinejob_name}"
584587
) from e
585588

586-
pipeline_job.pipeline_task_rerun_configs = pipeline_task_rerun_configs
589+
task_id_to_task_rerun_config = {}
590+
for task_rerun_config in pipeline_task_rerun_configs:
591+
task_id_to_task_rerun_config[task_rerun_config.task_id] = task_rerun_config
592+
593+
pipeline_job.pipeline_task_rerun_configs = []
594+
for task_detail in original_pipeline_task_details:
595+
if task_detail.task_id in task_id_to_task_rerun_config:
596+
task_rerun_config = task_id_to_task_rerun_config[task_detail.task_id]
597+
if task_detail.task_unique_name:
598+
task_rerun_config.task_name = task_detail.task_unique_name
599+
pipeline_job.pipeline_task_rerun_configs.append(task_rerun_config)
600+
else:
601+
pipeline_job.pipeline_task_rerun_configs.append(
602+
aiplatform_v1beta1.PipelineTaskRerunConfig(
603+
task_id=task_detail.task_id,
604+
task_name=task_detail.task_unique_name,
605+
skip_task=task_detail.state
606+
== aiplatform_v1beta1.PipelineTaskDetail.State.SUCCEEDED,
607+
)
608+
)
587609

588610
if parameter_values:
589611
runtime_config = self._v1_beta1_pipeline_job.runtime_config

tests/unit/aiplatform/test_pipeline_jobs.py

Lines changed: 93 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -74,6 +74,7 @@
7474
_TEST_PIPELINE_JOB_DISPLAY_NAME_2 = "sample-pipeline-job-display-name-2"
7575
_TEST_PIPELINE_JOB_ID = "sample-test-pipeline-202111111"
7676
_TEST_PIPELINE_JOB_ID_2 = "sample-test-pipeline-202111112"
77+
_TEST_PIPELINE_RERUN_JOB_ID = "sample-test-pipeline-rerun"
7778
_TEST_GCS_BUCKET_NAME = "my-bucket"
7879
_TEST_GCS_OUTPUT_DIRECTORY = f"gs://{_TEST_GCS_BUCKET_NAME}/output_artifacts/"
7980
_TEST_CREDENTIALS = auth_credentials.AnonymousCredentials()
@@ -339,6 +340,41 @@ def mock_pipeline_v1beta1_service_get():
339340
yield mock_get_pipeline_job
340341

341342

343+
@pytest.fixture
344+
def mock_pipeline_v1beta1_service_get_with_task_details():
345+
pipeline_job = make_v1beta1_pipeline_job(
346+
_TEST_PIPELINE_JOB_NAME,
347+
gca_pipeline_state.PipelineState.PIPELINE_STATE_FAILED,
348+
)
349+
pipeline_job.job_detail.task_details = [
350+
v1beta1_pipeline_job.PipelineTaskDetail(
351+
task_id=1,
352+
task_name="task-1",
353+
state=aiplatform_v1beta1.PipelineTaskDetail.State.SUCCEEDED,
354+
task_unique_name="task-1",
355+
),
356+
v1beta1_pipeline_job.PipelineTaskDetail(
357+
task_id=2,
358+
task_name="task-2",
359+
state=aiplatform_v1beta1.PipelineTaskDetail.State.FAILED,
360+
task_unique_name="task-2",
361+
),
362+
v1beta1_pipeline_job.PipelineTaskDetail(
363+
task_id=3,
364+
task_name="task-3",
365+
state=aiplatform_v1beta1.PipelineTaskDetail.State.SUCCEEDED,
366+
task_unique_name="task-3",
367+
),
368+
]
369+
370+
with mock.patch.object(
371+
v1beta1_pipeline_service.PipelineServiceClient, "get_pipeline_job"
372+
) as mock_get_pipeline_job:
373+
mock_get_pipeline_job.side_effect = [pipeline_job]
374+
375+
yield mock_get_pipeline_job
376+
377+
342378
@pytest.fixture
343379
def mock_pipeline_v1_service_batch_cancel():
344380
with patch.object(
@@ -2417,6 +2453,63 @@ def test_rerun_v1beta1_pipeline_job_returns_response(
24172453
assert mock_pipeline_v1beta1_service_get.call_count == 1
24182454
assert mock_pipeline_v1beta1_service_create.call_count == 2
24192455

2456+
@pytest.mark.usefixtures(
2457+
"mock_pipeline_v1beta1_service_create",
2458+
"mock_pipeline_v1beta1_service_get_with_task_details",
2459+
)
2460+
@pytest.mark.parametrize(
2461+
"job_spec",
2462+
[_TEST_PIPELINE_SPEC_JSON, _TEST_PIPELINE_SPEC_YAML, _TEST_PIPELINE_JOB],
2463+
)
2464+
def test_rerun_v1beta1_pipeline_job_returns_response_with_task_details(
2465+
self,
2466+
mock_load_yaml_and_json,
2467+
job_spec,
2468+
mock_pipeline_v1beta1_service_create,
2469+
mock_pipeline_v1beta1_service_get_with_task_details,
2470+
):
2471+
aiplatform.init(
2472+
project=_TEST_PROJECT,
2473+
staging_bucket=_TEST_GCS_BUCKET_NAME,
2474+
credentials=_TEST_CREDENTIALS,
2475+
)
2476+
2477+
job = preview_pipeline_jobs._PipelineJob(
2478+
display_name=_TEST_PIPELINE_JOB_DISPLAY_NAME,
2479+
template_path=_TEST_TEMPLATE_PATH,
2480+
job_id=_TEST_PIPELINE_JOB_ID,
2481+
)
2482+
2483+
job.submit()
2484+
2485+
job.rerun(
2486+
job_id=_TEST_PIPELINE_RERUN_JOB_ID,
2487+
original_pipelinejob_name=_TEST_PIPELINE_JOB_NAME,
2488+
pipeline_task_rerun_configs=[
2489+
aiplatform_v1beta1.PipelineTaskRerunConfig(
2490+
task_id=1,
2491+
skip_task=False,
2492+
)
2493+
],
2494+
parameter_values={"param-1": "value-1"},
2495+
)
2496+
2497+
assert mock_pipeline_v1beta1_service_get_with_task_details.call_count == 1
2498+
assert mock_pipeline_v1beta1_service_create.call_count == 2
2499+
2500+
rerun_config = mock_pipeline_v1beta1_service_create.call_args[1][
2501+
"request"
2502+
].pipeline_job.pipeline_task_rerun_configs
2503+
assert rerun_config[0].task_name == "task-1"
2504+
assert not rerun_config[0].skip_task
2505+
assert rerun_config[0].task_id == 1
2506+
assert rerun_config[1].task_name == "task-2"
2507+
assert not rerun_config[1].skip_task
2508+
assert rerun_config[1].task_id == 2
2509+
assert rerun_config[2].task_name == "task-3"
2510+
assert rerun_config[2].skip_task
2511+
assert rerun_config[2].task_id == 3
2512+
24202513
@pytest.mark.parametrize(
24212514
"job_spec",
24222515
[_TEST_PIPELINE_SPEC_JSON, _TEST_PIPELINE_SPEC_YAML, _TEST_PIPELINE_JOB],

0 commit comments

Comments
 (0)