Skip to content

Commit 415ef37

Browse files
Merge pull request AI-Hypercomputer#4276 from AI-Hypercomputer:xibin/ci
PiperOrigin-RevId: 940014403
2 parents f2917ee + e9b35a3 commit 415ef37

4 files changed

Lines changed: 83 additions & 6 deletions

File tree

.github/workflows/ci_pipeline.yml

Lines changed: 22 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -199,12 +199,31 @@ jobs:
199199
is_scheduled_run: ${{ github.event_name == 'schedule' }}
200200
maxtext_sha: ${{ needs.build_and_upload_maxtext_package.outputs.maxtext_sha }}
201201

202+
setup-pathways-parameters:
203+
name: Setup Pathways Parameters
204+
runs-on: ubuntu-latest
205+
outputs:
206+
total_workers: ${{ steps.set-params.outputs.total_workers }}
207+
worker_groups: ${{ steps.set-params.outputs.worker_groups }}
208+
steps:
209+
- id: set-params
210+
name: Set Pathways Worker Parameters
211+
run: |
212+
# TPU worker constants
213+
TPU_UNIT_TOTAL_WORKERS=2
214+
TPU_UNIT_WORKER_GROUPS='[1, 2]'
215+
216+
echo "total_workers=${TPU_UNIT_TOTAL_WORKERS}" >> "$GITHUB_OUTPUT"
217+
echo "worker_groups=${TPU_UNIT_WORKER_GROUPS}" >> "$GITHUB_OUTPUT"
218+
202219
maxtext_tpu_pathways_unit_tests:
203-
needs: build_and_upload_maxtext_package
220+
needs: [build_and_upload_maxtext_package, setup-pathways-parameters]
204221
if: needs.analyze_code_changes.outputs.run_tests == 'true'
205222
uses: ./.github/workflows/run_pathways_tests.yml
206223
strategy:
207224
fail-fast: false
225+
matrix:
226+
group: ${{ fromJSON(needs.setup-pathways-parameters.outputs.worker_groups) }}
208227
with:
209228
device_type: tpu
210229
device_name: v6e-4
@@ -217,6 +236,8 @@ jobs:
217236
container_resource_option: "--privileged"
218237
is_scheduled_run: ${{ github.event_name == 'schedule' }}
219238
maxtext_sha: ${{ needs.build_and_upload_maxtext_package.outputs.maxtext_sha }}
239+
total_workers: ${{ needs.setup-pathways-parameters.outputs.total_workers }}
240+
worker_group: ${{ matrix.group }}
220241

221242
maxtext_tpu_pathways_integration_tests:
222243
needs: build_and_upload_maxtext_package

.github/workflows/run_pathways_tests.yml

Lines changed: 16 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -53,6 +53,15 @@ on:
5353
maxtext_sha:
5454
required: true
5555
type: string
56+
total_workers:
57+
required: false
58+
type: string
59+
default: '1'
60+
worker_group:
61+
required: false
62+
type: string
63+
default: '1'
64+
5665

5766
permissions:
5867
contents: read
@@ -95,6 +104,12 @@ jobs:
95104
run : gcloud storage cp gs://maxtext-test-assets/* tests/assets
96105
- name: Run Tests
97106
run: |
107+
if [ "${{ inputs.total_workers }}" -gt 1 ]; then
108+
.venv/bin/python3 -m pip install --quiet pytest-split
109+
SPLIT_ARGS="--splits ${{ inputs.total_workers }} --group ${{ inputs.worker_group }}"
110+
else
111+
SPLIT_ARGS=""
112+
fi
98113
if [ "${{ inputs.is_scheduled_run }}" = "true" ]; then
99114
FINAL_PYTEST_MARKER="${{ inputs.pytest_marker }}"
100115
else
@@ -105,7 +120,7 @@ jobs:
105120
export MAXTEXT_TEST_ASSETS_ROOT=$(pwd)/tests/assets
106121
export MAXTEXT_PKG_DIR=$(pwd)/src/maxtext
107122
# TODO(b/454659463): Enable test_default_hlo_match after volume mount is supported.
108-
.venv/bin/python3 -m pytest ${{ inputs.pytest_addopts }} -v -m "${FINAL_PYTEST_MARKER}" -k "not AotHloIdenticalTest and not CompileThenLoad" --durations=0
123+
.venv/bin/python3 -m pytest ${{ inputs.pytest_addopts }} -v -m "${FINAL_PYTEST_MARKER}" -k "not AotHloIdenticalTest and not CompileThenLoad" --durations=0 ${SPLIT_ARGS}
109124
env:
110125
PYTHONPATH: "${{ github.workspace }}/src"
111126
services:

.github/workflows/run_tests_against_package.yml

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -176,8 +176,13 @@ jobs:
176176
fi
177177
fi
178178
if [ "${INPUTS_TOTAL_WORKERS}" -gt 1 ]; then
179-
$PYTHON_EXE -m pip install --quiet pytest-split pytest-xdist
180-
SPLIT_ARGS="--splits ${INPUTS_TOTAL_WORKERS} --group ${INPUTS_WORKER_GROUP} -n auto"
179+
$PYTHON_EXE -m pip install --quiet pytest-split
180+
if [ "${INPUTS_DEVICE_TYPE}" = "tpu" ]; then
181+
SPLIT_ARGS="--splits ${INPUTS_TOTAL_WORKERS} --group ${INPUTS_WORKER_GROUP}"
182+
else
183+
$PYTHON_EXE -m pip install --quiet pytest-xdist
184+
SPLIT_ARGS="--splits ${INPUTS_TOTAL_WORKERS} --group ${INPUTS_WORKER_GROUP} -n auto"
185+
fi
181186
else
182187
SPLIT_ARGS=""
183188
fi

.github/workflows/run_tests_coordinator.yml

Lines changed: 38 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,12 +61,48 @@ permissions:
6161
contents: read
6262

6363
jobs:
64+
setup-parameters:
65+
name: Setup Parameters
66+
runs-on: ubuntu-latest
67+
outputs:
68+
worker_groups: ${{ steps.set-params.outputs.worker_groups }}
69+
total_workers: ${{ steps.set-params.outputs.total_workers }}
70+
steps:
71+
- id: set-params
72+
name: Set Worker Parameters
73+
run: |
74+
# CPU worker constants
75+
CPU_UNIT_TOTAL_WORKERS=4
76+
CPU_UNIT_WORKER_GROUPS='[1, 2, 3, 4]'
77+
78+
# TPU worker constants
79+
TPU_UNIT_TOTAL_WORKERS=2
80+
TPU_UNIT_WORKER_GROUPS='[1, 2]'
81+
82+
# Default fallback constants
83+
DEFAULT_TOTAL_WORKERS=1
84+
DEFAULT_WORKER_GROUPS='[1]'
85+
86+
FLAVOR="${{ inputs.flavor }}"
87+
88+
if [[ "$FLAVOR" == "cpu-unit" || "$FLAVOR" == "cpu-post-training-unit" ]]; then
89+
echo "worker_groups=${CPU_UNIT_WORKER_GROUPS}" >> "$GITHUB_OUTPUT"
90+
echo "total_workers=${CPU_UNIT_TOTAL_WORKERS}" >> "$GITHUB_OUTPUT"
91+
elif [[ "$FLAVOR" == "tpu-unit" ]]; then
92+
echo "worker_groups=${TPU_UNIT_WORKER_GROUPS}" >> "$GITHUB_OUTPUT"
93+
echo "total_workers=${TPU_UNIT_TOTAL_WORKERS}" >> "$GITHUB_OUTPUT"
94+
else
95+
echo "worker_groups=${DEFAULT_WORKER_GROUPS}" >> "$GITHUB_OUTPUT"
96+
echo "total_workers=${DEFAULT_TOTAL_WORKERS}" >> "$GITHUB_OUTPUT"
97+
fi
98+
6499
execute-test-package:
100+
needs: setup-parameters
65101
name: ${{ inputs.flavor }}
66102
strategy:
67103
fail-fast: false
68104
matrix:
69-
worker_group: ${{ fromJSON(contains(inputs.flavor, 'cpu-unit') && '[1, 2, 3, 4]' || '[1]') }}
105+
worker_group: ${{ fromJSON(needs.setup-parameters.outputs.worker_groups) }}
70106

71107
uses: ./.github/workflows/run_tests_against_package.yml
72108
with:
@@ -158,6 +194,6 @@ jobs:
158194
is_scheduled_run: ${{ inputs.is_scheduled_run }}
159195
maxtext_installed: ${{ inputs.maxtext_installed }}
160196
worker_group: ${{ matrix.worker_group }}
161-
total_workers: ${{ contains(inputs.flavor, 'cpu-unit') && 4 || 1 }}
197+
total_workers: ${{ fromJSON(needs.setup-parameters.outputs.total_workers) }}
162198
maxtext_sha: ${{ inputs.maxtext_sha }}
163199
is_update_hlo: ${{ inputs.is_update_hlo }}

0 commit comments

Comments
 (0)